submission 311717
pagelessuptime · 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.
nvfp4_dual_gemm_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-311717?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:49367b25e3f3e95bbff0f87d593f10c2b217a2b999e5ebc15707e00e74ef6d90
license declaredunknown
license concludedunknown
authorspagelessuptime
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=8, num_stages=2),mma
acc1 += tl.dot(a_block, b1_block, out_dtype=tl.float32)num-warps = 8
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=8, num_stages=2),stages = 2
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=8, num_stages=2),Kernel source
nvfp4_dual_gemm_submission.py235 lines
#!POPCORN leaderboard nvfp4_dual_gemm
#!POPCORN gpu B200
import torch
try:
import triton
import triton.language as tl
_has_triton = True
except Exception:
_has_triton = False
def _ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
def _to_blocked(scale_matrix: torch.Tensor) -> torch.Tensor:
"""Convert (rows, cols) scale factors to the blocked layout expected by torch._scaled_mm."""
rows, cols = scale_matrix.shape
n_row_blocks = _ceil_div(rows, 128)
n_col_blocks = _ceil_div(cols, 4)
# Assumes rows are multiple of 128 and cols multiple of 4 (per problem constraints).
blocks = scale_matrix.reshape(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()
def _block_all_scales(sf: torch.Tensor) -> torch.Tensor:
"""
Block all scale tensors in one pass on the GPU.
sf: [rows, cols, L] -> returns [L, blocked_len]
"""
# Move batch to front and ensure contiguity for reshapes.
sf_l = sf.permute(2, 0, 1).contiguous()
# vmap applies _to_blocked across L.
return torch.vmap(_to_blocked)(sf_l)
_TRITON_CONFIGS = [
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=8, num_stages=2),
triton.Config({"BLOCK_M": 128, "BLOCK_N": 64, "BLOCK_K": 64}, num_warps=4, num_stages=2),
triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=4, num_stages=2),
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 128}, num_warps=8, num_stages=2),
]
@triton.autotune(configs=_TRITON_CONFIGS, key=["M", "N", "K"])
@triton.jit
def _dual_gemm_silu_kernel(
A, B1, B2, C,
M, N, K,
stride_aL, stride_am, stride_ak,
stride_b1L, stride_b1n, stride_b1k,
stride_b2L, stride_b2n, stride_b2k,
stride_cL, stride_cm, stride_cn,
**meta,
):
BLOCK_M = meta["BLOCK_M"]
BLOCK_N = meta["BLOCK_N"]
BLOCK_K = meta["BLOCK_K"]
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
pid_l = tl.program_id(2)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
a_ptr = A + pid_l * stride_aL
b1_ptr = B1 + pid_l * stride_b1L
b2_ptr = B2 + pid_l * stride_b2L
c_ptr = C + pid_l * stride_cL
acc1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k_iter = tl.cdiv(K, BLOCK_K)
for ki in range(0, k_iter):
k_idx = ki * BLOCK_K + offs_k
a_mask = (offs_m[:, None] < M) & (k_idx[None, :] < K)
b_mask = (offs_n[None, :] < N) & (k_idx[:, None] < K)
a_block = tl.load(
a_ptr + offs_m[:, None] * stride_am + k_idx[None, :] * stride_ak,
mask=a_mask,
other=0.0,
)
b1_block = tl.load(
b1_ptr + offs_n[None, :] * stride_b1n + k_idx[:, None] * stride_b1k,
mask=b_mask,
other=0.0,
)
b2_block = tl.load(
b2_ptr + offs_n[None, :] * stride_b2n + k_idx[:, None] * stride_b2k,
mask=b_mask,
other=0.0,
)
acc1 += tl.dot(a_block, b1_block, out_dtype=tl.float32)
acc2 += tl.dot(a_block, b2_block, out_dtype=tl.float32)
# fused silu(res1) * res2
silu = acc1 * tl.sigmoid(acc1)
out_block = silu * acc2
c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(
c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn,
out_block.to(tl.float16),
mask=c_mask,
)
def _nvfp4_dual_gemm_impl(data):
"""
Reference-style submission for nvfp4_dual_gemm.
Accepts the first 7 items of the tuple:
(a, b1, b2, sfa, sfb1, sfb2, c)
Extra items (e.g., permuted scales) are ignored to stay compatible with 10-input runners.
Returns: c_out with shape [M, N, L] in fp16.
"""
if len(data) < 7:
raise ValueError(f"Expected at least 7 inputs, got {len(data)}")
# Only the first 7 are used; any extras (permute variants) are ignored.
a, b1, b2, sfa, sfb1, sfb2, _c = data[:7]
device = a.device
# Ensure inputs are contiguous in their given layout to avoid hidden copies later.
a = a.contiguous()
b1 = b1.contiguous()
b2 = b2.contiguous()
sfa = sfa.contiguous()
sfb1 = sfb1.contiguous()
sfb2 = sfb2.contiguous()
m, k, l = a.shape
n, _, _ = b1.shape
# If triton is available and inputs are already fp16, run fused dual GEMM in one kernel.
# Otherwise, fall back to the scaled_mm reference path.
if _has_triton and a.dtype == torch.float16 and b1.dtype == torch.float16 and b2.dtype == torch.float16:
# Move batch to front: [L, M, K] etc.
a_b = a.permute(2, 0, 1).contiguous()
b1_b = b1.permute(2, 0, 1).contiguous()
b2_b = b2.permute(2, 0, 1).contiguous()
# Prepare output buffer
if (
isinstance(_c, torch.Tensor)
and _c.shape == (m, n, l)
and _c.dtype == torch.float16
and _c.device == device
):
out = _c
else:
out = torch.empty((m, n, l), device=device, dtype=torch.float16)
out_b = out.permute(2, 0, 1).contiguous()
# Per-shape launch: BLOCK sizes come from autotune.
grid = (triton.cdiv(m, 128), triton.cdiv(n, 128), l)
_dual_gemm_silu_kernel[grid](
a_b, b1_b, b2_b, out_b,
m, n, k,
a_b.stride(0), a_b.stride(1), a_b.stride(2),
b1_b.stride(0), b1_b.stride(1), b1_b.stride(2),
b2_b.stride(0), b2_b.stride(1), b2_b.stride(2),
out_b.stride(0), out_b.stride(1), out_b.stride(2),
)
return out
# Fallback: use torch._scaled_mm for correctness (supports FP4 + block scales).
scale_a_blocked = _block_all_scales(sfa).to(device)
scale_b1_blocked = _block_all_scales(sfb1).to(device)
scale_b2_blocked = _block_all_scales(sfb2).to(device)
a_b = a.permute(2, 0, 1).contiguous()
b1_b = b1.permute(2, 0, 1).contiguous()
b2_b = b2.permute(2, 0, 1).contiguous()
def _per_batch(a_l, b1_l, b2_l, sa_l, sb1_l, sb2_l):
res1 = torch._scaled_mm(
a_l,
b1_l.transpose(0, 1),
sa_l,
sb1_l,
bias=None,
out_dtype=torch.float32,
)
res2 = torch._scaled_mm(
a_l,
b2_l.transpose(0, 1),
sa_l,
sb2_l,
bias=None,
out_dtype=torch.float32,
)
return torch.nn.functional.silu(res1) * res2
if l == 1:
out_lmn = _per_batch(
a_b[0], b1_b[0], b2_b[0], scale_a_blocked[0], scale_b1_blocked[0], scale_b2_blocked[0]
).unsqueeze(0)
else:
out_lmn = torch.vmap(_per_batch)(
a_b,
b1_b,
b2_b,
scale_a_blocked,
scale_b1_blocked,
scale_b2_blocked,
)
if (
isinstance(_c, torch.Tensor)
and _c.shape == (m, n, l)
and _c.dtype == torch.float16
and _c.device == device
):
out = _c
else:
out = torch.empty((m, n, l), device=device, dtype=torch.float16)
out.copy_(out_lmn.permute(1, 2, 0).to(dtype=torch.float16))
return out
def custom_kernel(data):
"""Popcorn expects a custom_kernel entrypoint."""
return _nvfp4_dual_gemm_impl(data)
scrolls · 235 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