submission 327716
HayatoFujihara · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 268 lines, June 9 Researcher Reciprocity License v1.0.
nvfp4_dual_gemm_v12.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-327716?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:f7405e11b6e8341075a2dd1d45ffd98c5511eaf1eebb4fd1475dde646d67edb6
license declaredunknown
license concludedunknown
authorsHayatoFujihara
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
NVFP4 Dual GEMM + SiLU Fusion Kernel v12fused-epilogue
- SiLU fused in epiloguestages = 4
NUM_STAGES = 4 # Pipeline stages (optimal)tile-k = 256
BLOCK_K = 256 # K-dimension tile size (FP4 native K)tile-m = 128
BLOCK_M = 128 # M-dimension tile sizetile-n = 128
BLOCK_N = 128 # N-dimension tile sizeKernel source
nvfp4_dual_gemm_v12.py268 lines
"""
NVFP4 Dual GEMM + SiLU Fusion Kernel v12
============================================================================
Triton tl.dot_scaled Implementation - Host-side Optimization
Based on v11 (24.8us). Changes:
1. 2D grid instead of 1D (removes div/mod overhead in kernel)
2. L=1 specialization (removes batch loop overhead)
3. Inline scale conversion (reduces function call overhead)
4. Pre-computed constants outside batch loop
Target: 13.915us (1st place)
============================================================================
"""
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor
from task import input_t, output_t
# ============================================================================
# Configuration - Tuned for Blackwell B200
# ============================================================================
BLOCK_M = 128 # M-dimension tile size
BLOCK_N = 128 # N-dimension tile size
BLOCK_K = 256 # K-dimension tile size (FP4 native K)
VEC_SIZE = 16 # NVFP4: 16 elements per E4M3 scale
ELEM_PER_BYTE = 2 # FP4 packs 2 elements per byte
NUM_STAGES = 4 # Pipeline stages (optimal)
# ============================================================================
# Triton Kernel - Fused Dual GEMM + SiLU
# ============================================================================
@triton.jit
def fused_dual_gemm_silu_kernel(
# TMA descriptors
a_desc,
a_scale_desc,
b1_desc,
b1_scale_desc,
b2_desc,
b2_scale_desc,
c_desc,
# Shape parameters
M: tl.constexpr,
N: tl.constexpr,
K: tl.constexpr,
# Tile sizes
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
VEC_SIZE: tl.constexpr,
# Iteration parameters
rep_m: tl.constexpr,
rep_n: tl.constexpr,
rep_k: tl.constexpr,
NUM_STAGES: tl.constexpr,
ELEM_PER_BYTE: tl.constexpr,
):
"""
Fused Dual GEMM + SiLU kernel using TMA and tl.dot_scaled.
Computation: C = silu(A @ B1.T) * (A @ B2.T)
Optimizations:
- TMA for efficient memory access
- A matrix loaded once, reused for both B1 and B2 products
- Intermediate results kept in registers
- SiLU fused in epilogue
- 2D grid for direct program ID access (no div/mod)
"""
# Program ID (2D grid - direct access, no div/mod)
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
# Offset calculations
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
# Initialize accumulators (FP32 for precision)
acc1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# K-dimension loop with pipelining
num_k_iters = tl.cdiv(K, BLOCK_K)
for _ in tl.range(0, num_k_iters, num_stages=NUM_STAGES):
# Load A tile via TMA
a = a_desc.load([offs_am, offs_k])
# Load and transform A scale factors
raw_scale_a = a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
scale_a = raw_scale_a.reshape(rep_m, rep_k, 32, 4, 4)
scale_a = scale_a.trans(0, 3, 2, 1, 4)
scale_a = scale_a.reshape(BLOCK_M, BLOCK_K // VEC_SIZE)
# Load B1 tile
b1 = b1_desc.load([offs_bn, offs_k])
# Load and transform B1 scale factors
raw_scale_b1 = b1_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
scale_b1 = raw_scale_b1.reshape(rep_n, rep_k, 32, 4, 4)
scale_b1 = scale_b1.trans(0, 3, 2, 1, 4)
scale_b1 = scale_b1.reshape(BLOCK_N, BLOCK_K // VEC_SIZE)
# Load B2 tile
b2 = b2_desc.load([offs_bn, offs_k])
# Load and transform B2 scale factors
raw_scale_b2 = b2_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
scale_b2 = raw_scale_b2.reshape(rep_n, rep_k, 32, 4, 4)
scale_b2 = scale_b2.trans(0, 3, 2, 1, 4)
scale_b2 = scale_b2.reshape(BLOCK_N, BLOCK_K // VEC_SIZE)
# Dual GEMM with A reuse
# tl.dot_scaled handles FP4 E2M1 with E4M3 block scales natively
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)
# Update offsets for next iteration
offs_k += BLOCK_K // ELEM_PER_BYTE
offs_scale_k += rep_k
# Fused SiLU + multiply epilogue
# silu(x) = x * sigmoid(x)
result = (acc1 * tl.sigmoid(acc1)) * acc2
# Convert to FP16 for output
result = result.to(tl.float16)
# Store result via TMA
c_desc.store([offs_am, offs_bn], result)
# ============================================================================
# Host-side Entry Point
# ============================================================================
def custom_kernel(data: input_t) -> output_t:
"""
Triton Fused Dual GEMM + SiLU (TMA + tl.dot_scaled)
Computation: C = silu(A @ B1.T) * (A @ B2.T)
Optimizations applied:
- L=1 specialization: no batch loop overhead
- Inline scale conversion: no function call overhead
- Pre-computed shapes: avoid repeated calculations
"""
# Unpack inputs (use permuted scales for TMA compatibility)
a, b1, b2, _, _, _, sfa_perm, sfb1_perm, sfb2_perm, c = data
# Get dimensions
M, N, L = c.shape
K = a.shape[1] * 2 # FP4 packs 2 elements per byte
# Pre-computed constants (avoid recalculation in loop)
rep_m = BLOCK_M // 128 # = 1
rep_n = BLOCK_N // 128 # = 1
rep_k = BLOCK_K // VEC_SIZE // 4 # = 4
k_bytes = BLOCK_K // ELEM_PER_BYTE # = 128
# TMA block shapes (computed once)
a_block_shape = [BLOCK_M, k_bytes]
b_block_shape = [BLOCK_N, k_bytes]
c_block_shape = [BLOCK_M, BLOCK_N]
scale_block_shape = [1, rep_m, rep_k, 2, 256]
# Grid configuration
grid = (M // BLOCK_M, N // BLOCK_N)
# L=1 specialization: inline everything for minimal overhead
if L == 1:
# Direct slice without loop variable
a_u8 = a[:, :, 0].view(torch.uint8)
b1_u8 = b1[:, :, 0].view(torch.uint8)
b2_u8 = b2[:, :, 0].view(torch.uint8)
c_l = c[:, :, 0]
# Inline scale conversion for sfa
# (32, 4, rest_m, 4, rest_k, L) -> (1, rest_m, rest_k, 2, 256)
rest_m = sfa_perm.shape[2]
rest_k = sfa_perm.shape[4]
sfa_tma = sfa_perm[:, :, :, :, :, 0].permute(2, 4, 0, 1, 3).reshape(
1, rest_m, rest_k, 2, 256
).contiguous()
# Inline scale conversion for sfb1
rest_n = sfb1_perm.shape[2]
sfb1_tma = sfb1_perm[:, :, :, :, :, 0].permute(2, 4, 0, 1, 3).reshape(
1, rest_n, rest_k, 2, 256
).contiguous()
# Inline scale conversion for sfb2
sfb2_tma = sfb2_perm[:, :, :, :, :, 0].permute(2, 4, 0, 1, 3).reshape(
1, rest_n, rest_k, 2, 256
).contiguous()
# Create TMA tensor descriptors
a_desc = TensorDescriptor.from_tensor(a_u8, a_block_shape)
b1_desc = TensorDescriptor.from_tensor(b1_u8, b_block_shape)
b2_desc = TensorDescriptor.from_tensor(b2_u8, b_block_shape)
c_desc = TensorDescriptor.from_tensor(c_l, c_block_shape)
a_scale_desc = TensorDescriptor.from_tensor(sfa_tma, scale_block_shape)
b1_scale_desc = TensorDescriptor.from_tensor(sfb1_tma, scale_block_shape)
b2_scale_desc = TensorDescriptor.from_tensor(sfb2_tma, scale_block_shape)
# Launch kernel
fused_dual_gemm_silu_kernel[grid](
a_desc, a_scale_desc,
b1_desc, b1_scale_desc,
b2_desc, b2_scale_desc,
c_desc,
M=M, N=N, K=K,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
VEC_SIZE=VEC_SIZE,
rep_m=rep_m, rep_n=rep_n, rep_k=rep_k,
NUM_STAGES=NUM_STAGES, ELEM_PER_BYTE=ELEM_PER_BYTE,
)
else:
# Generic batch loop for L > 1
for l_idx in range(L):
a_u8 = a[:, :, l_idx].view(torch.uint8)
b1_u8 = b1[:, :, l_idx].view(torch.uint8)
b2_u8 = b2[:, :, l_idx].view(torch.uint8)
c_l = c[:, :, l_idx]
# Scale conversion
rest_m = sfa_perm.shape[2]
rest_k = sfa_perm.shape[4]
rest_n = sfb1_perm.shape[2]
sfa_tma = sfa_perm[:, :, :, :, :, l_idx].permute(2, 4, 0, 1, 3).reshape(
1, rest_m, rest_k, 2, 256
).contiguous()
sfb1_tma = sfb1_perm[:, :, :, :, :, l_idx].permute(2, 4, 0, 1, 3).reshape(
1, rest_n, rest_k, 2, 256
).contiguous()
sfb2_tma = sfb2_perm[:, :, :, :, :, l_idx].permute(2, 4, 0, 1, 3).reshape(
1, rest_n, rest_k, 2, 256
).contiguous()
a_desc = TensorDescriptor.from_tensor(a_u8, a_block_shape)
b1_desc = TensorDescriptor.from_tensor(b1_u8, b_block_shape)
b2_desc = TensorDescriptor.from_tensor(b2_u8, b_block_shape)
c_desc = TensorDescriptor.from_tensor(c_l, c_block_shape)
a_scale_desc = TensorDescriptor.from_tensor(sfa_tma, scale_block_shape)
b1_scale_desc = TensorDescriptor.from_tensor(sfb1_tma, scale_block_shape)
b2_scale_desc = TensorDescriptor.from_tensor(sfb2_tma, scale_block_shape)
fused_dual_gemm_silu_kernel[grid](
a_desc, a_scale_desc,
b1_desc, b1_scale_desc,
b2_desc, b2_scale_desc,
c_desc,
M=M, N=N, K=K,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
VEC_SIZE=VEC_SIZE,
rep_m=rep_m, rep_n=rep_n, rep_k=rep_k,
NUM_STAGES=NUM_STAGES, ELEM_PER_BYTE=ELEM_PER_BYTE,
)
return c
scrolls · 268 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