Skip to content
KernelIndex
Search⌘K

submission 188992

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

baseline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-188992?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
69.0µs
#328 of 420
2025-12-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e321230a6d657801862ac229c51405ee8306ba459a8a72aeeddf0581b0452b40
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15

Techniques

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

fp4PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("fused_dual_gemm", &fused_dual_gemm, "nvfp4 dual gemm"); }

Kernel source

baseline.py128 lines
import torch
from torch.utils.cpp_extension import load_inline


cpp_src = r"""
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <ATen/ops/_scaled_mm.h>

void silu_mul_cuda(const at::Half* g1, const at::Half* g2, at::Half* out, int64_t count, int64_t out_stride);

torch::Tensor to_blocked(torch::Tensor input) {
    auto rows = input.size(0);
    auto cols = input.size(1);
    auto n_row_blocks = (rows + 127) / 128;
    auto n_col_blocks = (cols + 3) / 4;
    auto padded = input.contiguous();
    auto blocks = padded.view({n_row_blocks, 128, n_col_blocks, 4}).permute({0, 2, 1, 3});
    auto rearranged = blocks.reshape({-1, 4, 32, 4}).transpose(1, 2).reshape({-1, 32, 16});
    return rearranged.reshape({-1}).contiguous();
}

torch::Tensor fused_dual_gemm(torch::Tensor a, torch::Tensor b1, torch::Tensor b2, torch::Tensor sfa, torch::Tensor sfb1, torch::Tensor sfb2, torch::Tensor c) {
    auto out = c.contiguous();
    auto l = out.size(2);
    auto out_stride = out.size(2);
    for (int64_t l_idx = 0; l_idx < l; ++l_idx) {
        auto a_l = a.select(2, l_idx);
        auto b1_l = b1.select(2, l_idx);
        auto b2_l = b2.select(2, l_idx);
        auto scale_a = to_blocked(sfa.select(2, l_idx));
        auto scale_b1 = to_blocked(sfb1.select(2, l_idx));
        auto scale_b2 = to_blocked(sfb2.select(2, l_idx));
        auto g1 = at::_scaled_mm(
            a_l,
            b1_l.transpose(0, 1),
            scale_a,
            scale_b1,
            c10::optional<at::Tensor>(),
            c10::optional<at::Tensor>(),
            c10::optional<at::ScalarType>(at::kHalf),
            false
        );
        auto g2 = at::_scaled_mm(
            a_l,
            b2_l.transpose(0, 1),
            scale_a,
            scale_b2,
            c10::optional<at::Tensor>(),
            c10::optional<at::Tensor>(),
            c10::optional<at::ScalarType>(at::kHalf),
            false
        );
        auto g1_c = g1.contiguous();
        auto g2_c = g2.contiguous();
        int64_t count = g1_c.numel();
        auto out_ptr = reinterpret_cast<at::Half*>(out.data_ptr()) + l_idx;
        silu_mul_cuda(
            reinterpret_cast<const at::Half*>(g1_c.data_ptr()),
            reinterpret_cast<const at::Half*>(g2_c.data_ptr()),
            out_ptr,
            count,
            out_stride
        );
    }
    return out;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("fused_dual_gemm", &fused_dual_gemm, "nvfp4 dual gemm"); }
"""


cuda_src = r"""
#include <ATen/ATen.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <math.h>
#include <stdint.h>

__global__ void silu_mul_kernel(const half* g1, const half* g2, half* out, int64_t count, int64_t out_stride) {
    int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (idx >= count) {
        return;
    }
    float x = __half2float(g1[idx]);
    float y = __half2float(g2[idx]);
    float silu = x / (1.0f + expf(-x));
    out[idx * out_stride] = __float2half(silu * y);
}

void silu_mul_cuda(const at::Half* g1, const at::Half* g2, at::Half* out, int64_t count, int64_t out_stride) {
    int threads = 256;
    int blocks = static_cast<int>((count + threads - 1) / threads);
    auto g1_ptr = reinterpret_cast<const half*>(g1);
    auto g2_ptr = reinterpret_cast<const half*>(g2);
    auto out_ptr = reinterpret_cast<half*>(out);
    silu_mul_kernel<<<blocks, threads>>>(g1_ptr, g2_ptr, out_ptr, count, out_stride);
}
"""


ext = load_inline(
    name="nvfp4_dual_gemm_ext",
    cpp_sources=cpp_src,
    cuda_sources=cuda_src,
    functions=None,
    with_cuda=True,
    extra_cflags=[
        "-O3",
        "-std=c++17",
    ],
    extra_cuda_cflags=[
        "-O3",
        "--use_fast_math",
        "-lineinfo",
    ],
    verbose=False,
)


def custom_kernel(data):
    a, b1, b2, sfa, sfb1, sfb2, sfa_p, sfb1_p, sfb2_p, c = data
    return ext.fused_dual_gemm(a, b1, b2, sfa, sfb1, sfb2, c)


__all__ = ["custom_kernel"]
scrolls · 128 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 188960.

