submission 491474
trvon · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 297 lines, June 9 Researcher Reciprocity License v1.0.
submission_candidate_latest.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-491474?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:1cf34ebdd2ae76bdeb4c58c63e1c69342160220c9c4e2e869decb140579a63cb
license declaredunknown
license concludedunknown
authorstrvon
imported2026-08-26
Kernel source
submission_candidate_latest.py297 lines
import torch
from task import input_t, output_t
from collections import OrderedDict
sf_vec_size = 16
_USE_NON_BLOCKING = False
_USE_BLOCK_CACHE = True
_USE_PAIR_CACHE = False
_BLOCK_CACHE_POLICY = "clear"
_PAIR_CACHE_POLICY = "lru"
_BLOCK_CACHE_MAX = 512
_PAIR_CACHE_MAX = 256
_CACHE_KEY_MODE = "ptr"
_REORDER_IMPL = "reshape_permute"
_FLATTEN_IMPL = "reshape"
_PAIR_PIPELINE = "on_demand"
_WARM_BLOCK_CACHE = False
_PIPELINE_STRUCTURE = "inline"
_GROUP_EXEC_ORDER = "input"
_PAIR_COMPUTE_ORDER = "ab"
_PAIR_CHUNK_SIZE = 2
_BLOCK_CACHE_ADMIT = "always"
_PAIR_CACHE_ADMIT = "large_only"
_PREFETCH_DEPTH = 4
_PREFETCH_MIN_L = 1
_CHUNK_STRATEGY = "fixed"
_CACHE_FLUSH_POLICY = "per_group"
_WARMUP_SCALE_PASS = True
if _BLOCK_CACHE_POLICY == "lru":
_BLOCKED_CACHE = OrderedDict()
else:
_BLOCKED_CACHE = {}
if _PAIR_CACHE_POLICY == "lru":
_SCALE_PAIR_CACHE = OrderedDict()
else:
_SCALE_PAIR_CACHE = {}
def ceil_div(a, b):
return (a + b - 1) // b
def _tensor_key(tensor: torch.Tensor):
if _CACHE_KEY_MODE == "ptr":
return (tensor.data_ptr(),)
return (
tensor.data_ptr(),
tuple(tensor.shape),
tuple(tensor.stride()),
str(tensor.dtype),
tensor.device.index if tensor.device.index is not None else -1,
)
def _pair_key(a: torch.Tensor, b: torch.Tensor):
return (_tensor_key(a), _tensor_key(b))
def _cache_get(cache, key, policy: str):
value = cache.get(key)
if value is not None and policy == "lru":
cache.move_to_end(key)
return value
def _cache_put(cache, key, value, max_items: int, policy: str):
cache[key] = value
if policy == "lru":
cache.move_to_end(key)
if len(cache) > max_items:
cache.popitem(last=False)
else:
if len(cache) > max_items:
cache.clear()
def _allow_block_cache(tensor: torch.Tensor) -> bool:
if not _USE_BLOCK_CACHE:
return False
if _BLOCK_CACHE_ADMIT == "never":
return False
if _BLOCK_CACHE_ADMIT == "large_only":
return tensor.numel() >= 8192
return True
def _allow_pair_cache(a: torch.Tensor, b: torch.Tensor) -> bool:
if not _USE_PAIR_CACHE:
return False
if _PAIR_CACHE_ADMIT == "never":
return False
if _PAIR_CACHE_ADMIT == "large_only":
return (a.numel() + b.numel()) >= 16384
return True
def to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
allow_cache = _allow_block_cache(input_matrix)
if allow_cache:
key = _tensor_key(input_matrix)
cached = _cache_get(_BLOCKED_CACHE, key, _BLOCK_CACHE_POLICY)
if cached is not None:
return cached
rows, cols = input_matrix.shape
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
padded_rows = n_row_blocks * 128
padded_cols = n_col_blocks * 4
if padded_rows != rows or padded_cols != cols:
padded = torch.nn.functional.pad(
input_matrix,
(0, padded_cols - cols, 0, padded_rows - rows),
mode="constant",
value=0,
)
else:
padded = input_matrix
if _REORDER_IMPL == "reshape_permute":
blocks = padded.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)
elif _REORDER_IMPL == "contiguous_permute":
blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3).contiguous()
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
else:
blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
chunked = []
for block in blocks.reshape(-1, 128, 4):
chunked.append(block.view(4, 32, 4).transpose(0, 1))
rearranged = torch.stack(chunked, dim=0).reshape(-1, 32, 16)
if _FLATTEN_IMPL == "reshape":
out = rearranged.reshape(-1)
else:
out = rearranged.flatten()
if allow_cache:
_cache_put(_BLOCKED_CACHE, key, out, _BLOCK_CACHE_MAX, _BLOCK_CACHE_POLICY)
return out
def _get_pair_scales(sfa: torch.Tensor, sfb: torch.Tensor):
allow_pair_cache = _allow_pair_cache(sfa, sfb)
if not allow_pair_cache:
if _PAIR_COMPUTE_ORDER == "ba":
scale_b = to_blocked(sfb)
scale_a = to_blocked(sfa)
return scale_a, scale_b
return to_blocked(sfa), to_blocked(sfb)
key = _pair_key(sfa, sfb)
cached = _cache_get(_SCALE_PAIR_CACHE, key, _PAIR_CACHE_POLICY)
if cached is not None:
return cached
if _PAIR_COMPUTE_ORDER == "ba":
scale_b = to_blocked(sfb)
scale_a = to_blocked(sfa)
value = (scale_a, scale_b)
else:
value = (to_blocked(sfa), to_blocked(sfb))
_cache_put(_SCALE_PAIR_CACHE, key, value, _PAIR_CACHE_MAX, _PAIR_CACHE_POLICY)
return value
def _prefetch_window_indices(length: int, start: int):
depth = max(int(_PREFETCH_DEPTH), 1)
end = min(start + depth, length)
return range(start, end)
def _should_prefetch(length: int) -> bool:
return int(length) >= max(int(_PREFETCH_MIN_L), 1)
def _effective_chunk_size(length: int) -> int:
base = max(int(_PAIR_CHUNK_SIZE), 1)
if _CHUNK_STRATEGY == "l_adaptive":
if length >= 16:
return max(base * 2, 2)
if length <= 2:
return 1
return base
def _flush_caches():
_BLOCKED_CACHE.clear()
_SCALE_PAIR_CACHE.clear()
def _group_indices(problem_sizes):
n = len(problem_sizes)
if _GROUP_EXEC_ORDER == "reverse":
return list(range(n - 1, -1, -1))
if _GROUP_EXEC_ORDER == "mn_desc":
pairs = [(i, problem_sizes[i][0] * problem_sizes[i][1]) for i in range(n)]
pairs.sort(key=lambda x: x[1], reverse=True)
return [i for i, _ in pairs]
return list(range(n))
def _run_mm(a_ref, b_ref, c_ref, l_idx: int, scale_a: torch.Tensor, scale_b: torch.Tensor):
c_ref[:, :, l_idx] = torch._scaled_mm(
a_ref[:, :, l_idx].view(torch.float4_e2m1fn_x2),
b_ref[:, :, l_idx].transpose(0, 1).view(torch.float4_e2m1fn_x2),
scale_a,
scale_b,
bias=None,
out_dtype=torch.float16,
)
@torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
abc_tensors, sfasfb_tensors, _, problem_sizes = data
results = [None] * len(abc_tensors)
if _CACHE_FLUSH_POLICY == "per_call":
_flush_caches()
for group_idx in _group_indices(problem_sizes):
if _CACHE_FLUSH_POLICY == "per_group":
_flush_caches()
a_ref, b_ref, c_ref = abc_tensors[group_idx]
sfa_ref, sfb_ref = sfasfb_tensors[group_idx]
_, _, _, l = problem_sizes[group_idx]
if sfa_ref.is_cuda:
sfa_cuda = sfa_ref
else:
sfa_cuda = sfa_ref.cuda(non_blocking=_USE_NON_BLOCKING)
if sfb_ref.is_cuda:
sfb_cuda = sfb_ref
else:
sfb_cuda = sfb_ref.cuda(non_blocking=_USE_NON_BLOCKING)
if _WARMUP_SCALE_PASS:
for l_idx in range(l):
_ = _get_pair_scales(sfa_cuda[:, :, l_idx], sfb_cuda[:, :, l_idx])
if _PIPELINE_STRUCTURE == "group_prefetch_then_mm" and _should_prefetch(l):
if _PREFETCH_DEPTH >= l:
scales = [_get_pair_scales(sfa_cuda[:, :, l_idx], sfb_cuda[:, :, l_idx]) for l_idx in range(l)]
for l_idx, (scale_a, scale_b) in enumerate(scales):
_run_mm(a_ref, b_ref, c_ref, l_idx, scale_a, scale_b)
else:
for start in range(0, l, max(int(_PREFETCH_DEPTH), 1)):
indices = list(_prefetch_window_indices(l, start))
scales = [_get_pair_scales(sfa_cuda[:, :, l_idx], sfb_cuda[:, :, l_idx]) for l_idx in indices]
for offset, l_idx in enumerate(indices):
scale_a, scale_b = scales[offset]
_run_mm(a_ref, b_ref, c_ref, l_idx, scale_a, scale_b)
elif _PIPELINE_STRUCTURE == "staged_chunks":
chunk = _effective_chunk_size(l)
for start in range(0, l, chunk):
end = min(start + chunk, l)
scales = []
for l_idx in range(start, end):
scales.append(_get_pair_scales(sfa_cuda[:, :, l_idx], sfb_cuda[:, :, l_idx]))
for offset, l_idx in enumerate(range(start, end)):
scale_a, scale_b = scales[offset]
_run_mm(a_ref, b_ref, c_ref, l_idx, scale_a, scale_b)
elif _PAIR_PIPELINE == "prefetch_all" and _should_prefetch(l):
if _PREFETCH_DEPTH >= l:
scales = []
for l_idx in range(l):
scales.append(_get_pair_scales(sfa_cuda[:, :, l_idx], sfb_cuda[:, :, l_idx]))
for l_idx, (scale_a, scale_b) in enumerate(scales):
_run_mm(a_ref, b_ref, c_ref, l_idx, scale_a, scale_b)
else:
for start in range(0, l, max(int(_PREFETCH_DEPTH), 1)):
indices = list(_prefetch_window_indices(l, start))
scales = []
for l_idx in indices:
scales.append(_get_pair_scales(sfa_cuda[:, :, l_idx], sfb_cuda[:, :, l_idx]))
for offset, l_idx in enumerate(indices):
scale_a, scale_b = scales[offset]
_run_mm(a_ref, b_ref, c_ref, l_idx, scale_a, scale_b)
else:
if _WARM_BLOCK_CACHE:
for l_idx in range(l):
_ = _get_pair_scales(sfa_cuda[:, :, l_idx], sfb_cuda[:, :, l_idx])
for l_idx in range(l):
scale_a, scale_b = _get_pair_scales(sfa_cuda[:, :, l_idx], sfb_cuda[:, :, l_idx])
_run_mm(a_ref, b_ref, c_ref, l_idx, scale_a, scale_b)
results[group_idx] = c_ref
return results
scrolls · 297 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