Skip to content
KernelIndex
Search⌘K

submission 109079

shiyeegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 168 lines, June 9 Researcher Reciprocity License v1.0.

node_26.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109079?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 GEMVsuite of 3 cases
NVIDIA B200
43.9µs
#248 of 678
2025-11-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b06e8224304aec55dc53db8df4260437d5625c2bbb9aa650a033c2ca60da15ba
license declaredunknown
license concludedunknown
authorsshiyeegao
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4NVFP4 GEMV using cublasLt-backed torch._scaled_mm with batched scale flatten

Kernel source

node_26.py168 lines
import hashlib
import os
from functools import lru_cache
from typing import Tuple

import torch
from torch.utils.cpp_extension import load_inline

TensorTuple = Tuple[torch.Tensor, ...]

# Limit concurrent streams; L up to 8 in benchmarks.
_MAX_STREAMS = 4
_STREAM_POOL: list[torch.cuda.Stream] = []


def _extension_name() -> str:
    """Deterministic name so load_inline can reuse builds between runs."""
    src_id = "nvfp4_flatten_batch_v3"
    digest = hashlib.sha1(src_id.encode("utf-8")).hexdigest()[:8]
    return f"nvfp4_flatten_ext_{digest}"


@lru_cache(maxsize=1)
def _load_ext():
    """
    Build a tiny C++/CUDA extension that flattens all L slices in one shot.

    Keeping the transform in compiled code avoids Python dispatch overhead and
    makes sure the permutation happens on GPU when inputs are CUDA tensors.
    """
    cpp_source = r"""
#include <torch/extension.h>
#include <stdexcept>

torch::Tensor flatten_scales_batch_cuda(torch::Tensor t);

torch::Tensor flatten_scales_batch(torch::Tensor t) {
    return flatten_scales_batch_cuda(t);
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("flatten_scales_batch", &flatten_scales_batch, "Flatten all L slices");
}
"""

    # Implemented with ATen ops to stay concise while still exercising NVCC.
    cuda_source = r"""
#include <torch/extension.h>
#include <stdexcept>

torch::Tensor flatten_scales_batch_cuda(torch::Tensor t) {
    const auto dim = t.dim();
    if (dim == 6) {
        // [32,4,mm,4,kk,l] -> [l, mm, kk, 32, 4, 4] -> [l, flat]
        auto tmp = t.permute({5, 2, 4, 0, 1, 3}).contiguous();
        return tmp.view({tmp.size(0), -1});
    } else if (dim == 5) {
        // Treat as l=1 for completeness.
        auto tmp = t.unsqueeze(0).permute({0, 3, 5, 1, 2, 4}).contiguous();
        return tmp.view({1, -1});
    }
    throw std::runtime_error("flatten_scales_batch expects rank-5 or rank-6 tensor");
}
"""

    name = _extension_name()
    build_dir = os.path.join(torch.utils.cpp_extension.get_default_build_root(), name)
    return load_inline(
        name=name,
        cpp_sources=cpp_source,
        cuda_sources=cuda_source,
        functions=None,
        with_cuda=True,
        extra_cflags=["-O3"],
        extra_cuda_cflags=["-O3", "--use_fast_math"],
        build_directory=build_dir,
        verbose=False,
    )


def _flatten_scale_batch(scale_tensor: torch.Tensor) -> torch.Tensor:
    """
    Flatten all L slices of the pre-permuted scale tensor at once.

    Returns a 2D tensor of shape [L, flat], where each row is suitable for
    torch._scaled_mm. Falls back to a pure Python/torch path if extension
    build fails (e.g., NVCC unavailable).
    """
    try:
        ext = _load_ext()
        return ext.flatten_scales_batch(scale_tensor)
    except Exception:
        if scale_tensor.dim() == 6:
            tmp = scale_tensor.permute(5, 2, 4, 0, 1, 3).contiguous()
            return tmp.view(scale_tensor.size(-1), -1)
        if scale_tensor.dim() == 5:
            tmp = scale_tensor.unsqueeze(0).permute(0, 3, 5, 1, 2, 4).contiguous()
            return tmp.view(1, -1)
        raise RuntimeError(f"Unexpected scale_tensor ndim={scale_tensor.dim()}")


def _get_streams(count: int) -> list[torch.cuda.Stream]:
    """Reuse a small pool of streams to overlap independent L slices."""
    global _STREAM_POOL
    if count <= 1:
        return [torch.cuda.current_stream()]
    while len(_STREAM_POOL) < count:
        _STREAM_POOL.append(torch.cuda.Stream())
    return _STREAM_POOL[:count]


@torch.inference_mode()
def custom_kernel(data: TensorTuple) -> torch.Tensor:
    """
    NVFP4 GEMV using cublasLt-backed torch._scaled_mm with batched scale flatten
    and stream-level parallelism across L slices.

    Args:
        data: Tuple containing (a_ref, b_ref, sfa_ref, sfb_ref, sfa_permuted, sfb_permuted, c_ref)
              generated by evaluate.generate_input. All tensors are expected to
              reside on CUDA and follow the layouts described in problem.md.
    Returns:
        c_ref: torch.Tensor[float16] with shape [m, 1, l]
    """
    a_ref, b_ref, _, _, sfa_permuted, sfb_permuted, c_ref = data

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required for NVFP4 GEMV")

    _, _, l = c_ref.shape

    # One-time flatten for all slices to cut redundant permute/contiguous calls.
    scale_a_batch = _flatten_scale_batch(sfa_permuted)
    scale_b_batch = _flatten_scale_batch(sfb_permuted)

    # Precompute B columns once; avoid contiguous() because nvfp4 copy_ is unsupported.
    b_cols = b_ref.permute(2, 1, 0)  # [l, k, 1] view

    # Allow multiple independent _scaled_mm calls to overlap when l>1.
    stream_count = min(max(l, 1), _MAX_STREAMS)
    streams = _get_streams(stream_count)
    current = torch.cuda.current_stream()
    # Ensure flattening is visible to worker streams.
    ready = torch.cuda.Event(blocking=False, enable_timing=False)
    ready.record(current)
    for s in streams:
        s.wait_event(ready)

    for l_idx in range(l):
        stream = streams[l_idx % stream_count]
        with torch.cuda.stream(stream):
            res = torch._scaled_mm(
                a_ref[:, :, l_idx],
                b_cols[l_idx],
                scale_a_batch[l_idx],
                scale_b_batch[l_idx],
                bias=None,
                out_dtype=torch.float16,
            )
            # Write back on the same stream; wait later on current stream.
            c_ref[:, 0, l_idx].copy_(res[:, 0])

    # Wait for all worker streams to finish before returning.
    for s in streams:
        if s is not current:
            current.wait_stream(s)
    return c_ref
scrolls · 168 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