submission 125299
shiyeegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 283 lines, June 9 Researcher Reciprocity License v1.0.
node_159.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-125299?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:d6d555ea8b22d41a99a0a2e5734fe581287583f7e7636116fdae8a277d226d26
license declaredunknown
license concludedunknown
authorsshiyeegao
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 8
num_warps = 8 if block_n == 256 else 4persistent-kernel
num_pid = tl.num_programs(axis=0)split-k
SPLIT_K: tl.constexpr,stages = 5
num_stages = 5tile-m = 128
BLOCK_M = 128Kernel source
node_159.py283 lines
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor
# -------------------------------------------------------------------------
# Kernel Metadata
# -------------------------------------------------------------------------
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,
}
# -------------------------------------------------------------------------
# Block-Scaled NVFP4 GEMM with TMA (Persistent)
# -------------------------------------------------------------------------
@triton.jit(launch_metadata=_matmul_launch_metadata)
def bmm_fp4_tma_kernel(
a_desc, # TMA descriptor for A: [M, L, K/2]
a_scale_desc, # TMA descriptor for packed scales of A
b_desc, # TMA descriptor for B: [N, L, K/2]
b_scale_desc, # TMA descriptor for packed scales of B
c_ptr, # Output pointer: [M, N, L]
# Strides
stride_cm, stride_cn, stride_cl,
# Dimensions
M, N, K, L,
# Constants
ELEM_PER_BYTE: tl.constexpr,
GROUP_SZ: tl.constexpr,
SPLIT_K: tl.constexpr,
# Block Tuning
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
REP_M: tl.constexpr,
REP_N: tl.constexpr,
REP_K: tl.constexpr,
# Compiler Hints
NUM_STAGES: tl.constexpr,
):
output_dtype: tl.constexpr = tl.float16 if SPLIT_K == 1 else tl.float32
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
tl.static_assert(BLOCK_M == 128)
tl.static_assert(BLOCK_N % 128 == 0)
tl.static_assert(BLOCK_K % (GROUP_SZ * 4) == 0)
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)
num_tiles_per_batch = num_pid_m * num_pid_n * SPLIT_K
total_tiles = num_tiles_per_batch * L
k_tiles_total = tl.cdiv(K, BLOCK_K)
part_k = tl.cdiv(k_tiles_total, SPLIT_K)
# Persistent loop keeps blocks active across tiles
for linear_id in tl.range(pid, total_tiles, num_pid, num_stages=NUM_STAGES):
tile_split = linear_id % SPLIT_K
tmp = linear_id // SPLIT_K
tile_n = tmp % num_pid_n
tmp = tmp // num_pid_n
tile_m = tmp % num_pid_m
tile_l = tmp // num_pid_m
offs_am = tile_m * BLOCK_M
offs_bn = tile_n * BLOCK_N
offs_scale_m = tile_m * REP_M
offs_scale_n = tile_n * REP_N
k_start = tile_split * part_k
k_end = tl.minimum(k_tiles_total, k_start + part_k)
if k_start < k_end:
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=acc_dtype)
for k_idx in tl.range(k_start, k_end):
offs_k = k_idx * BLOCK_K_ELEM_PER_BYTE
offs_scale_k = k_idx * REP_K
# TMA loads for packed FP4 tiles
a_tile = a_desc.load([offs_am, tile_l, offs_k])
b_tile = b_desc.load([offs_bn, tile_l, offs_k])
a_tile = a_tile.reshape(BLOCK_M, BLOCK_K_ELEM_PER_BYTE)
b_tile = b_tile.reshape(BLOCK_N, BLOCK_K_ELEM_PER_BYTE)
# Load block scales (already swizzled for WGMMA layout)
scale_a_pack = a_scale_desc.load([tile_l, offs_scale_m, offs_scale_k, 0, 0])
scale_b_pack = b_scale_desc.load([tile_l, offs_scale_n, offs_scale_k, 0, 0])
scale_a = (
scale_a_pack.reshape(REP_M, REP_K, 32, 4, 4)
.trans(0, 3, 2, 1, 4)
.reshape(BLOCK_M, BLOCK_K_GROUP_SZ)
)
scale_b = (
scale_b_pack.reshape(REP_N, REP_K, 32, 4, 4)
.trans(0, 3, 2, 1, 4)
.reshape(BLOCK_N, BLOCK_K_GROUP_SZ)
)
accumulator = tl.dot_scaled(
a_tile,
scale_a,
"e2m1",
b_tile.T,
scale_b,
"e2m1",
accumulator,
)
offs_cm = tile_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_cn = tile_n * BLOCK_N + tl.arange(0, BLOCK_N)
c_offset = (
offs_cm[:, None] * stride_cm
+ offs_cn[None, :] * stride_cn
+ tile_l * stride_cl
)
mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
if SPLIT_K == 1:
tl.store(c_ptr + c_offset, accumulator.to(output_dtype), mask=mask)
else:
tl.atomic_add(c_ptr + c_offset, accumulator.to(output_dtype), mask=mask)
# -------------------------------------------------------------------------
# Host API
# -------------------------------------------------------------------------
def _select_config(M: int, N: int, K: int, num_sms: int):
"""Shape-aware tile selection tuned for leaderboard shapes."""
wide_n = N >= 4096
block_n = 256 if (wide_n and K < 2048) else 128
block_k = 256
if K >= 4096:
num_stages = 5
elif K >= 2048:
num_stages = 4
else:
num_stages = 3
split_k = 1
num_warps = 8 if block_n == 256 else 4
# Persistent grid multiplier; slightly conservative to keep occupancy.
sm_mult = 8
return {
"BLOCK_N": block_n,
"BLOCK_K": block_k,
"NUM_STAGES": num_stages,
"NUM_WARPS": num_warps,
"SM_MULT": sm_mult,
"SPLIT_K": split_k,
}
@torch.inference_mode()
def custom_kernel(data):
"""Entry point expected by evaluator."""
a_tensor, b_tensor, _, _, sfa_permuted, sfb_permuted, c_tensor = data
BLOCK_M = 128
GROUP_SZ = 16
M, K_half, L = a_tensor.shape
N = b_tensor.shape[0]
K = 2 * K_half
ELEM_PER_BYTE = 2
num_sms = torch.cuda.get_device_properties(a_tensor.device).multi_processor_count
cfg = _select_config(M, N, K, num_sms)
BLOCK_N = cfg["BLOCK_N"]
BLOCK_K = cfg["BLOCK_K"]
NUM_STAGES = cfg["NUM_STAGES"]
NUM_WARPS = cfg["NUM_WARPS"]
SM_MULT = cfg["SM_MULT"]
split_k = cfg["SPLIT_K"]
REP_M = BLOCK_M // 128
REP_N = BLOCK_N // 128
REP_K = BLOCK_K // GROUP_SZ // 4
# Reorder A/B for TMA: place K as innermost for contiguous loads
a_tma = a_tensor.view(torch.uint8).permute(0, 2, 1).contiguous()
b_tma = b_tensor.view(torch.uint8).permute(0, 2, 1).contiguous()
a_desc = TensorDescriptor.from_tensor(
a_tma,
block_shape=[BLOCK_M, 1, BLOCK_K // ELEM_PER_BYTE],
)
b_desc = TensorDescriptor.from_tensor(
b_tma,
block_shape=[BLOCK_N, 1, BLOCK_K // ELEM_PER_BYTE],
)
rest_m = M // 128
rest_n = N // 128
rest_k = triton.cdiv(K, GROUP_SZ) // 4
sfa_packed = (
sfa_permuted.permute(5, 2, 4, 0, 1, 3)
.contiguous()
.view(L, rest_m, rest_k, 2, 256)
)
sfb_packed = (
sfb_permuted.permute(5, 2, 4, 0, 1, 3)
.contiguous()
.view(L, rest_n, rest_k, 2, 256)
)
a_scale_desc = TensorDescriptor.from_tensor(
sfa_packed,
block_shape=[1, REP_M, REP_K, 2, 256],
)
b_scale_desc = TensorDescriptor.from_tensor(
sfb_packed,
block_shape=[1, REP_N, REP_K, 2, 256],
)
base_tiles = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N) * L
num_tiles = base_tiles * split_k
c_buffer = c_tensor
if split_k > 1:
c_buffer = torch.zeros_like(c_tensor, dtype=torch.float32)
stride_cm, stride_cn, stride_cl = c_buffer.stride()
grid_target = num_sms * SM_MULT
grid = (max(1, min(num_tiles, grid_target)),)
bmm_fp4_tma_kernel[grid](
a_desc,
a_scale_desc,
b_desc,
b_scale_desc,
c_buffer,
stride_cm,
stride_cn,
stride_cl,
M,
N,
K,
L,
ELEM_PER_BYTE,
GROUP_SZ,
split_k,
BLOCK_M,
BLOCK_N,
BLOCK_K,
REP_M,
REP_N,
REP_K,
NUM_STAGES,
num_warps=NUM_WARPS,
num_stages=NUM_STAGES,
)
if split_k > 1:
c_tensor.copy_(c_buffer.to(dtype=c_tensor.dtype))
return c_tensor
__all__ = ["custom_kernel"]
scrolls · 283 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 121413.
- from __future__ import annotations+ import torch+ import triton+ import triton.language as tl+ from triton.tools.tensor_descriptor import TensorDescriptor- from typing import Tuple- import torch+ # -------------------------------------------------------------------------+ # Kernel Metadata+ # -------------------------------------------------------------------------+ 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 _prepare_scales_batch(scale_perm: torch.Tensor) -> torch.Tensor:- permuted = scale_perm.permute(5, 2, 4, 0, 1, 3)- return permuted.contiguous().view(scale_perm.size(-1), -1)+ # -------------------------------------------------------------------------+ # Block-Scaled NVFP4 GEMM with TMA (Persistent)+ # -------------------------------------------------------------------------+ @triton.jit(launch_metadata=_matmul_launch_metadata)+ def bmm_fp4_tma_kernel(+ a_desc, # TMA descriptor for A: [M, L, K/2]+ a_scale_desc, # TMA descriptor for packed scales of A+ b_desc, # TMA descriptor for B: [N, L, K/2]+ b_scale_desc, # TMA descriptor for packed scales of B+ c_ptr, # Output pointer: [M, N, L]+ # Strides+ stride_cm, stride_cn, stride_cl,+ # Dimensions+ M, N, K, L,+ # Constants+ ELEM_PER_BYTE: tl.constexpr,+ GROUP_SZ: tl.constexpr,+ SPLIT_K: tl.constexpr,+ # Block Tuning+ BLOCK_M: tl.constexpr,+ BLOCK_N: tl.constexpr,+ BLOCK_K: tl.constexpr,+ REP_M: tl.constexpr,+ REP_N: tl.constexpr,+ REP_K: tl.constexpr,+ # Compiler Hints+ NUM_STAGES: tl.constexpr,+ ):+ output_dtype: tl.constexpr = tl.float16 if SPLIT_K == 1 else tl.float32+ 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+ tl.static_assert(BLOCK_M == 128)+ tl.static_assert(BLOCK_N % 128 == 0)+ tl.static_assert(BLOCK_K % (GROUP_SZ * 4) == 0)++ 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)+ num_tiles_per_batch = num_pid_m * num_pid_n * SPLIT_K+ total_tiles = num_tiles_per_batch * L++ k_tiles_total = tl.cdiv(K, BLOCK_K)+ part_k = tl.cdiv(k_tiles_total, SPLIT_K)++ # Persistent loop keeps blocks active across tiles+ for linear_id in tl.range(pid, total_tiles, num_pid, num_stages=NUM_STAGES):+ tile_split = linear_id % SPLIT_K+ tmp = linear_id // SPLIT_K+ tile_n = tmp % num_pid_n+ tmp = tmp // num_pid_n+ tile_m = tmp % num_pid_m+ tile_l = tmp // num_pid_m++ offs_am = tile_m * BLOCK_M+ offs_bn = tile_n * BLOCK_N+ offs_scale_m = tile_m * REP_M+ offs_scale_n = tile_n * REP_N++ k_start = tile_split * part_k+ k_end = tl.minimum(k_tiles_total, k_start + part_k)++ if k_start < k_end:+ accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=acc_dtype)++ for k_idx in tl.range(k_start, k_end):+ offs_k = k_idx * BLOCK_K_ELEM_PER_BYTE+ offs_scale_k = k_idx * REP_K++ # TMA loads for packed FP4 tiles+ a_tile = a_desc.load([offs_am, tile_l, offs_k])+ b_tile = b_desc.load([offs_bn, tile_l, offs_k])++ a_tile = a_tile.reshape(BLOCK_M, BLOCK_K_ELEM_PER_BYTE)+ b_tile = b_tile.reshape(BLOCK_N, BLOCK_K_ELEM_PER_BYTE)++ # Load block scales (already swizzled for WGMMA layout)+ scale_a_pack = a_scale_desc.load([tile_l, offs_scale_m, offs_scale_k, 0, 0])+ scale_b_pack = b_scale_desc.load([tile_l, offs_scale_n, offs_scale_k, 0, 0])++ scale_a = (+ scale_a_pack.reshape(REP_M, REP_K, 32, 4, 4)+ .trans(0, 3, 2, 1, 4)+ .reshape(BLOCK_M, BLOCK_K_GROUP_SZ)+ )+ scale_b = (+ scale_b_pack.reshape(REP_N, REP_K, 32, 4, 4)+ .trans(0, 3, 2, 1, 4)+ .reshape(BLOCK_N, BLOCK_K_GROUP_SZ)+ )++ accumulator = tl.dot_scaled(+ a_tile,+ scale_a,+ "e2m1",+ b_tile.T,+ scale_b,+ "e2m1",+ accumulator,+ )++ offs_cm = tile_m * BLOCK_M + tl.arange(0, BLOCK_M)+ offs_cn = tile_n * BLOCK_N + tl.arange(0, BLOCK_N)++ c_offset = (+ offs_cm[:, None] * stride_cm+ + offs_cn[None, :] * stride_cn+ + tile_l * stride_cl+ )++ mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)++ if SPLIT_K == 1:+ tl.store(c_ptr + c_offset, accumulator.to(output_dtype), mask=mask)+ else:+ tl.atomic_add(c_ptr + c_offset, accumulator.to(output_dtype), mask=mask)+++ # -------------------------------------------------------------------------+ # Host API+ # -------------------------------------------------------------------------+ def _select_config(M: int, N: int, K: int, num_sms: int):+ """Shape-aware tile selection tuned for leaderboard shapes."""+ wide_n = N >= 4096++ block_n = 256 if (wide_n and K < 2048) else 128+ block_k = 256++ if K >= 4096:+ num_stages = 5+ elif K >= 2048:+ num_stages = 4+ else:+ num_stages = 3++ split_k = 1++ num_warps = 8 if block_n == 256 else 4++ # Persistent grid multiplier; slightly conservative to keep occupancy.+ sm_mult = 8++ return {+ "BLOCK_N": block_n,+ "BLOCK_K": block_k,+ "NUM_STAGES": num_stages,+ "NUM_WARPS": num_warps,+ "SM_MULT": sm_mult,+ "SPLIT_K": split_k,+ }++@torch.inference_mode()- def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:- a, b, _, _, sfa_perm, sfb_perm, c = data+ def custom_kernel(data):+ """Entry point expected by evaluator."""+ a_tensor, b_tensor, _, _, sfa_permuted, sfb_permuted, c_tensor = data- _, _, l = c.shape+ BLOCK_M = 128+ GROUP_SZ = 16- scale_a_batch = _prepare_scales_batch(sfa_perm)- scale_b_batch = _prepare_scales_batch(sfb_perm)+ M, K_half, L = a_tensor.shape+ N = b_tensor.shape[0]+ K = 2 * K_half+ ELEM_PER_BYTE = 2- for i in range(l):- res = torch._scaled_mm(- a[:, :, i],- b[:, :, i].transpose(0, 1),- scale_a_batch[i],- scale_b_batch[i],- bias=None,- out_dtype=torch.float16,- )- c[:, :, i].copy_(res)+ num_sms = torch.cuda.get_device_properties(a_tensor.device).multi_processor_count+ cfg = _select_config(M, N, K, num_sms)+ BLOCK_N = cfg["BLOCK_N"]+ BLOCK_K = cfg["BLOCK_K"]+ NUM_STAGES = cfg["NUM_STAGES"]+ NUM_WARPS = cfg["NUM_WARPS"]+ SM_MULT = cfg["SM_MULT"]+ split_k = cfg["SPLIT_K"]- return c+ REP_M = BLOCK_M // 128+ REP_N = BLOCK_N // 128+ REP_K = BLOCK_K // GROUP_SZ // 4+ # Reorder A/B for TMA: place K as innermost for contiguous loads+ a_tma = a_tensor.view(torch.uint8).permute(0, 2, 1).contiguous()+ b_tma = b_tensor.view(torch.uint8).permute(0, 2, 1).contiguous()- __all__ = ["custom_kernel"]No newline at end of file+ a_desc = TensorDescriptor.from_tensor(+ a_tma,+ block_shape=[BLOCK_M, 1, BLOCK_K // ELEM_PER_BYTE],+ )+ b_desc = TensorDescriptor.from_tensor(+ b_tma,+ block_shape=[BLOCK_N, 1, BLOCK_K // ELEM_PER_BYTE],+ )++ rest_m = M // 128+ rest_n = N // 128+ rest_k = triton.cdiv(K, GROUP_SZ) // 4++ sfa_packed = (+ sfa_permuted.permute(5, 2, 4, 0, 1, 3)+ .contiguous()+ .view(L, rest_m, rest_k, 2, 256)+ )+ sfb_packed = (+ sfb_permuted.permute(5, 2, 4, 0, 1, 3)+ .contiguous()+ .view(L, rest_n, rest_k, 2, 256)+ )++ a_scale_desc = TensorDescriptor.from_tensor(+ sfa_packed,+ block_shape=[1, REP_M, REP_K, 2, 256],+ )+ b_scale_desc = TensorDescriptor.from_tensor(+ sfb_packed,+ block_shape=[1, REP_N, REP_K, 2, 256],+ )++ base_tiles = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N) * L+ num_tiles = base_tiles * split_k++ c_buffer = c_tensor+ if split_k > 1:+ c_buffer = torch.zeros_like(c_tensor, dtype=torch.float32)++ stride_cm, stride_cn, stride_cl = c_buffer.stride()++ grid_target = num_sms * SM_MULT+ grid = (max(1, min(num_tiles, grid_target)),)++ bmm_fp4_tma_kernel[grid](+ a_desc,+ a_scale_desc,+ b_desc,+ b_scale_desc,+ c_buffer,+ stride_cm,+ stride_cn,+ stride_cl,+ M,+ N,+ K,+ L,+ ELEM_PER_BYTE,+ GROUP_SZ,+ split_k,+ BLOCK_M,+ BLOCK_N,+ BLOCK_K,+ REP_M,+ REP_N,+ REP_K,+ NUM_STAGES,+ num_warps=NUM_WARPS,+ num_stages=NUM_STAGES,+ )++ if split_k > 1:+ c_tensor.copy_(c_buffer.to(dtype=c_tensor.dtype))++ return c_tensor+++ __all__ = ["custom_kernel"]
scrolls · 306 diff lines total
Best evidence level for this revision: reported
JSON