submission 488917
kathsucurry · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 509 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-488917?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:c5801011f6a2840317ba7221cb24cf545629e27c7fa5d65caeb70b533ee7289a
license declaredunknown
license concludedunknown
authorskathsucurry
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
constexpr int MMA_K = 64; // FP4 MMA K-dimension size.mbarrier
__device__ inline void mbarrier_init(int mbar_addr, int count) {shared-memory
extern __shared__ __align__(1024) char smem[];stages = 4
constexpr int NUM_STAGES = 4;tcgen05
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128tma
TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);vector-width = half2
half2 out[4];Kernel source
submission.py509 lines
import os
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
cuda_source = r"""
#include <cuda_fp16.h>
#include <cudaTypedefs.h>
#include <torch/extension.h>
#include <torch/library.h>
#define WARP_SIZE 32
void check_cu_error(CUresult error) {
if (error == CUDA_SUCCESS) return;
const char *error_msg_ptr;
if (cuGetErrorString(error, &error_msg_ptr) != CUDA_SUCCESS)
error_msg_ptr = "unable to get error string";
TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);
}
template <const int NUM_ELEMENTS>
inline void create_tmap_descriptor(
CUtensorMap *tmap,
const char *ptr,
uint64_t global_height, uint64_t global_width,
uint32_t shared_height, uint32_t shared_width,
CUtensorMapSwizzle swizzle_type
) {
/*
The goal is to transfer multiple of [shared_height, NUM_ELEMENTS] spanning
[shared_height, shared_width] --> [shared_width / NUM_ELEMENTS, shared_height, NUM_ELEMENTS].
Code taken and modified from:
- https://docs.nvidia.com/cuda/cuda-programming-guide/04-special-topics/async-copies.html#using-tma-to-transfer-multi-dimensional-arrays.
- https://gau-nernst.github.io/tcgen05/
*/
constexpr int rank{3};
uint64_t global_dim[rank] = {NUM_ELEMENTS, global_height, global_width / (uint64_t) NUM_ELEMENTS};
// 4 bits would be 1/2 bytes.
uint64_t global_strides[rank - 1] = {global_width / 2, NUM_ELEMENTS / 2};
uint32_t box_dim[rank] = {NUM_ELEMENTS, shared_height, shared_width / NUM_ELEMENTS};
uint32_t element_strides[rank] = {1, 1, 1};
auto error = cuTensorMapEncodeTiled(
tmap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank,
(void *)ptr,
global_dim,
global_strides,
box_dim,
element_strides,
// Interleave patterns can be used to accelerate loading of values that
// are less than 4 bytes long.
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
swizzle_type,
// L2 Promotion can be used to widen the effect of a cache-policy to a wider
// set of L2 cache lines.
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
// Any element that is outside of bounds will be set to zero by the TMA transfer.
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
check_cu_error(error);
}
// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
__device__ inline uint32_t elect_sync() {
uint32_t pred = 0;
asm volatile(
"{\n\t"
".reg .pred %%px;\n\t"
"elect.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));
}
// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cutlass/arch/barrier.h#L408
__device__ inline void mbarrier_wait(int mbar_addr, int phase) {
uint32_t ticks = 0x989680; // arbitrarily large timer value.
asm volatile(
"{\n\t"
".reg .pred P1;\n\t"
"LAB_WAIT:\n\t"
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
"@P1 bra.uni DONE;\n\t"
"bra.uni LAB_WAIT;\n\t"
"DONE:\n\t"
"}" ::"r"(mbar_addr),
"r"(phase), "r"(ticks));
}
__device__ inline
void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr) {
asm volatile(
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"
:: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr));
}
template <int CTA_GROUP = 1>
__device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr)
{
// when CTA_GROUP=1, we can use .shared::cta instead.
// but .shared::cluster doesn't seem to be slower, so always use it unconditionally here.
// .cta_group::2 allows mbar_addr and dst to be in different CTA's smem.
asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%6 "
"[%0], [%1, {%2, %3, %4}], [%5];" ::"r"(dst),
"l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "n"(CTA_GROUP)
: "memory");
}
// Encodes the matrix descriptor and ensures 64 bits.
__device__ inline
constexpr uint64_t encode_descriptor(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }
// Copy scale factors from shared memory to tensor memory.
// .32x128b = 32 rows x 16 bytes = one scale factor tile for one MMA.
// .warpx4 duplicates data across all 32-lane groups.
__device__ inline
void copy_sf_smem2tmem(int taddr, uint64_t s_desc) {
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
}
// Issue FP4 MMA instruction with block scaling.
// d_tmem=0: accumulator always starts at TMEM column 0.
// enable_input_d: 0 = clear accumulator, nonzero = accumulate.
__device__ inline
void run_mma_nvfp4(
uint64_t a_desc,
uint64_t b_desc,
uint32_t i_desc,
int scale_A_tmem,
int scale_B_tmem,
int enable_input_d
) {
const int d_tmem = 0;
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p, %6, 0;\n\t"
"tcgen05.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_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
);
}
constexpr int MMA_K = 64; // FP4 MMA K-dimension size.
constexpr int NUM_STAGES = 4;
template <const int NUM_THREADS, const int BLOCK_M, const int BLOCK_N, const int BLOCK_K>
__global__
__launch_bounds__(NUM_THREADS) void kernel_v04_pipeline(
const __grid_constant__ CUtensorMap A_tmap,
const __grid_constant__ CUtensorMap B_tmap,
const char *SFA,
const char *SFB,
half *C,
int M,
int N,
int K
) {
const int thread_idx{static_cast<int>(threadIdx.x)};
const int block_idx{static_cast<int>(blockIdx.x)};
const int warp_idx{thread_idx / WARP_SIZE};
const int grid_dim_n{N / BLOCK_N};
const int block_idx_m{block_idx / grid_dim_n};
const int block_idx_n{block_idx % grid_dim_n};
const int offset_m{block_idx_m * BLOCK_M};
const int offset_n{block_idx_n * BLOCK_N};
// Double-buffered shared memory layout:
// [buf0: A | B | SFA | SFB | buf1: A | B | SFA | SFB]
constexpr int SF_size = 512 * BLOCK_K / MMA_K;
constexpr int BUF_SIZE = BLOCK_M * BLOCK_K / 2 + BLOCK_N * BLOCK_K / 2 + 2 * SF_size;
extern __shared__ __align__(1024) char smem[];
const int smem_base{static_cast<int>(__cvta_generic_to_shared(smem))};
int A_smem[NUM_STAGES], B_smem[NUM_STAGES], SFA_smem[NUM_STAGES], SFB_smem[NUM_STAGES];
for (int s{0}; s < NUM_STAGES; ++s) {
A_smem[s] = smem_base + s * BUF_SIZE;
B_smem[s] = A_smem[s] + BLOCK_M * BLOCK_K / 2;
SFA_smem[s] = B_smem[s] + BLOCK_N * BLOCK_K / 2;
SFB_smem[s] = SFA_smem[s] + SF_size;
}
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ uint64_t mbars[NUM_STAGES];
int mbar_addrs[NUM_STAGES];
for (int s{0}; s < NUM_STAGES; ++s)
mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&mbars[s]));
__shared__ int tmem_addr[1];
// TMEM layout:
// Columns [0, BLOCK_N) : accumulator D
// Columns [BLOCK_N, BLOCK_N + 4*BLOCK_K/MMA_K) : SFA
// Columns [BLOCK_N + 4*BLOCK_K/MMA_K, ...) : SFB
constexpr int SFA_tmem_start_col = BLOCK_N;
constexpr int SFB_tmem_start_col = SFA_tmem_start_col + 4 * (BLOCK_K / MMA_K);
constexpr int TMEM_COLS = BLOCK_N * 2; // BLOCK_N + 8 * (BLOCK_K / MMA_K), but it has to be a power of 2.
if (warp_idx == 0 && elect_sync()) {
for (int s{0}; s < NUM_STAGES; ++s)
mbarrier_init(mbar_addrs[s], 1);
asm volatile("fence.mbarrier_init.release.cluster;");
} else if (warp_idx == 1) {
// Allocate TMEM for accumulator + scale factors.
const int addr{static_cast<int>(__cvta_generic_to_shared(tmem_addr))};
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
::"r"(addr), "r"(TMEM_COLS));
}
__syncthreads();
const int taddr{tmem_addr[0]};
int phase[NUM_STAGES] = {0, 0};
// Instruction descriptor for tcgen05.mma.kind::mxf4nvf4
// atype=E2M1 (1), btype=E2M1 (1), MMA_N and MMA_M encoded in upper bits
constexpr uint32_t i_desc = (1U << 7U) // atype=E2M1
| (1U << 10U) // btype=E2M1
| ((uint32_t)BLOCK_N >> 3U << 17U) // MMA_N
| ((uint32_t)BLOCK_M >> 7U << 27U) // MMA_M
;
const int num_iters{K / BLOCK_K};
const int rest_k{K / 16 / 4}; // number of K-atoms (each covers 64 K-elements)
constexpr int cp_size = (BLOCK_M + BLOCK_N) * BLOCK_K / 2 + 2 * SF_size;
// === Prologue: load first tile into buf[0] ===
if (warp_idx == 0 && elect_sync()) {
const int off_k{0};
tma_3d_gmem2smem(A_smem[0], &A_tmap, 0, offset_m, off_k / 256, mbar_addrs[0]);
tma_3d_gmem2smem(B_smem[0], &B_tmap, 0, offset_n, off_k / 256, mbar_addrs[0]);
const char *SFA_src = SFA + ((offset_m / 128) * rest_k + off_k / (16 * 4)) * 512;
const char *SFB_src = SFB + ((offset_n / 128) * rest_k + off_k / (16 * 4)) * 512;
tma_gmem2smem(SFA_smem[0], SFA_src, SF_size, mbar_addrs[0]);
tma_gmem2smem(SFB_smem[0], SFB_src, SF_size, mbar_addrs[0]);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
::"r"(mbar_addrs[0]), "r"(cp_size) : "memory");
}
// Wait for first TMA to complete.
mbarrier_wait(mbar_addrs[0], phase[0]);
phase[0] ^= 1;
// === Main loop: overlap TMA load of next tile with MMA compute of current tile ===
for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
const int s{iter_k % NUM_STAGES};
if (warp_idx == 0 && elect_sync()) {
// Prefetch next tile into the other buffer (overlaps with MMA below).
if (iter_k + 1 < num_iters) {
const int ns{(iter_k + 1) % NUM_STAGES};
const int next_off_k{(iter_k + 1) * BLOCK_K};
tma_3d_gmem2smem(A_smem[ns], &A_tmap, 0, offset_m, next_off_k / 256, mbar_addrs[ns]);
tma_3d_gmem2smem(B_smem[ns], &B_tmap, 0, offset_n, next_off_k / 256, mbar_addrs[ns]);
const char *SFA_src = SFA + ((offset_m / 128) * rest_k + next_off_k / (16 * 4)) * 512;
const char *SFB_src = SFB + ((offset_n / 128) * rest_k + next_off_k / (16 * 4)) * 512;
tma_gmem2smem(SFA_smem[ns], SFA_src, SF_size, mbar_addrs[ns]);
tma_gmem2smem(SFB_smem[ns], SFB_src, SF_size, mbar_addrs[ns]);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
::"r"(mbar_addrs[ns]), "r"(cp_size) : "memory");
}
// MMA from current buffer.
auto make_desc_AB = [](int addr) -> uint64_t
{
constexpr int SBO = 8 * 256 / 2; // in bytes.
return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
auto make_desc_SF = [](int addr) -> uint64_t
{
const int SBO = 8 * 16; // = 128 bytes
return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL);
};
// Copy scale factors from shared memory to tensor memory.
for (int k{0}; k < BLOCK_K / MMA_K; ++k) {
uint64_t sfa_desc = make_desc_SF(SFA_smem[s] + k * 512);
uint64_t sfb_desc = make_desc_SF(SFB_smem[s] + k * 512);
copy_sf_smem2tmem(SFA_tmem_start_col + k * 4, sfa_desc);
copy_sf_smem2tmem(SFB_tmem_start_col + k * 4, sfb_desc);
}
// One swizzle tile = 256 FP4 elements = 128 bytes.
for (int k1{0}; k1 < BLOCK_K / 256; ++k1) {
for (int k2{0}; k2 < 256 / MMA_K; ++k2) {
uint64_t a_desc{make_desc_AB(A_smem[s] + k1 * BLOCK_M * 128 + k2 * MMA_K / 2)};
uint64_t b_desc{make_desc_AB(B_smem[s] + k1 * BLOCK_N * 128 + k2 * MMA_K / 2)};
int k{k1 * 256 / MMA_K + k2};
const int scale_A_tmem{SFA_tmem_start_col + k * 4};
const int scale_B_tmem{SFB_tmem_start_col + k * 4};
const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
run_mma_nvfp4(a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
}
}
// Signal MMA completion on the current buffer's mbarrier.
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
::"r"(mbar_addrs[s]) : "memory");
}
// Wait for MMA on current buffer to complete.
// This also ensures buf[s] is free to reuse for TMA 2 iterations later.
mbarrier_wait(mbar_addrs[s], phase[s]);
phase[s] ^= 1;
// Wait for next tile's TMA to complete (so next iteration can MMA from it).
if (iter_k + 1 < num_iters) {
const int ns{(iter_k + 1) % NUM_STAGES};
mbarrier_wait(mbar_addrs[ns], phase[ns]);
phase[ns] ^= 1;
}
}
// === Epilogue: Read accumulator from TMEM and store to global memory ===
// PTX docs require this fence before tcgen05.ld, after tcgen05.mma.
asm volatile("tcgen05.fence::after_thread_sync;");
// Each thread handles one row. 4 warps * 32 threads = 128 rows = BLOCK_M.
// Load 8 columns at a time from TMEM.
for (int n{0}; n < BLOCK_N / 8; ++n) {
float tmp[8];
// TMEM address: 16 MSBs = row offset, 16 LSBs = column offset.
const int addr = taddr + ((warp_idx * 32) << 16) + (n * 8);
asm volatile("tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0, %1, %2, %3, %4, %5, %6, %7}, [%8];"
: "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]),
"=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7])
: "r"(addr));
asm volatile("tcgen05.wait::ld.sync.aligned;");
// Convert f32 pairs to f16 pairs and write to global memory.
half2 out[4];
for (int i{0}; i < 4; ++i)
out[i] = __float22half2_rn({tmp[i * 2], tmp[i * 2 + 1]});
// Each thread writes 16 bytes (8 half values) to its row.
half *out_ptr = C + (offset_m + thread_idx) * N + (offset_n + n * 8);
reinterpret_cast<int4 *>(out_ptr)[0] = reinterpret_cast<int4 *>(out)[0];
}
__syncthreads();
if (warp_idx == 0) {
// Deallocate TMEM (accumulator + scale factor columns).
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(taddr), "r"(TMEM_COLS));
}
}
torch::Tensor launch_kernel_pipeline_04(
const torch::Tensor& A,
const torch::Tensor& B,
const torch::Tensor& sfa,
const torch::Tensor& sfb,
int M,
int N,
int K
) {
auto C = torch::empty({M, N}, torch::dtype(torch::kFloat16).device(A.device()));
// Tile sizes.
// BLOCK_M=128; 1 CTA for .kind::mxf4nvf4.
// BLOCK_N=128; ensure one SF atom covers all N rows in the tile.
constexpr int BLOCK_M{128};
constexpr int BLOCK_N{128};
constexpr int BLOCK_K{256};
constexpr int NUM_THREADS{4 * WARP_SIZE}; // 4 warps, 128 threads
auto A_ptr{reinterpret_cast<const char *>(A.data_ptr())};
auto B_ptr{reinterpret_cast<const char *>(B.data_ptr())};
auto SFA_ptr{reinterpret_cast<const char *>(sfa.data_ptr())};
auto SFB_ptr{reinterpret_cast<const char *>(sfb.data_ptr())};
// Create 3D TMA tensor maps for A and B.
CUtensorMap A_tmap{}, B_tmap{};
create_tmap_descriptor<256>(&A_tmap, A_ptr, M, K, BLOCK_M, BLOCK_K,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);
create_tmap_descriptor<256>(&B_tmap, B_ptr, N, K, BLOCK_N, BLOCK_K,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);
// Shared memory: 2x (A tile + B tile + SFA tile + SFB tile) for double buffering.
constexpr int AB_SHARED_SIZE{(BLOCK_M + BLOCK_N) * BLOCK_K / 2};
constexpr int SF_SHARED_SIZE{2 * 512 * BLOCK_K / MMA_K};
constexpr int SHARED_SIZE{NUM_STAGES * (AB_SHARED_SIZE + SF_SHARED_SIZE)};
dim3 num_threads(NUM_THREADS);
dim3 num_blocks((M / BLOCK_M) * (N / BLOCK_N));
auto kernel{kernel_v04_pipeline<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>};
if (SHARED_SIZE > 48'000)
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SHARED_SIZE);
kernel<<<num_blocks, num_threads, SHARED_SIZE>>>(
A_tmap,
B_tmap,
SFA_ptr,
SFB_ptr,
reinterpret_cast<half *>(C.data_ptr<at::Half>()),
M, N, K
);
return C;
}
"""
cpp_source = """
#include <torch/extension.h>
torch::Tensor launch_kernel_pipeline_04(
const torch::Tensor& A,
const torch::Tensor& B,
const torch::Tensor& sfa,
const torch::Tensor& sfb,
int M,
int N,
int K);
"""
module = load_inline(
name='kernel',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['launch_kernel_pipeline_04'],
verbose=True,
is_python_module=True,
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",
"-Xptxas=-v",
# "--keep",
# "--keep-dir",
# f"{Path(__file__).parent}/tmp",
],
extra_ldflags=["-lcuda"],
)
def custom_kernel(data: input_t) -> output_t:
abc_tensors, _, sfasfb_reordered_tensors, problem_sizes = data
BLOCK_M = 128
BLOCK_N = 128
results = []
for (a, b, c), (sfa_reordered, sfb_reordered), (m, n, k, l) in zip(
abc_tensors, sfasfb_reordered_tensors, problem_sizes
):
for l_idx in range(l):
a_ptr = a[:, :, l_idx].contiguous() # [m, k//2]
b_ptr = b[:, :, l_idx].contiguous() # [n, k//2]
# Pad M and N to multiples of BLOCK_M and BLOCK_N.
padded_m = ((m + BLOCK_M - 1) // BLOCK_M) * BLOCK_M
padded_n = ((n + BLOCK_N - 1) // BLOCK_N) * BLOCK_N
if padded_m != m:
a_padded = torch.zeros(padded_m, k // 2, dtype=torch.uint8, device=a_ptr.device).view(a_ptr.dtype)
a_padded[:m] = a_ptr
a_ptr = a_padded
if padded_n != n:
b_padded = torch.zeros(padded_n, k // 2, dtype=torch.uint8, device=b_ptr.device).view(b_ptr.dtype)
b_padded[:n] = b_ptr
b_ptr = b_padded
c_out = module.launch_kernel_pipeline_04(
a_ptr, b_ptr, sfa_reordered, sfb_reordered, padded_m, padded_n, k
)
# Trim back to original size.
c[:, :, l_idx] = c_out[:m, :n]
results.append(c)
return results
scrolls · 509 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 488357.
⋯ 162 unchanged lines}constexpr int MMA_K = 64; // FP4 MMA K-dimension size.+ constexpr int NUM_STAGES = 4;+template <const int NUM_THREADS, const int BLOCK_M, const int BLOCK_N, const int BLOCK_K>__global__- __launch_bounds__(NUM_THREADS) void kernel_v02_swizzling(+ __launch_bounds__(NUM_THREADS) void kernel_v04_pipeline(const __grid_constant__ CUtensorMap A_tmap,const __grid_constant__ CUtensorMap B_tmap,const char *SFA,⋯ 16 unchanged linesconst int offset_m{block_idx_m * BLOCK_M};const int offset_n{block_idx_n * BLOCK_N};- // Set up shared memory.- // Layout: [A tile | B tile | SFA tile | SFB tile]- extern __shared__ __align__(1024) char smem[];- const int A_smem{static_cast<int>(__cvta_generic_to_shared(smem))};- const int B_smem{A_smem + BLOCK_M * BLOCK_K / 2};+ // Double-buffered shared memory layout:+ // [buf0: A | B | SFA | SFB | buf1: A | B | SFA | SFB]constexpr int SF_size = 512 * BLOCK_K / MMA_K;- const int SFA_smem{B_smem + BLOCK_N * BLOCK_K / 2};- const int SFB_smem{SFA_smem + SF_size};+ constexpr int BUF_SIZE = BLOCK_M * BLOCK_K / 2 + BLOCK_N * BLOCK_K / 2 + 2 * SF_size;+ extern __shared__ __align__(1024) char smem[];+ const int smem_base{static_cast<int>(__cvta_generic_to_shared(smem))};++ int A_smem[NUM_STAGES], B_smem[NUM_STAGES], SFA_smem[NUM_STAGES], SFB_smem[NUM_STAGES];+ for (int s{0}; s < NUM_STAGES; ++s) {+ A_smem[s] = smem_base + s * BUF_SIZE;+ B_smem[s] = A_smem[s] + BLOCK_M * BLOCK_K / 2;+ SFA_smem[s] = B_smem[s] + BLOCK_N * BLOCK_K / 2;+ SFB_smem[s] = SFA_smem[s] + SF_size;+ }+#pragma nv_diag_suppress static_var_with_dynamic_init- __shared__ uint64_t mbars[1];- const int mbar_addr{static_cast<int>(__cvta_generic_to_shared(mbars))};+ __shared__ uint64_t mbars[NUM_STAGES];+ int mbar_addrs[NUM_STAGES];+ for (int s{0}; s < NUM_STAGES; ++s)+ mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&mbars[s]));__shared__ int tmem_addr[1];// TMEM layout:⋯ 2 unchanged lines// Columns [BLOCK_N + 4*BLOCK_K/MMA_K, ...) : SFBconstexpr int SFA_tmem_start_col = BLOCK_N;constexpr int SFB_tmem_start_col = SFA_tmem_start_col + 4 * (BLOCK_K / MMA_K);- constexpr int TMEM_COLS = BLOCK_N + 8 * (BLOCK_K / MMA_K); // = BLOCK_N + BLOCK_K/8+ constexpr int TMEM_COLS = BLOCK_N * 2; // BLOCK_N + 8 * (BLOCK_K / MMA_K), but it has to be a power of 2.if (warp_idx == 0 && elect_sync()) {- mbarrier_init(mbar_addr, 1);+ for (int s{0}; s < NUM_STAGES; ++s)+ mbarrier_init(mbar_addrs[s], 1);asm volatile("fence.mbarrier_init.release.cluster;");} else if (warp_idx == 1) {// Allocate TMEM for accumulator + scale factors.const int addr{static_cast<int>(__cvta_generic_to_shared(tmem_addr))};asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"- ::"r"(addr), "r"(BLOCK_N * 2));+ ::"r"(addr), "r"(TMEM_COLS));}__syncthreads();const int taddr{tmem_addr[0]};- int phase{0};+ int phase[NUM_STAGES] = {0, 0};// Instruction descriptor for tcgen05.mma.kind::mxf4nvf4// atype=E2M1 (1), btype=E2M1 (1), MMA_N and MMA_M encoded in upper bits⋯ 3 unchanged lines| ((uint32_t)BLOCK_M >> 7U << 27U) // MMA_M;- for (int iter_k{0}; iter_k < (K / BLOCK_K); ++iter_k) {- // === Phase 1: Load from global memory to shared memory via TMA ===- if (warp_idx == 0 && elect_sync()) {- const int off_k{iter_k * BLOCK_K};+ const int num_iters{K / BLOCK_K};+ const int rest_k{K / 16 / 4}; // number of K-atoms (each covers 64 K-elements)+ constexpr int cp_size = (BLOCK_M + BLOCK_N) * BLOCK_K / 2 + 2 * SF_size;- // Load A and B tiles via 3D TMA tensor map.- // z-coordinate = off_k / 256 because NUM_ELEMENTS=256 (T for FP4) in the tensor map.- tma_3d_gmem2smem(A_smem, &A_tmap, 0, offset_m, off_k / 256, mbar_addr);- tma_3d_gmem2smem(B_smem, &B_tmap, 0, offset_n, off_k / 256, mbar_addr);+ // === Prologue: load first tile into buf[0] ===+ if (warp_idx == 0 && elect_sync()) {+ const int off_k{0};+ tma_3d_gmem2smem(A_smem[0], &A_tmap, 0, offset_m, off_k / 256, mbar_addrs[0]);+ tma_3d_gmem2smem(B_smem[0], &B_tmap, 0, offset_n, off_k / 256, mbar_addrs[0]);+ const char *SFA_src = SFA + ((offset_m / 128) * rest_k + off_k / (16 * 4)) * 512;+ const char *SFB_src = SFB + ((offset_n / 128) * rest_k + off_k / (16 * 4)) * 512;+ tma_gmem2smem(SFA_smem[0], SFA_src, SF_size, mbar_addrs[0]);+ tma_gmem2smem(SFB_smem[0], SFB_src, SF_size, mbar_addrs[0]);+ asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"+ ::"r"(mbar_addrs[0]), "r"(cp_size) : "memory");+ }- // Load SFA/SFB via 1D bulk copy.- // Underlying storage order is (L, M/128, rest_k, 32, 4, 4) — each atom is 512 bytes.- const int rest_k = K / 16 / 4; // number of K-atoms (each covers 64 K-elements)- const char *SFA_src = SFA + ((offset_m / 128) * rest_k + off_k / (16 * 4)) * 512;- const char *SFB_src = SFB + ((offset_n / 128) * rest_k + off_k / (16 * 4)) * 512;- tma_gmem2smem(SFA_smem, SFA_src, SF_size, mbar_addr);- tma_gmem2smem(SFB_smem, SFB_src, SF_size, mbar_addr);+ // Wait for first TMA to complete.+ mbarrier_wait(mbar_addrs[0], phase[0]);+ phase[0] ^= 1;- // Signal expected number of bytes for all TMA transfers.- constexpr int cp_size = (BLOCK_M + BLOCK_N) * BLOCK_K / 2 + 2 * SF_size;- asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"- ::"r"(mbar_addr), "r"(cp_size) : "memory");- }+ // === Main loop: overlap TMA load of next tile with MMA compute of current tile ===+ for (int iter_k{0}; iter_k < num_iters; ++iter_k) {+ const int s{iter_k % NUM_STAGES};- // Wait for TMA to complete.- mbarrier_wait(mbar_addr, phase);- phase ^= 1;+ if (warp_idx == 0 && elect_sync()) {+ // Prefetch next tile into the other buffer (overlaps with MMA below).+ if (iter_k + 1 < num_iters) {+ const int ns{(iter_k + 1) % NUM_STAGES};+ const int next_off_k{(iter_k + 1) * BLOCK_K};+ tma_3d_gmem2smem(A_smem[ns], &A_tmap, 0, offset_m, next_off_k / 256, mbar_addrs[ns]);+ tma_3d_gmem2smem(B_smem[ns], &B_tmap, 0, offset_n, next_off_k / 256, mbar_addrs[ns]);+ const char *SFA_src = SFA + ((offset_m / 128) * rest_k + next_off_k / (16 * 4)) * 512;+ const char *SFB_src = SFB + ((offset_n / 128) * rest_k + next_off_k / (16 * 4)) * 512;+ tma_gmem2smem(SFA_smem[ns], SFA_src, SF_size, mbar_addrs[ns]);+ tma_gmem2smem(SFB_smem[ns], SFB_src, SF_size, mbar_addrs[ns]);+ asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"+ ::"r"(mbar_addrs[ns]), "r"(cp_size) : "memory");+ }- // === Phase 2: Copy scale factors to TMEM, then perform MMA ===- if (warp_idx == 0 && elect_sync())- {+ // MMA from current buffer.auto make_desc_AB = [](int addr) -> uint64_t{constexpr int SBO = 8 * 256 / 2; // in bytes.return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);};- // No-swizzle shared memory descriptor for scale factors.- // SBO = stride between 8-row groups = 8 rows * 16 bytes/row.auto make_desc_SF = [](int addr) -> uint64_t{const int SBO = 8 * 16; // = 128 bytes⋯ 1 unchanged lines};// Copy scale factors from shared memory to tensor memory.- // tcgen05.cp and tcgen05.mma are pipelined correctly per PTX docs.for (int k{0}; k < BLOCK_K / MMA_K; ++k) {- uint64_t sfa_desc = make_desc_SF(SFA_smem + k * 512); // 512 bytes per MMA_K atom- uint64_t sfb_desc = make_desc_SF(SFB_smem + k * 512);+ uint64_t sfa_desc = make_desc_SF(SFA_smem[s] + k * 512);+ uint64_t sfb_desc = make_desc_SF(SFB_smem[s] + k * 512);copy_sf_smem2tmem(SFA_tmem_start_col + k * 4, sfa_desc);copy_sf_smem2tmem(SFB_tmem_start_col + k * 4, sfb_desc);}⋯ 1 unchanged lines// One swizzle tile = 256 FP4 elements = 128 bytes.for (int k1{0}; k1 < BLOCK_K / 256; ++k1) {for (int k2{0}; k2 < 256 / MMA_K; ++k2) {- uint64_t a_desc{make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * MMA_K / 2)};- uint64_t b_desc{make_desc_AB(B_smem + k1 * BLOCK_N * 128 + k2 * MMA_K / 2)};+ uint64_t a_desc{make_desc_AB(A_smem[s] + k1 * BLOCK_M * 128 + k2 * MMA_K / 2)};+ uint64_t b_desc{make_desc_AB(B_smem[s] + k1 * BLOCK_N * 128 + k2 * MMA_K / 2)};int k{k1 * 256 / MMA_K + k2};const int scale_A_tmem{SFA_tmem_start_col + k * 4};⋯ 4 unchanged lines}}- // Signal MMA completion on the mbarrier.+ // Signal MMA completion on the current buffer's mbarrier.asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"- ::"r"(mbar_addr) : "memory");+ ::"r"(mbar_addrs[s]) : "memory");}- // Wait for MMA to complete.- mbarrier_wait(mbar_addr, phase);- phase ^= 1;+ // Wait for MMA on current buffer to complete.+ // This also ensures buf[s] is free to reuse for TMA 2 iterations later.+ mbarrier_wait(mbar_addrs[s], phase[s]);+ phase[s] ^= 1;++ // Wait for next tile's TMA to complete (so next iteration can MMA from it).+ if (iter_k + 1 < num_iters) {+ const int ns{(iter_k + 1) % NUM_STAGES};+ mbarrier_wait(mbar_addrs[ns], phase[ns]);+ phase[ns] ^= 1;+ }}// === Epilogue: Read accumulator from TMEM and store to global memory ===⋯ 25 unchanged linesif (warp_idx == 0) {// Deallocate TMEM (accumulator + scale factor columns).- asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(taddr), "r"(BLOCK_N * 2));+ asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(taddr), "r"(TMEM_COLS));}}- torch::Tensor launch_kernel_swizzling_02(- torch::Tensor A,- torch::Tensor B,- torch::Tensor sfa,- torch::Tensor sfb,+ torch::Tensor launch_kernel_pipeline_04(+ const torch::Tensor& A,+ const torch::Tensor& B,+ const torch::Tensor& sfa,+ const torch::Tensor& sfb,int M,int N,int K⋯ 20 unchanged linescreate_tmap_descriptor<256>(&B_tmap, B_ptr, N, K, BLOCK_N, BLOCK_K,CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);- // Shared memory: A tile + B tile + SFA tile + SFB tile.+ // Shared memory: 2x (A tile + B tile + SFA tile + SFB tile) for double buffering.constexpr int AB_SHARED_SIZE{(BLOCK_M + BLOCK_N) * BLOCK_K / 2};- constexpr int SF_SHARED_SIZE{2 * 512 * BLOCK_K / MMA_K}; // Each SF atom is 512 bytes.- constexpr int SHARED_SIZE{AB_SHARED_SIZE + SF_SHARED_SIZE};+ constexpr int SF_SHARED_SIZE{2 * 512 * BLOCK_K / MMA_K};+ constexpr int SHARED_SIZE{NUM_STAGES * (AB_SHARED_SIZE + SF_SHARED_SIZE)};dim3 num_threads(NUM_THREADS);dim3 num_blocks((M / BLOCK_M) * (N / BLOCK_N));- auto kernel{kernel_v02_swizzling<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>};+ auto kernel{kernel_v04_pipeline<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>};if (SHARED_SIZE > 48'000)cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SHARED_SIZE);⋯ 15 unchanged linescpp_source = """#include <torch/extension.h>- torch::Tensor launch_kernel_swizzling_02(- torch::Tensor a_bytes,- torch::Tensor b_bytes,- torch::Tensor sfa,- torch::Tensor sfb,+ torch::Tensor launch_kernel_pipeline_04(+ const torch::Tensor& A,+ const torch::Tensor& B,+ const torch::Tensor& sfa,+ const torch::Tensor& sfb,int M,int N,int K);⋯ 3 unchanged linesname='kernel',cpp_sources=cpp_source,cuda_sources=cuda_source,- functions=['launch_kernel_swizzling_02'],+ functions=['launch_kernel_pipeline_04'],verbose=True,is_python_module=True,no_implicit_headers=True,⋯ 41 unchanged linesb_padded[:n] = b_ptrb_ptr = b_padded- c_out = module.launch_kernel_swizzling_02(+ c_out = module.launch_kernel_pipeline_04(a_ptr, b_ptr, sfa_reordered, sfb_reordered, padded_m, padded_n, k)
scrolls · 290 diff lines total
Best evidence level for this revision: reported
JSON