File size: 2,621 Bytes
88b8ef2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
// SPDX-License-Identifier: Apache-2.0
#include "dequantize_fp4_sfa.cuh"

#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 fused_fp4 {

#if FV_HAVE_CUTLASS

using CfgDequant = cutlass::detail::Sm1xxBlockScaledConfig<16>;

__device__ __forceinline__ float e2m1_to_fp32_dequant(uint8_t value) {
  static constexpr float mags[8] = {0.f, 0.5f, 1.f, 1.5f, 2.f, 3.f, 4.f, 6.f};
  float mag = mags[value & 0x7];
  return (value & 0x8) ? -mag : mag;
}

template <class LayoutSF>
__global__ void dequantize_fp4_sfa_kernel(
    const uint8_t* __restrict__ packed,
    const uint8_t* __restrict__ sfa,
    __half* __restrict__ out,
    LayoutSF layout,
    int dim) {
  int row = blockIdx.y;
  int block_idx = blockIdx.x * blockDim.x + threadIdx.x;
  int n_blocks = dim / 16;
  if (block_idx >= n_blocks) return;

  int col_base = block_idx * 16;
  int sfa_off = layout(row, col_base, 0);
  __nv_fp8_e4m3 scale_q;
  *reinterpret_cast<uint8_t*>(&scale_q) = sfa[sfa_off];
  float scale = static_cast<float>(scale_q);

  const uint8_t* packed_block = packed + row * (dim / 2) + block_idx * 8;
  __half* out_block = out + row * dim + col_base;
#pragma unroll
  for (int p = 0; p < 8; ++p) {
    uint8_t byte = packed_block[p];
    out_block[2 * p] = __float2half(e2m1_to_fp32_dequant(byte & 0xF) * scale);
    out_block[2 * p + 1] = __float2half(e2m1_to_fp32_dequant(byte >> 4) * scale);
  }
}

#endif

void dequantize_fp4_sfa_fp16(
    const uint8_t* packed,
    const uint8_t* sfa,
    __half* out,
    int rows,
    int dim,
    bool is_sfb,
    cudaStream_t stream) {
#if FV_HAVE_CUTLASS
  int n_blocks = dim / 16;
  dim3 block(256);
  dim3 grid((n_blocks + block.x - 1) / block.x, rows);
  auto shape = cute::make_shape(is_sfb ? 1 : rows, is_sfb ? rows : 1, dim, 1);
  if (is_sfb) {
    auto layout = CfgDequant::tile_atom_to_shape_SFB(shape);
    dequantize_fp4_sfa_kernel<<<grid, block, 0, stream>>>(
        packed, sfa, out, layout, dim);
  } else {
    auto layout = CfgDequant::tile_atom_to_shape_SFA(shape);
    dequantize_fp4_sfa_kernel<<<grid, block, 0, stream>>>(
        packed, sfa, out, layout, dim);
  }
#else
  (void)packed; (void)sfa; (void)out; (void)rows; (void)dim; (void)is_sfb; (void)stream;
#endif
}

}  // namespace fused_fp4
}  // namespace flash_rt