submission 272244
boo · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 290 lines, June 9 Researcher Reciprocity License v1.0.
reference7_fast_v20.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-272244?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:4b298a4cfbe0440d38e5922e3506e41513d19b19908f5b48db5d24f72485bf92
license declaredunknown
license concludedunknown
authorsboo
imported2026-08-26
Kernel source
reference7_fast_v20.py290 lines
import os
import weakref
import torch
from task import input_t, output_t
from utils import make_match_reference
# -------------------------
# Fast, safe building blocks
# -------------------------
def _cuda_dev_index(t: torch.Tensor) -> int:
if not t.is_cuda:
return -1
return int(t.device.index) if t.device.index is not None else 0
def _tensor_ver(t: torch.Tensor) -> int:
# PyTorch tensors carry a version counter updated on in-place ops.
try:
return int(t._version) # type: ignore[attr-defined]
except Exception:
return 0
# -------------------------
# fp8 scale packing for _scaled_mm
# -------------------------
def to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
"""
Reference mapping used by the baseline:
input: [rows, cols] fp8 (e4m3fnuz)
output: flat tensor in a blocked order expected by torch._scaled_mm
Assumes rows % 128 == 0 and cols % 4 == 0 for the hot path.
"""
rows, cols = input_matrix.shape
# Fallback (should not trigger for contest sizes, but keep it safe).
if (rows % 128) != 0 or (cols % 4) != 0 or (not input_matrix.is_contiguous()):
n_row_blocks = (rows + 127) // 128
n_col_blocks = (cols + 3) // 4
# NOTE: For contest shapes, rows/cols are divisible and view() is valid.
blocks = input_matrix.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()
n_row_blocks = rows // 128
n_col_blocks = cols // 4
view5 = input_matrix.as_strided(
size=(n_row_blocks, n_col_blocks, 32, 4, 4),
stride=(128 * cols, 4, cols, 32 * cols, 1),
)
return view5.contiguous().view(-1)
def to_blocked_out(input_matrix: torch.Tensor, out_flat: torch.Tensor) -> torch.Tensor:
"""Same mapping as to_blocked(), but writes into a preallocated flat buffer."""
rows, cols = input_matrix.shape
if (rows % 128) != 0 or (cols % 4) != 0 or (not input_matrix.is_contiguous()):
tmp = to_blocked(input_matrix)
out_flat.copy_(tmp)
return out_flat
n_row_blocks = rows // 128
n_col_blocks = cols // 4
view5 = input_matrix.as_strided(
size=(n_row_blocks, n_col_blocks, 32, 4, 4),
stride=(128 * cols, 4, cols, 32 * cols, 1),
)
out5 = out_flat.view(n_row_blocks, n_col_blocks, 32, 4, 4)
out5.copy_(view5)
return out_flat
# -------------------------
# Buffer reuse (no data_ptr based signatures)
# -------------------------
_SCALE_BUF: dict[tuple, torch.Tensor] = {}
_OUT_BUF: dict[tuple, torch.Tensor] = {}
_MAT_BUF: dict[tuple, torch.Tensor] = {}
def _get_buf(cache: dict, key: tuple, make: callable) -> torch.Tensor:
buf = cache.get(key)
if buf is None:
buf = make()
cache[key] = buf
# avoid unbounded growth
if len(cache) > 64:
cache.clear()
return buf
def _get_scale_buf(slot: int, numel: int, dtype: torch.dtype, like: torch.Tensor) -> torch.Tensor:
dev = _cuda_dev_index(like)
key = (slot, numel, dtype, dev)
return _get_buf(_SCALE_BUF, key, lambda: torch.empty((numel,), device=like.device, dtype=dtype))
def _get_out_buf(slot: int, shape: tuple[int, int], dtype: torch.dtype, like: torch.Tensor) -> torch.Tensor:
dev = _cuda_dev_index(like)
key = (slot, shape[0], shape[1], dtype, dev)
return _get_buf(_OUT_BUF, key, lambda: torch.empty(shape, device=like.device, dtype=dtype))
def _get_mat_buf(slot: int, shape: tuple[int, int], dtype: torch.dtype, like: torch.Tensor) -> torch.Tensor:
dev = _cuda_dev_index(like)
key = (slot, shape[0], shape[1], dtype, dev)
return _get_buf(_MAT_BUF, key, lambda: torch.empty(shape, device=like.device, dtype=dtype))
# -------------------------
# _scaled_mm wrapper
# -------------------------
def _scaled_mm_out(mat1, mat2, scale_a_flat, scale_b_flat, out, out_dtype):
"""
Use aten._scaled_mm.out if available (saves an allocation).
IMPORTANT: in this contest, mat2 must be transposed: shape (K, N).
"""
op_out = getattr(torch.ops.aten, "_scaled_mm", None)
if op_out is not None and hasattr(torch.ops.aten._scaled_mm, "out"):
# use_fast_accum is not supported for FP4 inputs in the provided environment
return torch.ops.aten._scaled_mm.out(mat1, mat2, scale_a_flat, scale_b_flat, None, None, out_dtype, False, out=out)
# fallback
res = torch._scaled_mm(mat1, mat2, scale_a_flat, scale_b_flat, bias=None, out_dtype=out_dtype)
out.copy_(res)
return out
# -------------------------
# Modes
# -------------------------
# 0: safe baseline (2x scaled_mm)
# 1: packed 2N (1x scaled_mm on concatenated B), useful to experiment
_MODE = int(os.environ.get("NVFP4_DUAL_GEMM_MODE", "0"))
# -------------------------
# Kernels
# -------------------------
def _kernel_2x(a, b1, b2, sfa, sfb1, sfb2, c):
"""Fast & stable: two scaled_mm calls; mixed output dtypes."""
m, n, L = c.shape
cols_a = sfa.shape[1]
cols_b = sfb1.shape[1]
# Hot path for L==1
if L == 1:
a0 = a[:, :, 0]
b10 = b1[:, :, 0].transpose(0, 1) # (K, N)
b20 = b2[:, :, 0].transpose(0, 1) # (K, N)
scale_a = to_blocked_out(sfa[:, :, 0], _get_scale_buf(0, m * cols_a, sfa.dtype, sfa))
scale_b1 = to_blocked_out(sfb1[:, :, 0], _get_scale_buf(1, n * cols_b, sfb1.dtype, sfb1))
scale_b2 = to_blocked_out(sfb2[:, :, 0], _get_scale_buf(2, n * cols_b, sfb2.dtype, sfb2))
out1 = _get_out_buf(0, (m, n), torch.float32, c)
out2 = _get_out_buf(1, (m, n), torch.float16, c)
_scaled_mm_out(a0, b10, scale_a, scale_b1, out=out1, out_dtype=torch.float32)
_scaled_mm_out(a0, b20, scale_a, scale_b2, out=out2, out_dtype=torch.float16)
torch.nn.functional.silu(out1, inplace=True)
torch.mul(out1, out2, out=c[:, :, 0])
return c
# Generic L
out1 = _get_out_buf(0, (m, n), torch.float32, c)
out2 = _get_out_buf(1, (m, n), torch.float16, c)
for l_idx in range(L):
scale_a = to_blocked_out(sfa[:, :, l_idx], _get_scale_buf(0, m * cols_a, sfa.dtype, sfa))
scale_b1 = to_blocked_out(sfb1[:, :, l_idx], _get_scale_buf(1, n * cols_b, sfb1.dtype, sfb1))
scale_b2 = to_blocked_out(sfb2[:, :, l_idx], _get_scale_buf(2, n * cols_b, sfb2.dtype, sfb2))
_scaled_mm_out(a[:, :, l_idx], b1[:, :, l_idx].transpose(0, 1), scale_a, scale_b1, out=out1, out_dtype=torch.float32)
_scaled_mm_out(a[:, :, l_idx], b2[:, :, l_idx].transpose(0, 1), scale_a, scale_b2, out=out2, out_dtype=torch.float16)
torch.nn.functional.silu(out1, inplace=True)
torch.mul(out1, out2, out=c[:, :, l_idx])
return c
def _kernel_packed_2n(a, b1, b2, sfa, sfb1, sfb2, c):
"""
One _scaled_mm by concatenating B along N (so output is M x 2N), then SiLU(x)*y.
FIX: mat2 must be (K, 2N), so we pack B as (2N, K) and pass transpose view.
"""
m, n, L = c.shape
k = a.shape[1]
cols_a = sfa.shape[1]
cols_b = sfb1.shape[1]
# Allocate reusable buffers
bcat = _get_mat_buf(0, (2 * n, k), b1.dtype, b1)
bcat_t = bcat.transpose(0, 1) # (K, 2N) view
scale_a_buf = _get_scale_buf(0, m * cols_a, sfa.dtype, sfa)
scalecat_buf = _get_scale_buf(3, (2 * n) * cols_b, sfb1.dtype, sfb1)
tmp = _get_out_buf(2, (m, 2 * n), torch.float32, c)
# Optional: avoid allocating sfb_cat when rows align to 128-blocks (bench sizes do).
fast_scale_cat = (n % 128) == 0 and (cols_b % 4) == 0
if L == 1:
# pack B in (2N, K) with contiguous copies, then transpose view
bcat[:n].copy_(b1[:, :, 0])
bcat[n:].copy_(b2[:, :, 0])
# scale_a
scale_a = to_blocked_out(sfa[:, :, 0], scale_a_buf)
if fast_scale_cat:
half = n * cols_b
to_blocked_out(sfb1[:, :, 0], scalecat_buf[:half])
to_blocked_out(sfb2[:, :, 0], scalecat_buf[half:2 * half])
scale_b = scalecat_buf
else:
# safe path: build (2N, cols_b) then to_blocked_out once
sfb_cat = _get_mat_buf(1, (2 * n, cols_b), sfb1.dtype, sfb1)
sfb_cat[:n].copy_(sfb1[:, :, 0])
sfb_cat[n:].copy_(sfb2[:, :, 0])
scale_b = to_blocked_out(sfb_cat, scalecat_buf)
_scaled_mm_out(a[:, :, 0], bcat_t, scale_a, scale_b, out=tmp, out_dtype=torch.float32)
x = tmp[:, :n]
y = tmp[:, n:]
torch.nn.functional.silu(x, inplace=True)
torch.mul(x, y, out=c[:, :, 0])
return c
# Generic L
for l_idx in range(L):
bcat[:n].copy_(b1[:, :, l_idx])
bcat[n:].copy_(b2[:, :, l_idx])
scale_a = to_blocked_out(sfa[:, :, l_idx], scale_a_buf)
if fast_scale_cat:
half = n * cols_b
to_blocked_out(sfb1[:, :, l_idx], scalecat_buf[:half])
to_blocked_out(sfb2[:, :, l_idx], scalecat_buf[half:2 * half])
scale_b = scalecat_buf
else:
sfb_cat = _get_mat_buf(1, (2 * n, cols_b), sfb1.dtype, sfb1)
sfb_cat[:n].copy_(sfb1[:, :, l_idx])
sfb_cat[n:].copy_(sfb2[:, :, l_idx])
scale_b = to_blocked_out(sfb_cat, scalecat_buf)
_scaled_mm_out(a[:, :, l_idx], bcat_t, scale_a, scale_b, out=tmp, out_dtype=torch.float32)
x = tmp[:, :n]
y = tmp[:, n:]
torch.nn.functional.silu(x, inplace=True)
torch.mul(x, y, out=c[:, :, l_idx])
return c
def custom_kernel(data: input_t) -> output_t:
# The runner provides extra tensors; last one is output C.
a, b1, b2, sfa, sfb1, sfb2 = data[0], data[1], data[2], data[3], data[4], data[5]
c = data[-1]
if _MODE == 1:
return _kernel_packed_2n(a, b1, b2, sfa, sfb1, sfb2, c)
return _kernel_2x(a, b1, b2, sfa, sfb1, sfb2, c)
# Reference kernel (for correctness checking)
def ref_kernel(data: input_t) -> output_t:
a, b1, b2, sfa, sfb1, sfb2, _, _, _, c = data
m, n, L = c.shape
out1 = torch.empty((m, n, L), device=a.device, dtype=torch.float32)
out2 = torch.empty((m, n, L), device=a.device, dtype=torch.float32)
for l_idx in range(L):
scale_a = to_blocked(sfa[:, :, l_idx])
scale_b1 = to_blocked(sfb1[:, :, l_idx])
scale_b2 = to_blocked(sfb2[:, :, l_idx])
out1[:, :, l_idx] = torch._scaled_mm(
a[:, :, l_idx],
b1[:, :, l_idx].transpose(0, 1),
scale_a,
scale_b1,
bias=None,
out_dtype=torch.float32,
)
out2[:, :, l_idx] = torch._scaled_mm(
a[:, :, l_idx],
b2[:, :, l_idx].transpose(0, 1),
scale_a,
scale_b2,
bias=None,
out_dtype=torch.float32,
)
return (torch.nn.functional.silu(out1) * out2).to(torch.float16)
check_implementation = make_match_reference(ref_kernel, rtol=1e-3, atol=1e-3)
scrolls · 290 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