submission 126350
nrehiew · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 124 lines, June 9 Researcher Reciprocity License v1.0.
submission_naive.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-126350?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:ed18112f91f0171c5cd41c3cffde3dada0ee907b6f9e6075a004e0b451ac9c30
license declaredunknown
license concludedunknown
authorsnrehiew
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4,stages = 3
num_stages=3,tile-k = 256
BLOCK_SIZE_K = 256tile-m = 128
BLOCK_SIZE_M = 128tile-n = 128
BLOCK_SIZE_N = 128Kernel source
submission_naive.py124 lines
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor
from task import input_t, output_t
sf_vec_size = 16
elements_per_byte = 2
ACC_DTYPE = tl.float32
@triton.jit
def triton_kernel(
a_desc,
b_desc,
sfa_desc,
sfb_desc,
c_ptr,
c_m_stride,
c_n_stride,
c_l_stride,
M,
K,
BLOCK_SIZE_K: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
rep_m: tl.constexpr,
rep_n: tl.constexpr,
rep_k: tl.constexpr,
elements_per_byte: tl.constexpr = elements_per_byte,
sf_vec_size: tl.constexpr = sf_vec_size,
):
pid = tl.program_id(0)
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
pid_m = pid % num_pid_m
pid_n = pid // num_pid_m
offs_am = pid_m * BLOCK_SIZE_M
offs_bn = pid_n * BLOCK_SIZE_N
offs_k_a = 0
offs_k_b = 0
offs_scale_m = pid_m * rep_m
offs_scale_n = pid_n * rep_n
offs_scale_k = 0
packed_k = (K + elements_per_byte - 1) // elements_per_byte
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=ACC_DTYPE)
for k_byte in tl.range(0, packed_k, BLOCK_SIZE_K // elements_per_byte, num_stages=num_stages, disallow_acc_multi_buffer=True, flatten=True):
a_val_uint8 = a_desc.load([0, offs_am, offs_k_a]).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // elements_per_byte)
b_val_uint8 = b_desc.load([0, offs_bn, offs_k_b]).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // elements_per_byte)
a_scale_val = sfa_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
b_scale_val = sfb_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
a_scale_val = a_scale_val.reshape(rep_m, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // sf_vec_size)
b_scale_val = b_scale_val.reshape(rep_n, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // sf_vec_size)
acc = tl.dot_scaled(a_val_uint8, a_scale_val, "e2m1", b_val_uint8.T, b_scale_val, "e2m1", acc)
offs_k_a += BLOCK_SIZE_K // elements_per_byte
offs_k_b += BLOCK_SIZE_K // elements_per_byte
offs_scale_k += rep_k
out_offsets = (offs_am + tl.arange(0, BLOCK_SIZE_M))[:, None] * c_m_stride + (offs_bn + tl.arange(0, BLOCK_SIZE_N))[None, :] * c_n_stride
tl.store(c_ptr + out_offsets, acc.to(tl.float16))
def custom_kernel(data: input_t) -> output_t:
# c: [m, n, l] is pre-allocated memory to avoid timing allocation overhead.
a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data
m, k_packed, l = a.shape
k = k_packed * elements_per_byte
m, n, l = c.shape
# a: [128, K, 1]
# b: [N, K, 1]
# sfa: [128, K // 16, 1]
# sfb: [N, K // 16, 1]
# c: [M, N, 1]
BLOCK_SIZE_M = 128
BLOCK_SIZE_K = 256
BLOCK_SIZE_N = 128
rep_m = BLOCK_SIZE_M // 128
rep_n = BLOCK_SIZE_N // 128
# rep_n = n // 128
rep_k = BLOCK_SIZE_K // sf_vec_size // 4
a = a.view(torch.uint8).permute(2, 0, 1) # [l, m, k]
b = b.view(torch.uint8).permute(2, 0, 1) # [l, n, k]
sfa_permuted_permute = sfa_permuted.permute(5, 2, 4, 0, 1, 3)
sfb_permuted_permute = sfb_permuted.permute(5, 2, 4, 0, 1, 3)
sfa_5d = sfa_permuted_permute.reshape(l, sfa_permuted_permute.shape[1], sfa_permuted_permute.shape[2], 2, 256)
sfb_5d = sfb_permuted_permute.reshape(l, sfb_permuted_permute.shape[1], sfb_permuted_permute.shape[2], 2, 256)
a_desc = TensorDescriptor.from_tensor(a, [1, BLOCK_SIZE_M, BLOCK_SIZE_K // elements_per_byte])
b_desc = TensorDescriptor.from_tensor(b, [1, BLOCK_SIZE_N, BLOCK_SIZE_K // elements_per_byte])
a_scale_desc = TensorDescriptor.from_tensor(sfa_5d, [1, rep_m, rep_k, 2, 256])
b_scale_desc = TensorDescriptor.from_tensor(sfb_5d, [1, rep_n, rep_k, 2, 256])
grid = lambda meta: (triton.cdiv(m, meta["BLOCK_SIZE_M"]) * triton.cdiv(n, meta["BLOCK_SIZE_N"]), l)
triton_kernel[grid](
a_desc,
b_desc,
a_scale_desc,
b_scale_desc,
c,
c.stride(0),
c.stride(1),
c.stride(2),
m,
k,
BLOCK_SIZE_K=BLOCK_SIZE_K,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
num_warps=4,
num_stages=3,
rep_m=rep_m,
rep_n=rep_n,
rep_k=rep_k,
)
return c
scrolls · 124 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