Skip to content
KernelIndex
Search⌘K

submission 364761

Natalie Gollahon · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

bs_improved.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-364761?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
82.3µs
#386 of 420
2026-01-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a016f95a609b95d09d576dba25dddb26bcc255d3c373062f8a6e0622b4f827a3
license declaredunknown
license concludedunknown
authorsNatalie Gollahon
imported2026-08-26

Techniques

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

fused-epilogue_CUDA_EPILOGUE = r"""

Kernel source

bs_improved.py211 lines
import os

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t
from utils import make_match_reference

# Scaling factor vector size
sf_vec_size = 16


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


def to_blocked(input_matrix):
    rows, cols = input_matrix.shape
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)
    blocks = input_matrix.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()


# -----------------------------------------------------------------------------
# Fast baseline: use Tensor-Core-backed torch._scaled_mm for the two GEMMs, then
# fuse only the epilogue (SiLU * mul + fp16 cast) in a small CUDA kernel.
# -----------------------------------------------------------------------------

_CUDA_EPILOGUE = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>

__device__ __forceinline__ float sigmoidf_fast(float x) {
    return 1.0f / (1.0f + __expf(-x));
}

__global__ void fused_silu_mul_f16_kernel_f32(
    const float* __restrict__ x,
    const float* __restrict__ y,
    half* __restrict__ out,
    int64_t numel
) {
    int64_t idx = (int64_t)blockIdx.x * blockDim.x + threadIdx.x;
    if (idx >= numel) return;
    float xv = x[idx];
    float silu = xv * sigmoidf_fast(xv);
    out[idx] = __float2half_rn(silu * y[idx]);
}

__global__ void fused_silu_mul_f16_kernel_f16(
    const half* __restrict__ x,
    const half* __restrict__ y,
    half* __restrict__ out,
    int64_t numel
) {
    int64_t idx = (int64_t)blockIdx.x * blockDim.x + threadIdx.x;
    if (idx >= numel) return;
    float xv = __half2float(x[idx]);
    float silu = xv * sigmoidf_fast(xv);
    float yv = __half2float(y[idx]);
    out[idx] = __float2half_rn(silu * yv);
}

torch::Tensor fused_silu_mul_f16(torch::Tensor x, torch::Tensor y) {
    TORCH_CHECK(x.is_cuda() && y.is_cuda(), "x/y must be CUDA");
    TORCH_CHECK(x.sizes() == y.sizes(), "x/y shape mismatch");
    TORCH_CHECK(x.is_contiguous() && y.is_contiguous(), "x/y must be contiguous");

    TORCH_CHECK(
        (x.scalar_type() == torch::kFloat32 && y.scalar_type() == torch::kFloat32) ||
        (x.scalar_type() == torch::kFloat16 && y.scalar_type() == torch::kFloat16),
        "x/y must both be float32 or both be float16"
    );

    auto out = torch::empty_like(x, x.options().dtype(torch::kFloat16));
    int64_t numel = x.numel();
    constexpr int threads = 256;
    int blocks = (int)((numel + threads - 1) / threads);

    if (x.scalar_type() == torch::kFloat32) {
        fused_silu_mul_f16_kernel_f32<<<blocks, threads>>>(
            (const float*)x.data_ptr<float>(),
            (const float*)y.data_ptr<float>(),
            (half*)out.data_ptr<at::Half>(),
            numel
        );
    } else {
        fused_silu_mul_f16_kernel_f16<<<blocks, threads>>>(
            (const half*)x.data_ptr<at::Half>(),
            (const half*)y.data_ptr<at::Half>(),
            (half*)out.data_ptr<at::Half>(),
            numel
        );
    }
    return out;
}
"""

_CPP_EPILOGUE = r"""
#include <torch/extension.h>
torch::Tensor fused_silu_mul_f16(torch::Tensor x, torch::Tensor y);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("fused_silu_mul_f16", &fused_silu_mul_f16, "fused silu(x) * y -> fp16 (CUDA)");
}
"""

