submission 148348
Arseni Ivanov · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 205 lines, June 9 Researcher Reciprocity License v1.0.
triton_hardcoded_not_persistent.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-148348?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:b9be5c47f759afde7483770b7d130c214ef47ba3888dbd5009faa4eddef17307
license declaredunknown
license concludedunknown
authorsArseni Ivanov
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps = 4stages = 3
num_stages = 3tile-k = 512
BLOCK_K = 512tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128warp-specialization
WARP_SPECIALIZE_INNER: tl.constexpr,Kernel source
triton_hardcoded_not_persistent.py205 lines
#!POPCORN leaderboard nvfp4_gemm
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor
@triton.jit
def block_scaled_batched_gemm_kernel(
a_desc,
a_scale_desc,
b_desc,
b_scale_desc,
c_ptr,
stride_cm,
stride_cn,
M,
N,
K,
ELEM_PER_BYTE: tl.constexpr,
GROUP_SZ: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
REP_K: tl.constexpr,
NUM_INNER_STAGES: tl.constexpr,
WARP_SPECIALIZE_INNER: tl.constexpr,
):
BLOCK_K_ELEM_PER_BYTE: tl.constexpr = BLOCK_K // ELEM_PER_BYTE
BLOCK_K_GROUP_SZ: tl.constexpr = BLOCK_K // GROUP_SZ
pid = tl.program_id(axis=0)
num_pid_n = tl.cdiv(N, BLOCK_N)
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
# Base offsets for this tile
offs_am = pid_m * BLOCK_M
offs_bn = pid_n * BLOCK_N
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for i in tl.range(
0,
tl.cdiv(K, BLOCK_K),
num_stages=NUM_INNER_STAGES,
warp_specialize=WARP_SPECIALIZE_INNER,
):
offs_k = i * BLOCK_K_ELEM_PER_BYTE
offs_scale_k = i * REP_K
# A: [BLOCK_M, BLOCK_K/2]
# B: [BLOCK_N, BLOCK_K/2]
a = a_desc.load([offs_am, offs_k])
b = b_desc.load([offs_bn, offs_k])
scale_a = (
a_scale_desc.load([pid_m, offs_scale_k, 0, 0])
.reshape(REP_K, 32, 4, 4)
.trans(2, 1, 0, 3)
.reshape(BLOCK_M, BLOCK_K_GROUP_SZ)
)
scale_b = (
b_scale_desc.load([pid_n, offs_scale_k, 0, 0])
.reshape(REP_K, 32, 4, 4)
.trans(2, 1, 0, 3)
.reshape(BLOCK_N, BLOCK_K_GROUP_SZ)
)
accumulator = tl.dot_scaled(
a,
scale_a,
"e2m1",
b.T,
scale_b,
"e2m1",
accumulator,
)
# Calculate output pointers
offset_m = offs_am + tl.arange(0, BLOCK_M)
offset_n = offs_bn + tl.arange(0, BLOCK_N)
c_off = (
offset_m[:, None] * stride_cm
+ offset_n[None, :] * stride_cn
)
c_mask = (offset_m[:, None] < M) & (offset_n[None, :] < N)
tl.store(c_ptr + c_off, accumulator.to(tl.float16), mask=c_mask)
def custom_kernel(data):
a_tensor, b_tensor, _, _, sfa_tensor, sfb_tensor, c_tensor = data
#We only have a single batch every time
a_tensor = a_tensor.squeeze(-1)
b_tensor = b_tensor.squeeze(-1)
sfa_tensor = sfa_tensor.squeeze(-1)
sfb_tensor = sfb_tensor.squeeze(-1)
# Input Shapes
M, K_half = a_tensor.shape
N = b_tensor.shape[0]
K = K_half * 2
# Configuration constants
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 512
GROUP_SZ = 16
ELEM_PER_BYTE = 2
REP_K = BLOCK_K // GROUP_SZ // 4
# --- Manual Configuration Selection ---
# Default config (fallback)
num_inner_stages = 2
warp_specialize_inner = False
num_stages = 3
num_warps = 4
# Match specific shapes
if M == 128 and N == 7168 and K == 16384:
# Config [128x7168x16384]: BK:512 OS:2 IS:3 WSO:True WSI:True Stg:2 Wrp:4
num_inner_stages = 3
warp_specialize_inner = True
num_stages = 2
num_warps = 4
elif M == 128 and N == 4096 and K == 7168:
# Config [128x4096x7168]: BK:512 OS:2 IS:3 WSO:True WSI:True Stg:3 Wrp:4
num_inner_stages = 3
warp_specialize_inner = True
num_stages = 2
num_warps = 4
elif M == 128 and N == 7168 and K == 2048:
# Config [128x7168x2048]: BK:512 OS:3 IS:2 WSO:False WSI:False Stg:2 Wrp:4
num_inner_stages = 3
warp_specialize_inner = False
num_stages = 2
num_warps = 4
a_tma = a_tensor.view(torch.uint8) # [M, K/2]
a_desc = TensorDescriptor.from_tensor(
a_tma,
block_shape=[BLOCK_M, BLOCK_K // ELEM_PER_BYTE],
)
b_tma = b_tensor.view(torch.uint8) # [N, K/2]
b_desc = TensorDescriptor.from_tensor(
b_tma,
block_shape=[BLOCK_N, BLOCK_K // ELEM_PER_BYTE],
)
rest_m = M // 128
rest_n = N // 128
rest_k = triton.cdiv(K, GROUP_SZ) // 4
# sfa_permuted: [32, 4, rest_m, 4, rest_k]
# sfb_permuted: [32, 4, rest_n, 4, rest_k]
# Permute to [rest_m or rest_n, rest_k, 32, 4, 4]
sfa_back = sfa_tensor.permute(2, 4, 0, 1, 3)
sfb_back = sfb_tensor.permute(2, 4, 0, 1, 3)
# Pack final three dims: (rest_m, rest_k, 32, 4, 4) -> (rest_m, rest_k, 2, 256)
a_scale_packed = sfa_back.view(rest_m, rest_k, 2, 256)
b_scale_packed = sfb_back.view(rest_n, rest_k, 2, 256)
a_scale_desc = TensorDescriptor.from_tensor(
a_scale_packed,
block_shape=[1, REP_K, 2, 256],
)
b_scale_desc = TensorDescriptor.from_tensor(
b_scale_packed,
block_shape=[1, REP_K, 2, 256],
)
stride_cm, stride_cn, _ = c_tensor.stride()
# Launch Grid
num_tiles = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N)
grid=(num_tiles,)
block_scaled_batched_gemm_kernel[grid](
a_desc,
a_scale_desc,
b_desc,
b_scale_desc,
c_tensor,
stride_cm,
stride_cn,
M,
N,
K,
ELEM_PER_BYTE,
GROUP_SZ,
BLOCK_M,
BLOCK_N,
BLOCK_K,
REP_K,
NUM_INNER_STAGES=num_inner_stages,
WARP_SPECIALIZE_INNER=warp_specialize_inner,
num_warps=num_warps,
num_stages=num_stages
)
return c_tensor
scrolls · 205 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 145523.
⋯ 21 unchanged linesBLOCK_N: tl.constexpr,BLOCK_K: tl.constexpr,REP_K: tl.constexpr,- NUM_OUTER_STAGES: tl.constexpr,NUM_INNER_STAGES: tl.constexpr,- WARP_SPECIALIZE_OUTER: tl.constexpr,WARP_SPECIALIZE_INNER: tl.constexpr,- FLATTEN: tl.constexpr,):BLOCK_K_ELEM_PER_BYTE: tl.constexpr = BLOCK_K // ELEM_PER_BYTEBLOCK_K_GROUP_SZ: tl.constexpr = BLOCK_K // GROUP_SZpid = tl.program_id(axis=0)- num_pid = tl.num_programs(axis=0)-- num_pid_m = tl.cdiv(M, BLOCK_M)num_pid_n = tl.cdiv(N, BLOCK_N)- total_tiles = num_pid_m * num_pid_n+ pid_m = pid // num_pid_n+ pid_n = pid % num_pid_n+ # Base offsets for this tile+ offs_am = pid_m * BLOCK_M+ offs_bn = pid_n * BLOCK_N+ accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)- for linear in tl.range(- pid,- total_tiles,- num_pid,- num_stages=NUM_OUTER_STAGES,- flatten=FLATTEN,- warp_specialize=WARP_SPECIALIZE_OUTER,+ for i in tl.range(+ 0,+ tl.cdiv(K, BLOCK_K),+ num_stages=NUM_INNER_STAGES,+ warp_specialize=WARP_SPECIALIZE_INNER,):- tile_id = linear % (num_pid_m * num_pid_n)- pid_m = tile_id // num_pid_n- pid_n = tile_id % num_pid_n+ offs_k = i * BLOCK_K_ELEM_PER_BYTE+ offs_scale_k = i * REP_K- # Base offsets for this tile- offs_am = pid_m * BLOCK_M- offs_bn = pid_n * BLOCK_N+ # A: [BLOCK_M, BLOCK_K/2]+ # B: [BLOCK_N, BLOCK_K/2]+ a = a_desc.load([offs_am, offs_k])+ b = b_desc.load([offs_bn, offs_k])- accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)+ scale_a = (+ a_scale_desc.load([pid_m, offs_scale_k, 0, 0])+ .reshape(REP_K, 32, 4, 4)+ .trans(2, 1, 0, 3)+ .reshape(BLOCK_M, BLOCK_K_GROUP_SZ)+ )- for i in tl.range(- 0,- tl.cdiv(K, BLOCK_K),- num_stages=NUM_INNER_STAGES,- warp_specialize=WARP_SPECIALIZE_INNER,- ):- offs_k = i * BLOCK_K_ELEM_PER_BYTE- offs_scale_k = i * REP_K+ scale_b = (+ b_scale_desc.load([pid_n, offs_scale_k, 0, 0])+ .reshape(REP_K, 32, 4, 4)+ .trans(2, 1, 0, 3)+ .reshape(BLOCK_N, BLOCK_K_GROUP_SZ)+ )- # A: [BLOCK_M, BLOCK_K/2]- # B: [BLOCK_N, BLOCK_K/2]- a = a_desc.load([offs_am, offs_k])- b = b_desc.load([offs_bn, offs_k])-- scale_a = (- a_scale_desc.load([pid_m, offs_scale_k, 0, 0])- .reshape(REP_K, 32, 4, 4)- .trans(2, 1, 0, 3)- .reshape(BLOCK_M, BLOCK_K_GROUP_SZ)- )-- scale_b = (- b_scale_desc.load([pid_n, offs_scale_k, 0, 0])- .reshape(REP_K, 32, 4, 4)- .trans(2, 1, 0, 3)- .reshape(BLOCK_N, BLOCK_K_GROUP_SZ)- )-- accumulator = tl.dot_scaled(- a,- scale_a,- "e2m1",- b.T,- scale_b,- "e2m1",- accumulator,- )-- # Calculate output pointers- offset_m = offs_am + tl.arange(0, BLOCK_M)- offset_n = offs_bn + tl.arange(0, BLOCK_N)-- c_off = (- offset_m[:, None] * stride_cm- + offset_n[None, :] * stride_cn+ accumulator = tl.dot_scaled(+ a,+ scale_a,+ "e2m1",+ b.T,+ scale_b,+ "e2m1",+ accumulator,)-- c_mask = (offset_m[:, None] < M) & (offset_n[None, :] < N)- tl.store(c_ptr + c_off, accumulator.to(tl.float16), mask=c_mask)+ # Calculate output pointers+ offset_m = offs_am + tl.arange(0, BLOCK_M)+ offset_n = offs_bn + tl.arange(0, BLOCK_N)++ c_off = (+ offset_m[:, None] * stride_cm+ + offset_n[None, :] * stride_cn+ )++ c_mask = (offset_m[:, None] < M) & (offset_n[None, :] < N)+ tl.store(c_ptr + c_off, accumulator.to(tl.float16), mask=c_mask)+def custom_kernel(data):a_tensor, b_tensor, _, _, sfa_tensor, sfb_tensor, c_tensor = data⋯ 18 unchanged lines# --- Manual Configuration Selection ---# Default config (fallback)- num_outer_stages = 2num_inner_stages = 2- warp_specialize_outer = Truewarp_specialize_inner = Falsenum_stages = 3num_warps = 4⋯ 1 unchanged lines# Match specific shapesif M == 128 and N == 7168 and K == 16384:# Config [128x7168x16384]: BK:512 OS:2 IS:3 WSO:True WSI:True Stg:2 Wrp:4- num_outer_stages = 2num_inner_stages = 3- warp_specialize_outer = Truewarp_specialize_inner = Truenum_stages = 2num_warps = 4elif M == 128 and N == 4096 and K == 7168:# Config [128x4096x7168]: BK:512 OS:2 IS:3 WSO:True WSI:True Stg:3 Wrp:4- num_outer_stages = 2num_inner_stages = 3- warp_specialize_outer = Truewarp_specialize_inner = True- num_stages = 3+ num_stages = 2num_warps = 4elif M == 128 and N == 7168 and K == 2048:# Config [128x7168x2048]: BK:512 OS:3 IS:2 WSO:False WSI:False Stg:2 Wrp:4- num_outer_stages = 3- num_inner_stages = 2- warp_specialize_outer = False+ num_inner_stages = 3warp_specialize_inner = Falsenum_stages = 2num_warps = 4⋯ 56 unchanged linesBLOCK_N,BLOCK_K,REP_K,- NUM_OUTER_STAGES=num_outer_stages,NUM_INNER_STAGES=num_inner_stages,- WARP_SPECIALIZE_OUTER=warp_specialize_outer,WARP_SPECIALIZE_INNER=warp_specialize_inner,- FLATTEN=True,num_warps=num_warps,num_stages=num_stages)
scrolls · 191 diff lines total
Best evidence level for this revision: reported
JSON