submission 315436
Gusarich · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 425 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-315436?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:ad593749a1091870c4bd4377c7557a5d794589fd821c3a35e3d26ee541c887e2
license declaredunknown
license concludedunknown
authorsGusarich
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 = 8
num_warps=8,tile-k = 256
BLOCK_K = 256tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128Kernel source
submission.py425 lines
import torch
from task import input_t, output_t
from utils import make_match_reference
# Scaling factor vector size for NVFP4 block scaling
sf_vec_size = 16
# ---------------------------
# Optional Triton fast path
# ---------------------------
_TRITON_AVAILABLE = False
try:
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor
_TRITON_AVAILABLE = True
except Exception:
_TRITON_AVAILABLE = False
# Helper function for ceiling division
def ceil_div(a, b):
return (a + b - 1) // b
# Helper function to convert scale factor tensor to blocked format
# Used only by the reference kernel (slow CPU-side path).
def to_blocked(input_matrix):
rows, cols = input_matrix.shape
# Please ensure rows and cols are multiples of 128 and 4 respectively
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
padded = input_matrix
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()
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
)
# Get dimensions from MxNxL layout
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):
# Convert the scale factor tensor to blocked format
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])
# (m, k) @ (n, k).T -> (m, n)
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
c_ref = (torch.nn.functional.silu(ref1) * ref2).to(torch.float16)
return c_ref
# ---------------------------
# Triton fused kernel
# ---------------------------
if _TRITON_AVAILABLE:
@triton.jit
def _dual_nvfp4_scaled_gemm_silu_kernel(
a_desc, # packed uint8 FP4: (M, K//2)
a_scale_desc, # FP8 scales in TMA-friendly layout
b1_desc, # packed uint8 FP4: (N, K//2)
b1_scale_desc, # FP8 scales
b2_desc, # packed uint8 FP4: (N, K//2)
b2_scale_desc, # FP8 scales
c_desc, # FP16 output: (M, N)
M: tl.constexpr,
N: tl.constexpr,
K: tl.constexpr, # logical K (elements, not bytes)
ELEM_PER_BYTE: tl.constexpr,
VEC_SIZE: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
rep_m: tl.constexpr,
rep_n: tl.constexpr,
rep_k: tl.constexpr,
NUM_STAGES: tl.constexpr,
):
pid = tl.program_id(axis=0)
num_pid_m = tl.cdiv(M, BLOCK_M)
pid_m = pid % num_pid_m
pid_n = pid // num_pid_m
offs_am = pid_m * BLOCK_M
offs_bn = pid_n * BLOCK_N
offs_k = 0
offs_scale_m = pid_m * rep_m
offs_scale_n = pid_n * rep_n
offs_scale_k = 0
acc1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for _ in tl.range(0, tl.cdiv(K, BLOCK_K), num_stages=NUM_STAGES):
a = a_desc.load([offs_am, offs_k])
b1 = b1_desc.load([offs_bn, offs_k])
b2 = b2_desc.load([offs_bn, offs_k])
scale_a = a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
scale_b1 = b1_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
scale_b2 = b2_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
# Unpack into the 2D scale layout expected by tl.dot_scaled
scale_a = (
scale_a.reshape(rep_m, rep_k, 32, 4, 4)
.trans(0, 3, 2, 1, 4)
.reshape(BLOCK_M, BLOCK_K // VEC_SIZE)
)
scale_b1 = (
scale_b1.reshape(rep_n, rep_k, 32, 4, 4)
.trans(0, 3, 2, 1, 4)
.reshape(BLOCK_N, BLOCK_K // VEC_SIZE)
)
scale_b2 = (
scale_b2.reshape(rep_n, rep_k, 32, 4, 4)
.trans(0, 3, 2, 1, 4)
.reshape(BLOCK_N, BLOCK_K // VEC_SIZE)
)
# NVFP4 x NVFP4
acc1 = tl.dot_scaled(a, scale_a, "e2m1", b1.T, scale_b1, "e2m1", acc1)
acc2 = tl.dot_scaled(a, scale_a, "e2m1", b2.T, scale_b2, "e2m1", acc2)
offs_k += BLOCK_K // ELEM_PER_BYTE
offs_scale_k += rep_k
# SiLU: x * sigmoid(x)
sigmoid = 1.0 / (1.0 + tl.exp(-acc1))
out = (acc1 * sigmoid) * acc2
c_desc.store([offs_am, offs_bn], out.to(tl.float16))
def custom_kernel(data: input_t) -> output_t:
"""
Optimized kernel:
C = silu(A @ B1) * (A @ B2)
Uses a fused Triton kernel on Blackwell (CC 10.x/11.x) when available.
Falls back to torch._scaled_mm otherwise.
"""
a, b1, b2, sfa, sfb1, sfb2, sfa_perm, sfb1_perm, sfb2_perm, c = data
# Shapes
# a/b are torch.float4_e2m1fn_x2 with physical K dimension stored as K//2 bytes
m = a.shape[0]
n = b1.shape[0]
l = a.shape[2]
k_bytes = a.shape[1]
k = k_bytes * 2 # logical K elements
# Fast path: Triton fused kernel (Blackwell+)
if _TRITON_AVAILABLE and torch.cuda.is_available():
major, _minor = torch.cuda.get_device_capability()
if major >= 10:
# Views for packed FP4
a_u8 = a.view(torch.uint8)
b1_u8 = b1.view(torch.uint8)
b2_u8 = b2.view(torch.uint8)
# Convert preshuffled scale tensors into the exact TMA-friendly layout used by Triton tutorial:
# shape: (L, M//128, K//64, 2, 256) for A
# shape: (L, N//128, K//64, 2, 256) for B
#
# Input s*_perm is a view with shape (32, 4, rest_m, 4, rest_k, L) and underlying storage
# corresponds to (L, rest_m, rest_k, 32, 4, 4) which is already the cublas packed layout.
rest_m = m // 128
rest_n = n // 128
rest_k = k // sf_vec_size // 4 # K/16/4 = K/64
# Inverse permute back to contiguous (L, rest_m, rest_k, 32, 4, 4), then reshape 32*4*4=512 -> 2*256
a_scale_tma = sfa_perm.permute(5, 2, 4, 0, 1, 3).reshape(
l, rest_m, rest_k, 2, 256
)
b1_scale_tma = sfb1_perm.permute(5, 2, 4, 0, 1, 3).reshape(
l, rest_n, rest_k, 2, 256
)
b2_scale_tma = sfb2_perm.permute(5, 2, 4, 0, 1, 3).reshape(
l, rest_n, rest_k, 2, 256
)
# Kernel params (safe for dual-accumulator register pressure)
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 256
VEC_SIZE = 16
ELEM_PER_BYTE = 2
rep_m = BLOCK_M // 128 # 1
rep_n = BLOCK_N // 128 # 1
rep_k = BLOCK_K // VEC_SIZE // 4 # 4
# Grid over (M,N) tiles
grid = (triton.cdiv(m, BLOCK_M) * triton.cdiv(n, BLOCK_N), 1)
# Launch per-batch (L is small in the benchmark, typically 1)
for l_idx in range(l):
a_desc = TensorDescriptor.from_tensor(
a_u8[:, :, l_idx], [BLOCK_M, BLOCK_K // ELEM_PER_BYTE]
)
b1_desc = TensorDescriptor.from_tensor(
b1_u8[:, :, l_idx], [BLOCK_N, BLOCK_K // ELEM_PER_BYTE]
)
b2_desc = TensorDescriptor.from_tensor(
b2_u8[:, :, l_idx], [BLOCK_N, BLOCK_K // ELEM_PER_BYTE]
)
a_scale_desc = TensorDescriptor.from_tensor(
a_scale_tma[l_idx : l_idx + 1],
block_shape=[1, rep_m, rep_k, 2, 256],
)
b1_scale_desc = TensorDescriptor.from_tensor(
b1_scale_tma[l_idx : l_idx + 1],
block_shape=[1, rep_n, rep_k, 2, 256],
)
b2_scale_desc = TensorDescriptor.from_tensor(
b2_scale_tma[l_idx : l_idx + 1],
block_shape=[1, rep_n, rep_k, 2, 256],
)
c_desc = TensorDescriptor.from_tensor(
c[:, :, l_idx], [BLOCK_M, BLOCK_N]
)
_dual_nvfp4_scaled_gemm_silu_kernel[grid](
a_desc,
a_scale_desc,
b1_desc,
b1_scale_desc,
b2_desc,
b2_scale_desc,
c_desc,
m,
n,
k,
ELEM_PER_BYTE,
VEC_SIZE,
BLOCK_M,
BLOCK_N,
BLOCK_K,
rep_m,
rep_n,
rep_k,
4, # NUM_STAGES
num_warps=8,
)
return c
# Fallback: torch._scaled_mm (correct, slower)
# Use the already-provided preshuffled scales to build cublas scale vectors without CPU loops.
# For nvfp4, cublas expects FP8 E4M3 scales in a packed cublas layout flattened.
a_scale_flat_all = (
sfa_perm.permute(5, 2, 4, 0, 1, 3)
.reshape(l, m // 128, k // 64, 32, 16)
.contiguous()
)
b1_scale_flat_all = (
sfb1_perm.permute(5, 2, 4, 0, 1, 3)
.reshape(l, n // 128, k // 64, 32, 16)
.contiguous()
)
b2_scale_flat_all = (
sfb2_perm.permute(5, 2, 4, 0, 1, 3)
.reshape(l, n // 128, k // 64, 32, 16)
.contiguous()
)
for l_idx in range(l):
scale_a = a_scale_flat_all[l_idx].flatten()
scale_b1 = b1_scale_flat_all[l_idx].flatten()
scale_b2 = b2_scale_flat_all[l_idx].flatten()
r1 = torch._scaled_mm(
a[:, :, l_idx],
b1[:, :, l_idx].transpose(0, 1),
scale_a,
scale_b1,
bias=None,
out_dtype=torch.float32,
)
r2 = torch._scaled_mm(
a[:, :, l_idx],
b2[:, :, l_idx].transpose(0, 1),
scale_a,
scale_b2,
bias=None,
out_dtype=torch.float32,
)
c[:, :, l_idx] = (torch.nn.functional.silu(r1) * r2).to(torch.float16)
return c
def generate_input(m: int, n: int, k: int, l: int, seed: int):
"""
Generate input tensors for NVFP4 block-scaled dual GEMM with silu activation,
C = silu(A @ B1) * (A @ B2).
"""
torch.manual_seed(seed)
def create_fp4_tensors(l, mn, k):
# generate uint8 tensor, then convert to float4e2m1fn_x2 data type
# generate all bit patterns
ref_i8 = torch.randint(
255, size=(l, mn, k // 2), dtype=torch.uint8, device="cuda"
)
# for each nibble, only keep the sign bit and 2 LSBs
# the possible values are [-1.5, -1, -0.5, 0, +0.5, +1, +1.5]
ref_i8 = ref_i8 & 0b1011_1011
return ref_i8.permute(1, 2, 0).view(torch.float4_e2m1fn_x2)
# FP4 inputs (packed, 2 elems per byte along K)
a_ref = create_fp4_tensors(l, m, k).view(torch.float4_e2m1fn_x2)
b1_ref = create_fp4_tensors(l, n, k).view(torch.float4_e2m1fn_x2)
b2_ref = create_fp4_tensors(l, n, k).view(torch.float4_e2m1fn_x2)
# Output buffer
c_ref = torch.randn((l, m, n), dtype=torch.float16, device="cuda").permute(1, 2, 0)
def create_scale_factor_tensors(l, mn, sf_k):
# Create the reference scale factor tensor (mn, sf_k, l) on GPU, then also return a preshuffled layout.
ref_shape = (l, mn, sf_k)
ref_permute_order = (1, 2, 0)
ref_f8_random_fp32 = torch.rand(ref_shape, dtype=torch.float32, device="cuda")
ref_f8_torch_tensor = ref_f8_random_fp32.to(dtype=torch.float8_e4m3fn)
ref_f8_torch_tensor_permuted = ref_f8_torch_tensor.permute(*ref_permute_order)
atom_m = (32, 4)
atom_k = 4
mma_shape = (
l, # batch size
ceil_div(mn, atom_m[0] * atom_m[1]),
ceil_div(sf_k, atom_k),
atom_m[0],
atom_m[1],
atom_k,
)
# Create backing storage and a view with the "CuTe" axis order.
mma_permute_order = (3, 4, 1, 5, 2, 0)
rand_int_tensor = torch.empty(mma_shape, dtype=torch.int8, device="cuda")
reordered_f8_torch_tensor = rand_int_tensor.to(
dtype=torch.float8_e4m3fn
).permute(*mma_permute_order)
# Vectorized reordering on GPU
i_idx = torch.arange(mn, device="cuda")
j_idx = torch.arange(sf_k, device="cuda")
b_idx = torch.arange(l, device="cuda")
i_grid, j_grid, b_grid = torch.meshgrid(i_idx, j_idx, b_idx, indexing="ij")
mm = i_grid // (atom_m[0] * atom_m[1])
mm32 = i_grid % atom_m[0]
mm4 = (i_grid % 128) // atom_m[0]
kk = j_grid // atom_k
kk4 = j_grid % atom_k
reordered_f8_torch_tensor[mm32, mm4, mm, kk4, kk, b_grid] = (
ref_f8_torch_tensor_permuted[i_grid, j_grid, b_grid]
)
return ref_f8_torch_tensor_permuted.cpu(), reordered_f8_torch_tensor
sf_k = ceil_div(k, sf_vec_size)
sfa_ref_cpu, sfa_ref_permuted = create_scale_factor_tensors(l, m, sf_k)
sfb1_ref_cpu, sfb1_ref_permuted = create_scale_factor_tensors(l, n, sf_k)
sfb2_ref_cpu, sfb2_ref_permuted = create_scale_factor_tensors(l, n, sf_k)
return (
a_ref,
b1_ref,
b2_ref,
sfa_ref_cpu.to("cuda"),
sfb1_ref_cpu.to("cuda"),
sfb2_ref_cpu.to("cuda"),
sfa_ref_permuted,
sfb1_ref_permuted,
sfb2_ref_permuted,
c_ref,
)
check_implementation = make_match_reference(ref_kernel, rtol=1e-03, atol=1e-03)
scrolls · 425 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