# Optional compilation probe: if BS_TRY_CUTLASS=1, we also compile a tiny TU that
# includes CUTLASS headers. This is just to test availability in the evaluator
# image. If CUTLASS isn't present, the extension compilation will fail with a
# clear "cutlass/..." include error.
_CUDA_CUTLASS_PROBE = r"""
#include <cutlass/cutlass.h>
__global__ void cutlass_probe_kernel() {}
"""

_ep = None


def _get_ep():
    global _ep
    if _ep is not None:
        return _ep
    try_cutlass = os.environ.get("BS_TRY_CUTLASS", "0") == "1"
    cuda_sources = [_CUDA_EPILOGUE]
    if try_cutlass:
        cuda_sources.append(_CUDA_CUTLASS_PROBE)
    _ep = load_inline(
        name=f"bs_improved_ep{os.environ.get('KERNELGEN_EXT_SUFFIX','')}",
        cpp_sources=[_CPP_EPILOGUE],
        cuda_sources=cuda_sources,
        functions=None,
        with_cuda=True,
        verbose=True,
        extra_cflags=["-O3", "-std=c++17"],
        extra_cuda_cflags=["-O3", "--use_fast_math"],
    )
    return _ep


def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, sfa, sfb1, sfb2, _, _, _, c = data
    m, n, l = c.shape

    # Default to fp32 intermediates to match reference numerics.
    # You can try fp16 intermediates by setting BS_MM_OUT_DTYPE=fp16 (may fail correctness).
    mm_out_dtype = os.environ.get("BS_MM_OUT_DTYPE", "fp32").lower()
    if mm_out_dtype == "fp16":
        out_dtype = torch.float16
    else:
        out_dtype = torch.float32

    # Optional experiment: do one GEMM with Bcat=[B1;B2] to reduce launch/scale overhead
    # (still computes 2N columns, so FLOPs are the same; may help for small-M regimes).
    one_gemm = os.environ.get("BS_ONE_GEMM", "0") == "1"

    out1 = torch.empty((m, n, l), dtype=out_dtype, device="cuda")
    out2 = torch.empty((m, n, l), dtype=out_dtype, device="cuda")

    if one_gemm:
        # Concatenate along N (dim 0): (2N, K, L)
        bcat = torch.cat([b1, b2], dim=0)
        sfbcat = torch.cat([sfb1, sfb2], dim=0)
        for l_idx in range(l):
            # (No caching here because sfbcat is a fresh tensor)
            scale_a = to_blocked(sfa[:, :, l_idx]).contiguous()
            scale_bcat = to_blocked(sfbcat[:, :, l_idx]).contiguous()
            tmp = torch._scaled_mm(
                a[:, :, l_idx],
                bcat[:, :, l_idx].transpose(0, 1),  # (K, 2N)
                scale_a,
                scale_bcat,
                bias=None,
                out_dtype=out_dtype,
            )
            out1[:, :, l_idx] = tmp[:, :n]
            out2[:, :, l_idx] = tmp[:, n:]
    else:
        for l_idx in range(l):
            scale_a = to_blocked(sfa[:, :, l_idx]).contiguous()
            scale_b1 = to_blocked(sfb1[:, :, l_idx]).contiguous()
            scale_b2 = to_blocked(sfb2[:, :, l_idx]).contiguous()

            out1[:, :, l_idx] = torch._scaled_mm(
                a[:, :, l_idx],
                b1[:, :, l_idx].transpose(0, 1),
                scale_a,
                scale_b1,
                bias=None,
                out_dtype=out_dtype,
            )
            out2[:, :, l_idx] = torch._scaled_mm(
                a[:, :, l_idx],
                b2[:, :, l_idx].transpose(0, 1),
                scale_a,
                scale_b2,
                bias=None,
                out_dtype=out_dtype,
            )

    ep = _get_ep()
    out = ep.fused_silu_mul_f16(out1.contiguous(), out2.contiguous()).view(m, n, l)
    return out


check_implementation = make_match_reference(custom_kernel, rtol=1e-03, atol=1e-03)

scrolls · 211 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