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
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-memory
extern __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 = 64
constexpr int BN = 64;tma
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;"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