submission 260929
phuc9702 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 248 lines, June 9 Researcher Reciprocity License v1.0.
nvfp4_dual_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-260929?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:4de8aa5c9c3f9f0922819ee8555ca64133e72d80a1efd061e9f67082b20c52a9
license declaredunknown
license concludedunknown
authorsphuc9702
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,num-warps = 4
num_warps = 4stages = 3
num_stages = 3tile-k = 256
BLOCK_K = 256tile-m = 128
BLOCK_M = 128tile-n = 64
BLOCK_N = 64Kernel source
nvfp4_dual_gemm.py248 lines
import torch
import triton
import triton.language as tl
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: torch.Tensor) -> torch.Tensor:
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()
def ref_kernel(data: input_t) -> output_t:
"""
PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,
C = silu(A @ B1) * (A @ B2).
"""
a_ref, b1_ref, b2_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu, _, _, _, c_ref = data
m, n, l = c_ref.shape
ref1 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)
ref2 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)
for l_idx in range(l):
scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
scale_b1 = to_blocked(sfb1_ref_cpu[:, :, l_idx])
scale_b2 = to_blocked(sfb2_ref_cpu[:, :, l_idx])
res1 = torch._scaled_mm(
a_ref[:, :, l_idx],
b1_ref[:, :, l_idx].transpose(0, 1),
scale_a.cuda(),
scale_b1.cuda(),
bias=None,
out_dtype=torch.float32,
)
ref1[:, :, l_idx] = res1
res2 = torch._scaled_mm(
a_ref[:, :, l_idx],
b2_ref[:, :, l_idx].transpose(0, 1),
scale_a.cuda(),
scale_b2.cuda(),
bias=None,
out_dtype=torch.float32,
)
ref2[:, :, l_idx] = res2
return (torch.nn.functional.silu(ref1) * ref2).to(torch.float16)
@triton.jit
def _dual_gemm_silu_kernel_opt(
a_ptr, # uint8 packed fp4, [M, K_bytes, L]
b1_ptr, # uint8 packed fp4, [N, K_bytes, L]
b2_ptr, # uint8 packed fp4, [N, K_bytes, L]
sfa_ptr, # fp8 scale, [M, K//VEC_SIZE, L]
sfb1_ptr, # fp8 scale, [N, K//VEC_SIZE, L]
sfb2_ptr, # fp8 scale, [N, K//VEC_SIZE, L]
c_ptr, # fp16 output, [M, N, L]
stride_am,
stride_akb,
stride_al,
stride_b1n,
stride_b1kb,
stride_b1l,
stride_b2n,
stride_b2kb,
stride_b2l,
stride_sfam,
stride_sfak,
stride_sfal,
stride_sfb1n,
stride_sfb1k,
stride_sfb1l,
stride_sfb2n,
stride_sfb2k,
stride_sfb2l,
stride_cm,
stride_cn,
stride_cl,
M: tl.constexpr,
N: tl.constexpr,
K: tl.constexpr,
L: tl.constexpr,
ELEM_PER_BYTE: tl.constexpr,
VEC_SIZE: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
NUM_STAGES: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
pid_l = tl.program_id(axis=2)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
acc1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k_bytes_total = K // ELEM_PER_BYTE
k_tiles = tl.cdiv(K, BLOCK_K)
kb_tile = BLOCK_K // ELEM_PER_BYTE
ks_tile = BLOCK_K // VEC_SIZE
base_a_l = a_ptr + pid_l * stride_al
base_b1_l = b1_ptr + pid_l * stride_b1l
base_b2_l = b2_ptr + pid_l * stride_b2l
base_sfa_l = sfa_ptr + pid_l * stride_sfal
base_sfb1_l = sfb1_ptr + pid_l * stride_sfb1l
base_sfb2_l = sfb2_ptr + pid_l * stride_sfb2l
for kt in tl.range(0, k_tiles, num_stages=NUM_STAGES):
offs_kb = kt * kb_tile + tl.arange(0, BLOCK_K // ELEM_PER_BYTE)
offs_ks = kt * ks_tile + tl.arange(0, BLOCK_K // VEC_SIZE)
a_ptrs = base_a_l + offs_m[:, None] * stride_am + offs_kb[None, :] * stride_akb
a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_kb[None, :] < k_bytes_total), other=0).to(tl.uint8)
sfa_ptrs = base_sfa_l + offs_m[:, None] * stride_sfam + offs_ks[None, :] * stride_sfak
scale_a = tl.load(sfa_ptrs, mask=(offs_m[:, None] < M) & (offs_ks[None, :] < (K // VEC_SIZE)), other=0.0)
b_ptrs = base_b1_l + offs_n[:, None] * stride_b1n + offs_kb[None, :] * stride_b1kb
b = tl.load(b_ptrs, mask=(offs_n[:, None] < N) & (offs_kb[None, :] < k_bytes_total), other=0).to(tl.uint8)
sfb_ptrs = base_sfb1_l + offs_n[:, None] * stride_sfb1n + offs_ks[None, :] * stride_sfb1k
scale_b = tl.load(sfb_ptrs, mask=(offs_n[:, None] < N) & (offs_ks[None, :] < (K // VEC_SIZE)), other=0.0)
acc1 = tl.dot_scaled(
a, scale_a, "e2m1",
b.T, scale_b, "e2m1",
acc=acc1,
fast_math=True,
lhs_k_pack=True,
rhs_k_pack=True,
out_dtype=tl.float32,
)
b_ptrs = base_b2_l + offs_n[:, None] * stride_b2n + offs_kb[None, :] * stride_b2kb
b = tl.load(b_ptrs, mask=(offs_n[:, None] < N) & (offs_kb[None, :] < k_bytes_total), other=0).to(tl.uint8)
sfb_ptrs = base_sfb2_l + offs_n[:, None] * stride_sfb2n + offs_ks[None, :] * stride_sfb2k
scale_b = tl.load(sfb_ptrs, mask=(offs_n[:, None] < N) & (offs_ks[None, :] < (K // VEC_SIZE)), other=0.0)
acc2 = tl.dot_scaled(
a, scale_a, "e2m1",
b.T, scale_b, "e2m1",
acc=acc2,
fast_math=True,
lhs_k_pack=True,
rhs_k_pack=True,
out_dtype=tl.float32,
)
sig = 1.0 / (1.0 + tl.exp(-acc1))
acc1 = acc1 * sig
out = acc1 * acc2
out = out.to(tl.float16)
c_ptrs = c_ptr + pid_l * stride_cl + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptrs, out, mask=mask)
def custom_kernel(data: input_t) -> output_t:
a, b1, b2, sfa, sfb1, sfb2, _, _, _, c = data
m, k_bytes, l = a.shape
n = b1.shape[0]
k = k_bytes * 2
a_u8 = a.view(torch.uint8)
b1_u8 = b1.view(torch.uint8)
b2_u8 = b2.view(torch.uint8)
out = c
BLOCK_M = 128
BLOCK_N = 64
BLOCK_K = 256
VEC_SIZE = 16
ELEM_PER_BYTE = 2
num_warps = 4
num_stages = 3
grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N), l)
_dual_gemm_silu_kernel_opt[grid](
a_u8,
b1_u8,
b2_u8,
sfa,
sfb1,
sfb2,
out,
a_u8.stride(0),
a_u8.stride(1),
a_u8.stride(2),
b1_u8.stride(0),
b1_u8.stride(1),
b1_u8.stride(2),
b2_u8.stride(0),
b2_u8.stride(1),
b2_u8.stride(2),
sfa.stride(0),
sfa.stride(1),
sfa.stride(2),
sfb1.stride(0),
sfb1.stride(1),
sfb1.stride(2),
sfb2.stride(0),
sfb2.stride(1),
sfb2.stride(2),
out.stride(0),
out.stride(1),
out.stride(2),
M=m,
N=n,
K=k,
L=l,
ELEM_PER_BYTE=ELEM_PER_BYTE,
VEC_SIZE=VEC_SIZE,
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
NUM_STAGES=num_stages,
num_warps=num_warps,
)
return out
check_implementation = make_match_reference(custom_kernel, rtol=1e-03, atol=1e-03)
scrolls · 248 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