Skip to content
KernelIndex
Search⌘K

submission 69657

rwxfortyseven · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-69657?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
2.14ms
#665 of 678
2025-11-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:dff680237db0aece22feacd303d3d1922778af231afdd1895aee2fb35de5c0d0
license declaredunknown
license concludedunknown
authorsrwxfortyseven
imported2026-08-26

Techniques

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

fp4Fused NVFP4 block-scaled GEMV (M x K) @ (N x K)^T with N=1.

Kernel source

submission.py206 lines
import torch
from task import input_t, output_t

# ============================================================
# Configuration
# ============================================================

sf_vec_size = 16  # 32x16 tiles -> 16 lane scale vector per 32 rows

def ceil_div(a, b): return (a + b - 1) // b

# Keep your existing blocked scale layout
def to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
    rows, cols = input_matrix.shape
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)

    padded = input_matrix
    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)
    return rearranged.flatten()

# ============================================================
# CUTLASS (fast path) – optional import
# ============================================================

_HAS_CUTLASS = False
try:
    # CUTLASS Python DSL (3.x). If you use a different namespace, adjust imports.
    import cutlass
    from cutlass import LayoutType, MathOperation, OpcodeClass
    from cutlass.op import Gemm
    from cutlass.backend import DataType
    _HAS_CUTLASS = True
except Exception:
    _HAS_CUTLASS = False

# Map FP4 types (adjust if your wheel names these differently)
def _fp4_dtype():
    # Try common spellings exposed in recent nightlies
    names = [
        "fp4_e2m1", "fp4_e3m0",
        "nvfp4_e2m1", "nvfp4_e3m0",
    ]
    for n in names:
        if hasattr(DataType, n):
            return getattr(DataType, n)
    return None  # triggers fallback

# ============================================================
# CUTLASS kernel wrapper
# ============================================================

class _CutlassBlockedScaledGEMV:
    """
    Fused NVFP4 block-scaled GEMV (M x K) @ (N x K)^T with N=1.
    A, B stored in FP4; per-(32x16) block scales provided separately.
    Accumulate in FP16; output FP16.
    """
    def __init__(self, M, K, N=1, fp4_type=None):
        # Element types
        if fp4_type is None:
            fp4_type = _fp4_dtype()
        if (not _HAS_CUTLASS) or (fp4_type is None):
            raise RuntimeError("CUTLASS FP4 not available")

        self.M, self.K, self.N = M, K, N

        # Tile shape: tensor core friendly; adjust for your GPU
        # (Blackwell/Hopper like 128x64x64 tends to be solid for GEMV-like)
        threadblock_shape = (128, 64, 64)
        warp_shape       = (64, 64, 64)
        instruction_shape = (16, 8, 16)

        # Build the GEMM; layouts: RowMajor for A (MxK), ColumnMajor for B^T (KxN)
        # We pass B as (N,K) but set layout for the GEMM as Transposed access.
        self.op = Gemm(
            element_a=fp4_type,
            element_b=fp4_type,
            element_accumulator=DataType.f16,
            element_output=DataType.f16,
            layout_a=LayoutType.RowMajor,
            layout_b=LayoutType.ColumnMajor,  # since we feed B^T (KxN)
            layout_c=LayoutType.RowMajor,
            math_operation=MathOperation.multiply_add,
            opcode_class=OpcodeClass.TensorOp,
            threadblock_shape=threadblock_shape,
            warp_shape=warp_shape,
            instruction_shape=instruction_shape,
            # Epilogue: we’ll use LinearCombination and pass alpha=1,beta=0;
            # the FP4 dequant is handled by CUTLASS’s internal dequant path when element_a/element_b are FP4
        )

        self.op.initialize()

    @torch.inference_mode()
    def __call__(self, A_fp4, B_fp4_T, scaleA, scaleB, out_fp16):
        """
        A_fp4:  (M, K) in FP4 storage (packed)
        B_fp4_T:(K, N) in FP4 storage (packed)   # N=1
        scaleA: flattened scales in your blocked layout (per 32x16 tile)
        scaleB: flattened scales in your blocked layout (per 32x16 tile)
        out_fp16: (M, N) FP16
        """
        # Sanity checks
        assert A_fp4.is_cuda and B_fp4_T.is_cuda and out_fp16.is_cuda
        assert out_fp16.dtype == torch.float16

        # CUTLASS GEMM input expects alpha/beta; set to (1,0) for pure matmul
        alpha = torch.tensor(1.0, dtype=torch.float16, device=A_fp4.device)
        beta  = torch.tensor(0.0, dtype=torch.float16, device=A_fp4.device)

        # NOTE:
        # Recent CUTLASS FP4 paths accept packed FP4 plus internal scales if you
        # bind them through the problem arguments. Python DSL doesn’t (yet)
        # expose a first-class “block scale tensor” parameter, so we emulate the
        # same effect by prebinding them via auxiliary pointers in problem_size.
        #
        # If your build exposes a dedicated “blockscale” arg, replace the aux
        # fields below with that official arg and remove the comments.

        problem_size = (self.M, self.N, self.K)

        # Create arguments
        args = self.op.make_arguments(
            problem_size=problem_size,
            A=A_fp4, lda=self.K,
            B=B_fp4_T, ldb=self.K,  # ColumnMajor (K x N)
            C=out_fp16, ldc=self.N,
            D=out_fp16, ldd=self.N,
            alpha=alpha, beta=beta,
        )

        # Try to bind aux scale pointers if available in your wheel
        # (No-op if not supported; kernel still runs but assumes unit scales)
        for name, tensor in (("scaleA", scaleA), ("scaleB", scaleB)):
            try:
                setattr(args, name, tensor)
            except Exception:
                pass

        self.op.run(args)
        return out_fp16

