| |
| |
| |
| |
| |
| |
| |
| |
| #include "quantize_fp4_sfa.cuh" |
|
|
| #include <cuda_bf16.h> |
| #include <cuda_fp16.h> |
| #include <cuda_fp8.h> |
|
|
| #ifndef CUTLASS_ARCH_MMA_SM100_SUPPORTED |
| # define CUTLASS_ARCH_MMA_SM100_SUPPORTED 1 |
| #endif |
| #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED) || defined(__CUDA_ARCH__) |
| # include "cutlass/cutlass.h" |
| # include "cutlass/detail/sm100_blockscaled_layout.hpp" |
| # include "cute/tensor.hpp" |
| # define FV_HAVE_CUTLASS 1 |
| #else |
| # define FV_HAVE_CUTLASS 0 |
| #endif |
|
|
| namespace flash_rt { |
| namespace fp4 { |
|
|
| #if FV_HAVE_CUTLASS |
|
|
| using Cfg = cutlass::detail::Sm1xxBlockScaledConfig<16>; |
|
|
| |
| |
| __device__ __forceinline__ uint8_t fp32_to_e2m1_sfa(float x) { |
| uint8_t sign = (x < 0.f) ? 0x8u : 0x0u; |
| float ax = fabsf(x); |
| uint8_t mant; |
| if (ax <= 0.25f) mant = 0u; |
| else if (ax <= 0.75f) mant = 1u; |
| else if (ax <= 1.25f) mant = 2u; |
| else if (ax <= 1.75f) mant = 3u; |
| else if (ax <= 2.5f) mant = 4u; |
| else if (ax <= 3.5f) mant = 5u; |
| else if (ax <= 5.0f) mant = 6u; |
| else mant = 7u; |
| return sign | mant; |
| } |
|
|
| __device__ __forceinline__ __nv_fp8_e4m3 quantize_ue4m3_sfa(float x) { |
| float v = fmaxf(x, 0.f); |
| return __nv_fp8_e4m3(v); |
| } |
|
|
| __device__ __forceinline__ float dequantize_ue4m3_sfa(__nv_fp8_e4m3 s) { |
| return static_cast<float>(s); |
| } |
|
|
| __device__ __forceinline__ float input_to_float(__half value) { |
| return __half2float(value); |
| } |
|
|
| __device__ __forceinline__ float input_to_float(__nv_bfloat16 value) { |
| return __bfloat162float(value); |
| } |
|
|
| |
| |
| |
| template <typename Input, class LayoutSF> |
| __global__ void kernel_quantize_fp4_sfa( |
| const Input* __restrict__ src, |
| uint8_t* __restrict__ dst_packed, |
| uint8_t* __restrict__ dst_sfa, |
| LayoutSF layout, |
| int N, int D) { |
| const int block_idx = blockIdx.x * blockDim.x + threadIdx.x; |
| const int row = blockIdx.y; |
| const int n_blocks = D / 16; |
| if (row >= N || block_idx >= n_blocks) return; |
|
|
| const int base = row * D + block_idx * 16; |
| float vals[16]; |
| float amax = 0.f; |
| #pragma unroll |
| for (int i = 0; i < 16; ++i) { |
| vals[i] = input_to_float(src[base + i]); |
| float a = fabsf(vals[i]); |
| if (a > amax) amax = a; |
| } |
|
|
| float desired = amax / 6.f; |
| if (desired < 1e-12f) desired = 1e-12f; |
| __nv_fp8_e4m3 bs_q = quantize_ue4m3_sfa(desired); |
| float bs_dq = dequantize_ue4m3_sfa(bs_q); |
|
|
| |
| |
| |
| |
| int sfa_off = layout(row, block_idx * 16, 0); |
| dst_sfa[sfa_off] = *reinterpret_cast<uint8_t*>(&bs_q); |
|
|
| |
| const int out_base = row * (D / 2) + block_idx * 8; |
| const float inv_bs = 1.f / bs_dq; |
| #pragma unroll |
| for (int p = 0; p < 8; ++p) { |
| float v_lo = vals[2 * p ] * inv_bs; |
| float v_hi = vals[2 * p + 1] * inv_bs; |
| uint8_t lo = fp32_to_e2m1_sfa(v_lo); |
| uint8_t hi = fp32_to_e2m1_sfa(v_hi); |
| dst_packed[out_base + p] = lo | (hi << 4); |
| } |
| } |
|
|
| #endif |
|
|
| int quantize_fp4_dynamic_sfa_fp16( |
| const void* src_fp16, void* dst_packed, void* dst_sfa, |
| int N, int D, bool is_sfb, cudaStream_t stream) { |
| #if FV_HAVE_CUTLASS |
| if (D % 16 != 0) return -1; |
| const int n_blocks = D / 16; |
| const int threads = 128; |
| dim3 grid((n_blocks + threads - 1) / threads, N); |
| dim3 block(threads); |
|
|
| |
| auto shape = cute::make_shape( |
| is_sfb ? 1 : N, |
| is_sfb ? N : 1, |
| D, 1); |
|
|
| if (is_sfb) { |
| auto layout = Cfg::tile_atom_to_shape_SFB(shape); |
| kernel_quantize_fp4_sfa<__half><<<grid, block, 0, stream>>>( |
| reinterpret_cast<const __half*>(src_fp16), |
| reinterpret_cast<uint8_t*>(dst_packed), |
| reinterpret_cast<uint8_t*>(dst_sfa), |
| layout, N, D); |
| } else { |
| auto layout = Cfg::tile_atom_to_shape_SFA(shape); |
| kernel_quantize_fp4_sfa<__half><<<grid, block, 0, stream>>>( |
| reinterpret_cast<const __half*>(src_fp16), |
| reinterpret_cast<uint8_t*>(dst_packed), |
| reinterpret_cast<uint8_t*>(dst_sfa), |
| layout, N, D); |
| } |
| cudaError_t e = cudaGetLastError(); |
| return (e == cudaSuccess) ? 0 : -static_cast<int>(e); |
| #else |
| (void)src_fp16; (void)dst_packed; (void)dst_sfa; |
| (void)N; (void)D; (void)is_sfb; (void)stream; |
| return -2; |
| #endif |
| } |
|
|
| int quantize_fp4_dynamic_sfa_bf16( |
| const void* src_bf16, void* dst_packed, void* dst_sfa, |
| int N, int D, bool is_sfb, cudaStream_t stream) { |
| #if FV_HAVE_CUTLASS |
| if (D % 16 != 0) return -1; |
| const int n_blocks = D / 16; |
| const int threads = 128; |
| dim3 grid((n_blocks + threads - 1) / threads, N); |
| dim3 block(threads); |
| auto shape = cute::make_shape(is_sfb ? 1 : N, is_sfb ? N : 1, D, 1); |
|
|
| if (is_sfb) { |
| auto layout = Cfg::tile_atom_to_shape_SFB(shape); |
| kernel_quantize_fp4_sfa<__nv_bfloat16><<<grid, block, 0, stream>>>( |
| reinterpret_cast<const __nv_bfloat16*>(src_bf16), |
| reinterpret_cast<uint8_t*>(dst_packed), |
| reinterpret_cast<uint8_t*>(dst_sfa), layout, N, D); |
| } else { |
| auto layout = Cfg::tile_atom_to_shape_SFA(shape); |
| kernel_quantize_fp4_sfa<__nv_bfloat16><<<grid, block, 0, stream>>>( |
| reinterpret_cast<const __nv_bfloat16*>(src_bf16), |
| reinterpret_cast<uint8_t*>(dst_packed), |
| reinterpret_cast<uint8_t*>(dst_sfa), layout, N, D); |
| } |
| cudaError_t e = cudaGetLastError(); |
| return (e == cudaSuccess) ? 0 : -static_cast<int>(e); |
| #else |
| (void)src_bf16; (void)dst_packed; (void)dst_sfa; |
| (void)N; (void)D; (void)is_sfb; (void)stream; |
| return -2; |
| #endif |
| } |
|
|
| } |
| } |
|
|