fix: use unsigned short for BF16 storage, inline PTX for conversions
This commit is contained in:
@@ -11,8 +11,10 @@
|
||||
#include <cstdint>
|
||||
#include <cmath>
|
||||
|
||||
// BF16: use __bf16 (CUDA 13+ built-in) to avoid cuda_bf16.h C++17 bug
|
||||
typedef __bf16 bf16_t;
|
||||
// BF16 storage type: use uint16_t (2 bytes) to avoid __bf16 ICE on CUDA 13.2
|
||||
// Conversions via inline PTX asm.
|
||||
|
||||
typedef unsigned short bf16_t;
|
||||
|
||||
__device__ __forceinline__ bf16_t f32_to_bf16(float f) {
|
||||
bf16_t h; asm("cvt.rn.bf16.f32 %0, %1;" : "=h"(h) : "f"(f)); return h;
|
||||
@@ -82,10 +84,10 @@ __device__ uint64_t make_umma_desc(const void* smem, uint32_t ld_bytes) {
|
||||
template<int HD, bool CAUSAL=false, bool SINK=false>
|
||||
__global__ void __launch_bounds__(NTHREADS)
|
||||
fmha_decode(
|
||||
const bf16_t* __restrict__ q, // (B, H, T, HD)
|
||||
const bf16_t* __restrict__ k, // (B, sk, HD)
|
||||
const bf16_t* __restrict__ v, // (B, HD, sk)
|
||||
bf16_t* __restrict__ o, // (B, H, T, HD)
|
||||
const uint16_t* __restrict__ q, // (B, H, T, HD)
|
||||
const uint16_t* __restrict__ k, // (B, sk, HD)
|
||||
const uint16_t* __restrict__ v, // (B, HD, sk)
|
||||
uint16_t* __restrict__ o, // (B, H, T, HD)
|
||||
int bstride_q, int bstride_kv, int bstride_o,
|
||||
int s_k, int n_comp, int swa_len,
|
||||
float scale, // 1/sqrt(HD)
|
||||
@@ -111,17 +113,17 @@ fmha_decode(
|
||||
// SMEM
|
||||
const int kvs = (HD > 128) ? 1 : 2;
|
||||
extern __shared__ char sbuf[];
|
||||
bf16_t* sQ = (bf16_t*)sbuf;
|
||||
uint16_t* sQ = (uint16_t*)sbuf;
|
||||
int off = TILE_M * HD;
|
||||
bf16_t* sK = (bf16_t*)(sbuf + off * sizeof(bf16_t)); off += TILE_K * HD * kvs;
|
||||
bf16_t* sV = (bf16_t*)(sbuf + off * sizeof(bf16_t)); off += TILE_K * HD * kvs;
|
||||
bf16_t* sC = (bf16_t*)(sbuf + off * sizeof(bf16_t));
|
||||
uint16_t* sK = (uint16_t*)(sbuf + off * sizeof(uint16_t)); off += TILE_K * HD * kvs;
|
||||
uint16_t* sV = (uint16_t*)(sbuf + off * sizeof(uint16_t)); off += TILE_K * HD * kvs;
|
||||
uint16_t* sC = (uint16_t*)(sbuf + off * sizeof(uint16_t));
|
||||
|
||||
// Pointers for this head/batch
|
||||
const bf16_t* qh = q + batch*bstride_q + head*HD;
|
||||
const bf16_t* kb = k + batch*bstride_kv;
|
||||
const bf16_t* vb = v + batch*bstride_kv;
|
||||
bf16_t* oh = o + batch*bstride_o + head*HD;
|
||||
const uint16_t* qh = q + batch*bstride_q + head*HD;
|
||||
const uint16_t* kb = k + batch*bstride_kv;
|
||||
const uint16_t* vb = v + batch*bstride_kv;
|
||||
uint16_t* oh = o + batch*bstride_o + head*HD;
|
||||
|
||||
const float scale_log2 = scale * 1.4426950408889634f;
|
||||
const int nkt = (s_k + TILE_K - 1) / TILE_K;
|
||||
@@ -150,8 +152,8 @@ fmha_decode(
|
||||
// --- MMA warp: QK GEMM (S = Q @ K^T) ---
|
||||
// tcgen05.mma.cta_group::1.kind::f16 [tmem_c], desc_a, desc_b, idescE, scaleC, mask, pred
|
||||
if (is_mma && lane == 0) {
|
||||
uint64_t desc_a = make_umma_desc(sQ, HD * sizeof(bf16_t));
|
||||
uint64_t desc_b = make_umma_desc(sK, HD * sizeof(bf16_t));
|
||||
uint64_t desc_a = make_umma_desc(sQ, HD * sizeof(uint16_t));
|
||||
uint64_t desc_b = make_umma_desc(sK, HD * sizeof(uint16_t));
|
||||
uint32_t idescE = 0; // no E descriptor
|
||||
uint32_t tmem_c = ts; // S accumulator starts at TMEM_S
|
||||
int pred = 1;
|
||||
|
||||
@@ -15,7 +15,7 @@ static int compute_smem(int D) {
|
||||
int k = 128 * D * kvs; // K
|
||||
int v = 128 * D * kvs; // V
|
||||
int c = 128 * D; // C (epilogue)
|
||||
return (q + k + v + c) * sizeof(bf16_t);
|
||||
return (q + k + v + c) * sizeof(uint16_t);
|
||||
}
|
||||
|
||||
std::tuple<torch::Tensor, torch::Tensor> fmha_decode_cuda(
|
||||
@@ -37,10 +37,10 @@ std::tuple<torch::Tensor, torch::Tensor> fmha_decode_cuda(
|
||||
const float* sp = attn_sink.has_value() ? attn_sink->data_ptr<float>() : nullptr;
|
||||
|
||||
#define L(D, C, S) fmha_decode<D,C,S><<<grid,block,smem>>>( \
|
||||
(bf16_t*)q.data_ptr<at::BFloat16>(), \
|
||||
(bf16_t*)k.data_ptr<at::BFloat16>(), \
|
||||
(bf16_t*)v.data_ptr<at::BFloat16>(), \
|
||||
(bf16_t*)o.data_ptr<at::BFloat16>(), \
|
||||
(uint16_t*)q.data_ptr<at::BFloat16>(), \
|
||||
(uint16_t*)k.data_ptr<at::BFloat16>(), \
|
||||
(uint16_t*)v.data_ptr<at::BFloat16>(), \
|
||||
(uint16_t*)o.data_ptr<at::BFloat16>(), \
|
||||
q.stride(0), k.stride(0), o.stride(0), \
|
||||
sk, n_comp, swa_len, (float)scale, sp, lse.data_ptr<float>())
|
||||
|
||||
|
||||
Reference in New Issue
Block a user