Skip to content
KernelIndex
Search⌘K

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
NVFP4 group GEMMsuite of 4 cases
NVIDIA B200
890.9µs
#135 of 145
2026-02-14

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