From 6f7449ce71ba9dc20f3c9d8151e12c00f9f325f8 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Thu, 28 May 2026 05:25:59 +0000 Subject: [PATCH] =?UTF-8?q?FMHA=20SM100:=20Fix=20tcgen05.mma=20PTX=20synta?= =?UTF-8?q?x=20=E2=80=94=20correct=20register=20constraints?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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) --- dsv4/kernels/attention/fmha_sm100.cuh | 54 +++++++++++++++------------ 1 file changed, 30 insertions(+), 24 deletions(-) diff --git a/dsv4/kernels/attention/fmha_sm100.cuh b/dsv4/kernels/attention/fmha_sm100.cuh index c7692c14..9bb8f5d2 100644 --- a/dsv4/kernels/attention/fmha_sm100.cuh +++ b/dsv4/kernels/attention/fmha_sm100.cuh @@ -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();