submission 246321
zyvren · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 204 lines, June 9 Researcher Reciprocity License v1.0.
optimized_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-246321?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:0b01413fc6e116d089f8b136179aa8b4a8e42e2f3d79ba03be1a83d9de492cd9
license declaredunknown
license concludedunknown
authorszyvren
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
_SF_VEC_SIZE = 16 # NVFP4 block scale is per 16 K elementsKernel source
optimized_submission.py204 lines
# optimized_submission.py
import torch
from task import input_t, output_t
_SF_VEC_SIZE = 16 # NVFP4 block scale is per 16 K elements
@torch.no_grad()
def _blocked_flat_from_permuted_all_l(sf_perm: torch.Tensor) -> torch.Tensor:
"""
sf_perm: [32, 4, rest_mn, 4, rest_k, L] (fp8, cuda)
returns: [L, rest_mn*rest_k*32*16] (fp8, cuda)
"""
# -> [rest_mn, rest_k, 32, 4, 4, L]
t = sf_perm.permute(2, 4, 0, 1, 3, 5).contiguous()
L = t.shape[-1]
# -> [rest_mn*rest_k, 32, 16, L]
t = t.view(-1, 32, 16, L)
# -> [L, rest_mn*rest_k*32*16]
return t.permute(3, 0, 1, 2).contiguous().view(L, -1)
@torch.no_grad()
def _blocked_flat_from_logical_all_l(sf: torch.Tensor) -> torch.Tensor:
"""
Fallback if permuted scales are not provided.
sf: [mn, sf_k, L] (fp8, cuda)
returns [L, rest_mn*rest_k*32*16] (fp8, cuda) matching torch._scaled_mm.
"""
assert sf.dim() == 3
mn, sf_k, L = sf.shape
# pad rows to multiple of 128, cols to multiple of 4
rest_mn = (mn + 127) // 128
rest_k = (sf_k + 3) // 4
pad_m = rest_mn * 128
pad_k = rest_k * 4
if pad_m != mn or pad_k != sf_k:
padded = torch.zeros((pad_m, pad_k, L), device=sf.device, dtype=sf.dtype)
padded[:mn, :sf_k, :] = sf
else:
padded = sf
# reference to_blocked does (mn,sf_k)-> view(rest_mn,128, rest_k,4).permute(0,2,1,3)
# then reshape(-1,4,32,4).transpose(1,2).reshape(-1,32,16).flatten()
# We do it for all L at once.
x = padded.permute(2, 0, 1).contiguous() # [L, pad_m, pad_k]
x = x.view(L, rest_mn, 128, rest_k, 4).permute(0, 1, 3, 2, 4).contiguous() # [L, rest_mn, rest_k, 128, 4]
x = x.view(L, -1, 4, 32, 4).transpose(2, 3).contiguous() # [L, -1, 32, 4, 4]
x = x.view(L, -1, 32, 16) # [L, -1, 32, 16]
return x.reshape(L, -1) # [L, flat]
def _parse_inputs(data):
"""
Support both 7-tensor and 10-tensor tuples.
7: (a, b1, b2, sfa, sfb1, sfb2, c)
10: may include logical and permuted scales; order can vary in some harnesses.
We detect by tensor dimensionality.
Returns:
a, b1, b2, sfa_perm_or_none, sfb1_perm_or_none, sfb2_perm_or_none,
sfa_logical, sfb1_logical, sfb2_logical, c
"""
if not isinstance(data, (tuple, list)):
raise TypeError("custom_kernel expects tuple/list inputs")
if len(data) == 7:
a, b1, b2, sfa, sfb1, sfb2, c = data
return a, b1, b2, None, None, None, sfa, sfb1, sfb2, c
if len(data) == 10:
# Common reference order:
# (a, b1, b2, sfa, sfb1, sfb2, sfa_perm, sfb1_perm, sfb2_perm, c)
# But to be robust, detect:
a = data[0]
b1 = data[1]
b2 = data[2]
# Remaining 7 tensors
rem = list(data[3:])
# Output c is 3D fp16
c = None
for i, t in enumerate(rem):
if isinstance(t, torch.Tensor) and t.dim() == 3 and t.dtype == torch.float16:
c = t
rem.pop(i)
break
if c is None:
# fallback: last tensor
c = data[-1]
rem = list(data[3:-1])
# Permuted scales are 6D with leading dims (32,4,...,4,...,L)
perms = []
logicals = []
for t in rem:
if isinstance(t, torch.Tensor) and t.dim() == 6 and t.shape[0] == 32 and t.shape[1] == 4 and t.shape[3] == 4:
perms.append(t)
else:
logicals.append(t)
# Assign perms by matching their "rest_mn" dim to M or N when possible
# We'll fill later once we know M,N from a/b shapes.
return a, b1, b2, perms, logicals, c
raise ValueError(f"Unexpected input tuple length: {len(data)}")
@torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
parsed = _parse_inputs(data)
if len(parsed) == 10:
a, b1, b2, _, _, _, sfa, sfb1, sfb2, c = parsed
perms = None
else:
# robust path for len==10 with detection
a, b1, b2, perms, logicals, c = parsed
# logicals expected to be [sfa, sfb1, sfb2] (3D fp8); if missing we still proceed
sfa = logicals[0] if len(logicals) > 0 else None
sfb1 = logicals[1] if len(logicals) > 1 else None
sfb2 = logicals[2] if len(logicals) > 2 else None
# Shapes
M, K, L = a.shape
N = b1.shape[0]
assert b1.shape == (N, K, L)
assert b2.shape == (N, K, L)
# Prepare output
if not (isinstance(c, torch.Tensor) and c.is_cuda and c.dtype == torch.float16 and c.shape == (M, N, L)):
c_out = torch.empty((M, N, L), device=a.device, dtype=torch.float16)
else:
c_out = c
# Build blocked-flat scales
# Prefer permuted if present and identifiable; else fallback to logical.
blocked_a = blocked_b1 = blocked_b2 = None
if perms is not None and len(perms) >= 3:
# Identify which perm corresponds to A (rest_m matches ceil(M/128)) and which to B (rest_n matches ceil(N/128))
rest_m = (M + 127) // 128
rest_n = (N + 127) // 128
# split perms by their 3rd dim
a_perm = None
b_perms = []
for t in perms:
if t.shape[2] == rest_m and a_perm is None:
a_perm = t
elif t.shape[2] == rest_n:
b_perms.append(t)
if a_perm is None:
# fallback: just take first as A
a_perm = perms[0]
b_perms = perms[1:]
# b_perms should have 2 entries
if len(b_perms) < 2:
# fallback: remaining in order
b_perms = [p for p in perms if p is not a_perm]
if len(b_perms) < 2:
b_perms = (b_perms + b_perms)[:2]
b1_perm, b2_perm = b_perms[0], b_perms[1]
blocked_a = _blocked_flat_from_permuted_all_l(a_perm)
blocked_b1 = _blocked_flat_from_permuted_all_l(b1_perm)
blocked_b2 = _blocked_flat_from_permuted_all_l(b2_perm)
else:
# Fallback to logical scales if permuted aren't available.
if sfa is None or sfb1 is None or sfb2 is None:
raise RuntimeError("Permuted scales not found and logical scales missing; cannot run.")
blocked_a = _blocked_flat_from_logical_all_l(sfa)
blocked_b1 = _blocked_flat_from_logical_all_l(sfb1)
blocked_b2 = _blocked_flat_from_logical_all_l(sfb2)
# Compute per L slice, write fp16
for l_idx in range(L):
scale_a = blocked_a[l_idx]
scale_b1 = blocked_b1[l_idx]
scale_b2 = blocked_b2[l_idx]
x1 = torch._scaled_mm(
a[:, :, l_idx],
b1[:, :, l_idx].transpose(0, 1),
scale_a, scale_b1,
bias=None,
out_dtype=torch.float32,
)
x2 = torch._scaled_mm(
a[:, :, l_idx],
b2[:, :, l_idx].transpose(0, 1),
scale_a, scale_b2,
bias=None,
out_dtype=torch.float32,
)
c_out[:, :, l_idx] = (torch.nn.functional.silu(x1) * x2).to(torch.float16)
return c_out
scrolls · 204 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