| #include "smallm_gemv.h" |
|
|
| #include <ATen/cuda/CUDAContext.h> |
| #include <c10/cuda/CUDAGuard.h> |
| #include <cuda_bf16.h> |
| #include <cuda_fp8.h> |
| #include <cuda_runtime.h> |
|
|
| #include <cstdint> |
| #include <limits> |
| #include <optional> |
| #include <vector> |
|
|
| namespace { |
|
|
| constexpr int kWarpSize = 32; |
| constexpr int kWarpsPerBlock = 8; |
| constexpr int kThreads = kWarpSize * kWarpsPerBlock; |
| constexpr int kFp4BlockElements = 16; |
| constexpr int kScaleTileOuter = 128; |
|
|
| __device__ __forceinline__ size_t scale_offset( |
| int outer, |
| int inner_scale, |
| int scale_inner_dim) { |
| const int outer_tile = outer / kScaleTileOuter; |
| const int local_outer = outer % kScaleTileOuter; |
| const int local_inner = inner_scale & 3; |
| const int inner_tile_start = inner_scale - local_inner; |
| const size_t tile_base = |
| static_cast<size_t>( |
| inner_tile_start + outer_tile * scale_inner_dim) * |
| kScaleTileOuter; |
| return tile_base + static_cast<size_t>(local_outer & 31) * 16 + |
| static_cast<size_t>(local_outer >> 5) * 4 + local_inner; |
| } |
|
|
| __device__ __forceinline__ float e4m3_to_float(uint8_t raw) { |
| const __half_raw half_raw = __nv_cvt_fp8_to_halfraw(raw, __NV_E4M3); |
| return __half2float(static_cast<__half>(half_raw)); |
| } |
|
|
| __device__ __forceinline__ float2 e2m1x2_to_float2(uint8_t packed) { |
| const __half2_raw raw = |
| __nv_cvt_fp4x2_to_halfraw2(packed, __NV_E2M1); |
| const __half2 converted(raw); |
| return __half22float2(converted); |
| } |
|
|
| __global__ void smallm_nvfp4_gemv_kernel( |
| const __nv_bfloat16* __restrict__ input, |
| const uint8_t* __restrict__ packed_weight, |
| const uint8_t* __restrict__ weight_block_scales, |
| const float* __restrict__ weight_tensor_scale, |
| const __nv_bfloat16* __restrict__ bias, |
| __nv_bfloat16* __restrict__ output, |
| int m, |
| int n, |
| int k, |
| int scale_inner_dim) { |
| const int lane = threadIdx.x & (kWarpSize - 1); |
| const int warp_in_block = threadIdx.x / kWarpSize; |
| const int64_t output_linear = |
| static_cast<int64_t>(blockIdx.x) * kWarpsPerBlock + warp_in_block; |
| const int64_t output_count = static_cast<int64_t>(m) * n; |
| if (output_linear >= output_count) { |
| return; |
| } |
|
|
| const int row_m = static_cast<int>(output_linear / n); |
| const int row_n = static_cast<int>( |
| output_linear - static_cast<int64_t>(row_m) * n); |
| const int packed_k = k / 2; |
| const __nv_bfloat16* input_row = |
| input + static_cast<int64_t>(row_m) * k; |
| const uint8_t* weight_row = |
| packed_weight + static_cast<int64_t>(row_n) * packed_k; |
| float accumulator = 0.0f; |
|
|
| for (int pair = lane; pair < packed_k; pair += kWarpSize) { |
| const uint8_t packed = weight_row[pair]; |
| const int scale_block = pair / (kFp4BlockElements / 2); |
| const uint8_t scale_raw = weight_block_scales[ |
| scale_offset(row_n, scale_block, scale_inner_dim)]; |
| const float scale = |
| e4m3_to_float(scale_raw) * weight_tensor_scale[0]; |
| const float2 weights = e2m1x2_to_float2(packed); |
| const int input_index = pair * 2; |
| accumulator = fmaf( |
| __bfloat162float(input_row[input_index]), |
| weights.y * scale, |
| accumulator); |
| accumulator = fmaf( |
| __bfloat162float(input_row[input_index + 1]), |
| weights.x * scale, |
| accumulator); |
| } |
|
|
| #pragma unroll |
| for (int offset = 16; offset > 0; offset >>= 1) { |
| accumulator += __shfl_down_sync(0xFFFFFFFF, accumulator, offset); |
| } |
| if (lane == 0) { |
| if (bias != nullptr) { |
| accumulator += __bfloat162float(bias[row_n]); |
| } |
| output[output_linear] = __float2bfloat16_rn(accumulator); |
| } |
| } |
|
|
| } |
|
|
| torch::Tensor smallm_nvfp4_linear_cuda( |
| const torch::Tensor& input, |
| const torch::Tensor& packed_weight, |
| const torch::Tensor& weight_block_scales, |
| const torch::Tensor& weight_tensor_scale, |
| const std::optional<torch::Tensor>& bias) { |
| const auto device = input.device(); |
| c10::cuda::CUDAGuard guard(device); |
| const int64_t out_features64 = packed_weight.size(0); |
| const int64_t in_features64 = packed_weight.size(1) * 2; |
| const int64_t logical_m64 = input.numel() / in_features64; |
| TORCH_CHECK( |
| logical_m64 > 0 && |
| logical_m64 <= static_cast<int64_t>(std::numeric_limits<int>::max()), |
| "small-M GEMV M is out of range"); |
| TORCH_CHECK( |
| out_features64 <= static_cast<int64_t>(std::numeric_limits<int>::max()) && |
| in_features64 <= static_cast<int64_t>(std::numeric_limits<int>::max()), |
| "small-M GEMV N or K is out of range"); |
| const int m = static_cast<int>(logical_m64); |
| const int n = static_cast<int>(out_features64); |
| const int k = static_cast<int>(in_features64); |
|
|
| std::vector<int64_t> output_shape = input.sizes().vec(); |
| output_shape.back() = out_features64; |
| torch::Tensor output = torch::empty( |
| output_shape, |
| input.options().dtype(at::kBFloat16)); |
| const int64_t output_count = logical_m64 * out_features64; |
| const int64_t block_count64 = |
| (output_count + kWarpsPerBlock - 1) / kWarpsPerBlock; |
| TORCH_CHECK( |
| block_count64 <= static_cast<int64_t>(std::numeric_limits<int>::max()), |
| "small-M GEMV grid is too large"); |
|
|
| const auto stream = |
| at::cuda::getCurrentCUDAStream(device.index()).stream(); |
| const __nv_bfloat16* bias_pointer = |
| bias.has_value() |
| ? reinterpret_cast<const __nv_bfloat16*>(bias->data_ptr()) |
| : nullptr; |
| smallm_nvfp4_gemv_kernel<<< |
| static_cast<int>(block_count64), |
| kThreads, |
| 0, |
| stream>>>( |
| reinterpret_cast<const __nv_bfloat16*>(input.data_ptr()), |
| reinterpret_cast<const uint8_t*>(packed_weight.data_ptr()), |
| reinterpret_cast<const uint8_t*>(weight_block_scales.data_ptr()), |
| reinterpret_cast<const float*>(weight_tensor_scale.data_ptr()), |
| bias_pointer, |
| reinterpret_cast<__nv_bfloat16*>(output.data_ptr()), |
| m, |
| n, |
| k, |
| static_cast<int>(weight_block_scales.size(1))); |
| C10_CUDA_KERNEL_LAUNCH_CHECK(); |
| return output; |
| } |
|
|