Skip to content
KernelIndex
Search⌘K

submission 241202

snowclipsed · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 551 lines, June 9 Researcher Reciprocity License v1.0.

solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-241202?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 dual GEMMsuite of 4 cases
NVIDIA B200
14.8µs
#79 of 420
2025-12-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4b8ea174c0a69d01f2e6c14e65536cd671632b653f1c2500c116d195d8de11aa
license declaredunknown
license concludedunknown
authorssnowclipsed
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

mbarrier__device__ inline void mbar_init(int a, int c) { asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(a), "r"(c)); }
shared-memoryextern __shared__ __align__(1024) char smem[];
tcgen05__device__ inline void tcgen05_cp(int t, uint64_t d) { asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(t), "l"(d)); }
tile-n = 64constexpr int BN = 64;
tmaasm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint [%0], [%1, {%2, %3, %4}], [%5], %6;"

Kernel source

solution.py551 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

CUDA_SRC = r"""
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
#include <cutlass/cutlass.h>
#include <cute/tensor.hpp>

constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST  = 0x14F0000000000000ULL;

__device__ inline constexpr uint64_t enc(uint64_t x) { return (x & 0x3FFFFULL) >> 4ULL; }

__device__ uint32_t elect() {
    uint32_t p = 0;
    asm volatile("{\n.reg .pred %%px;\nelect.sync _|%%px, %1;\n@%%px mov.s32 %0, 1;\n}" : "+r"(p) : "r"(0xFFFFFFFF));
    return p;
}

__device__ inline void mbar_init(int a, int c) { asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(a), "r"(c)); }
__device__ inline void mbar_wait(int a, int p) {
    asm volatile("{\n.reg .pred P1;\nLAB_WAIT:\nmbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1;\n@P1 bra.uni DONE;\nbra.uni LAB_WAIT;\nDONE:\n}" :: "r"(a), "r"(p));
}
__device__ inline void mbar_arrive_tx(int a, uint32_t b) {
    asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;" :: "r"(a), "r"(b) : "memory");
}
__device__ inline void tma_load_cached(int d, const void* t, int x, int y, int z, int m, uint64_t cache) {
    asm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint [%0], [%1, {%2, %3, %4}], [%5], %6;" 
                 :: "r"(d), "l"(t), "r"(x), "r"(y), "r"(z), "r"(m), "l"(cache) : "memory");
}
__device__ inline void bulk_cp_cached(int d, const void* s, int sz, int m, uint64_t cache) {
    asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;" 
                 :: "r"(d), "l"(s), "r"(sz), "r"(m), "l"(cache) : "memory");
}
__device__ inline void tcgen05_cp(int t, uint64_t d) { asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(t), "l"(d)); }

template<int D>
__device__ inline void tcgen05_mma(uint64_t a, uint64_t b, uint32_t i, int sa, int sb, int en) {
    asm volatile("{\n.reg .pred p;\nsetp.ne.b32 p, %6, 0;\n"
                 "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\n}"
        :: "r"(D), "l"(a), "l"(b), "r"(i), "r"(sa), "r"(sb), "r"(en) : "memory");
}

__device__ inline void tcgen05_ld_16x256b_x16_addr(float* d, uint32_t addr) {
    asm volatile(
        "tcgen05.ld.sync.aligned.16x256b.x16.b32 "
        "{%0,%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];"
        : "=f"(d[0]),"=f"(d[1]),"=f"(d[2]),"=f"(d[3]),"=f"(d[4]),"=f"(d[5]),"=f"(d[6]),"=f"(d[7]),
          "=f"(d[8]),"=f"(d[9]),"=f"(d[10]),"=f"(d[11]),"=f"(d[12]),"=f"(d[13]),"=f"(d[14]),"=f"(d[15]),
          "=f"(d[16]),"=f"(d[17]),"=f"(d[18]),"=f"(d[19]),"=f"(d[20]),"=f"(d[21]),"=f"(d[22]),"=f"(d[23]),
          "=f"(d[24]),"=f"(d[25]),"=f"(d[26]),"=f"(d[27]),"=f"(d[28]),"=f"(d[29]),"=f"(d[30]),"=f"(d[31]),
          "=f"(d[32]),"=f"(d[33]),"=f"(d[34]),"=f"(d[35]),"=f"(d[36]),"=f"(d[37]),"=f"(d[38]),"=f"(d[39]),
          "=f"(d[40]),"=f"(d[41]),"=f"(d[42]),"=f"(d[43]),"=f"(d[44]),"=f"(d[45]),"=f"(d[46]),"=f"(d[47]),
          "=f"(d[48]),"=f"(d[49]),"=f"(d[50]),"=f"(d[51]),"=f"(d[52]),"=f"(d[53]),"=f"(d[54]),"=f"(d[55]),
          "=f"(d[56]),"=f"(d[57]),"=f"(d[58]),"=f"(d[59]),"=f"(d[60]),"=f"(d[61]),"=f"(d[62]),"=f"(d[63])
        : "r"(addr));
}

__device__ inline void tcgen05_ld_16x256b_x8_addr(float* d, uint32_t addr) {
    asm volatile(
        "tcgen05.ld.sync.aligned.16x256b.x8.b32 "
        "{%0,%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];"
        : "=f"(d[0]),"=f"(d[1]),"=f"(d[2]),"=f"(d[3]),"=f"(d[4]),"=f"(d[5]),"=f"(d[6]),"=f"(d[7]),
          "=f"(d[8]),"=f"(d[9]),"=f"(d[10]),"=f"(d[11]),"=f"(d[12]),"=f"(d[13]),"=f"(d[14]),"=f"(d[15]),
          "=f"(d[16]),"=f"(d[17]),"=f"(d[18]),"=f"(d[19]),"=f"(d[20]),"=f"(d[21]),"=f"(d[22]),"=f"(d[23]),
          "=f"(d[24]),"=f"(d[25]),"=f"(d[26]),"=f"(d[27]),"=f"(d[28]),"=f"(d[29]),"=f"(d[30]),"=f"(d[31])
        : "r"(addr));
}

__device__ inline float silu(float x) { return x / (1.0f + __expf(-x)); }

template<int BM, int BN, int BK, int NS>
__global__ void __launch_bounds__(BM + 2*WARP_SIZE)
cutlass_kernel_bn128(
    const __grid_constant__ CUtensorMap At, const __grid_constant__ CUtensorMap B1t, const __grid_constant__ CUtensorMap B2t,
    const char* SFA, const char* SFB1, const char* SFB2, half* C, int M, int N, int K
) {
    constexpr int A_SZ = BM*BK/2, B_SZ = BN*BK/2, SFA_SZ = 128*(BK/16), SFB_SZ = 128*(BK/16);
    constexpr int G1_SZ = A_SZ + B_SZ + SFA_SZ + SFB_SZ, G2_SZ = B_SZ + SFB_SZ;
    constexpr int STAGE_SZ = G1_SZ + G2_SZ;
    constexpr int D0 = 0, D1 = BN, SFA_T = 2*BN, SFB1_T = SFA_T + 4*(BK/MMA_K), SFB2_T = SFB1_T + 4*(BK/MMA_K);
    constexpr int NW = BM/WARP_SIZE + 2;
    constexpr int SMEM_STRIDE = 136;
    
    extern __shared__ __align__(1024) char smem[];
    const int sm = (int)__cvta_generic_to_shared(smem);
    const int tid = threadIdx.x, wid = tid/WARP_SIZE, lid = tid%WARP_SIZE;
    const int bid = blockIdx.x, off_m = (bid/(N/BN))*BM, off_n = (bid%(N/BN))*BN;
    
    auto stage_ptr = [&](int s) { return sm + s*STAGE_SZ; };
    const int mbar_tma1 = sm + NS*STAGE_SZ, mbar_tma2 = mbar_tma1 + NS*8;
    const int mbar_mma = mbar_tma2 + NS*8, mbar_done = mbar_mma + NS*8;
    
    if (wid == 0 && elect()) {
        for (int s = 0; s < NS; s++) { mbar_init(mbar_tma1+s*8,1); mbar_init(mbar_tma2+s*8,1); mbar_init(mbar_mma+s*8,1); }
        mbar_init(mbar_done, 1);
    } else if (wid == 1) {
        asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(sm), "r"(512));
    }
    __syncthreads();
    
    const int num_k = K/BK, rest_k = K/64;
    
    // --- TMA PRODUCER ---
    if (wid == NW-2 && elect()) {
        constexpr uint64_t cache_A = EVICT_LAST, cache_B = EVICT_FIRST;
        auto issue_tma = [&](int ik, int s) {
            int base = stage_ptr(s), off_k = ik*BK;
            int A_sm = base, B1_sm = base+A_SZ, SFA_sm = base+A_SZ+B_SZ, SFB1_sm = SFA_sm+SFA_SZ;
            int B2_sm = base+G1_SZ, SFB2_sm = B2_sm+B_SZ;
            int mb1 = mbar_tma1 + s*8, mb2 = mbar_tma2 + s*8;
            tma_load_cached(A_sm, &At, 0, off_m, ik, mb1, cache_A);
            tma_load_cached(B1_sm, &B1t, 0, off_n, ik, mb1, cache_B);
            bulk_cp_cached(SFA_sm, SFA + ((off_m/128)*rest_k + off_k/64)*512, SFA_SZ, mb1, cache_A);
            bulk_cp_cached(SFB1_sm, SFB1 + ((off_n/128)*rest_k + off_k/64)*512, SFB_SZ, mb1, cache_B);
            mbar_arrive_tx(mb1, G1_SZ);
            tma_load_cached(B2_sm, &B2t, 0, off_n, ik, mb2, cache_B);
            bulk_cp_cached(SFB2_sm, SFB2 + ((off_n/128)*rest_k + off_k/64)*512, SFB_SZ, mb2, cache_B);
            mbar_arrive_tx(mb2, G2_SZ);
        };
        for (int i = 0; i < NS && i < num_k; i++) issue_tma(i, i);
        for (int ik = NS; ik < num_k; ik++) { mbar_wait(mbar_mma + (ik%NS)*8, (ik/NS-1)%2); issue_tma(ik, ik%NS); }
    }
    
    // --- MMA CONSUMER ---
    if (wid == NW-1 && elect()) {
        constexpr uint32_t idesc = (1U<<7)|(1U<<10)|((uint32_t)BN>>3<<17)|(1U<<27);
        auto desc_ab = [](int a) -> uint64_t { return enc(a)|(enc(8*128)<<32)|(1ULL<<46)|(2ULL<<61); };
        auto desc_sf = [](int a) -> uint64_t { return enc(a)|(enc(8*16)<<32)|(1ULL<<46); };
        
        for (int ik = 0; ik < num_k; ik++) {
            int s = ik % NS, ph = (ik/NS) % 2;
            mbar_wait(mbar_tma1 + s*8, ph);
            int base = stage_ptr(s);
            int A_sm = base, B1_sm = base+A_SZ, SFA_sm = base+A_SZ+B_SZ, SFB1_sm = SFA_sm+SFA_SZ;
            int B2_sm = base+G1_SZ, SFB2_sm = B2_sm+B_SZ;
            uint64_t sfa_d = desc_sf(0)+((uint64_t)SFA_sm>>4), sfb1_d = desc_sf(0)+((uint64_t)SFB1_sm>>4);
            
            #pragma unroll
            for (int k = 0; k < BK/MMA_K; k++) { tcgen05_cp(SFA_T+k*4, sfa_d+k*(512ULL>>4)); tcgen05_cp(SFB1_T+k*4, sfb1_d+k*(512ULL>>4)); }
            
            {
                uint64_t ad = desc_ab(A_sm), b1d = desc_ab(B1_sm);
                int en = (ik == 0) ? 0 : 1;
                tcgen05_mma<D0>(ad, b1d, idesc, SFA_T, SFB1_T, en);
            }
            #pragma unroll
            for (int k = 1; k < BK/MMA_K; k++) {
                uint64_t ad = desc_ab(A_sm + k*32), b1d = desc_ab(B1_sm + k*32);
                tcgen05_mma<D0>(ad, b1d, idesc, SFA_T+k*4, SFB1_T+k*4, 1);
            }

            mbar_wait(mbar_tma2 + s*8, ph);
            uint64_t sfb2_d = desc_sf(0)+((uint64_t)SFB2_sm>>4);
            
            #pragma unroll
            for (int k = 0; k < BK/MMA_K; k++) tcgen05_cp(SFB2_T+k*4, sfb2_d+k*(512ULL>>4));
            
            {
                uint64_t ad = desc_ab(A_sm), b2d = desc_ab(B2_sm);
                int en = (ik == 0) ? 0 : 1;
                tcgen05_mma<D1>(ad, b2d, idesc, SFA_T, SFB2_T, en);
            }
            #pragma unroll
            for (int k = 1; k < BK/MMA_K; k++) {
                uint64_t ad = desc_ab(A_sm + k*32), b2d = desc_ab(B2_sm + k*32);
                tcgen05_mma<D1>(ad, b2d, idesc, SFA_T+k*4, SFB2_T+k*4, 1);
            }
            
            asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];" :: "r"(mbar_mma + s*8) : "memory");
        }
        asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];" :: "r"(mbar_done) : "memory");
    }
    
    // ============================================
    // EPILOGUE WITH ALL PRE-WORK HOISTED
    // ============================================
    if (tid < BM) {
        // ========== PRE-WORK (While MMA is running) ==========
        
        // 1. SMEM pointers
        half* epi_smem = reinterpret_cast<half*>(smem);
        half* my_smem = epi_smem + wid * 32 * SMEM_STRIDE;
        const int col_offset = lid * 4;
        
        // 2. Precompute TMEM load addresses (avoid shifts in critical path)
        const uint32_t tmem_d0_lo = ((wid * 32 + 0) << 16) | D0;
        const uint32_t tmem_d0_hi = ((wid * 32 + 16) << 16) | D0;
        const uint32_t tmem_d1_lo = ((wid * 32 + 0) << 16) | D1;
        const uint32_t tmem_d1_hi = ((wid * 32 + 16) << 16) | D1;
        
        // 3. Precompute ALL 32 global store row pointers
        const long long row_base = (long long)(off_m + wid * 32) * N + off_n + col_offset;
        half* C_rows[32];
        #pragma unroll
        for (int r = 0; r < 32; r++) {
            C_rows[r] = C + row_base + (long long)r * N;
        }
        
        // 4. Precompute ALL 32 SMEM read row pointers
        half* S_rows[32];
        #pragma unroll
        for (int r = 0; r < 32; r++) {
            S_rows[r] = my_smem + r * SMEM_STRIDE + col_offset;
        }
        
        // ========== WAIT ==========
        mbar_wait(mbar_done, 0);
        
        // ========== POST-WAIT (Critical Path - minimal work) ==========
        float d0[BN/2], d1[BN/2];
        
        // Use precomputed TMEM addresses
        tcgen05_ld_16x256b_x16_addr(d0, tmem_d0_lo);
        tcgen05_ld_16x256b_x16_addr(d1, tmem_d1_lo);
        asm volatile("tcgen05.wait::ld.sync.aligned;");
        
        #pragma unroll
        for (int m = 0; m < 2; m++) {
            const int row_off = m * 16;
            half* my_smem_m = my_smem + row_off * SMEM_STRIDE;
            
            // Math + SMEM Write
            #pragma unroll
            for (int i = 0; i < BN/8; i++) {
                int col_base = i * 8 + (lid % 4) * 2;
                int row0 = lid / 4;
                int row8 = lid / 4 + 8;
                
                my_smem_m[row0 * SMEM_STRIDE + col_base + 0] = __float2half(silu(d0[i*4+0]) * d1[i*4+0]);
                my_smem_m[row0 * SMEM_STRIDE + col_base + 1] = __float2half(silu(d0[i*4+1]) * d1[i*4+1]);
                my_smem_m[row8 * SMEM_STRIDE + col_base + 0] = __float2half(silu(d0[i*4+2]) * d1[i*4+2]);
                my_smem_m[row8 * SMEM_STRIDE + col_base + 1] = __float2half(silu(d0[i*4+3]) * d1[i*4+3]);
            }
            
            // Issue next TMEM load early (use precomputed addresses)
            if (m == 0) {
                tcgen05_ld_16x256b_x16_addr(d0, tmem_d0_hi);
                tcgen05_ld_16x256b_x16_addr(d1, tmem_d1_hi);
            }
            
            // Store using precomputed pointers (no address math here)
            #pragma unroll
            for (int row = 0; row < 16; row++) {
                const int r = row_off + row;
                *reinterpret_cast<float2*>(C_rows[r]) = *reinterpret_cast<float2*>(S_rows[r]);
            }
            
            if (m == 0) asm volatile("tcgen05.wait::ld.sync.aligned;");
        }
        
        if (wid == 0 && lid == 0) 
            asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(512));
    }
}

template<int BM, int BK, int NS>
__global__ void __launch_bounds__(BM + 2*WARP_SIZE)
cutlass_kernel_bn64(
    const __grid_constant__ CUtensorMap At, const __grid_constant__ CUtensorMap B1t, const __grid_constant__ CUtensorMap B2t,
    const char* SFA, const char* SFB1, const char* SFB2, half* C, int M, int N, int K
) {
    constexpr int BN = 64;
    constexpr int A_SZ = BM*BK/2, B_SZ = BN*BK/2, SFA_SZ = 128*(BK/16), SFB_SZ = 128*(BK/16);
    constexpr int G1_SZ = A_SZ + B_SZ + SFA_SZ + SFB_SZ, G2_SZ = B_SZ + SFB_SZ;
    constexpr int STAGE_SZ = G1_SZ + G2_SZ;
    constexpr int D0 = 0, D1 = BN, SFA_T = 2*BN, SFB1_T = SFA_T + 4*(BK/MMA_K), SFB2_T = SFB1_T + 4*(BK/MMA_K);
    constexpr int NW = BM/WARP_SIZE + 2;
    constexpr int SMEM_STRIDE = 72;
    
    extern __shared__ __align__(1024) char smem[];
    const int sm = (int)__cvta_generic_to_shared(smem);
    const int tid = threadIdx.x, wid = tid/WARP_SIZE, lid = tid%WARP_SIZE;
    const int bid = blockIdx.x, off_m = (bid/(N/BN))*BM, off_n = (bid%(N/BN))*BN;
    const int sfb_off = (off_n % 128) / 64 * 2;
    
    auto stage_ptr = [&](int s) { return sm + s*STAGE_SZ; };
    const int mbar_tma1 = sm + NS*STAGE_SZ, mbar_tma2 = mbar_tma1 + NS*8;
    const int mbar_mma = mbar_tma2 + NS*8, mbar_done = mbar_mma + NS*8;
    
    if (wid == 0 && elect()) {
        for (int s = 0; s < NS; s++) { mbar_init(mbar_tma1+s*8,1); mbar_init(mbar_tma2+s*8,1); mbar_init(mbar_mma+s*8,1); }
        mbar_init(mbar_done, 1);
    } else if (wid == 1) {
        asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(sm), "r"(512));
    }
    __syncthreads();
    
    const int num_k = K/BK, rest_k = K/64;
    
    // --- TMA PRODUCER ---
    if (wid == NW-2 && elect()) {
        constexpr uint64_t cache_A = EVICT_LAST, cache_B = EVICT_FIRST;
        auto issue_tma = [&](int ik, int s) {
            int base = stage_ptr(s), off_k = ik*BK;
            int A_sm = base, B1_sm = base+A_SZ, SFA_sm = base+A_SZ+B_SZ, SFB1_sm = SFA_sm+SFA_SZ;
            int B2_sm = base+G1_SZ, SFB2_sm = B2_sm+B_SZ;
            int mb1 = mbar_tma1 + s*8, mb2 = mbar_tma2 + s*8;
            tma_load_cached(A_sm, &At, 0, off_m, ik, mb1, cache_A);
            tma_load_cached(B1_sm, &B1t, 0, off_n, ik, mb1, cache_B);
            bulk_cp_cached(SFA_sm, SFA + ((off_m/128)*rest_k + off_k/64)*512, SFA_SZ, mb1, cache_A);
            bulk_cp_cached(SFB1_sm, SFB1 + ((off_n/128)*rest_k + off_k/64)*512, SFB_SZ, mb1, cache_B);
            mbar_arrive_tx(mb1, G1_SZ);
            tma_load_cached(B2_sm, &B2t, 0, off_n, ik, mb2, cache_B);
            bulk_cp_cached(SFB2_sm, SFB2 + ((off_n/128)*rest_k + off_k/64)*512, SFB_SZ, mb2, cache_B);
            mbar_arrive_tx(mb2, G2_SZ);
        };
        for (int i = 0; i < NS && i < num_k; i++) issue_tma(i, i);
        for (int ik = NS; ik < num_k; ik++) { mbar_wait(mbar_mma + (ik%NS)*8, (ik/NS-1)%2); issue_tma(ik, ik%NS); }
    }
    
    // --- MMA CONSUMER ---
    if (wid == NW-1 && elect()) {
        constexpr uint32_t idesc = (1U<<7)|(1U<<10)|((uint32_t)BN>>3<<17)|(1U<<27);
        auto desc_ab = [](int a) -> uint64_t { return enc(a)|(enc(8*128)<<32)|(1ULL<<46)|(2ULL<<61); };
        auto desc_sf = [](int a) -> uint64_t { return enc(a)|(enc(8*16)<<32)|(1ULL<<46); };
        for (int ik = 0; ik < num_k; ik++) {
            int s = ik % NS, ph = (ik/NS) % 2;
            mbar_wait(mbar_tma1 + s*8, ph);
            int base = stage_ptr(s);
            int A_sm = base, B1_sm = base+A_SZ, SFA_sm = base+A_SZ+B_SZ, SFB1_sm = SFA_sm+SFA_SZ;
            int B2_sm = base+G1_SZ, SFB2_sm = B2_sm+B_SZ;
            uint64_t sfa_d = desc_sf(0)+((uint64_t)SFA_sm>>4), sfb1_d = desc_sf(0)+((uint64_t)SFB1_sm>>4);
            
            #pragma unroll
            for (int k = 0; k < BK/MMA_K; k++) { tcgen05_cp(SFA_T+k*4, sfa_d+k*(512ULL>>4)); tcgen05_cp(SFB1_T+k*4, sfb1_d+k*(512ULL>>4)); }
            
            {
                constexpr int k1 = 0;
                {
                    constexpr int k2 = 0;
                    uint64_t ad = desc_ab(A_sm + k1*BM*128 + k2*32), b1d = desc_ab(B1_sm + k1*BN*128 + k2*32);
                    int ksf = k1*4+k2;
                    int en = (ik == 0) ? 0 : 1;
                    tcgen05_mma<D0>(ad, b1d, idesc, SFA_T+ksf*4, SFB1_T+ksf*4+sfb_off, en);
                }
                #pragma unroll
                for (int k2 = 1; k2 < 256/MMA_K; k2++) {
                    uint64_t ad = desc_ab(A_sm + k1*BM*128 + k2*32), b1d = desc_ab(B1_sm + k1*BN*128 + k2*32);
                    int ksf = k1*4+k2;
                    tcgen05_mma<D0>(ad, b1d, idesc, SFA_T+ksf*4, SFB1_T+ksf*4+sfb_off, 1);
                }
            }
            #pragma unroll
            for (int k1 = 1; k1 < BK/256; k1++) {
                #pragma unroll
                for (int k2 = 0; k2 < 256/MMA_K; k2++) {
                    uint64_t ad = desc_ab(A_sm + k1*BM*128 + k2*32), b1d = desc_ab(B1_sm + k1*BN*128 + k2*32);
                    int ksf = k1*4+k2;
                    tcgen05_mma<D0>(ad, b1d, idesc, SFA_T+ksf*4, SFB1_T+ksf*4+sfb_off, 1);
                }
            }

            mbar_wait(mbar_tma2 + s*8, ph);
            uint64_t sfb2_d = desc_sf(0)+((uint64_t)SFB2_sm>>4);
            
            #pragma unroll
            for (int k = 0; k < BK/MMA_K; k++) tcgen05_cp(SFB2_T+k*4, sfb2_d+k*(512ULL>>4));
            
            {
                constexpr int k1 = 0;
                {
                    constexpr int k2 = 0;
                    uint64_t ad = desc_ab(A_sm + k1*BM*128 + k2*32), b2d = desc_ab(B2_sm + k1*BN*128 + k2*32);
                    int ksf = k1*4+k2;
                    int en = (ik == 0) ? 0 : 1;
                    tcgen05_mma<D1>(ad, b2d, idesc, SFA_T+ksf*4, SFB2_T+ksf*4+sfb_off, en);
                }
                #pragma unroll
                for (int k2 = 1; k2 < 256/MMA_K; k2++) {
                    uint64_t ad = desc_ab(A_sm + k1*BM*128 + k2*32), b2d = desc_ab(B2_sm + k1*BN*128 + k2*32);
                    int ksf = k1*4+k2;
                    tcgen05_mma<D1>(ad, b2d, idesc, SFA_T+ksf*4, SFB2_T+ksf*4+sfb_off, 1);
                }
            }
            #pragma unroll
            for (int k1 = 1; k1 < BK/256; k1++) {
                #pragma unroll
                for (int k2 = 0; k2 < 256/MMA_K; k2++) {
                    uint64_t ad = desc_ab(A_sm + k1*BM*128 + k2*32), b2d = desc_ab(B2_sm + k1*BN*128 + k2*32);
                    int ksf = k1*4+k2;
                    tcgen05_mma<D1>(ad, b2d, idesc, SFA_T+ksf*4, SFB2_T+ksf*4+sfb_off, 1);
                }
            }

            asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];" :: "r"(mbar_mma + s*8) : "memory");
        }
        asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];" :: "r"(mbar_done) : "memory");
    }
    
    // ============================================
    // EPILOGUE WITH ALL PRE-WORK HOISTED (BN64)
    // ============================================
    if (tid < BM) {
        // ========== PRE-WORK (While MMA is running) ==========
        
        // 1. SMEM pointers
        half* epi_smem = reinterpret_cast<half*>(smem);
        half* my_smem = epi_smem + wid * 32 * SMEM_STRIDE;
        const int col_offset = lid * 2;
        
        // 2. Precompute TMEM load addresses
        const uint32_t tmem_d0_lo = ((wid * 32 + 0) << 16) | D0;
        const uint32_t tmem_d0_hi = ((wid * 32 + 16) << 16) | D0;
        const uint32_t tmem_d1_lo = ((wid * 32 + 0) << 16) | D1;
        const uint32_t tmem_d1_hi = ((wid * 32 + 16) << 16) | D1;
        
        // 3. Precompute ALL 32 global store row pointers
        const long long row_base = (long long)(off_m + wid * 32) * N + off_n + col_offset;
        half* C_rows[32];
        #pragma unroll
        for (int r = 0; r < 32; r++) {
            C_rows[r] = C + row_base + (long long)r * N;
        }
        
        // 4. Precompute ALL 32 SMEM read row pointers
        half* S_rows[32];
        #pragma unroll
        for (int r = 0; r < 32; r++) {
            S_rows[r] = my_smem + r * SMEM_STRIDE + col_offset;
        }
        
        // ========== WAIT ==========
        mbar_wait(mbar_done, 0);
        
        // ========== POST-WAIT (Critical Path - minimal work) ==========
        float d0[BN/2], d1[BN/2];
        
        // Use precomputed TMEM addresses
        tcgen05_ld_16x256b_x8_addr(d0, tmem_d0_lo);
        tcgen05_ld_16x256b_x8_addr(d1, tmem_d1_lo);
        asm volatile("tcgen05.wait::ld.sync.aligned;");

        #pragma unroll
        for (int m = 0; m < 2; m++) {
            const int row_off = m * 16;
            half* my_smem_m = my_smem + row_off * SMEM_STRIDE;
            
            // Math + SMEM Write
            #pragma unroll
            for (int i = 0; i < BN/8; i++) {
                int col_base = i * 8 + (lid % 4) * 2;
                int row0 = lid / 4;
                int row8 = lid / 4 + 8;
                
                my_smem_m[row0 * SMEM_STRIDE + col_base + 0] = __float2half(silu(d0[i*4+0]) * d1[i*4+0]);
                my_smem_m[row0 * SMEM_STRIDE + col_base + 1] = __float2half(silu(d0[i*4+1]) * d1[i*4+1]);
                my_smem_m[row8 * SMEM_STRIDE + col_base + 0] = __float2half(silu(d0[i*4+2]) * d1[i*4+2]);
                my_smem_m[row8 * SMEM_STRIDE + col_base + 1] = __float2half(silu(d0[i*4+3]) * d1[i*4+3]);
            }
            
            // Issue next TMEM load early
            if (m == 0) {
                tcgen05_ld_16x256b_x8_addr(d0, tmem_d0_hi);
                tcgen05_ld_16x256b_x8_addr(d1, tmem_d1_hi);
            }
            
            // Store using precomputed pointers
            #pragma unroll
            for (int row = 0; row < 16; row++) {
                const int r = row_off + row;
                *reinterpret_cast<__half2*>(C_rows[r]) = *reinterpret_cast<__half2*>(S_rows[r]);
            }

            if (m == 0) asm volatile("tcgen05.wait::ld.sync.aligned;");
        }
        
        if (wid == 0 && lid == 0) 
            asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(512));
    }
}

void check_cu(CUresult e) { if (e == CUDA_SUCCESS) return; const char* m; cuGetErrorString(e, &m); TORCH_CHECK(false, "CUDA: ", m); }
void init_tmap(CUtensorMap* t, const char* p, int dim, int K, int BD, int BK) {
    uint64_t gd[3] = {256, (uint64_t)dim, (uint64_t)(K/256)};
    uint64_t gs[2] = {(uint64_t)(K/2), 128};
    uint32_t bd[3] = {256, (uint32_t)BD, (uint32_t)(BK/256)};
    uint32_t es[3] = {1, 1, 1};
    check_cu(cuTensorMapEncodeTiled(t, CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B, 3, (void*)p, gd, gs, bd, es,
        CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B, CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE));
}

template<int BM, int BN, int BK, int NS>
at::Tensor launch_pipe(const at::Tensor& A, const at::Tensor& B1, const at::Tensor& B2,
                        const at::Tensor& SFA, const at::Tensor& SFB1, const at::Tensor& SFB2, at::Tensor& C) {
    int M = A.size(0), N = B1.size(0), K = A.size(1)*2;
    CUtensorMap At, B1t, B2t;
    init_tmap(&At, (const char*)A.data_ptr(), M, K, BM, BK);
    init_tmap(&B1t, (const char*)B1.data_ptr(), N, K, BN, BK);
    init_tmap(&B2t, (const char*)B2.data_ptr(), N, K, BN, BK);
    constexpr int A_SZ = BM*BK/2, B_SZ = BN*BK/2, SFA_SZ = 128*(BK/16), SFB_SZ = 128*(BK/16);
    constexpr int STAGE_SZ = A_SZ + 2*B_SZ + SFA_SZ + 2*SFB_SZ;
    int smem = NS*STAGE_SZ + (3*NS+1)*8 + 256;
    auto kern = cutlass_kernel_bn128<BM, BN, BK, NS>;
    if (smem > 48000) cudaFuncSetAttribute(kern, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
    kern<<<(M/BM)*(N/BN), BM+2*WARP_SIZE, smem>>>(At, B1t, B2t,
        (const char*)SFA.data_ptr(), (const char*)SFB1.data_ptr(), (const char*)SFB2.data_ptr(), (half*)C.data_ptr(), M, N, K);
    return C;
}

template<int BM, int BK, int NS>
at::Tensor launch_bn64(const at::Tensor& A, const at::Tensor& B1, const at::Tensor& B2,
                        const at::Tensor& SFA, const at::Tensor& SFB1, const at::Tensor& SFB2, at::Tensor& C) {
    int M = A.size(0), N = B1.size(0), K = A.size(1)*2;
    constexpr int BN = 64;
    CUtensorMap At, B1t, B2t;
    init_tmap(&At, (const char*)A.data_ptr(), M, K, BM, BK);
    init_tmap(&B1t, (const char*)B1.data_ptr(), N, K, BN, BK);
    init_tmap(&B2t, (const char*)B2.data_ptr(), N, K, BN, BK);
    constexpr int A_SZ = BM*BK/2, B_SZ = BN*BK/2, SFA_SZ = 128*(BK/16), SFB_SZ = 128*(BK/16);
    constexpr int STAGE_SZ = A_SZ + 2*B_SZ + SFA_SZ + 2*SFB_SZ;
    int smem = NS*STAGE_SZ + (3*NS+1)*8 + 256;
    auto kern = cutlass_kernel_bn64<BM, BK, NS>;
    if (smem > 48000) cudaFuncSetAttribute(kern, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
    kern<<<(M/BM)*(N/BN), BM+2*WARP_SIZE, smem>>>(At, B1t, B2t,
        (const char*)SFA.data_ptr(), (const char*)SFB1.data_ptr(), (const char*)SFB2.data_ptr(), (half*)C.data_ptr(), M, N, K);
    return C;
}

at::Tensor dual_gemm(const at::Tensor& A, const at::Tensor& B1, const at::Tensor& B2,
                     const at::Tensor& SFA, const at::Tensor& SFB1, const at::Tensor& SFB2, at::Tensor& C) {
    int M = A.size(0);
    if (M == 256) {
        return launch_bn64<128, 256, 5>(A, B1, B2, SFA, SFB1, SFB2, C);
    }
    return launch_pipe<128, 128, 256, 4>(A, B1, B2, SFA, SFB1, SFB2, C);
}

TORCH_LIBRARY(dual_gemm_module, m) {
    m.def("dual_gemm(Tensor A, Tensor B1, Tensor B2, Tensor SFA, Tensor SFB1, Tensor SFB2, Tensor(a!) C) -> Tensor");
    m.impl("dual_gemm", &dual_gemm);
}
"""
load_inline(
    "dual_gemm_ext", cpp_sources="", cuda_sources=CUDA_SRC, verbose=True, is_python_module=False, no_implicit_headers=True,
    extra_cuda_cflags=["-O3", "-gencode=arch=compute_100a,code=sm_100a", "--use_fast_math", "--expt-relaxed-constexpr", "-lineinfo", "-maxrregcount=128"],
    extra_ldflags=["-lcuda"],
)
def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, sfa, sfb1, sfb2, sfa_p, sfb1_p, sfb2_p, c = data
    return torch.ops.dual_gemm_module.dual_gemm(a, b1, b2, sfa_p, sfb1_p, sfb2_p, c)
scrolls · 551 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON