Skip to content
KernelIndex
Search⌘K

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
NVFP4 dual GEMMsuite of 4 cases
NVIDIA B200
30.7µs
#254 of 420
2025-12-22

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.

fp4const int64_t K = A.size(1) * 2; // FP4 is packed 2 per byte

Kernel 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