Remove obsolete sgl-kernel legacy paths (#21528)
This commit is contained in:
@@ -1,55 +0,0 @@
|
||||
// Adapted from
|
||||
// https://github.com/flashinfer-ai/flashinfer/blob/55576c626421b5ee7e7ebe74afd26465c8ae863f/csrc/cascade.cu
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include <flashinfer/attention/cascade.cuh>
|
||||
|
||||
#include "pytorch_extension_utils.h"
|
||||
|
||||
using namespace flashinfer;
|
||||
|
||||
void merge_state(
|
||||
at::Tensor v_a, at::Tensor s_a, at::Tensor v_b, at::Tensor s_b, at::Tensor v_merged, at::Tensor s_merged) {
|
||||
CHECK_INPUT(v_a);
|
||||
CHECK_INPUT(s_a);
|
||||
CHECK_INPUT(v_b);
|
||||
CHECK_INPUT(s_b);
|
||||
auto device = v_a.device();
|
||||
CHECK_EQ(s_a.device(), device);
|
||||
CHECK_EQ(v_b.device(), device);
|
||||
CHECK_EQ(s_b.device(), device);
|
||||
CHECK_DIM(3, v_a);
|
||||
CHECK_DIM(2, s_a);
|
||||
CHECK_DIM(3, v_b);
|
||||
CHECK_DIM(2, s_b);
|
||||
CHECK_SHAPE(v_a, v_b);
|
||||
CHECK_SHAPE(s_a, s_b);
|
||||
CHECK_EQ(v_a.size(0), s_a.size(0));
|
||||
CHECK_EQ(v_a.size(1), s_b.size(1));
|
||||
unsigned int seq_len = v_a.size(0);
|
||||
unsigned int num_heads = v_a.size(1);
|
||||
unsigned int head_dim = v_a.size(2);
|
||||
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(v_a.device());
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
bool success = DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(v_a.scalar_type(), c_type, [&] {
|
||||
cudaError_t status = MergeState(
|
||||
static_cast<c_type*>(v_a.data_ptr()),
|
||||
static_cast<float*>(s_a.data_ptr()),
|
||||
static_cast<c_type*>(v_b.data_ptr()),
|
||||
static_cast<float*>(s_b.data_ptr()),
|
||||
static_cast<c_type*>(v_merged.data_ptr()),
|
||||
static_cast<float*>(s_merged.data_ptr()),
|
||||
seq_len,
|
||||
num_heads,
|
||||
head_dim,
|
||||
stream);
|
||||
TORCH_CHECK(status == cudaSuccess, "MergeState kernel launch failed: ", cudaGetErrorString(status));
|
||||
return true;
|
||||
});
|
||||
|
||||
TORCH_CHECK(success, "MergeState kernel launch failed: unsupported data type");
|
||||
}
|
||||
@@ -50,8 +50,6 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
/*
|
||||
* From csrc/attention
|
||||
*/
|
||||
m.def("merge_state(Tensor v_a, Tensor s_a, Tensor v_b, Tensor s_b, Tensor! v_merged, Tensor! s_merged) -> ()");
|
||||
m.impl("merge_state", torch::kCUDA, &merge_state);
|
||||
m.def("merge_state_v2(Tensor v_a, Tensor s_a, Tensor v_b, Tensor s_b, Tensor! v_merged, Tensor! s_merged) -> ()");
|
||||
m.impl("merge_state_v2", torch::kCUDA, &merge_state_v2);
|
||||
m.def(
|
||||
@@ -90,11 +88,6 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
" Tensor cos_sin_cache, bool is_neox) -> ()");
|
||||
m.impl("rotary_embedding", torch::kCUDA, &rotary_embedding);
|
||||
|
||||
m.def(
|
||||
"downcast_fp8(Tensor k, Tensor v, Tensor k_out, Tensor v_out, Tensor k_scale, Tensor v_scale, Tensor loc, "
|
||||
"int mult, int offset) -> ()");
|
||||
m.impl("downcast_fp8", torch::kCUDA, &downcast_fp8);
|
||||
|
||||
m.def("copy_to_gpu_no_ce(Tensor input, Tensor! output) -> ()");
|
||||
m.impl("copy_to_gpu_no_ce", torch::kCUDA, ©_to_gpu_no_ce);
|
||||
m.def("concat_mla_k(Tensor! k, Tensor k_nope, Tensor k_rope) -> ()");
|
||||
@@ -364,9 +357,6 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
m.def("top_p_renorm_probs(Tensor probs, Tensor! renorm_probs, Tensor? maybe_top_p_arr, float top_p_val) -> ()");
|
||||
m.impl("top_p_renorm_probs", torch::kCUDA, &top_p_renorm_probs);
|
||||
|
||||
m.def("top_k_mask_logits(Tensor logits, Tensor mask_logits, Tensor? maybe_top_k_arr, int top_k_val) -> ()");
|
||||
m.impl("top_k_mask_logits", torch::kCUDA, &top_k_mask_logits);
|
||||
|
||||
/*
|
||||
* From Sparse Flash Attention
|
||||
*/
|
||||
|
||||
@@ -43,9 +43,6 @@ TORCH_LIBRARY_EXPAND(sgl_kernel, m) {
|
||||
"top_k_top_p_sampling_from_probs(Tensor probs, Tensor output, Tensor? maybe_indices, Tensor? maybe_top_k_arr, "
|
||||
"float top_k_val, Tensor? maybe_top_p_arr, float top_p_val, bool deterministic, Generator? gen) -> ()");
|
||||
m.impl("top_k_top_p_sampling_from_probs", torch::kMUSA, &top_k_top_p_sampling_from_probs);
|
||||
|
||||
m.def("top_k_mask_logits(Tensor logits, Tensor mask_logits, Tensor? maybe_top_k_arr, int top_k_val) -> ()");
|
||||
m.impl("top_k_mask_logits", torch::kMUSA, &top_k_mask_logits);
|
||||
}
|
||||
|
||||
REGISTER_EXTENSION(common_ops)
|
||||
|
||||
@@ -1,172 +0,0 @@
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
|
||||
#include "utils.h"
|
||||
|
||||
template <typename T>
|
||||
struct ConvertToFP8 {
|
||||
static __device__ __nv_fp8_storage_t convert_to_fp8(T value) {
|
||||
return 0;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct ConvertToFP8<__nv_bfloat16> {
|
||||
static __device__ __nv_fp8_storage_t convert_to_fp8(__nv_bfloat16 value) {
|
||||
return __nv_cvt_bfloat16raw_to_fp8(value, __NV_SATFINITE, __NV_E4M3);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct ConvertToFP8<half> {
|
||||
static __device__ __nv_fp8_storage_t convert_to_fp8(half value) {
|
||||
return __nv_cvt_halfraw_to_fp8(value, __NV_SATFINITE, __NV_E4M3);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct ConvertFromFloat {
|
||||
static __device__ T convert_from_float(float value) {
|
||||
return 0;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct ConvertFromFloat<__nv_bfloat16> {
|
||||
static __device__ __nv_bfloat16 convert_from_float(float value) {
|
||||
return __float2bfloat16(value);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct ConvertFromFloat<half> {
|
||||
static __device__ half convert_from_float(float value) {
|
||||
return __float2half(value);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
__global__ void fused_downcast_kernel(
|
||||
const T* cache_k,
|
||||
const T* cache_v,
|
||||
const float* k_scale,
|
||||
const float* v_scale,
|
||||
__nv_fp8_storage_t* output_k,
|
||||
__nv_fp8_storage_t* output_v,
|
||||
const int input_sl,
|
||||
const int head,
|
||||
const int dim,
|
||||
const T max_fp8,
|
||||
const T min_fp8,
|
||||
const int64_t mult,
|
||||
const int64_t offset,
|
||||
const int64_t* loc) {
|
||||
// TODO: change name
|
||||
int token_idx = blockIdx.x;
|
||||
int thread_idx = threadIdx.x;
|
||||
int total_threads = blockDim.x;
|
||||
|
||||
T k_scale_val = ConvertFromFloat<T>::convert_from_float(k_scale[0]);
|
||||
T v_scale_val = ConvertFromFloat<T>::convert_from_float(v_scale[0]);
|
||||
|
||||
T k_scale_inv = static_cast<T>(1.f) / k_scale_val;
|
||||
T v_scale_inv = static_cast<T>(1.f) / v_scale_val;
|
||||
|
||||
auto clamp = [&](T val) { return val > max_fp8 ? max_fp8 : (min_fp8 > val ? min_fp8 : val); };
|
||||
|
||||
if (token_idx < input_sl) {
|
||||
int out_seq_idx = loc[token_idx];
|
||||
|
||||
#pragma unroll
|
||||
for (int i = thread_idx; i < head * dim; i += total_threads) {
|
||||
int in_idx = token_idx * head * dim + i;
|
||||
int out_idx = (out_seq_idx * mult + offset) * head * dim + i;
|
||||
|
||||
T k_val = cache_k[in_idx] * k_scale_inv;
|
||||
k_val = clamp(k_val);
|
||||
output_k[out_idx] = ConvertToFP8<T>::convert_to_fp8(k_val);
|
||||
|
||||
T v_val = cache_v[in_idx] * v_scale_inv;
|
||||
v_val = clamp(v_val);
|
||||
output_v[out_idx] = ConvertToFP8<T>::convert_to_fp8(v_val);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void downcast_fp8_impl(
|
||||
at::Tensor& k,
|
||||
at::Tensor& v,
|
||||
at::Tensor& k_out,
|
||||
at::Tensor& v_out,
|
||||
at::Tensor& k_scale,
|
||||
at::Tensor& v_scale,
|
||||
at::Tensor& loc,
|
||||
int64_t mult,
|
||||
int64_t offset,
|
||||
cudaStream_t stream) {
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
CHECK_INPUT(k_out);
|
||||
CHECK_INPUT(v_out);
|
||||
CHECK_INPUT(k_scale);
|
||||
CHECK_INPUT(v_scale);
|
||||
CHECK_INPUT(loc);
|
||||
|
||||
int64_t input_sl = k.size(0);
|
||||
int64_t head = k.size(1);
|
||||
int64_t dim = k.size(2);
|
||||
|
||||
dim3 grid(input_sl * head);
|
||||
int vec_size = 8;
|
||||
dim3 block(std::min(int(dim) / vec_size, 1024));
|
||||
|
||||
const T max_fp8 = static_cast<T>(FP8_E4M3_MAX);
|
||||
const T min_fp8 = static_cast<T>(-FP8_E4M3_MAX);
|
||||
|
||||
fused_downcast_kernel<T><<<grid, block, 0, stream>>>(
|
||||
static_cast<const T*>(k.data_ptr()),
|
||||
static_cast<const T*>(v.data_ptr()),
|
||||
static_cast<const float*>(k_scale.data_ptr()),
|
||||
static_cast<const float*>(v_scale.data_ptr()),
|
||||
static_cast<__nv_fp8_storage_t*>(k_out.data_ptr()),
|
||||
static_cast<__nv_fp8_storage_t*>(v_out.data_ptr()),
|
||||
input_sl,
|
||||
head,
|
||||
dim,
|
||||
max_fp8,
|
||||
min_fp8,
|
||||
mult,
|
||||
offset,
|
||||
static_cast<const int64_t*>(loc.data_ptr()));
|
||||
|
||||
cudaError_t status = cudaGetLastError();
|
||||
TORCH_CHECK(status == cudaSuccess, "Kernel launch failed: " + std::string(cudaGetErrorString(status)));
|
||||
}
|
||||
|
||||
void downcast_fp8(
|
||||
at::Tensor& k,
|
||||
at::Tensor& v,
|
||||
at::Tensor& k_out,
|
||||
at::Tensor& v_out,
|
||||
at::Tensor& k_scale,
|
||||
at::Tensor& v_scale,
|
||||
at::Tensor& loc,
|
||||
int64_t mult,
|
||||
int64_t offset) {
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
CHECK_INPUT(k_out);
|
||||
CHECK_INPUT(v_out);
|
||||
|
||||
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
switch (k.scalar_type()) {
|
||||
case at::ScalarType::BFloat16:
|
||||
downcast_fp8_impl<__nv_bfloat16>(k, v, k_out, v_out, k_scale, v_scale, loc, mult, offset, stream);
|
||||
break;
|
||||
case at::ScalarType::Half:
|
||||
downcast_fp8_impl<__half>(k, v, k_out, v_out, k_scale, v_scale, loc, mult, offset, stream);
|
||||
break;
|
||||
default:
|
||||
TORCH_CHECK(false, "Unsupported input type for downcast_fp8. Expected bfloat16 or float16.");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user