Skip to content
KernelIndex
Search⌘K

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
NVFP4 dual GEMMsuite of 4 cases
NVIDIA B200
43.5µs
#292 of 420
2026-01-01

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 elements

Kernel 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