submission 145523
Arseni Ivanov · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 235 lines, June 9 Researcher Reciprocity License v1.0.
triton_hardcoded.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-145523?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:653ba9d0b6f071d23d00cf80faaba2e1893aaf2b35ea00e34142044fe79af8b7
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 = 4persistent-kernel
num_pid = tl.num_programs(axis=0)stages = 3
num_stages = 3tile-k = 512
BLOCK_K = 512tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128warp-specialization
WARP_SPECIALIZE_OUTER: tl.constexpr,Kernel source
triton_hardcoded.py235 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_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_BYTE
BLOCK_K_GROUP_SZ: tl.constexpr = BLOCK_K // GROUP_SZ
pid = 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
for linear in tl.range(
pid,
total_tiles,
num_pid,
num_stages=NUM_OUTER_STAGES,
flatten=FLATTEN,
warp_specialize=WARP_SPECIALIZE_OUTER,
):
tile_id = linear % (num_pid_m * num_pid_n)
pid_m = tile_id // num_pid_n
pid_n = tile_id % 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_outer_stages = 2
num_inner_stages = 2
warp_specialize_outer = True
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_outer_stages = 2
num_inner_stages = 3
warp_specialize_outer = True
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_outer_stages = 2
num_inner_stages = 3
warp_specialize_outer = True
warp_specialize_inner = True
num_stages = 3
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_outer_stages = 3
num_inner_stages = 2
warp_specialize_outer = False
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_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
)
return c_tensor
scrolls · 235 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 144525.
#!POPCORN leaderboard nvfp4_gemm- import functoolsimport torchimport tritonimport triton.language as tlfrom triton.tools.tensor_descriptor import TensorDescriptor- def _matmul_launch_metadata(grid, kernel, args):- M, N, K = args["M"], args["N"], args["K"]- return {- "name": f"{kernel.name} [M={M}, N={N}, K={K}]",- "flops": 2.0 * M * N * K,- }--- def _config(**autotune_kwargs):- class inner:- def __init__(self, fn):- self.fn = fn-- def __getitem__(self, s):- return functools.partial(self.fn[s], **autotune_kwargs)-- return inner--- @_config(- NUM_OUTER_STAGES=None,- NUM_INNER_STAGES=None,- WARP_SPECIALIZE_OUTER=True,- WARP_SPECIALIZE_INNER=False,- FLATTEN=True,- num_warps=4,- num_stages=3,- num_ctas=1,- )- @triton.jit(launch_metadata=_matmul_launch_metadata)+ @triton.jitdef block_scaled_batched_gemm_kernel(a_desc,a_scale_desc,⋯ 72 unchanged lines.trans(2, 1, 0, 3).reshape(BLOCK_N, BLOCK_K_GROUP_SZ))+accumulator = tl.dot_scaled(a,scale_a,⋯ 31 unchanged linesN = b_tensor.shape[0]K = K_half * 2- # Configuration+ # Configuration constantsBLOCK_M = 128BLOCK_N = 128BLOCK_K = 512GROUP_SZ = 16ELEM_PER_BYTE = 2- SM_MULT = 1REP_K = BLOCK_K // GROUP_SZ // 4+ # --- Manual Configuration Selection ---+ # Default config (fallback)+ num_outer_stages = 2+ num_inner_stages = 2+ warp_specialize_outer = True+ 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_outer_stages = 2+ num_inner_stages = 3+ warp_specialize_outer = True+ 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_outer_stages = 2+ num_inner_stages = 3+ warp_specialize_outer = True+ warp_specialize_inner = True+ num_stages = 3+ 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_outer_stages = 3+ num_inner_stages = 2+ warp_specialize_outer = False+ 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,⋯ 33 unchanged lines# Launch Gridnum_tiles = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N)- num_sms = torch.cuda.get_device_properties(a_tensor.device).multi_processor_count- # Persistent kernel grid size- grid = (min(num_tiles, num_sms * SM_MULT),)-+ grid=(num_tiles,)+block_scaled_batched_gemm_kernel[grid](a_desc,a_scale_desc,⋯ 11 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)return c_tensor
scrolls · 128 diff lines total
Best evidence level for this revision: reported
JSON