FMHA SM100: Fix tcgen05.mma PTX syntax — correct register constraints
- tcgen05.mma.cta_group::1.kind::f16 [tmem_c], desc_a, desc_b, idescE_hi, scaleC, {mask0..3}, pred
- idescE is upper 32 bits of the E descriptor
- scaleC is a float (1.0 for accumulate)
- mask is 4 uint32 values (0xFFFFFFFF for no masking)
This commit is contained in:
@@ -150,31 +150,37 @@ fmha_decode(
|
||||
__syncthreads();
|
||||
|
||||
// --- 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
|
||||
// tcgen05.mma.cta_group::1.kind::f16 [tmem_c], desc_a, desc_b, idescE_hi, scaleC, {mask0..3}, pred
|
||||
if (is_mma && lane == 0) {
|
||||
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;
|
||||
// QK: A=Q (128×HD), B=K (128×HD), C=S (128×128)
|
||||
// For hd=64: single MMA pass (K=64 fits in one instruction)
|
||||
// For hd=128: K-dim needs 2 sub-tiles (k_sub=0,1)
|
||||
const int k_tiles = (HD + 63) / 64; // sub-tiles along K
|
||||
for (int ksub = 0; ksub < k_tiles; ksub++) {
|
||||
int accumulate = (ksub > 0) ? 1 : 0;
|
||||
// The MMA PTX for bf16 S→S:
|
||||
// tcgen05.mma.cta_group::1.kind::f16 [tmem_c], desc_a, desc_b, idescE, 1.0, mask, p
|
||||
// scaleC=1.0 for first, scaleC=1.0 (accumulate) for subsequent
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %6, 0;\n\t"
|
||||
"tcgen05.mma.cta_group::1.kind::f16 [%0], %1, %2, %3, 1.0, 0, p;\n\t"
|
||||
"}\n\t"
|
||||
:: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(idescE), "f"(1.0f), "l"(0ULL), "r"(pred)
|
||||
);
|
||||
}
|
||||
uint64_t desc_a = make_umma_desc(sQ, HD * sizeof(bf16_t));
|
||||
uint64_t desc_b = make_umma_desc(sK, HD * sizeof(bf16_t));
|
||||
uint32_t tmem_c = ts; // S accumulator
|
||||
uint32_t idescE_hi = 0; // no E descriptor
|
||||
float scaleC = 1.0f; // accumulate
|
||||
uint32_t mask[4] = {0xFFFFFFFF, 0xFFFFFFFF, 0xFFFFFFFF, 0xFFFFFFFF};
|
||||
uint32_t pred = 1;
|
||||
|
||||
// QK GEMM: A=Q (128×HD BF16, SMEM), B=K (128×HD BF16, SMEM), C=S (128×128 FP32, TMEM)
|
||||
// For hd=64: single MMA instruction
|
||||
// For hd=128: need to handle K-dim sub-tiling
|
||||
//
|
||||
// The MMA instruction processes a 128×128 output tile.
|
||||
// A-desc encodes Q layout, B-desc encodes K layout.
|
||||
// The descriptor's leading dimension tells the MMA how to iterate over K.
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %9, 0;\n\t"
|
||||
"tcgen05.mma.cta_group::1.kind::f16 [%0], %1, %2, %3, %4, {%5, %6, %7, %8}, p;\n\t"
|
||||
"}\n\t"
|
||||
:: "r"(tmem_c),
|
||||
"l"(desc_a),
|
||||
"l"(desc_b),
|
||||
"r"(idescE_hi),
|
||||
"f"(scaleC),
|
||||
"r"(mask[0]), "r"(mask[1]), "r"(mask[2]), "r"(mask[3]),
|
||||
"r"(pred)
|
||||
);
|
||||
}
|
||||
|
||||
tmem_fence();
|
||||
|
||||
Reference in New Issue
Block a user