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:
2026-05-28 05:25:59 +00:00
parent a11a245307
commit 6f7449ce71

View File

@@ -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();