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
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.
fp4
PYBIND11_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