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
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