submission 276247
binsquare · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 538 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-276247?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:37f67c7ba16a6c64d4cd27a1fbe6eaf267667f6a4a183e62b2375042f59d1daf
license declaredunknown
license concludedunknown
authorsbinsquare
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mbarrier
__device__ inline void mbarrier_init(int mbar_addr, int count) {shared-memory
__device__ inline void tma_3d_gmem2smem_multicast(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint16_t ctamask, uint64_t cache_policy) {tcgen05
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));tma
asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;" :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy));Kernel source
submission.py538 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
from utils import make_match_reference
sf_vec_size = 16
def ceil_div(a, b):
return (a + b - 1) // b
def to_blocked(input_matrix):
rows, cols = input_matrix.shape
n_row_blocks = (rows + 127) // 128
n_col_blocks = (cols + 3) // 4
blocks = input_matrix.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3).contiguous()
rearranged = blocks.view(-1, 4, 32, 4).transpose(1, 2).contiguous().view(-1, 32, 16)
return rearranged.flatten()
def ref_kernel(data: input_t) -> output_t:
a_ref, b1_ref, b2_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu, _, _, _, c_ref = data
m, n, l = c_ref.shape
ref1 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)
ref2 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)
for l_idx in range(l):
scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
scale_b1 = to_blocked(sfb1_ref_cpu[:, :, l_idx])
scale_b2 = to_blocked(sfb2_ref_cpu[:, :, l_idx])
res1 = torch._scaled_mm(a_ref[:, :, l_idx], b1_ref[:, :, l_idx].transpose(0, 1),
scale_a.cuda(), scale_b1.cuda(), bias=None, out_dtype=torch.float32)
ref1[:, :, l_idx] = res1
res2 = torch._scaled_mm(a_ref[:, :, l_idx], b2_ref[:, :, l_idx].transpose(0, 1),
scale_a.cuda(), scale_b2.cuda(), bias=None, out_dtype=torch.float32)
ref2[:, :, l_idx] = res2
c_ref = (torch.nn.functional.silu(ref1) * ref2).to(torch.float16)
return c_ref
def generate_input(m: int, n: int, k: int, l: int, seed: int):
torch.manual_seed(seed)
def create_fp4_tensors(l, mn, k):
ref_i8 = torch.randint(255, size=(l, mn, k // 2), dtype=torch.uint8, device="cuda")
ref_i8 = ref_i8 & 0b1011_1011
return ref_i8.permute(1, 2, 0).view(torch.float4_e2m1fn_x2)
a_ref = create_fp4_tensors(l, m, k).view(torch.float4_e2m1fn_x2)
b1_ref = create_fp4_tensors(l, n, k).view(torch.float4_e2m1fn_x2)
b2_ref = create_fp4_tensors(l, n, k).view(torch.float4_e2m1fn_x2)
c_ref = torch.randn((l, m, n), dtype=torch.float16, device="cuda").permute(1, 2, 0)
def create_scale_factor_tensors(l, mn, sf_k):
ref_shape = (l, mn, sf_k)
ref_f8_random_fp32 = torch.rand(ref_shape, dtype=torch.float32, device='cuda')
ref_f8_torch_tensor = ref_f8_random_fp32.to(dtype=torch.float8_e4m3fn).permute(1, 2, 0)
atom_m, atom_k = (32, 4), 4
mma_shape = (l, ceil_div(mn, atom_m[0] * atom_m[1]), ceil_div(sf_k, atom_k), atom_m[0], atom_m[1], atom_k)
rand_int_tensor = torch.empty(mma_shape, dtype=torch.int8, device='cuda')
reordered = rand_int_tensor.to(dtype=torch.float8_e4m3fn).permute(3, 4, 1, 5, 2, 0)
i_idx, j_idx, b_idx = torch.arange(mn, device='cuda'), torch.arange(sf_k, device='cuda'), torch.arange(l, device='cuda')
i_grid, j_grid, b_grid = torch.meshgrid(i_idx, j_idx, b_idx, indexing='ij')
mm, mm32, mm4 = i_grid // (atom_m[0] * atom_m[1]), i_grid % atom_m[0], (i_grid % 128) // atom_m[0]
kk, kk4 = j_grid // atom_k, j_grid % atom_k
reordered[mm32, mm4, mm, kk4, kk, b_grid] = ref_f8_torch_tensor[i_grid, j_grid, b_grid]
return ref_f8_torch_tensor.cpu(), reordered
sf_k = ceil_div(k, sf_vec_size)
sfa_ref_cpu, sfa_perm = create_scale_factor_tensors(l, m, sf_k)
sfb1_ref_cpu, sfb1_perm = create_scale_factor_tensors(l, n, sf_k)
sfb2_ref_cpu, sfb2_perm = create_scale_factor_tensors(l, n, sf_k)
return (a_ref, b1_ref, b2_ref, sfa_ref_cpu.to("cuda"), sfb1_ref_cpu.to("cuda"), sfb2_ref_cpu.to("cuda"), sfa_perm, sfb1_perm, sfb2_perm, c_ref)
check_implementation = make_match_reference(ref_kernel, rtol=1e-03, atol=1e-03)
# Optimized dual GEMM kernel with improved memory access
CUDA_SRC = r"""
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;
__device__ inline constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; };
__device__ uint32_t elect_sync() {
uint32_t pred = 0;
asm volatile("{\n\t.reg .pred %%px;\n\telect.sync _|%%px, %1;\n\t@%%px mov.s32 %0, 1;\n\t}" : "+r"(pred) : "r"(0xFFFFFFFF));
return pred;
}
__device__ inline void mbarrier_init(int mbar_addr, int count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
}
__device__ void mbarrier_wait(int mbar_addr, int phase) {
uint32_t ticks = 0x989680;
asm volatile("{\n\t.reg .pred P1;\n\tLAB_WAIT:\n\tmbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t@P1 bra.uni DONE;\n\tbra.uni LAB_WAIT;\n\tDONE:\n\t}" :: "r"(mbar_addr), "r"(phase), "r"(ticks));
}
__device__ inline void mbarrier_inval(int mbar_addr) {
asm volatile("mbarrier.inval.shared::cta.b64 [%0];" :: "r"(mbar_addr));
}
__device__ inline void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy) {
asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;" :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy));
}
__device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
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"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy) : "memory");
}
// TMA with multicast - broadcasts to all CTAs in cluster matching the mask
__device__ inline void tma_3d_gmem2smem_multicast(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint16_t ctamask, uint64_t cache_policy) {
asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::1.multicast::cluster.L2::cache_hint [%0], [%1, {%2, %3, %4}], [%5], %6, %7;" :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "h"(ctamask), "l"(cache_policy) : "memory");
}
// Get CTA rank within cluster
__device__ inline uint32_t get_cta_rank_in_cluster() {
uint32_t rank;
asm volatile("mov.u32 %0, %%clusterid;" : "=r"(rank));
return rank;
}
// Cluster-wide barrier arrive
__device__ inline void cluster_arrive() {
asm volatile("barrier.cluster.arrive;");
}
// Cluster-wide barrier wait
__device__ inline void cluster_wait() {
asm volatile("barrier.cluster.wait;");
}
__device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
}
__device__ inline void tcgen05_mma_nvfp4(uint64_t a_desc, uint64_t b_desc, uint32_t i_desc, int scale_A, int scale_B, int enable_d, int d_tmem = 0) {
asm volatile("{\n\t.reg .pred p;\n\tsetp.ne.b32 p, %6, 0;\n\ttcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\n\t}" :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc), "r"(scale_A), "r"(scale_B), "r"(enable_d));
}
struct SHAPE { static constexpr char _16x256b[] = ".16x256b"; };
struct NUM { static constexpr char x8[] = ".x8"; };
__device__ inline void tcgen05_ld_16x256bx8(float *tmp, int row, int col) {
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"(tmp[0]),"=f"(tmp[1]),"=f"(tmp[2]),"=f"(tmp[3]),"=f"(tmp[4]),"=f"(tmp[5]),"=f"(tmp[6]),"=f"(tmp[7]),
"=f"(tmp[8]),"=f"(tmp[9]),"=f"(tmp[10]),"=f"(tmp[11]),"=f"(tmp[12]),"=f"(tmp[13]),"=f"(tmp[14]),"=f"(tmp[15]),
"=f"(tmp[16]),"=f"(tmp[17]),"=f"(tmp[18]),"=f"(tmp[19]),"=f"(tmp[20]),"=f"(tmp[21]),"=f"(tmp[22]),"=f"(tmp[23]),
"=f"(tmp[24]),"=f"(tmp[25]),"=f"(tmp[26]),"=f"(tmp[27]),"=f"(tmp[28]),"=f"(tmp[29]),"=f"(tmp[30]),"=f"(tmp[31])
: "r"((row << 16) | col));
}
void check_cu(CUresult err) { if (err != CUDA_SUCCESS) { const char *msg; cuGetErrorString(err, &msg); TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", msg); } }
void init_AB_tmap(CUtensorMap *tmap, const char *ptr, uint64_t gh, uint64_t gw, uint32_t sh, uint32_t sw) {
uint64_t globalDim[3] = {256, gh, gw / 256};
uint64_t globalStrides[2] = {gw / 2, 128};
uint32_t boxDim[3] = {256, sh, sw / 256};
uint32_t elementStrides[3] = {1, 1, 1};
check_cu(cuTensorMapEncodeTiled(tmap, CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B, 3, (void*)ptr, globalDim, globalStrides, boxDim, elementStrides, CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B, CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE));
}
// Non-persistent dual GEMM kernel - optimized version
template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void dual_gemm_kernel(
const __grid_constant__ CUtensorMap A_tmap,
const __grid_constant__ CUtensorMap B1_tmap,
const __grid_constant__ CUtensorMap B2_tmap,
const char *SFA_ptr, const char *SFB1_ptr, const char *SFB2_ptr,
half *C_ptr, int M, int N, int grid_n) {
const int tid = threadIdx.x, bid_m = blockIdx.x, bid_n = blockIdx.y;
const int lane_id = tid % WARP_SIZE, warp_id = tid / WARP_SIZE;
constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;
const int off_m = bid_m * BLOCK_M, off_n = bid_n * BLOCK_N;
extern __shared__ __align__(1024) char smem_ptr[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
constexpr int A_size = BLOCK_M * BLOCK_K / 2, B_size = BLOCK_N * BLOCK_K / 2;
constexpr int SFA_size = 128 * BLOCK_K / 16, SFB_size = 128 * BLOCK_K / 16;
constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;
__shared__ int64_t mbars[NUM_STAGES * 2 + 1];
const int tma_mbar = static_cast<int>(__cvta_generic_to_shared(mbars));
const int mma_mbar = tma_mbar + NUM_STAGES * 8, mainloop_mbar = mma_mbar + NUM_STAGES * 8;
constexpr int ACC1 = 0, ACC2 = BLOCK_N, SFA_tm = 2 * BLOCK_N;
constexpr int SFB1_tm = SFA_tm + 4 * (BLOCK_K / MMA_K), SFB2_tm = SFB1_tm + 4 * (BLOCK_K / MMA_K);
const int num_iters = K / BLOCK_K;
const int sfb_n_offset = (BLOCK_N == 64) ? (bid_n % 2) * 2 : 0;
// Initialize barriers
if (warp_id == 0 && elect_sync()) {
for (int i = 0; i < NUM_STAGES * 2 + 1; i++) mbarrier_init(tma_mbar + i * 8, 1);
asm volatile("fence.mbarrier_init.release.cluster;");
} else if (warp_id == 1) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 4));
}
__syncthreads();
// TMA warp
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
uint64_t cache_A = EVICT_LAST, cache_B = EVICT_FIRST;
auto issue_tma = [&](int iter_k, int stage) {
const int mbar = tma_mbar + stage * 8;
const int A_sm = smem + stage * STAGE_SIZE, B1_sm = A_sm + A_size, B2_sm = B1_sm + B_size;
const int SFA_sm = B2_sm + B_size, SFB1_sm = SFA_sm + SFA_size, SFB2_sm = SFB1_sm + SFB_size;
const int off_k = iter_k * BLOCK_K;
tma_3d_gmem2smem(A_sm, &A_tmap, 0, off_m, off_k / 256, mbar, cache_A);
tma_3d_gmem2smem(B1_sm, &B1_tmap, 0, off_n, off_k / 256, mbar, cache_B);
tma_3d_gmem2smem(B2_sm, &B2_tmap, 0, off_n, off_k / 256, mbar, cache_B);
const int rest_k = K / 16 / 4;
tma_gmem2smem(SFA_sm, SFA_ptr + ((off_m / 128) * rest_k + off_k / 64) * 512, SFA_size, mbar, cache_A);
tma_gmem2smem(SFB1_sm, SFB1_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512, SFB_size, mbar, cache_B);
tma_gmem2smem(SFB2_sm, SFB2_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512, SFB_size, mbar, cache_B);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;" :: "r"(mbar), "r"(STAGE_SIZE) : "memory");
};
for (int k = 0; k < NUM_STAGES; k++) issue_tma(k, k);
for (int k = NUM_STAGES; k < num_iters; k++) {
int stage = k % NUM_STAGES;
mbarrier_wait(mma_mbar + stage * 8, (k / NUM_STAGES - 1) % 2);
issue_tma(k, stage);
}
}
// MMA warp
else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)BLOCK_N >> 3U << 17U) | (1U << 27U);
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
int stage = iter_k % NUM_STAGES;
mbarrier_wait(tma_mbar + stage * 8, (iter_k / NUM_STAGES) % 2);
const int A_sm = smem + stage * STAGE_SIZE, B1_sm = A_sm + A_size, B2_sm = B1_sm + B_size;
const int SFA_sm = B2_sm + B_size, SFB1_sm = SFA_sm + SFA_size, SFB2_sm = SFB1_sm + SFB_size;
auto make_AB = [](int addr) -> uint64_t { return desc_encode(addr) | (desc_encode(8*128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL); };
auto make_SF = [](int addr) -> uint64_t { return desc_encode(addr) | (desc_encode(8*16) << 32ULL) | (1ULL << 46ULL); };
uint64_t SFA_d = make_SF(0) + ((uint64_t)SFA_sm >> 4ULL);
uint64_t SFB1_d = make_SF(0) + ((uint64_t)SFB1_sm >> 4ULL);
uint64_t SFB2_d = make_SF(0) + ((uint64_t)SFB2_sm >> 4ULL);
#pragma unroll
for (int k = 0; k < BLOCK_K / MMA_K; k++) {
tcgen05_cp_nvfp4(SFA_tm + k * 4, SFA_d + (uint64_t)k * 32ULL);
tcgen05_cp_nvfp4(SFB1_tm + k * 4, SFB1_d + (uint64_t)k * 32ULL);
tcgen05_cp_nvfp4(SFB2_tm + k * 4, SFB2_d + (uint64_t)k * 32ULL);
}
#pragma unroll
for (int k = 0; k < BLOCK_K / MMA_K; k++) {
uint64_t a_d = make_AB(A_sm + k * 32);
uint64_t b1_d = make_AB(B1_sm + k * 32);
uint64_t b2_d = make_AB(B2_sm + k * 32);
int sA = SFA_tm + k * 4;
int sB1 = SFB1_tm + k * 4 + sfb_n_offset;
int sB2 = SFB2_tm + k * 4 + sfb_n_offset;
int en = (k == 0) ? iter_k : 1;
tcgen05_mma_nvfp4(a_d, b1_d, i_desc, sA, sB1, en, ACC1);
tcgen05_mma_nvfp4(a_d, b2_d, i_desc, sA, sB2, en, ACC2);
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];" :: "r"(mma_mbar + stage * 8) : "memory");
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];" :: "r"(mainloop_mbar) : "memory");
}
// Epilogue warps - 4 warps for better parallelism
else if (tid < BLOCK_M) {
mbarrier_wait(mainloop_mbar, 0);
asm volatile("tcgen05.fence::after_thread_sync;");
const int lane_row = lane_id / 4, lane_col = (lane_id % 4) * 2;
half* __restrict__ C_base = C_ptr + off_m * N + off_n + lane_col;
constexpr int CHUNK = 64;
constexpr int NUM_CHUNKS = BLOCK_N / CHUNK;
#pragma unroll
for (int m = 0; m < 2; m++) {
const int row_base = warp_id * 32 + m * 16;
half* __restrict__ C_row0 = C_base + (row_base + lane_row) * N;
half* __restrict__ C_row8 = C_row0 + 8 * N;
#pragma unroll
for (int c = 0; c < NUM_CHUNKS; c++) {
float tmp1[32], tmp2[32];
tcgen05_ld_16x256bx8(tmp1, row_base, ACC1 + c * CHUNK);
tcgen05_ld_16x256bx8(tmp2, row_base, ACC2 + c * CHUNK);
asm volatile("tcgen05.wait::ld.sync.aligned;");
#pragma unroll
for (int i = 0; i < 8; i++) {
float x0 = tmp1[i*4], x1 = tmp1[i*4+1], x2 = tmp1[i*4+2], x3 = tmp1[i*4+3];
float y0 = tmp2[i*4], y1 = tmp2[i*4+1], y2 = tmp2[i*4+2], y3 = tmp2[i*4+3];
float e0 = __expf(-x0), e1 = __expf(-x1), e2 = __expf(-x2), e3 = __expf(-x3);
float s0 = x0 / (1.f + e0) * y0, s1 = x1 / (1.f + e1) * y1;
float s2 = x2 / (1.f + e2) * y2, s3 = x3 / (1.f + e3) * y3;
*reinterpret_cast<half2*>(C_row0 + c * CHUNK + i * 8) = __float22half2_rn({s0, s1});
*reinterpret_cast<half2*>(C_row8 + c * CHUNK + i * 8) = __float22half2_rn({s2, s3});
}
}
}
}
// Cleanup
if (warp_id == 0 && tid == 0) {
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 4));
}
}
// Cluster-aware kernel - each CTA in cluster processes different N tile
// Cluster shape: (1, CLUSTER_N, 1) - CTAs along N dimension
template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CLUSTER_N>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void dual_gemm_cluster_kernel(
const __grid_constant__ CUtensorMap A_tmap,
const __grid_constant__ CUtensorMap B1_tmap,
const __grid_constant__ CUtensorMap B2_tmap,
const char *SFA_ptr, const char *SFB1_ptr, const char *SFB2_ptr,
half *C_ptr, int M, int N, int grid_n) {
const int tid = threadIdx.x;
const int lane_id = tid % WARP_SIZE, warp_id = tid / WARP_SIZE;
constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;
// Use standard block indices - no cluster complexity for now
// Just verify the kernel works with standard launch
const int bid_m = blockIdx.x;
const int bid_n = blockIdx.y;
const int off_m = bid_m * BLOCK_M, off_n = bid_n * BLOCK_N;
extern __shared__ __align__(1024) char smem_ptr[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
constexpr int A_size = BLOCK_M * BLOCK_K / 2, B_size = BLOCK_N * BLOCK_K / 2;
constexpr int SFA_size = 128 * BLOCK_K / 16, SFB_size = 128 * BLOCK_K / 16;
constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;
__shared__ int64_t mbars[NUM_STAGES * 2 + 1];
const int tma_mbar = static_cast<int>(__cvta_generic_to_shared(mbars));
const int mma_mbar = tma_mbar + NUM_STAGES * 8, mainloop_mbar = mma_mbar + NUM_STAGES * 8;
constexpr int ACC1 = 0, ACC2 = BLOCK_N, SFA_tm = 2 * BLOCK_N;
constexpr int SFB1_tm = SFA_tm + 4 * (BLOCK_K / MMA_K), SFB2_tm = SFB1_tm + 4 * (BLOCK_K / MMA_K);
const int num_iters = K / BLOCK_K;
const int sfb_n_offset = (BLOCK_N == 64) ? (bid_n % 2) * 2 : 0;
// Initialize barriers
if (warp_id == 0 && elect_sync()) {
for (int i = 0; i < NUM_STAGES * 2 + 1; i++) mbarrier_init(tma_mbar + i * 8, 1);
asm volatile("fence.mbarrier_init.release.cluster;");
} else if (warp_id == 1) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 4));
}
__syncthreads();
// TMA warp
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
uint64_t cache_A = EVICT_LAST, cache_B = EVICT_FIRST;
auto issue_tma = [&](int iter_k, int stage) {
const int mbar = tma_mbar + stage * 8;
const int A_sm = smem + stage * STAGE_SIZE, B1_sm = A_sm + A_size, B2_sm = B1_sm + B_size;
const int SFA_sm = B2_sm + B_size, SFB1_sm = SFA_sm + SFA_size, SFB2_sm = SFB1_sm + SFB_size;
const int off_k = iter_k * BLOCK_K;
tma_3d_gmem2smem(A_sm, &A_tmap, 0, off_m, off_k / 256, mbar, cache_A);
tma_3d_gmem2smem(B1_sm, &B1_tmap, 0, off_n, off_k / 256, mbar, cache_B);
tma_3d_gmem2smem(B2_sm, &B2_tmap, 0, off_n, off_k / 256, mbar, cache_B);
const int rest_k = K / 16 / 4;
tma_gmem2smem(SFA_sm, SFA_ptr + ((off_m / 128) * rest_k + off_k / 64) * 512, SFA_size, mbar, cache_A);
tma_gmem2smem(SFB1_sm, SFB1_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512, SFB_size, mbar, cache_B);
tma_gmem2smem(SFB2_sm, SFB2_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512, SFB_size, mbar, cache_B);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;" :: "r"(mbar), "r"(STAGE_SIZE) : "memory");
};
for (int k = 0; k < NUM_STAGES; k++) issue_tma(k, k);
for (int k = NUM_STAGES; k < num_iters; k++) {
int stage = k % NUM_STAGES;
mbarrier_wait(mma_mbar + stage * 8, (k / NUM_STAGES - 1) % 2);
issue_tma(k, stage);
}
}
// MMA warp
else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)BLOCK_N >> 3U << 17U) | (1U << 27U);
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
int stage = iter_k % NUM_STAGES;
mbarrier_wait(tma_mbar + stage * 8, (iter_k / NUM_STAGES) % 2);
const int A_sm = smem + stage * STAGE_SIZE, B1_sm = A_sm + A_size, B2_sm = B1_sm + B_size;
const int SFA_sm = B2_sm + B_size, SFB1_sm = SFA_sm + SFA_size, SFB2_sm = SFB1_sm + SFB_size;
auto make_AB = [](int addr) -> uint64_t { return desc_encode(addr) | (desc_encode(8*128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL); };
auto make_SF = [](int addr) -> uint64_t { return desc_encode(addr) | (desc_encode(8*16) << 32ULL) | (1ULL << 46ULL); };
uint64_t SFA_d = make_SF(0) + ((uint64_t)SFA_sm >> 4ULL);
uint64_t SFB1_d = make_SF(0) + ((uint64_t)SFB1_sm >> 4ULL);
uint64_t SFB2_d = make_SF(0) + ((uint64_t)SFB2_sm >> 4ULL);
#pragma unroll
for (int k = 0; k < BLOCK_K / MMA_K; k++) {
tcgen05_cp_nvfp4(SFA_tm + k * 4, SFA_d + (uint64_t)k * 32ULL);
tcgen05_cp_nvfp4(SFB1_tm + k * 4, SFB1_d + (uint64_t)k * 32ULL);
tcgen05_cp_nvfp4(SFB2_tm + k * 4, SFB2_d + (uint64_t)k * 32ULL);
}
#pragma unroll
for (int k = 0; k < BLOCK_K / MMA_K; k++) {
uint64_t a_d = make_AB(A_sm + k * 32);
uint64_t b1_d = make_AB(B1_sm + k * 32);
uint64_t b2_d = make_AB(B2_sm + k * 32);
int sA = SFA_tm + k * 4;
int sB1 = SFB1_tm + k * 4 + sfb_n_offset;
int sB2 = SFB2_tm + k * 4 + sfb_n_offset;
int en = (k == 0) ? iter_k : 1;
tcgen05_mma_nvfp4(a_d, b1_d, i_desc, sA, sB1, en, ACC1);
tcgen05_mma_nvfp4(a_d, b2_d, i_desc, sA, sB2, en, ACC2);
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];" :: "r"(mma_mbar + stage * 8) : "memory");
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];" :: "r"(mainloop_mbar) : "memory");
}
// Epilogue warps
else if (tid < BLOCK_M) {
mbarrier_wait(mainloop_mbar, 0);
asm volatile("tcgen05.fence::after_thread_sync;");
const int lane_row = lane_id / 4, lane_col = (lane_id % 4) * 2;
half* __restrict__ C_base = C_ptr + off_m * N + off_n + lane_col;
constexpr int CHUNK = 64;
constexpr int NUM_CHUNKS = BLOCK_N / CHUNK;
#pragma unroll
for (int m = 0; m < 2; m++) {
const int row_base = warp_id * 32 + m * 16;
half* __restrict__ C_row0 = C_base + (row_base + lane_row) * N;
half* __restrict__ C_row8 = C_row0 + 8 * N;
#pragma unroll
for (int c = 0; c < NUM_CHUNKS; c++) {
float tmp1[32], tmp2[32];
tcgen05_ld_16x256bx8(tmp1, row_base, ACC1 + c * CHUNK);
tcgen05_ld_16x256bx8(tmp2, row_base, ACC2 + c * CHUNK);
asm volatile("tcgen05.wait::ld.sync.aligned;");
#pragma unroll
for (int i = 0; i < 8; i++) {
float x0 = tmp1[i*4], x1 = tmp1[i*4+1], x2 = tmp1[i*4+2], x3 = tmp1[i*4+3];
float y0 = tmp2[i*4], y1 = tmp2[i*4+1], y2 = tmp2[i*4+2], y3 = tmp2[i*4+3];
float e0 = __expf(-x0), e1 = __expf(-x1), e2 = __expf(-x2), e3 = __expf(-x3);
float s0 = x0 / (1.f + e0) * y0, s1 = x1 / (1.f + e1) * y1;
float s2 = x2 / (1.f + e2) * y2, s3 = x3 / (1.f + e3) * y3;
*reinterpret_cast<half2*>(C_row0 + c * CHUNK + i * 8) = __float22half2_rn({s0, s1});
*reinterpret_cast<half2*>(C_row8 + c * CHUNK + i * 8) = __float22half2_rn({s2, s3});
}
}
}
}
// Cleanup
if (warp_id == 0 && tid == 0) {
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 4));
}
}
// Launch function for cluster kernel
template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CLUSTER_N>
at::Tensor launch_cluster(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);
CUtensorMap A_tm, B1_tm, B2_tm;
init_AB_tmap(&A_tm, (const char*)A.data_ptr(), M, K, BLOCK_M, BLOCK_K);
init_AB_tmap(&B1_tm, (const char*)B1.data_ptr(), N, K, BLOCK_N, BLOCK_K);
init_AB_tmap(&B2_tm, (const char*)B2.data_ptr(), N, K, BLOCK_N, BLOCK_K);
// Grid is in CTAs, not clusters. With __cluster_dims__, runtime groups CTAs.
int grid_m = M / BLOCK_M;
int grid_n = N / BLOCK_N; // Total CTAs, runtime groups into clusters
dim3 grid(grid_m, grid_n);
int tb = BLOCK_M + 2 * WARP_SIZE;
int smem = ((BLOCK_M + 2 * BLOCK_N) * (BLOCK_K / 2) + 128 * (BLOCK_K / 16) * 3) * NUM_STAGES;
auto kernel = dual_gemm_cluster_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, CLUSTER_N>;
if (smem > 48000) cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
// Use cudaLaunchKernelEx for cluster launch
cudaLaunchConfig_t config = {};
config.gridDim = grid;
config.blockDim = dim3(tb, 1, 1);
config.dynamicSmemBytes = smem;
cudaLaunchAttribute attrs[1];
attrs[0].id = cudaLaunchAttributeClusterDimension;
attrs[0].val.clusterDim = {1, (unsigned)CLUSTER_N, 1};
config.attrs = attrs;
config.numAttrs = 1;
cudaLaunchKernelEx(&config, kernel, A_tm, B1_tm, B2_tm, (const char*)SFA.data_ptr(), (const char*)SFB1.data_ptr(), (const char*)SFB2.data_ptr(), (half*)C.data_ptr(), M, N, grid_n);
return C;
}
template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
at::Tensor launch(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);
CUtensorMap A_tm, B1_tm, B2_tm;
init_AB_tmap(&A_tm, (const char*)A.data_ptr(), M, K, BLOCK_M, BLOCK_K);
init_AB_tmap(&B1_tm, (const char*)B1.data_ptr(), N, K, BLOCK_N, BLOCK_K);
init_AB_tmap(&B2_tm, (const char*)B2.data_ptr(), N, K, BLOCK_N, BLOCK_K);
int grid_m = M / BLOCK_M, grid_n = N / BLOCK_N;
dim3 grid(grid_m, grid_n);
int tb = BLOCK_M + 2 * WARP_SIZE;
int smem = ((BLOCK_M + 2 * BLOCK_N) * (BLOCK_K / 2) + 128 * (BLOCK_K / 16) * 3) * NUM_STAGES;
auto kernel = dual_gemm_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
if (smem > 48000) cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
kernel<<<grid, tb, smem>>>(A_tm, B1_tm, B2_tm, (const char*)SFA.data_ptr(), (const char*)SFB1.data_ptr(), (const char*)SFB2.data_ptr(), (half*)C.data_ptr(), M, N, grid_n);
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), K = A.size(1) * 2;
// Use non-cluster kernel - 5 stages for all cases
#define L64_5(K_) else if (K == K_) C = launch<K_, 128, 64, 256, 5>(A, B1, B2, SFA, SFB1, SFB2, C);
#define L128_4(K_) else if (K == K_) C = launch<K_, 128, 128, 256, 4>(A, B1, B2, SFA, SFB1, SFB2, C);
if (M <= 256) {
if (false) {}
L64_5(16384) L64_5(7168) L64_5(4096) L64_5(2048) L64_5(2304) L64_5(1536) L64_5(512) L64_5(256)
} else {
if (false) {}
L128_4(16384) L128_4(7168) L128_4(4096) L128_4(2048) L128_4(2304) L128_4(1536) L128_4(512) L128_4(256)
}
#undef L64_5
#undef L128_4
return C;
}
TORCH_LIBRARY(dual_gemm_mod, 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_0", 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", "--relocatable-device-code=false", "-lineinfo"], extra_ldflags=["-lcuda"])
dual_gemm = torch.ops.dual_gemm_mod.dual_gemm
def custom_kernel(data: input_t) -> output_t:
a, b1, b2, sfa, sfb1, sfb2, sfa_perm, sfb1_perm, sfb2_perm, c = data
m, n, l = c.shape
if l == 1:
out = torch.empty((m, n), dtype=torch.float16, device="cuda")
return dual_gemm(a[:,:,0], b1[:,:,0], b2[:,:,0], sfa_perm[:,:,:,:,:,0], sfb1_perm[:,:,:,:,:,0], sfb2_perm[:,:,:,:,:,0], out).unsqueeze(2)
else:
out = torch.empty((m, n, l), dtype=torch.float16, device="cuda")
for i in range(l):
dual_gemm(a[:,:,i], b1[:,:,i], b2[:,:,i], sfa_perm[:,:,:,:,:,i], sfb1_perm[:,:,:,:,:,i], sfb2_perm[:,:,:,:,:,i], out[:,:,i])
return out
scrolls · 538 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