[JIT sgl-kernel] Jit support per tensor quant (#15709)

This commit is contained in:
Xiaoyu Zhang
2025-12-25 16:24:37 +08:00
committed by GitHub
parent a89e85e739
commit de2f2880b5
11 changed files with 497 additions and 7 deletions

View 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

View File

@@ -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>

View File

@@ -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 {

View File

@@ -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>