submission 317513
Julius T · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 131 lines, June 9 Researcher Reciprocity License v1.0.
solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-317513?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:db205ef4ef3b43e10e51fb663f16505b5cebca3238cde0f96279704de17264b4
license declaredunknown
license concludedunknown
authorsJulius T
imported2026-08-26
Kernel source
solution.py131 lines
"""
Block-Scaled Dual GEMM with SiLU Activation for NVIDIA B200
Operation: C = SiLU(A @ B1.T) * (A @ B2.T)
Pure PyTorch implementation (no Triton, no JIT).
"""
import torch
def ceil_div(a: int, b: int) -> int:
"""Ceiling division."""
return (a + b - 1) // b
def to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
"""
Convert scale factor tensor to blocked format for cuBLAS.
"""
rows = input_matrix.shape[0]
cols = input_matrix.shape[1]
n_row_blocks = (rows + 127) // 128
n_col_blocks = (cols + 3) // 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().contiguous()
def custom_kernel(data: tuple) -> torch.Tensor:
"""
Dual GEMM with SiLU activation.
C = SiLU(A @ B1.T) * (A @ B2.T)
"""
# Unpack inputs
a, b1, b2, sfa_cpu, sfb1_cpu, sfb2_cpu, _, _, _, c = data
# Get dimensions
m, n, l = c.shape
device = a.device
sfa = sfa_cpu.to(device=device, non_blocking=True) if sfa_cpu.device != device else sfa_cpu
sfb1 = sfb1_cpu.to(device=device, non_blocking=True) if sfb1_cpu.device != device else sfb1_cpu
sfb2 = sfb2_cpu.to(device=device, non_blocking=True) if sfb2_cpu.device != device else sfb2_cpu
# Single-batch path (L=1)
if l == 1:
# Convert scale factors to blocked format
scale_a = to_blocked(sfa[:, :, 0])
scale_b1 = to_blocked(sfb1[:, :, 0])
scale_b2 = to_blocked(sfb2[:, :, 0])
b1_t = b1[:, :, 0].t()
b2_t = b2[:, :, 0].t()
# First GEMM: A @ B1.T
gemm1 = torch._scaled_mm(
a[:, :, 0], b1_t,
scale_a, scale_b1,
bias=None, out_dtype=torch.float32,
)
# Second GEMM: A @ B2.T
gemm2 = torch._scaled_mm(
a[:, :, 0], b2_t,
scale_a, scale_b2,
bias=None, out_dtype=torch.float32,
)
# Fused SiLU and multiply using PyTorch
# SiLU(x) = x * sigmoid(x)
result = (torch.nn.functional.silu(gemm1) * gemm2).to(torch.float16)
# Reshape to M x N x 1
return result.unsqueeze(2)
# Multi-batch case (L > 1)
gemm1_results = []
gemm2_results = []
for l_idx in range(l):
# Convert scale factors to blocked format
scale_a = to_blocked(sfa[:, :, l_idx])
scale_b1 = to_blocked(sfb1[:, :, l_idx])
scale_b2 = to_blocked(sfb2[:, :, l_idx])
b1_t = b1[:, :, l_idx].t()
b2_t = b2[:, :, l_idx].t()
# First GEMM: A @ B1.T
gemm1 = torch._scaled_mm(
a[:, :, l_idx], b1_t,
scale_a, scale_b1,
bias=None, out_dtype=torch.float32,
)
gemm1_results.append(gemm1)
# Second GEMM: A @ B2.T
gemm2 = torch._scaled_mm(
a[:, :, l_idx], b2_t,
scale_a, scale_b2,
bias=None, out_dtype=torch.float32,
)
gemm2_results.append(gemm2)
# Stack results: list of (M, N) -> (L, M, N) -> (M, N, L)
gemm1_stacked = torch.stack(gemm1_results, dim=0).permute(1, 2, 0)
gemm2_stacked = torch.stack(gemm2_results, dim=0).permute(1, 2, 0)
# Fused SiLU and multiply
result = (torch.nn.functional.silu(gemm1_stacked) * gemm2_stacked).to(torch.float16)
return result
def kernel(data: tuple) -> torch.Tensor:
"""
Main entry point for the kernel.
Args:
data: Tuple of (a, b1, b2, sfa, sfb1, sfb2, sfa_perm, sfb1_perm, sfb2_perm, c)
Returns:
Output tensor C = SiLU(A @ B1.T) * (A @ B2.T) in float16
"""
return custom_kernel(data)
scrolls · 131 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