submission 115121
sahanp · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 148 lines, June 9 Researcher Reciprocity License v1.0.
gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-115121?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:e5b539dca3242b024d87e44245ff248e5f9ea1c8b40d9f6db3951f8e4cb9b918
license declaredunknown
license concludedunknown
authorssahanp
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
gemm.py148 lines
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor
# Fixed constants
NVFP4_VEC_SIZE = 16
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 256
ELEM_PER_BYTE = 2
ROWS_PER_SCALE_CHUNK = 128
K_PACKED = BLOCK_K // ELEM_PER_BYTE
def get_tma_configs():
configs = []
for num_stages in [2, 3, 4, 5]:
for num_warps in [4, 8]:
configs.append(
triton.Config(
{'NUM_STAGES': num_stages},
num_warps=num_warps,
num_stages=num_stages,
)
)
return configs
@triton.autotune(
configs=get_tma_configs(),
key=['M', 'N', 'K'],
)
@triton.jit
def batched_block_scaled_gemm_kernel(
a_desc, a_scale_desc, b_desc, b_scale_desc,
c_ptr,
stride_c_m, stride_c_n,
M: tl.constexpr, N: tl.constexpr, K: tl.constexpr,
D_chunks_m: tl.constexpr, D_chunks_n: tl.constexpr,
VEC_SIZE: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
rep_m: tl.constexpr, rep_n: tl.constexpr, rep_k: tl.constexpr,
NUM_STAGES: tl.constexpr,
num_m_blocks: tl.constexpr, num_n_blocks: tl.constexpr,
):
pid = tl.program_id(axis=0)
# 2D tiling over M and N
pid_m = pid % num_m_blocks
pid_n = pid // num_m_blocks
offs_am = pid_m * BLOCK_M
offs_bn = pid_n * BLOCK_N
offs_scale_m = pid_m * rep_m
offs_scale_n = pid_n * rep_n
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
K_STEP: tl.constexpr = BLOCK_K // 2
k_iters = tl.cdiv(K, BLOCK_K)
offs_k = 0
offs_scale_k = 0
for _ in tl.range(0, k_iters, num_stages=NUM_STAGES):
a = a_desc.load([offs_am, offs_k])
b = b_desc.load([offs_bn, offs_k])
scale_a = a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
scale_b = b_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
scale_a = scale_a.reshape(rep_m, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_M, BLOCK_K // VEC_SIZE)
scale_b = scale_b.reshape(rep_n, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_N, BLOCK_K // VEC_SIZE)
accumulator = tl.dot_scaled(a, scale_a, "e2m1", b.T, scale_b, "e2m1", accumulator)
offs_k += K_STEP
offs_scale_k += rep_k
# Store full BLOCK_M x BLOCK_N output
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
m_mask = offs_m[:, None] < M
n_mask = offs_n[None, :] < N
mask = m_mask & n_mask
c_ptrs = c_ptr + offs_m[:, None] * stride_c_m + offs_n[None, :] * stride_c_n
tl.store(c_ptrs, accumulator.to(tl.float16), mask=mask)
def custom_kernel(data):
a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, sfa_permuted, sfb_permuted, c_ref = data
device = a_ref.device
M, K_packed_dim, L = a_ref.shape
N_ext = b_ref.shape[0]
_, K_blocks, _ = sfa_ref_cpu.shape
K_real = K_blocks * NVFP4_VEC_SIZE
D_chunks_m = M // ROWS_PER_SCALE_CHUNK
D_chunks_n = N_ext // ROWS_PER_SCALE_CHUNK
K_chunks = K_blocks // 4
rep_m, rep_n, rep_k = 1, 1, 4
assert L == 1, "This kernel only supports L=1"
# A and B: slice off L dimension
a_stacked = a_ref[:, :, 0].view(torch.uint8)
b_stacked = b_ref[:, :, 0].view(torch.uint8)
# Scale factors - use sfa_permuted directly (already on GPU)
# sfa_permuted[:,:,:,:,:,0] has shape (32, 4, D_chunks_m, 4, K_chunks)
# Need to get to (1, D_chunks_m, K_chunks, 2, 256)
sfa_5d = (sfa_permuted[:, :, :, :, :, 0]
.permute(2, 4, 0, 1, 3) # (D_chunks_m, K_chunks, 32, 4, 4)
.reshape(1, D_chunks_m, K_chunks, 2, 256))
sfb_5d = (sfb_permuted[:, :, :, :, :, 0]
.permute(2, 4, 0, 1, 3)
.reshape(1, D_chunks_n, K_chunks, 2, 256))
num_m_blocks = triton.cdiv(M, BLOCK_M)
num_n_blocks = triton.cdiv(N_ext, BLOCK_N)
total_tiles = num_m_blocks * num_n_blocks
# TMA descriptors
a_desc = TensorDescriptor.from_tensor(a_stacked, [BLOCK_M, K_PACKED])
b_desc = TensorDescriptor.from_tensor(b_stacked, [BLOCK_N, K_PACKED])
a_scale_desc = TensorDescriptor.from_tensor(sfa_5d, [1, rep_m, rep_k, 2, 256])
b_scale_desc = TensorDescriptor.from_tensor(sfb_5d, [1, rep_n, rep_k, 2, 256])
# Output is M x N (slice L dimension)
c_out = c_ref[:, :, 0]
grid = (total_tiles,)
batched_block_scaled_gemm_kernel[grid](
a_desc, a_scale_desc, b_desc, b_scale_desc,
c_out, c_out.stride(0), c_out.stride(1),
M, N_ext, K_real,
D_chunks_m, D_chunks_n,
NVFP4_VEC_SIZE,
BLOCK_M, BLOCK_N, BLOCK_K,
rep_m, rep_n, rep_k,
num_m_blocks=num_m_blocks, num_n_blocks=num_n_blocks,
)
return c_refscrolls · 148 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