submission 189713
jiab_85281 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 223 lines, June 9 Researcher Reciprocity License v1.0.
sub_cub2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-189713?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:25a3a0edea502929005d6f202bc5a1219658334d5d6713b7d26da842dbce4205
license declaredunknown
license concludedunknown
authorsjiab_85281
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
const int64_t K = A.size(1) * 2; // FP4 is packed 2 per byteKernel source
sub_cub2.py223 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
cpp_src = """
#include <torch/extension.h>
torch::Tensor cuda_nvfp4_dual_gemm_cublaslt(
torch::Tensor A,
torch::Tensor B1,
torch::Tensor B2,
torch::Tensor SFA,
torch::Tensor SFB1,
torch::Tensor SFB2,
torch::Tensor SFA_perm,
torch::Tensor SFB1_perm,
torch::Tensor SFB2_perm,
torch::Tensor C);
"""
cuda_src = """
#include <torch/extension.h>
#include <cublasLt.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_fp4.h>
#include <stdexcept>
namespace {
inline cublasLtHandle_t get_handle() {
static cublasLtHandle_t h = [] {
cublasLtHandle_t t;
cublasLtCreate(&t);
return t;
}();
return h;
}
inline void check(cublasStatus_t s, const char* m) {
if (s != CUBLAS_STATUS_SUCCESS) throw std::runtime_error(m);
}
// SiLU activation kernel: out = silu(acc1) * acc2
// where silu(x) = x * sigmoid(x) = x / (1 + exp(-x))
__global__ void silu_mul_kernel(
const half* __restrict__ acc1,
const half* __restrict__ acc2,
half* __restrict__ out,
int64_t size)
{
int64_t idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx < size) {
float x = __half2float(acc1[idx]);
float y = __half2float(acc2[idx]);
// silu(x) = x * sigmoid(x) = x / (1 + exp(-x))
float silu_x = x / (1.0f + expf(-x));
out[idx] = __float2half(silu_x * y);
}
}
void run_gemm(
cublasLtHandle_t handle,
torch::Tensor A,
torch::Tensor B,
torch::Tensor SFA_perm,
torch::Tensor SFB_perm,
torch::Tensor C,
int64_t M, int64_t K, int64_t N)
{
const int64_t lda = K;
const int64_t ldb = K;
const int64_t ldc = N;
cublasLtMatmulDesc_t opDesc;
check(cublasLtMatmulDescCreate(&opDesc, CUBLAS_COMPUTE_32F, CUDA_R_32F), "desc");
cublasOperation_t transA = CUBLAS_OP_N;
cublasOperation_t transB = CUBLAS_OP_T;
check(cublasLtMatmulDescSetAttribute(opDesc, CUBLASLT_MATMUL_DESC_TRANSA,
&transA, sizeof(transA)), "ta");
check(cublasLtMatmulDescSetAttribute(opDesc, CUBLASLT_MATMUL_DESC_TRANSB,
&transB, sizeof(transB)), "tb");
cublasLtMatmulMatrixScale_t scale_mode = CUBLASLT_MATMUL_MATRIX_SCALE_VEC16_UE4M3;
check(cublasLtMatmulDescSetAttribute(opDesc, CUBLASLT_MATMUL_DESC_A_SCALE_MODE,
&scale_mode, sizeof(scale_mode)), "asm");
check(cublasLtMatmulDescSetAttribute(opDesc, CUBLASLT_MATMUL_DESC_B_SCALE_MODE,
&scale_mode, sizeof(scale_mode)), "bsm");
const void* a_scale_ptr = SFA_perm.data_ptr();
const void* b_scale_ptr = SFB_perm.data_ptr();
check(cublasLtMatmulDescSetAttribute(opDesc, CUBLASLT_MATMUL_DESC_A_SCALE_POINTER,
&a_scale_ptr, sizeof(a_scale_ptr)), "asp");
check(cublasLtMatmulDescSetAttribute(opDesc, CUBLASLT_MATMUL_DESC_B_SCALE_POINTER,
&b_scale_ptr, sizeof(b_scale_ptr)), "bsp");
cublasLtMatrixLayout_t Adesc, Bdesc, Cdesc, Ddesc;
check(cublasLtMatrixLayoutCreate(&Adesc, CUDA_R_4F_E2M1, M, K, lda), "Adesc");
check(cublasLtMatrixLayoutCreate(&Bdesc, CUDA_R_4F_E2M1, N, K, ldb), "Bdesc");
check(cublasLtMatrixLayoutCreate(&Cdesc, CUDA_R_16F, M, N, ldc), "Cdesc");
check(cublasLtMatrixLayoutCreate(&Ddesc, CUDA_R_16F, M, N, ldc), "Ddesc");
cublasLtOrder_t order = CUBLASLT_ORDER_ROW;
check(cublasLtMatrixLayoutSetAttribute(Adesc, CUBLASLT_MATRIX_LAYOUT_ORDER,
&order, sizeof(order)), "orderA");
check(cublasLtMatrixLayoutSetAttribute(Bdesc, CUBLASLT_MATRIX_LAYOUT_ORDER,
&order, sizeof(order)), "orderB");
check(cublasLtMatrixLayoutSetAttribute(Cdesc, CUBLASLT_MATRIX_LAYOUT_ORDER,
&order, sizeof(order)), "orderC");
check(cublasLtMatrixLayoutSetAttribute(Ddesc, CUBLASLT_MATRIX_LAYOUT_ORDER,
&order, sizeof(order)), "orderD");
cublasLtMatmulPreference_t pref;
check(cublasLtMatmulPreferenceCreate(&pref), "pref");
size_t ws = 0;
check(cublasLtMatmulPreferenceSetAttribute(
pref, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &ws, sizeof(ws)),
"pref_ws");
cublasLtMatmulHeuristicResult_t hres{};
int ret = 0;
check(cublasLtMatmulAlgoGetHeuristic(
handle, opDesc, Adesc, Bdesc, Cdesc, Ddesc, pref, 1, &hres, &ret),
"heuristic");
TORCH_CHECK(ret > 0, "no algo");
const float alpha = 1.0f;
const float beta = 0.0f;
check(cublasLtMatmul(
handle, opDesc, &alpha,
A.data_ptr(), Adesc,
B.data_ptr(), Bdesc,
&beta,
C.data_ptr(), Cdesc,
C.data_ptr(), Ddesc,
&hres.algo,
nullptr, 0,
nullptr),
"matmul");
cublasLtMatmulPreferenceDestroy(pref);
cublasLtMatrixLayoutDestroy(Adesc);
cublasLtMatrixLayoutDestroy(Bdesc);
cublasLtMatrixLayoutDestroy(Cdesc);
cublasLtMatrixLayoutDestroy(Ddesc);
cublasLtMatmulDescDestroy(opDesc);
}
} // namespace
torch::Tensor cuda_nvfp4_dual_gemm_cublaslt(
torch::Tensor A,
torch::Tensor B1,
torch::Tensor B2,
torch::Tensor SFA,
torch::Tensor SFB1,
torch::Tensor SFB2,
torch::Tensor SFA_perm,
torch::Tensor SFB1_perm,
torch::Tensor SFB2_perm,
torch::Tensor C)
{
const int64_t M = A.size(0);
const int64_t K = A.size(1) * 2; // FP4 is packed 2 per byte
const int64_t N = B1.size(0);
cublasLtHandle_t handle = get_handle();
// Allocate temporary buffer for second GEMM result
auto acc2 = torch::empty_like(C);
// Perform first GEMM: acc1 = A @ B1^T (stored in C temporarily)
run_gemm(handle, A, B1, SFA_perm, SFB1_perm, C, M, K, N);
// Perform second GEMM: acc2 = A @ B2^T
run_gemm(handle, A, B2, SFA_perm, SFB2_perm, acc2, M, K, N);
// Apply SiLU and multiply: C = silu(acc1) * acc2
int64_t size = M * N;
int threads = 256;
int blocks = (size + threads - 1) / threads;
silu_mul_kernel<<<blocks, threads>>>(
reinterpret_cast<const half*>(C.data_ptr()),
reinterpret_cast<const half*>(acc2.data_ptr()),
reinterpret_cast<half*>(C.data_ptr()),
size
);
return C;
}
"""
nvfp4_dual_gemm_module = load_inline(
name="nvfp4_dual_gemm_cublaslt",
cpp_sources=[cpp_src],
cuda_sources=[cuda_src],
functions=["cuda_nvfp4_dual_gemm_cublaslt"],
extra_cuda_cflags=[
"-std=c++17",
"-gencode=arch=compute_100a,code=sm_100a",
"--ptxas-options=--gpu-name=sm_100a",
"-O3",
"-w",
"-allow-unsupported-compiler",
],
extra_ldflags=["-lcuda"],
verbose=False,
)
def custom_kernel(data: input_t) -> output_t:
a, b1, b2, sfa, sfb1, sfb2, sfa_perm, sfb1_perm, sfb2_perm, c = data
return nvfp4_dual_gemm_module.cuda_nvfp4_dual_gemm_cublaslt(
a, b1, b2, sfa, sfb1, sfb2, sfa_perm, sfb1_perm, sfb2_perm, c
)
scrolls · 223 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