submission 118340
poornaravuri · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 164 lines, June 9 Researcher Reciprocity License v1.0.
submission_triton.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-118340?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:e93b7ac982ef292d8820a4de3d0b63eb30660498b9d80f00a4d4d92867121d7f
license declaredunknown
license concludedunknown
authorspoornaravuri
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
other=tl.zeros((BLOCK_M, num_vec), dtype=tl.float8e4nv),num-warps = 8
num_warps=8,stages = 2
num_stages=2,tile-k = 256
BLOCK_K = 256tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128Kernel source
submission_triton.py164 lines
import torch
try:
import triton
import triton.language as tl
except Exception:
triton = None
from task import input_t, output_t
from utils import make_match_reference
sf_vec_size = 16
def ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
if triton is not None:
@triton.jit
def _bsmm_kernel(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
M, N, K, sfK,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_sm, stride_sk,
stride_tn, stride_tk,
stride_cm, stride_cn,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
VEC_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
num_pid_m = tl.cdiv(M, BLOCK_M)
pid_m = pid % num_pid_m
pid_n = pid // num_pid_m
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 < M
n_mask = offs_n < N
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
K_packed = K // 2
BLOCK_K_PACKED: tl.constexpr = BLOCK_K // 2
for k0 in range(0, K, BLOCK_K):
offs_k_packed = (k0 // 2) + tl.arange(0, BLOCK_K_PACKED)
k_mask = offs_k_packed < K_packed
a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k_packed[None, :] * stride_ak
b_ptrs = b_ptr + offs_n[:, None] * stride_bn + offs_k_packed[None, :] * stride_bk
a_tile = tl.load(a_ptrs, mask=m_mask[:, None] & k_mask[None, :], other=0)
b_tile = tl.load(b_ptrs, mask=n_mask[:, None] & k_mask[None, :], other=0)
num_vec: tl.constexpr = BLOCK_K // VEC_SIZE
offs_vec = (k0 // VEC_SIZE) + tl.arange(0, num_vec)
vec_mask = offs_vec < sfK
sfa_ptrs = sfa_ptr + offs_m[:, None] * stride_sm + offs_vec[None, :] * stride_sk
sfb_ptrs = sfb_ptr + offs_n[:, None] * stride_tn + offs_vec[None, :] * stride_tk
scale_a = tl.load(
sfa_ptrs,
mask=m_mask[:, None] & vec_mask[None, :],
other=tl.zeros((BLOCK_M, num_vec), dtype=tl.float8e4nv),
)
scale_b = tl.load(
sfb_ptrs,
mask=n_mask[:, None] & vec_mask[None, :],
other=tl.zeros((BLOCK_N, num_vec), dtype=tl.float8e4nv),
)
acc = tl.dot_scaled(a_tile, scale_a, "e2m1", b_tile.T, scale_b, "e2m1", acc)
c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
out = acc.to(c_ptr.type.element_ty)
tl.store(c_ptrs, out, mask=m_mask[:, None] & n_mask[None, :])
def _fallback_scaled_mm(data: input_t) -> output_t:
a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_ref = data
_, _, L = c_ref.shape
device = a_ref.device
def to_blocked(mat: torch.Tensor) -> torch.Tensor:
rows, cols = mat.shape
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
blocks = mat.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return rearranged.flatten()
for l_idx in range(L):
scale_a = to_blocked(sfa_ref[:, :, l_idx]).to(device=device)
scale_b = to_blocked(sfb_ref[:, :, l_idx]).to(device=device)
res = torch._scaled_mm(
a_ref[:, :, l_idx],
b_ref[:, :, l_idx].transpose(0, 1),
scale_a,
scale_b,
bias=None,
out_dtype=torch.float16,
)
c_ref[:, :, l_idx] = res
return c_ref
def custom_kernel(data: input_t) -> output_t:
a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_ref = data
M, K_packed, L = a_ref.shape
N, _, _ = b_ref.shape
K = K_packed * 2
sfK = sfa_ref.shape[1]
use_triton = (
triton is not None
and torch.cuda.is_available()
and torch.cuda.get_device_capability()[0] >= 10
)
if not use_triton:
return _fallback_scaled_mm(data)
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 256
VEC_SIZE = sf_vec_size
for l_idx in range(L):
# packed NVFP4 must be uint8 for tl.dot_scaled("e2m1")
a_l = a_ref[:, :, l_idx].view(torch.uint8).contiguous()
b_l = b_ref[:, :, l_idx].view(torch.uint8).contiguous()
sfa_l = sfa_ref[:, :, l_idx].contiguous()
sfb_l = sfb_ref[:, :, l_idx].contiguous()
c_l = c_ref[:, :, l_idx]
grid = (triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N),)
_bsmm_kernel[grid](
a_l, b_l, sfa_l, sfb_l, c_l,
M, N, K, sfK,
a_l.stride(0), a_l.stride(1),
b_l.stride(0), b_l.stride(1),
sfa_l.stride(0), sfa_l.stride(1),
sfb_l.stride(0), sfb_l.stride(1),
c_l.stride(0), c_l.stride(1),
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
VEC_SIZE=VEC_SIZE,
num_warps=8,
num_stages=2,
)
return c_ref
check_implementation = make_match_reference(custom_kernel, rtol=1e-3, atol=1e-3)
scrolls · 164 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