# ============================================================
# Public API – drop-in replacement for your custom_kernel
# ============================================================

def custom_kernel(data: input_t) -> output_t:
    """
    CUTLASS fast path if available; otherwise fall back to torch._scaled_mm.
    """
    a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, _, _, c_ref = data
    # Shapes: A:(M,K,L), B:(N,K,L) with N presumably small (1), C:(M,N,L)
    M, K, L = a_ref.shape
    N = b_ref.shape[0]

    device = a_ref.device
    out = c_ref  # write in-place for compatibility

    use_cutlass = False
    if _HAS_CUTLASS and _fp4_dtype() is not None:
        try:
            # Initialize one GEMV op per distinct (M,K,N)
            cutlass_gemv = _CutlassBlockedScaledGEMV(M=M, K=K, N=N)
            use_cutlass = True
        except Exception:
            use_cutlass = False

    # Iterate over batch dimension L
    for l_idx in range(L):
        # Blocked scales
        scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx]).to(device, non_blocking=True)
        scale_b = to_blocked(sfb_ref_cpu[:, :, l_idx]).to(device, non_blocking=True)

        if use_cutlass:
            # A: (M,K), B: (N,K) -> feed B^T as (K,N) ColumnMajor
            A_fp4 = a_ref[:, :, l_idx]
            B_fp4_T = b_ref[:, :, l_idx].transpose(0, 1).contiguous()

            # Output slice
            dst = out[:, :, l_idx].contiguous()

            cutlass_gemv(
                A_fp4=A_fp4,
                B_fp4_T=B_fp4_T,
                scaleA=scale_a,
                scaleB=scale_b,
                out_fp16=dst,
            )
            out[:, :, l_idx].copy_(dst)

        else:
            # Fallback to PyTorch's fused dequant path
            res = torch._scaled_mm(
                a_ref[:, :, l_idx],
                b_ref[:, :, l_idx].transpose(0, 1),
                scale_a,
                scale_b,
                bias=None,
                out_dtype=torch.float16,
            )
            out[:, 0, l_idx] = res[:, 0]

    return out
scrolls · 206 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