import torch
- import torch.nn.functional as F
+ from torch.utils.cpp_extension import load_inline
- def _to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
- rows, cols = input_matrix.shape
- n_row_blocks = (rows + 127) // 128
- n_col_blocks = (cols + 3) // 4
- padded = input_matrix.contiguous()
- 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()
+ cpp_src = r"""
+ #include <torch/extension.h>
+ #include <ATen/ATen.h>
+ #include <ATen/ops/_scaled_mm.h>
+ void silu_mul_cuda(const at::Half* g1, const at::Half* g2, at::Half* out, int64_t count, int64_t out_stride);
- def custom_kernel(data):
- a, b1, b2, sfa, sfb1, sfb2, sfa_p, sfb1_p, sfb2_p, c = data
- out = c.contiguous()
- _, _, l = out.shape
- for l_idx in range(l):
- scale_a = _to_blocked(sfa[:, :, l_idx])
- scale_b1 = _to_blocked(sfb1[:, :, l_idx])
- scale_b2 = _to_blocked(sfb2[:, :, l_idx])
- g1 = torch._scaled_mm(
- a[:, :, l_idx],
- b1[:, :, l_idx].transpose(0, 1),
+ torch::Tensor to_blocked(torch::Tensor input) {
+ auto rows = input.size(0);
+ auto cols = input.size(1);
+ auto n_row_blocks = (rows + 127) / 128;
+ auto n_col_blocks = (cols + 3) / 4;
+ auto padded = input.contiguous();
+ auto blocks = padded.view({n_row_blocks, 128, n_col_blocks, 4}).permute({0, 2, 1, 3});
+ auto rearranged = blocks.reshape({-1, 4, 32, 4}).transpose(1, 2).reshape({-1, 32, 16});
+ return rearranged.reshape({-1}).contiguous();
+ }
+
+ torch::Tensor fused_dual_gemm(torch::Tensor a, torch::Tensor b1, torch::Tensor b2, torch::Tensor sfa, torch::Tensor sfb1, torch::Tensor sfb2, torch::Tensor c) {
+ auto out = c.contiguous();
+ auto l = out.size(2);
+ auto out_stride = out.size(2);
+ for (int64_t l_idx = 0; l_idx < l; ++l_idx) {
+ auto a_l = a.select(2, l_idx);
+ auto b1_l = b1.select(2, l_idx);
+ auto b2_l = b2.select(2, l_idx);
+ auto scale_a = to_blocked(sfa.select(2, l_idx));
+ auto scale_b1 = to_blocked(sfb1.select(2, l_idx));
+ auto scale_b2 = to_blocked(sfb2.select(2, l_idx));
+ auto g1 = at::_scaled_mm(
+ a_l,
+ b1_l.transpose(0, 1),
scale_a,
scale_b1,
- bias=None,
- out_dtype=torch.float16,
- )
- g2 = torch._scaled_mm(
- a[:, :, l_idx],
- b2[:, :, l_idx].transpose(0, 1),
+ c10::optional<at::Tensor>(),
+ c10::optional<at::Tensor>(),
+ c10::optional<at::ScalarType>(at::kHalf),
+ false
+ );
+ auto g2 = at::_scaled_mm(
+ a_l,
+ b2_l.transpose(0, 1),
scale_a,
scale_b2,
- bias=None,
- out_dtype=torch.float16,
- )
- out[:, :, l_idx] = F.silu(g1) * g2
- return out
+ c10::optional<at::Tensor>(),
+ c10::optional<at::Tensor>(),
+ c10::optional<at::ScalarType>(at::kHalf),
+ false
+ );
+ auto g1_c = g1.contiguous();
+ auto g2_c = g2.contiguous();
+ int64_t count = g1_c.numel();
+ auto out_ptr = reinterpret_cast<at::Half*>(out.data_ptr()) + l_idx;
+ silu_mul_cuda(
+ reinterpret_cast<const at::Half*>(g1_c.data_ptr()),
+ reinterpret_cast<const at::Half*>(g2_c.data_ptr()),
+ out_ptr,
+ count,
+ out_stride
+ );
+ }
+ return out;
+ }
+ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("fused_dual_gemm", &fused_dual_gemm, "nvfp4 dual gemm"); }
+ """
+
+ cuda_src = r"""
+ #include <ATen/ATen.h>
+ #include <cuda.h>
+ #include <cuda_fp16.h>
+ #include <cuda_runtime.h>
+ #include <math.h>
+ #include <stdint.h>
+
+ __global__ void silu_mul_kernel(const half* g1, const half* g2, half* out, int64_t count, int64_t out_stride) {
+ int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
+ if (idx >= count) {
+ return;
+ }
+ float x = __half2float(g1[idx]);
+ float y = __half2float(g2[idx]);
+ float silu = x / (1.0f + expf(-x));
+ out[idx * out_stride] = __float2half(silu * y);
+ }
+
+ void silu_mul_cuda(const at::Half* g1, const at::Half* g2, at::Half* out, int64_t count, int64_t out_stride) {
+ int threads = 256;
+ int blocks = static_cast<int>((count + threads - 1) / threads);
+ auto g1_ptr = reinterpret_cast<const half*>(g1);
+ auto g2_ptr = reinterpret_cast<const half*>(g2);
+ auto out_ptr = reinterpret_cast<half*>(out);
+ silu_mul_kernel<<<blocks, threads>>>(g1_ptr, g2_ptr, out_ptr, count, out_stride);
+ }
+ """
+
+
+ ext = load_inline(
+ name="nvfp4_dual_gemm_ext",
+ cpp_sources=cpp_src,
+ cuda_sources=cuda_src,
+ functions=None,
+ with_cuda=True,
+ extra_cflags=[
+ "-O3",
+ "-std=c++17",
+ ],
+ extra_cuda_cflags=[
+ "-O3",
+ "--use_fast_math",
+ "-lineinfo",
+ ],
+ verbose=False,
+ )
+
+
+ def custom_kernel(data):
+ a, b1, b2, sfa, sfb1, sfb2, sfa_p, sfb1_p, sfb2_p, c = data
+ return ext.fused_dual_gemm(a, b1, b2, sfa, sfb1, sfb2, c)
+
+
__all__ = ["custom_kernel"]
scrolls · 158 diff lines total

Best evidence level for this revision: reported

JSON