submission 102999
nrehiew · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 359 lines, June 9 Researcher Reciprocity License v1.0.
bestv2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-102999?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:c33d1e04a75532f7f511695f9dd26a2cb7e2d538e3b30cff9474207be839a8ad
license declaredunknown
license concludedunknown
authorsnrehiew
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
logical_k_size_scale = (K + sf_vec_size - 1) // sf_vec_size # sf_vec_size-wide FP4 groupswarp-specialization
warp_specialize: tl.constexpr,Kernel source
bestv2.py359 lines
import os
import torch
import triton
import triton.language as tl
from typing import Optional
from task import input_t, output_t
from reference import generate_input
from triton.tools.tensor_descriptor import TensorDescriptor
# print(triton.__version__) # 3.5.0
# os.environ["TRITON_PRINT_AUTOTUNING"] = "1"
sf_vec_size = 16
elements_per_byte = 2
ACC_DTYPE = tl.float32
CONFIGS = {
(7168, 16384, 1): {
"block_size_k": 4096,
"num_stages": 3,
"num_warps": 2,
"block_size_m": 2,
# "block_size_n": 64,what about
},
# (4096, 7168, 8): {
# "block_size_k": 512,
# "num_stages": 4,
# "num_warps": 1,
# "block_size_m": 4,
# },
# (7168, 2048, 4): {
# "block_size_k": 256,
# "num_stages": 4,
# "num_warps": 2,
# "block_size_m": 8,
# },
(4096, 7168, 8): {
"block_size_k": 128,
"num_stages": 4,
"num_warps": 4,
"block_size_m": 128,
"block_size_n": 64,
},
(7168, 2048, 4): {
"block_size_k": 128,
"num_stages": 4,
"num_warps": 4,
"block_size_m": 128,
"block_size_n": 64,
},
}
def get_config(m, k, l):
key = (m, k, l)
if key in CONFIGS:
return CONFIGS[key]
else:
return {
"block_size_k": 128,
"num_stages": 4,
"num_warps": 4,
"block_size_m": 128,
"block_size_n": 64,
} # need 128 because there is a really small sized test case
@triton.jit
def triton_kernel_tma(
a_desc,
b_desc,
out_desc,
sfa_desc,
sfb_desc,
M,
K,
BLOCK_SIZE_K: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K_SCALE: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
elements_per_byte: tl.constexpr = elements_per_byte,
sf_vec_size: tl.constexpr = sf_vec_size,
):
m_pid = tl.program_id(0)
l_pid = tl.program_id(1)
row_start = m_pid * BLOCK_SIZE_M
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, num_stages=num_stages):
a_val_uint8 = a_desc.load([l_pid, row_start, k_byte])
b_val_uint8 = b_desc.load([l_pid, 0, k_byte])
scale_offset = (k_byte * elements_per_byte) // sf_vec_size
a_scale_val = sfa_desc.load([l_pid, row_start, scale_offset])
b_scale_val = sfb_desc.load([l_pid, 0, scale_offset])
a_val_uint8 = tl.reshape(a_val_uint8, (BLOCK_SIZE_M, BLOCK_SIZE_K))
b_val_uint8 = tl.reshape(b_val_uint8, (BLOCK_SIZE_N, BLOCK_SIZE_K))
a_scale_val = tl.reshape(a_scale_val, (BLOCK_SIZE_M, BLOCK_SIZE_K_SCALE))
b_scale_val = tl.reshape(b_scale_val, (BLOCK_SIZE_N, BLOCK_SIZE_K_SCALE))
acc = tl.dot_scaled(a_val_uint8, a_scale_val, "e2m1", b_val_uint8.T, b_scale_val, "e2m1", acc)
out_desc.store([l_pid, row_start, 0], tl.reshape(acc, (1, BLOCK_SIZE_M, BLOCK_SIZE_N)))
@triton.jit
def triton_kernel_naive(
a_ptr,
b_ptr,
out_ptr,
M,
K,
L,
am_stride,
ak_stride,
al_stride,
bm_stride,
bk_stride,
bl_stride,
outm_stride,
outk_stride,
outl_stride, # outk_stride = 1
sfa_ptr,
sfb_ptr,
sfa_m_stride,
sfa_k_stride,
sfa_l_stride, # M x (K // sf_vec_size) x L
sfb_m_stride,
sfb_k_stride,
sfb_l_stride, # 1 x (K // sf_vec_size) x L
BLOCK_SIZE_K: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
warp_specialize: tl.constexpr,
elements_per_byte: tl.constexpr = elements_per_byte,
sf_vec_size: tl.constexpr = sf_vec_size,
):
# each iteration we want to load BLOCK_SIZE_K bytes of a and b = 2 * BLOCK_SIZE_K FP4 values for a and b each
# with 2 * BLOCK_SIZE_K FP4 values we need (2 * BLOCK_SIZE_K) // sf_vec_size scale factors (FP8) each load
m_pid = tl.program_id(0)
l_pid = tl.program_id(1)
row_start = m_pid * BLOCK_SIZE_M
row_offsets = row_start + tl.arange(0, BLOCK_SIZE_M)
a_row_ptr = a_ptr + row_offsets * am_stride + l_pid * al_stride
b_starting_ptr = b_ptr + bl_stride * l_pid
a_scale_row_ptr = sfa_ptr + row_offsets * sfa_m_stride + l_pid * sfa_l_stride
b_scale_starting_ptr = sfb_ptr + sfb_l_stride * l_pid
out_row_ptr = out_ptr + row_offsets * outm_stride + l_pid * outl_stride
logical_k_size_scale = (K + sf_vec_size - 1) // sf_vec_size # sf_vec_size-wide FP4 groups
NUM_LOGICAL_ITEMS_K_PER_ITERATION: tl.constexpr = BLOCK_SIZE_K * elements_per_byte
BLOCK_SIZE_K_SCALE: tl.constexpr = (NUM_LOGICAL_ITEMS_K_PER_ITERATION + sf_vec_size - 1) // sf_vec_size
scale_offsets = tl.arange(0, BLOCK_SIZE_K_SCALE)
val_offsets = tl.arange(0, BLOCK_SIZE_K)
val_idx = 0
acc_vec = tl.zeros((BLOCK_SIZE_M,), dtype=ACC_DTYPE)
ONES_MASK = 0xFFFF
for scale_idx in tl.range(0, logical_k_size_scale, BLOCK_SIZE_K_SCALE, num_stages=num_stages, warp_specialize=warp_specialize):
# each loop should process uint8 values so 2 values per iteration
scale_idx_offsets = scale_offsets + scale_idx
a_scale_ptr_curr = a_scale_row_ptr[:, None] + scale_idx_offsets[None, :] * sfa_k_stride
b_scale_ptr_curr = b_scale_starting_ptr + scale_idx_offsets * sfb_k_stride
# Load scales as f16 for SIMD optimization
a_scale_val = tl.load(a_scale_ptr_curr, cache_modifier=".cv", eviction_policy="evict_first").to(tl.float16)
b_scale_val = tl.load(b_scale_ptr_curr, cache_modifier=".cg").to(tl.float16)
val_idx_offsets = val_offsets + val_idx
a_ptr_curr = a_row_ptr[:, None] + val_idx_offsets[None, :] * ak_stride
b_ptr_curr = b_starting_ptr + val_idx_offsets * bk_stride
a_val_uint8 = tl.load(a_ptr_curr, cache_modifier=".cv", eviction_policy="evict_first") # [BLOCK_SIZE_M, BLOCK_SIZE_K]
b_val_uint8 = tl.load(b_ptr_curr, cache_modifier=".cg") # [BLOCK_SIZE_K]
# Source: https://github.com/triton-lang/triton/blob/main/python/triton_kernels/triton_kernels/numerics_details/mxfp_details/_upcast_from_mxfp.py
# Convert FP4 to f16x2 packed format (2 f16 values per u32)
a_packed_u32 = tl.inline_asm_elementwise(
asm="""
{
.reg .b8 in_8;
.reg .f16x2 out;
cvt.u8.u32 in_8, $1;
cvt.rn.f16x2.e2m1x2 out, in_8;
mov.b32 $0, out;
}
""",
constraints="=r,r",
args=[a_val_uint8],
dtype=tl.uint32,
is_pure=True,
pack=1,
)
b_packed_u32 = tl.inline_asm_elementwise(
asm="""
{
.reg .b8 in_8;
.reg .f16x2 out;
cvt.u8.u32 in_8, $1;
cvt.rn.f16x2.e2m1x2 out, in_8;
mov.b32 $0, out;
}
""",
constraints="=r,r",
args=[b_val_uint8],
dtype=tl.uint32,
is_pure=True,
pack=1,
)
# Unpack f16x2 to individual f16 values
a_lo_f16 = (a_packed_u32 & ONES_MASK).to(tl.uint16).to(tl.float16, bitcast=True)
a_hi_f16 = (a_packed_u32 >> 16).to(tl.uint16).to(tl.float16, bitcast=True)
b_lo_f16 = (b_packed_u32 & ONES_MASK).to(tl.uint16).to(tl.float16, bitcast=True)
b_hi_f16 = (b_packed_u32 >> 16).to(tl.uint16).to(tl.float16, bitcast=True)
# Interleave to get [M, K*2] for both a and b
a_val_f16 = tl.interleave(a_lo_f16, a_hi_f16) # [BLOCK_SIZE_M, BLOCK_SIZE_K * 2]
b_val_f16 = tl.interleave(b_lo_f16, b_hi_f16) # [BLOCK_SIZE_K * 2]
b_val_f16 = tl.broadcast_to(b_val_f16[None, :], (BLOCK_SIZE_M, NUM_LOGICAL_ITEMS_K_PER_ITERATION))
# Broadcast scales to match [M, K*2]
a_scale_broadcast = tl.broadcast_to(a_scale_val[:, :, None], (BLOCK_SIZE_M, BLOCK_SIZE_K_SCALE, sf_vec_size))
b_scale_broadcast = tl.broadcast_to(b_scale_val[None, :, None], (BLOCK_SIZE_M, BLOCK_SIZE_K_SCALE, sf_vec_size))
a_scale_f16 = a_scale_broadcast.reshape((BLOCK_SIZE_M, NUM_LOGICAL_ITEMS_K_PER_ITERATION))
b_scale_f16 = b_scale_broadcast.reshape((BLOCK_SIZE_M, NUM_LOGICAL_ITEMS_K_PER_ITERATION))
# # Use FMA in f16 for the multiply chain, convert to f32 for accumulation
# # FMA: result = ((a * b) * scale_a) * scale_b
temp1 = a_val_f16 * b_val_f16
temp2 = temp1 * a_scale_f16
result_f16 = temp2 * b_scale_f16
acc_vec += tl.sum(result_f16, axis=1).to(tl.float32)
val_idx += BLOCK_SIZE_K # in bytes
# # acc = tl.sum(acc_vec, axis=1)
# # tl.store(out_row_ptr, acc.to(tl.float16))
tl.store(out_row_ptr, acc_vec.to(tl.float16))
def custom_kernel(
data: input_t,
) -> output_t:
a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_ref = data
m, k_packed, l = a_ref.shape
k = k_packed * elements_per_byte
_, _, l = c_ref.shape
a_ref = a_ref.view(torch.uint8)
b_ref = b_ref.view(torch.uint8)
config = get_config(m, k, l)
if l == 1:
grid = lambda meta: (triton.cdiv(m, meta["BLOCK_SIZE_M"]), l)
triton_kernel_naive[grid](
a_ref,
b_ref,
c_ref,
m,
k_packed * elements_per_byte,
l,
a_ref.stride(0),
a_ref.stride(1),
a_ref.stride(2),
b_ref.stride(0),
b_ref.stride(1),
b_ref.stride(2),
c_ref.stride(0),
c_ref.stride(1),
c_ref.stride(2),
sfa_ref,
sfb_ref,
sfa_ref.stride(0),
sfa_ref.stride(1),
sfa_ref.stride(2),
sfb_ref.stride(0),
sfb_ref.stride(1),
sfb_ref.stride(2),
BLOCK_SIZE_K=config["block_size_k"],
BLOCK_SIZE_M=config["block_size_m"],
num_warps=config["num_warps"],
num_stages=config["num_stages"],
warp_specialize=False,
)
return c_ref
else:
block_m = config["block_size_m"]
block_k = config["block_size_k"]
block_n = config["block_size_n"]
block_k_scale = (block_k * elements_per_byte + sf_vec_size - 1) // sf_vec_size
a_uint8 = a_ref.view(torch.uint8).permute(2, 0, 1)
b_uint8 = b_ref.view(torch.uint8).permute(2, 0, 1)
sfa_perm = sfa_ref.permute(2, 0, 1)
sfb_perm = sfb_ref.permute(2, 0, 1)
a_desc = TensorDescriptor.from_tensor(a_uint8, [1, block_m, block_k])
b_desc = TensorDescriptor.from_tensor(b_uint8, [1, block_n, block_k])
sfa_desc = TensorDescriptor.from_tensor(sfa_perm, [1, block_m, block_k_scale])
sfb_desc = TensorDescriptor.from_tensor(sfb_perm, [1, block_n, block_k_scale])
c_buf = torch.empty((l, m, block_n), dtype=torch.float16, device="cuda")
out_desc = TensorDescriptor.from_tensor(c_buf, [1, block_m, block_n])
grid = (triton.cdiv(m, block_m), l)
triton_kernel_tma[grid](
a_desc,
b_desc,
out_desc,
sfa_desc,
sfb_desc,
m,
k,
BLOCK_SIZE_K=block_k,
BLOCK_SIZE_M=block_m,
BLOCK_SIZE_N=block_n,
BLOCK_SIZE_K_SCALE=block_k_scale,
num_warps=config["num_warps"],
num_stages=config["num_stages"],
)
out_first_col = c_buf[:, :, 0].permute(1, 0)
return out_first_col.unsqueeze(1)
shapes = [
{"m": 7168, "k": 16384, "l": 1, "seed": 1111},
{"m": 4096, "k": 7168, "l": 8, "seed": 1111},
{"m": 7168, "k": 2048, "l": 4, "seed": 1111},
# {"m": 128, "k": 256, "l": 1, "seed": 1111},
# {"m": 128, "k": 1536, "l": 1, "seed": 1111},
# {"m": 128, "k": 3072, "l": 1, "seed": 1111},
# {"m": 256, "k": 7168, "l": 1, "seed": 1111},
# {"m": 256, "k": 7168, "l": 1, "seed": 1111},
# {"m": 2432, "k": 4608, "l": 2, "seed": 1111},
# {"m": 512, "k": 1536, "l": 2, "seed": 1111},
]
for shape in shapes:
key_ = (shape["m"], shape["k"], shape["l"])
data_ = generate_input(**shape)
for _ in range(5):
out = custom_kernel(data_)
torch.cuda.synchronize()
torch.cuda.empty_cache()
scrolls · 359 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