submission 498891
kathsucurry · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 559 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-498891?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:3420c4ce60b3bbc2e96235a7eaf41071d6ffa4b473dfb7e716326115c80f3b8b
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.fused-epilogue
__launch_bounds__(NUM_THREADS) void kernel_v09_improve_epilogue(mbarrier
__device__ inline void mbarrier_init(int mbar_addr, int count) {shared-memory
void tma_gmem2smem_cluster(int dst, const void *src, int size, int mbar_addr) {stages = 4
constexpr int NUM_STAGES = 4;tcgen05
asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(NUM_CTA));tma
TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);vector-width = half2
half2 out[WIDTH / 2];Kernel source
submission.py559 lines
import os
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
cuda_common_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 zero-filled 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));
}
__device__ inline
void tma_gmem2smem_cluster(int dst, const void *src, int size, int mbar_addr) {
asm volatile(
"cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"
:: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr));
}
template <int NUM_CTA = 1>
__device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr)
{
// when NUM_CTA=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"(NUM_CTA)
: "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.
template <int NUM_CTA=1>
__device__ inline
void copy_sf_smem2tmem(int taddr, uint64_t s_desc) {
asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(NUM_CTA));
}
// 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.
template <int NUM_CTA=1>
__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::%7.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), "n"(NUM_CTA)
);
}
__device__ inline
void load_from_tmem_32x32b_x8(float *tmp, const int addr) {
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));
}
__device__ inline
void load_from_tmem_32x32b_x16(float *tmp, const int addr) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x16.b32 "
"{%0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15}, [%16];"
: "=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])
: "r"(addr));
}
"""
cuda_kernel_source = r"""
constexpr int MMA_K = 64; // FP4 MMA K-dimension size.
constexpr int NUM_STAGES = 4;
constexpr int MAX_GROUPS = 8;
struct GroupParams {
const char *SFA;
const char *SFB;
half *C;
int M, N, K;
int block_offset; // cumulative block count before this group
int grid_dim_n;
int rest_k; // K / 16 / 4
int num_iters; // K / BLOCK_K
};
struct GroupedKernelArgs {
CUtensorMap tmaps[MAX_GROUPS * 2]; // A_tmap, B_tmap per group
GroupParams params[MAX_GROUPS];
int num_groups;
};
template <const int NUM_THREADS, const int BLOCK_M, const int BLOCK_N, const int BLOCK_K>
__global__
__launch_bounds__(NUM_THREADS) void kernel_v09_improve_epilogue(
const __grid_constant__ GroupedKernelArgs args
) {
const int thread_idx{static_cast<int>(threadIdx.x)};
const int global_block_idx{static_cast<int>(blockIdx.x)};
const int warp_idx{thread_idx / WARP_SIZE};
// Find which group this block belongs to (linear scan, num_groups <= 8).
int group = 0;
for (int g = 1; g < args.num_groups; ++g) {
if (global_block_idx >= args.params[g].block_offset)
group = g;
}
// Load per-group parameters.
const GroupParams &gp = args.params[group];
const int block_idx = global_block_idx - gp.block_offset;
const int block_idx_m{block_idx / gp.grid_dim_n};
const int block_idx_n{block_idx % gp.grid_dim_n};
const int offset_m{block_idx_m * BLOCK_M};
const int offset_n{block_idx_n * BLOCK_N};
const int M = gp.M;
const int N = gp.N;
const int num_iters = gp.num_iters;
const int rest_k = gp.rest_k;
const char *SFA = gp.SFA;
const char *SFB = gp.SFB;
half *C = gp.C;
// Tensor maps for this group (in __grid_constant__ / .param space).
const CUtensorMap *A_tmap_ptr = &args.tmaps[group * 2];
const CUtensorMap *B_tmap_ptr = &args.tmaps[group * 2 + 1];
// Multi-buffered shared memory layout:
// [buf0: A | B | SFA | SFB | buf1: ... | buf2: ... | buf3: ...]
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 tma_mbars[NUM_STAGES];
__shared__ uint64_t mma_mbars[NUM_STAGES];
int tma_mbar_addrs[NUM_STAGES], mma_mbar_addrs[NUM_STAGES];
for (int s{0}; s < NUM_STAGES; ++s) {
tma_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&tma_mbars[s]));
mma_mbar_addrs[s] = static_cast<int>(__cvta_generic_to_shared(&mma_mbars[s]));
}
__shared__ int tmem_addr[1];
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;
if (warp_idx == 0 && elect_sync()) {
for (int s{0}; s < NUM_STAGES; ++s) {
mbarrier_init(tma_mbar_addrs[s], 1);
mbarrier_init(mma_mbar_addrs[s], 1);
}
asm volatile("fence.mbarrier_init.release.cluster;");
} else if (warp_idx == 1) {
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]};
constexpr uint32_t i_desc = (1U << 7U)
| (1U << 10U)
| ((uint32_t)BLOCK_N >> 3U << 17U)
| ((uint32_t)BLOCK_M >> 7U << 27U);
constexpr int cp_size = (BLOCK_M + BLOCK_N) * BLOCK_K / 2 + 2 * SF_size;
// =========================================================================
// Warp 0: TMA Producer
// =========================================================================
if (warp_idx == 0 && elect_sync()) {
int mma_prod_phase[NUM_STAGES] = {};
for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
const int s{iter_k % NUM_STAGES};
if (iter_k >= NUM_STAGES) {
mbarrier_wait(mma_mbar_addrs[s], mma_prod_phase[s]);
mma_prod_phase[s] ^= 1;
}
const int off_k{iter_k * BLOCK_K};
tma_3d_gmem2smem(A_smem[s], A_tmap_ptr, 0, offset_m, off_k / 256, tma_mbar_addrs[s]);
tma_3d_gmem2smem(B_smem[s], B_tmap_ptr, 0, offset_n, off_k / 256, tma_mbar_addrs[s]);
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[s], SFA_src, SF_size, tma_mbar_addrs[s]);
tma_gmem2smem(SFB_smem[s], SFB_src, SF_size, tma_mbar_addrs[s]);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
::"r"(tma_mbar_addrs[s]), "r"(cp_size) : "memory");
}
// =========================================================================
// Warp 1: MMA Consumer
// =========================================================================
} else if (warp_idx == 1 && elect_sync()) {
int tma_cons_phase[NUM_STAGES] = {};
int mma_done_phase[NUM_STAGES] = {};
for (int iter_k{0}; iter_k < num_iters; ++iter_k) {
const int s{iter_k % NUM_STAGES};
mbarrier_wait(tma_mbar_addrs[s], tma_cons_phase[s]);
tma_cons_phase[s] ^= 1;
auto make_desc_AB = [](int addr) -> uint64_t
{
constexpr int SBO = 8 * 256 / 2;
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;
return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL);
};
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);
}
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);
}
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
::"r"(mma_mbar_addrs[s]) : "memory");
mma_done_phase[s] ^= 1;
}
if (num_iters > 0) {
const int last_s{(num_iters - 1) % NUM_STAGES};
mbarrier_wait(mma_mbar_addrs[last_s], mma_done_phase[last_s] ^ 1);
}
}
__syncthreads();
// === Epilogue: Read accumulator from TMEM and store to global memory ===
asm volatile("tcgen05.fence::after_thread_sync;");
const int row{offset_m + thread_idx};
constexpr int WIDTH{16};
for (int n{0}; n < BLOCK_N / WIDTH; ++n) {
float tmp[WIDTH];
const int addr = taddr + ((warp_idx * 32) << 16) + (n * WIDTH);
load_from_tmem_32x32b_x16(tmp, addr);
asm volatile("tcgen05.wait::ld.sync.aligned;");
if (row >= M) continue;
const int col{offset_n + n * WIDTH};
half2 out[WIDTH / 2];
for (int i{0}; i < WIDTH / 2; ++i)
out[i] = __float22half2_rn({tmp[i * 2], tmp[i * 2 + 1]});
half *out_ptr = C + row * N + col;
if (col + WIDTH <= N) {
for (int i{0}; i < WIDTH / 8; ++i)
reinterpret_cast<int4 *>(out_ptr)[i] = reinterpret_cast<int4 *>(out)[i];
} else {
const half *out_half = reinterpret_cast<const half *>(out);
for (int i{0}; i < WIDTH && col + i < N; ++i)
out_ptr[i] = out_half[i];
}
}
__syncthreads();
if (warp_idx == 0) {
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(taddr), "r"(TMEM_COLS));
}
}
std::vector<torch::Tensor> launch_grouped_kernel(
std::vector<torch::Tensor> As,
std::vector<torch::Tensor> Bs,
std::vector<torch::Tensor> SFAs,
std::vector<torch::Tensor> SFBs,
std::vector<int64_t> Ms,
std::vector<int64_t> Ns,
std::vector<int64_t> Ks
) {
constexpr int BLOCK_M{128};
constexpr int BLOCK_N{128};
constexpr int BLOCK_K{256};
constexpr int NUM_THREADS{4 * WARP_SIZE};
int num_groups = As.size();
// Create output tensors.
std::vector<torch::Tensor> Cs;
for (int g = 0; g < num_groups; ++g) {
Cs.push_back(torch::empty({Ms[g], Ns[g]},
torch::dtype(torch::kFloat16).device(As[g].device())));
}
// Build kernel args on the host stack (~2.5 KB, fits in 4 KB param limit).
GroupedKernelArgs host_args{};
host_args.num_groups = num_groups;
int total_blocks = 0;
for (int g = 0; g < num_groups; ++g) {
auto A_ptr = reinterpret_cast<const char *>(As[g].data_ptr());
auto B_ptr = reinterpret_cast<const char *>(Bs[g].data_ptr());
create_tmap_descriptor<256>(&host_args.tmaps[g * 2], A_ptr, Ms[g], Ks[g],
BLOCK_M, BLOCK_K, CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);
create_tmap_descriptor<256>(&host_args.tmaps[g * 2 + 1], B_ptr, Ns[g], Ks[g],
BLOCK_N, BLOCK_K, CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B);
int grid_m = (Ms[g] + BLOCK_M - 1) / BLOCK_M;
int grid_n = (Ns[g] + BLOCK_N - 1) / BLOCK_N;
host_args.params[g].SFA = reinterpret_cast<const char *>(SFAs[g].data_ptr());
host_args.params[g].SFB = reinterpret_cast<const char *>(SFBs[g].data_ptr());
host_args.params[g].C = reinterpret_cast<half *>(Cs[g].data_ptr<at::Half>());
host_args.params[g].M = Ms[g];
host_args.params[g].N = Ns[g];
host_args.params[g].K = Ks[g];
host_args.params[g].block_offset = total_blocks;
host_args.params[g].grid_dim_n = grid_n;
host_args.params[g].rest_k = Ks[g] / 16 / 4;
host_args.params[g].num_iters = Ks[g] / BLOCK_K;
total_blocks += grid_m * grid_n;
}
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)};
auto kernel = kernel_v09_improve_epilogue<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>;
if (SHARED_SIZE > 48'000)
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SHARED_SIZE);
// Using __grid_constant__ allows passing host_args via constant memory, no need to do allocation.
kernel<<<total_blocks, NUM_THREADS, SHARED_SIZE>>>(host_args);
return Cs;
}
"""
cpp_source = """
#include <torch/extension.h>
std::vector<torch::Tensor> launch_grouped_kernel(
std::vector<torch::Tensor> As,
std::vector<torch::Tensor> Bs,
std::vector<torch::Tensor> SFAs,
std::vector<torch::Tensor> SFBs,
std::vector<int64_t> Ms,
std::vector<int64_t> Ns,
std::vector<int64_t> Ks);
"""
module = load_inline(
name='kernel',
cpp_sources=cpp_source,
cuda_sources=cuda_common_source + cuda_kernel_source,
functions=['launch_grouped_kernel'],
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:
As, Bs, SFAs, SFBs = [], [], [], []
Ms, Ns, Ks = [], [], []
cs = []
for (a, b, c), _, (sfa_reordered, sfb_reordered), (m, n, k, _) in zip(*data):
As.append(a[:, :, 0])
Bs.append(b[:, :, 0])
SFAs.append(sfa_reordered)
SFBs.append(sfb_reordered)
Ms.append(m)
Ns.append(n)
Ks.append(k)
cs.append(c)
outputs = module.launch_grouped_kernel(As, Bs, SFAs, SFBs, Ms, Ns, Ks)
for c, out in zip(cs, outputs):
c[:, :, 0] = out
return cs
scrolls · 559 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 489099.
⋯ 2 unchanged linesfrom torch.utils.cpp_extension import load_inlinefrom task import input_t, output_t- cuda_source = r"""+ cuda_common_source = r"""#include <cuda_fp16.h>#include <cudaTypedefs.h>⋯ 52 unchanged lines// 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.+ // Any element that is outside of bounds will be zero-filled by the TMA transfer.CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);check_cu_error(error);}⋯ 43 unchanged lines}- template <int CTA_GROUP = 1>+ __device__ inline+ void tma_gmem2smem_cluster(int dst, const void *src, int size, int mbar_addr) {+ asm volatile(+ "cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"+ :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr));+ }+++ template <int NUM_CTA = 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.+ // when NUM_CTA=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)+ "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "n"(NUM_CTA): "memory");}⋯ 6 unchanged lines// 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.+ template <int NUM_CTA=1>__device__ inlinevoid 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));+ asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(NUM_CTA));}// 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.+ template <int NUM_CTA=1>__device__ inlinevoid run_mma_nvfp4(uint64_t a_desc,⋯ 8 unchanged lines"{\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"+ "tcgen05.mma.cta_group::%7.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)+ "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d), "n"(NUM_CTA));}+ __device__ inline+ void load_from_tmem_32x32b_x8(float *tmp, const int addr) {+ 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));+ }+++ __device__ inline+ void load_from_tmem_32x32b_x16(float *tmp, const int addr) {+ asm volatile("tcgen05.ld.sync.aligned.32x32b.x16.b32 "+ "{%0, %1, %2, %3, %4, %5, %6, %7, "+ " %8, %9, %10, %11, %12, %13, %14, %15}, [%16];"+ : "=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])+ : "r"(addr));+ }+ """++ cuda_kernel_source = r"""constexpr int MMA_K = 64; // FP4 MMA K-dimension size.constexpr int NUM_STAGES = 4;constexpr int MAX_GROUPS = 8;⋯ 19 unchanged linestemplate <const int NUM_THREADS, const int BLOCK_M, const int BLOCK_N, const int BLOCK_K>__global__- __launch_bounds__(NUM_THREADS) void kernel_v05_grouped(+ __launch_bounds__(NUM_THREADS) void kernel_v09_improve_epilogue(const __grid_constant__ GroupedKernelArgs args) {const int thread_idx{static_cast<int>(threadIdx.x)};⋯ 84 unchanged lines// =========================================================================// Warp 0: TMA Producer// =========================================================================- if (warp_idx == 0) {+ if (warp_idx == 0 && elect_sync()) {int mma_prod_phase[NUM_STAGES] = {};for (int iter_k{0}; iter_k < num_iters; ++iter_k) {⋯ 4 unchanged linesmma_prod_phase[s] ^= 1;}- if (elect_sync()) {- const int off_k{iter_k * BLOCK_K};- tma_3d_gmem2smem(A_smem[s], A_tmap_ptr, 0, offset_m, off_k / 256, tma_mbar_addrs[s]);- tma_3d_gmem2smem(B_smem[s], B_tmap_ptr, 0, offset_n, off_k / 256, tma_mbar_addrs[s]);- 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[s], SFA_src, SF_size, tma_mbar_addrs[s]);- tma_gmem2smem(SFB_smem[s], SFB_src, SF_size, tma_mbar_addrs[s]);- asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"- ::"r"(tma_mbar_addrs[s]), "r"(cp_size) : "memory");- }+ const int off_k{iter_k * BLOCK_K};+ tma_3d_gmem2smem(A_smem[s], A_tmap_ptr, 0, offset_m, off_k / 256, tma_mbar_addrs[s]);+ tma_3d_gmem2smem(B_smem[s], B_tmap_ptr, 0, offset_n, off_k / 256, tma_mbar_addrs[s]);+ 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[s], SFA_src, SF_size, tma_mbar_addrs[s]);+ tma_gmem2smem(SFB_smem[s], SFB_src, SF_size, tma_mbar_addrs[s]);+ asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"+ ::"r"(tma_mbar_addrs[s]), "r"(cp_size) : "memory");}// =========================================================================// Warp 1: MMA Consumer// =========================================================================- } else if (warp_idx == 1) {+ } else if (warp_idx == 1 && elect_sync()) {int tma_cons_phase[NUM_STAGES] = {};int mma_done_phase[NUM_STAGES] = {};⋯ 3 unchanged linesmbarrier_wait(tma_mbar_addrs[s], tma_cons_phase[s]);tma_cons_phase[s] ^= 1;- if (elect_sync()) {- auto make_desc_AB = [](int addr) -> uint64_t- {- constexpr int SBO = 8 * 256 / 2;- return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);- };+ auto make_desc_AB = [](int addr) -> uint64_t+ {+ constexpr int SBO = 8 * 256 / 2;+ 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;- return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL);- };+ auto make_desc_SF = [](int addr) -> uint64_t+ {+ const int SBO = 8 * 16;+ return encode_descriptor(addr) | (encode_descriptor(SBO) << 32ULL) | (1ULL << 46ULL);+ };- 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);- }+ 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);+ }- 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)};+ 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};+ 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);- }+ 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);}-- asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"- ::"r"(mma_mbar_addrs[s]) : "memory");}++ asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"+ ::"r"(mma_mbar_addrs[s]) : "memory");mma_done_phase[s] ^= 1;}⋯ 9 unchanged linesasm volatile("tcgen05.fence::after_thread_sync;");const int row{offset_m + thread_idx};- for (int n{0}; n < BLOCK_N / 8; ++n) {- float tmp[8];- 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));+ constexpr int WIDTH{16};+ for (int n{0}; n < BLOCK_N / WIDTH; ++n) {+ float tmp[WIDTH];+ const int addr = taddr + ((warp_idx * 32) << 16) + (n * WIDTH);+ load_from_tmem_32x32b_x16(tmp, addr);asm volatile("tcgen05.wait::ld.sync.aligned;");if (row >= M) continue;- const int col{offset_n + n * 8};-- half2 out[4];- for (int i{0}; i < 4; ++i)+ const int col{offset_n + n * WIDTH};+ half2 out[WIDTH / 2];+ for (int i{0}; i < WIDTH / 2; ++i)out[i] = __float22half2_rn({tmp[i * 2], tmp[i * 2 + 1]});half *out_ptr = C + row * N + col;- if (col + 8 <= N) {- reinterpret_cast<int4 *>(out_ptr)[0] = reinterpret_cast<int4 *>(out)[0];+ if (col + WIDTH <= N) {+ for (int i{0}; i < WIDTH / 8; ++i)+ reinterpret_cast<int4 *>(out_ptr)[i] = reinterpret_cast<int4 *>(out)[i];} else {const half *out_half = reinterpret_cast<const half *>(out);- for (int i{0}; i < 8 && col + i < N; ++i)+ for (int i{0}; i < WIDTH && col + i < N; ++i)out_ptr[i] = out_half[i];}}+__syncthreads();if (warp_idx == 0) {⋯ 60 unchanged linesconstexpr int SF_SHARED_SIZE{2 * 512 * BLOCK_K / MMA_K};constexpr int SHARED_SIZE{NUM_STAGES * (AB_SHARED_SIZE + SF_SHARED_SIZE)};- auto kernel = kernel_v05_grouped<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>;+ auto kernel = kernel_v09_improve_epilogue<NUM_THREADS, BLOCK_M, BLOCK_N, BLOCK_K>;if (SHARED_SIZE > 48'000)cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SHARED_SIZE);+ // Using __grid_constant__ allows passing host_args via constant memory, no need to do allocation.kernel<<<total_blocks, NUM_THREADS, SHARED_SIZE>>>(host_args);return Cs;}--"""cpp_source = """⋯ 12 unchanged linesmodule = load_inline(name='kernel',cpp_sources=cpp_source,- cuda_sources=cuda_source,+ cuda_sources=cuda_common_source + cuda_kernel_source,functions=['launch_grouped_kernel'],verbose=True,is_python_module=True,
scrolls · 301 diff lines total
Best evidence level for this revision: reported
JSON