submission 120588
Arseni Ivanov · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 249 lines, June 9 Researcher Reciprocity License v1.0.
triton_improved.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-120588?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:594d1a1ecf35b0f459e890a84aec07b8c006705d545a9679dbfbd2658e100d0e
license declaredunknown
license concludedunknown
authorsArseni Ivanov
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
def _config(**autotune_kwargs):num-warps = 4
num_warps=4,persistent-kernel
num_pid = tl.num_programs(axis=0)stages = 3
num_stages=3,tile-k = 256
BLOCK_K = 256tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128warp-specialization
WARP_SPECIALIZE_OUTER=True,Kernel source
triton_improved.py249 lines
#!POPCORN leaderboard nvfp4_gemm
import functools
import torch
import triton
import triton.language as tl
from 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)
def block_scaled_batched_gemm_kernel(
a_desc,
a_scale_desc,
b_desc,
b_scale_desc,
c_ptr,
stride_cm,
stride_cn,
stride_cl,
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_M: tl.constexpr,
REP_N: 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,
):
output_dtype: tl.constexpr = tl.float16
acc_dtype: tl.constexpr = tl.float32
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,
):
# Decode linear index into (pid_m, pid_n, pid_b)
tile_id = linear % (num_pid_m * num_pid_n)
pid_b = 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
offs_scale_m = pid_m * REP_M
offs_scale_n = pid_n * REP_N
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=acc_dtype)
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])
# Reconstruct packed scales for A
scale_a = (
a_scale_desc.load([pid_b, offs_scale_m, offs_scale_k, 0, 0])
.reshape(REP_M, REP_K, 32, 4, 4)
.trans(0, 3, 2, 1, 4)
.reshape(BLOCK_M, BLOCK_K_GROUP_SZ)
)
# Reconstruct packed scales for B (using pid_n/REP_N)
scale_b = (
b_scale_desc.load([pid_b, offs_scale_n, offs_scale_k, 0, 0])
.reshape(REP_M, REP_K, 32, 4, 4)
.trans(0, 3, 2, 1, 4)
.reshape(BLOCK_N, BLOCK_K_GROUP_SZ)
)
# Scaled Dot Product: A * B.T
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
+ pid_b * stride_cl
)
c_mask = (offset_m[:, None] < M) & (offset_n[None, :] < N)
tl.store(c_ptr + c_off, accumulator.to(output_dtype), mask=c_mask, cache_modifier=".cg")
def custom_kernel(data):
a_tensor, b_tensor, _, _, sfa_tensor, sfb_tensor, c_tensor = data
a_tensor = a_tensor.squeeze(-1)
b_tensor = b_tensor.squeeze(-1)
# Input Shapes
# a: [M, K/2, L], b: [N, K/2, L], sfa: [M, K/16, L], sfb: [N, K/16, L]
M, K_half = a_tensor.shape
N = b_tensor.shape[0]
K = K_half * 2
# Configuration
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 256
GROUP_SZ = 16
ELEM_PER_BYTE = 2
SM_MULT = 1
REP_M = BLOCK_M // 128
REP_N = BLOCK_N // 128
REP_K = BLOCK_K // GROUP_SZ // 4
# Prepare TMA Descriptors
# View as uint8 for TMA to handle the 4-bit packed data correctly
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],
)
# Scales: invert CuTe layout and pack for TMA
# sfa_permuted: [32, 4, rest_m, 4, rest_k, L]
# sfb_permuted: [32, 4, rest_n, 4, rest_k, L]
rest_m = M // 128
rest_n = N // 128 # = 1
rest_k = triton.cdiv(K, GROUP_SZ) // 4
# Permute to [L, rest_m, rest_k, 32, 4, 4]
# Permute to [rest_m, rest_k, 32, 4, 4]
sfa_back = sfa_tensor.permute(5, 2, 4, 0, 1, 3)
sfb_back = sfb_tensor.permute(5, 2, 4, 0, 1, 3)
assert sfa_back.shape == (1, rest_m, rest_k, 32, 4, 4)
assert sfb_back.shape == (1, rest_n, rest_k, 32, 4, 4)
# Pack final three dims: (L, rest_m, rest_k, 32, 4, 4) -> (L, rest_m, rest_k, 2, 256)
a_scale_packed = sfa_back.view(1, rest_m, rest_k, 2, 256)
b_scale_packed = sfb_back.view(1, rest_n, rest_k, 2, 256)
a_scale_desc = TensorDescriptor.from_tensor(
a_scale_packed,
block_shape=[1, REP_M, REP_K, 2, 256],
)
b_scale_desc = TensorDescriptor.from_tensor(
b_scale_packed,
block_shape=[1, REP_N, REP_K, 2, 256],
)
stride_cm, stride_cn, stride_cl = c_tensor.stride()
# Launch Grid
num_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),)
block_scaled_batched_gemm_kernel[grid](
a_desc,
a_scale_desc,
b_desc,
b_scale_desc,
c_tensor,
stride_cm,
stride_cn,
stride_cl,
M,
N,
K,
ELEM_PER_BYTE,
GROUP_SZ,
BLOCK_M,
BLOCK_N,
BLOCK_K,
REP_M,
REP_N,
REP_K,
)
return c_tensor
scrolls · 249 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 120443.
⋯ 46 unchanged linesM,N,K,- L,ELEM_PER_BYTE: tl.constexpr,GROUP_SZ: tl.constexpr,BLOCK_M: tl.constexpr,⋯ 18 unchanged linesnum_pid_m = tl.cdiv(M, BLOCK_M)num_pid_n = tl.cdiv(N, BLOCK_N)- total_tiles = num_pid_m * num_pid_n * L+ total_tiles = num_pid_m * num_pid_nfor linear in tl.range(pid,⋯ 27 unchanged linesoffs_k = i * BLOCK_K_ELEM_PER_BYTEoffs_scale_k = i * REP_K- # A: [BLOCK_M, 1, BLOCK_K/2]- # B: [BLOCK_N, 1, BLOCK_K/2]- a = a_desc.load([offs_am, pid_b, offs_k])- b = b_desc.load([offs_bn, pid_b, offs_k])-- a = a.reshape(BLOCK_M, BLOCK_K_ELEM_PER_BYTE)- b = b.reshape(BLOCK_N, BLOCK_K_ELEM_PER_BYTE)+ # 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])# Reconstruct packed scales for Ascale_a = (⋯ 6 unchanged lines# Reconstruct packed scales for B (using pid_n/REP_N)scale_b = (b_scale_desc.load([pid_b, offs_scale_n, offs_scale_k, 0, 0])- .reshape(REP_N, REP_K, 32, 4, 4)+ .reshape(REP_M, REP_K, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_N, BLOCK_K_GROUP_SZ))⋯ 20 unchanged lines)c_mask = (offset_m[:, None] < M) & (offset_n[None, :] < N)- tl.store(c_ptr + c_off, accumulator.to(output_dtype), mask=c_mask)+ tl.store(c_ptr + c_off, accumulator.to(output_dtype), mask=c_mask, cache_modifier=".cg")def custom_kernel(data):a_tensor, b_tensor, _, _, sfa_tensor, sfb_tensor, c_tensor = data-++ a_tensor = a_tensor.squeeze(-1)+ b_tensor = b_tensor.squeeze(-1)# Input Shapes# a: [M, K/2, L], b: [N, K/2, L], sfa: [M, K/16, L], sfb: [N, K/16, L]- M, K_half, L = a_tensor.shape+ M, K_half = a_tensor.shapeN = b_tensor.shape[0]K = K_half * 2⋯ 11 unchanged lines# Prepare TMA Descriptors# View as uint8 for TMA to handle the 4-bit packed data correctly- a_tma = a_tensor.view(torch.uint8).permute(0, 2, 1) # [M, L, K/2]+ a_tma = a_tensor.view(torch.uint8) # [M, K/2]a_desc = TensorDescriptor.from_tensor(a_tma,- block_shape=[BLOCK_M, 1, BLOCK_K // ELEM_PER_BYTE],+ block_shape=[BLOCK_M, BLOCK_K // ELEM_PER_BYTE],)- b_tma = b_tensor.view(torch.uint8).permute(0, 2, 1) # [N, L, K/2]+ b_tma = b_tensor.view(torch.uint8)# [N, K/2]b_desc = TensorDescriptor.from_tensor(b_tma,- block_shape=[BLOCK_N, 1, BLOCK_K // ELEM_PER_BYTE],+ block_shape=[BLOCK_N, BLOCK_K // ELEM_PER_BYTE],)# Scales: invert CuTe layout and pack for TMA⋯ 4 unchanged linesrest_k = triton.cdiv(K, GROUP_SZ) // 4# Permute to [L, rest_m, rest_k, 32, 4, 4]+ # Permute to [rest_m, rest_k, 32, 4, 4]sfa_back = sfa_tensor.permute(5, 2, 4, 0, 1, 3)sfb_back = sfb_tensor.permute(5, 2, 4, 0, 1, 3)- assert sfa_back.shape == (L, rest_m, rest_k, 32, 4, 4)- assert sfb_back.shape == (L, rest_n, rest_k, 32, 4, 4)+ assert sfa_back.shape == (1, rest_m, rest_k, 32, 4, 4)+ assert sfb_back.shape == (1, rest_n, rest_k, 32, 4, 4)# Pack final three dims: (L, rest_m, rest_k, 32, 4, 4) -> (L, rest_m, rest_k, 2, 256)- a_scale_packed = sfa_back.view(L, rest_m, rest_k, 2, 256)- b_scale_packed = sfb_back.view(L, rest_n, rest_k, 2, 256)+ a_scale_packed = sfa_back.view(1, rest_m, rest_k, 2, 256)+ b_scale_packed = sfb_back.view(1, rest_n, rest_k, 2, 256)a_scale_desc = TensorDescriptor.from_tensor(a_scale_packed,block_shape=[1, REP_M, REP_K, 2, 256],⋯ 7 unchanged linesstride_cm, stride_cn, stride_cl = c_tensor.stride()# Launch Grid- num_tiles = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N) * L+ num_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 sizegrid = (min(num_tiles, num_sms * SM_MULT),)⋯ 10 unchanged linesM,N,K,- L,ELEM_PER_BYTE,GROUP_SZ,BLOCK_M,
scrolls · 123 diff lines total
Best evidence level for this revision: reported
JSON