[JIT sgl-kernel] Jit support per tensor quant (#15709)
This commit is contained in:
13
python/sglang/jit_kernel/include/sgl_kernel/fp8_utils.cuh
Normal file
13
python/sglang/jit_kernel/include/sgl_kernel/fp8_utils.cuh
Normal file
@@ -0,0 +1,13 @@
|
||||
#pragma once
|
||||
|
||||
#ifdef __CUDACC__
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_fp8.h>
|
||||
#endif
|
||||
|
||||
namespace device {
|
||||
|
||||
inline constexpr float FP8_E4M3_MAX = 448.0f;
|
||||
|
||||
} // namespace device
|
||||
@@ -23,6 +23,7 @@
|
||||
#ifdef __CUDACC__
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_fp8.h>
|
||||
#endif
|
||||
|
||||
namespace host {
|
||||
@@ -65,6 +66,10 @@ template <>
|
||||
struct dtype_trait<__nv_bfloat16> {
|
||||
inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLBfloat, .bits = 16, .lanes = 1};
|
||||
};
|
||||
template <>
|
||||
struct dtype_trait<__nv_fp8_e4m3> {
|
||||
inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLFloat8_e4m3fn, .bits = 8, .lanes = 1};
|
||||
};
|
||||
#endif
|
||||
|
||||
template <DLDeviceType Code>
|
||||
|
||||
@@ -12,6 +12,49 @@
|
||||
namespace device {
|
||||
|
||||
inline constexpr auto kWarpThreads = 32u;
|
||||
inline constexpr auto kFullMask = 0xffffffffu;
|
||||
|
||||
__device__ __forceinline__ float atomicMaxFloat(float* addr, float value) {
|
||||
#ifndef USE_ROCM
|
||||
float old;
|
||||
old = (value >= 0) ? __int_as_float(atomicMax((int*)addr, __float_as_int(value)))
|
||||
: __uint_as_float(atomicMin((unsigned int*)addr, __float_as_uint(value)));
|
||||
return old;
|
||||
#else
|
||||
int* addr_as_i = (int*)addr;
|
||||
int old = *addr_as_i, assumed;
|
||||
do {
|
||||
assumed = old;
|
||||
old = atomicCAS(addr_as_i, assumed, __float_as_int(fmaxf(value, __int_as_float(assumed))));
|
||||
} while (assumed != old);
|
||||
return __int_as_float(old);
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float warpReduceMax(float value) {
|
||||
value = fmaxf(value, __shfl_xor_sync(kFullMask, value, 16));
|
||||
value = fmaxf(value, __shfl_xor_sync(kFullMask, value, 8));
|
||||
value = fmaxf(value, __shfl_xor_sync(kFullMask, value, 4));
|
||||
value = fmaxf(value, __shfl_xor_sync(kFullMask, value, 2));
|
||||
value = fmaxf(value, __shfl_xor_sync(kFullMask, value, 1));
|
||||
return value;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float blockReduceMax(float value) {
|
||||
static __shared__ float warpLevelMaxs[kWarpThreads];
|
||||
const int laneId = threadIdx.x % kWarpThreads;
|
||||
const int warpId = threadIdx.x / kWarpThreads;
|
||||
|
||||
value = warpReduceMax(value);
|
||||
|
||||
if (laneId == 0) warpLevelMaxs[warpId] = value;
|
||||
__syncthreads();
|
||||
|
||||
value = (threadIdx.x < blockDim.x / kWarpThreads) ? warpLevelMaxs[laneId] : 0;
|
||||
if (warpId == 0) value = warpReduceMax(value);
|
||||
|
||||
return value;
|
||||
}
|
||||
|
||||
namespace pointer {
|
||||
|
||||
|
||||
@@ -2,6 +2,9 @@
|
||||
|
||||
// ref: https://forums.developer.nvidia.com/t/c-20s-source-location-compilation-error-when-using-nvcc-12-1/258026/3
|
||||
#ifdef __CUDACC__
|
||||
#include <cuda.h>
|
||||
#if CUDA_VERSION <= 12010
|
||||
|
||||
#pragma push_macro("__cpp_consteval")
|
||||
#pragma push_macro("_NODISCARD")
|
||||
#pragma push_macro("__builtin_LINE")
|
||||
@@ -23,7 +26,10 @@
|
||||
#undef consteval
|
||||
#pragma pop_macro("__cpp_consteval")
|
||||
#pragma pop_macro("_NODISCARD")
|
||||
#else
|
||||
#else // __CUDACC__ && CUDA_VERSION > 12010
|
||||
#include <source_location>
|
||||
#endif
|
||||
#else // no __CUDACC__
|
||||
#include <source_location>
|
||||
#endif
|
||||
|
||||
@@ -33,7 +39,6 @@
|
||||
#include <cstddef>
|
||||
#include <ostream>
|
||||
#include <ranges>
|
||||
#include <source_location>
|
||||
#include <sstream>
|
||||
#include <utility>
|
||||
|
||||
|
||||
Reference in New Issue
Block a user