submission 127650
gilsaia · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 390 lines, June 9 Researcher Reciprocity License v1.0.
triton_5.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-127650?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:dd8f3df409df9bcb4004a7668e1fb8d9100d2851d6fcf7155cd094def2ff9a2a
license declaredunknown
license concludedunknown
authorsgilsaia
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
output_dtype = tl.float8e4nvsplit-k
M, N, K, split_k = args["M"], args["N"], args["K"], args["SPLIT_K"]stages = 4
num_stages = 4tile-k = 256
BLOCK_K = 256tile-m = 128
BLOCK_M = 128tile-n = 128
BLOCK_N = 128vector-width = float4
k = k_half * 2 # Actual k dimension (float4 packs 2 elements per byte)Kernel source
triton_5.py390 lines
from numpy import c_
from task import input_t,output_t
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor
def _matmul_launch_metadata(grid, kernel, args):
ret = {}
M, N, K, split_k = args["M"], args["N"], args["K"], args["SPLIT_K"]
kernel_name = kernel.name
if "ELEM_PER_BYTE_A" and "ELEM_PER_BYTE_B" and "VEC_SIZE" in args:
if args["ELEM_PER_BYTE_A"] == 1 and args["ELEM_PER_BYTE_B"] == 1:
kernel_name += "_mxfp8"
elif args["ELEM_PER_BYTE_A"] == 1 and args["ELEM_PER_BYTE_B"] == 2:
kernel_name += "_mixed"
elif args["ELEM_PER_BYTE_A"] == 2 and args["ELEM_PER_BYTE_B"] == 2:
if args["VEC_SIZE"] == 16:
kernel_name += "_nvfp4"
elif args["VEC_SIZE"] == 32:
kernel_name += "_mxfp4"
ret["name"] = f"{kernel_name} [M={M}, N={N}, K={K}, SPLIT_K={split_k}]"
ret["flops"] = 2.0 * M * N * K
return ret
@triton.jit(launch_metadata=_matmul_launch_metadata)
def block_scaled_matmul_kernel( #
a_desc, #
a_scale_desc, #
b_desc, #
b_scale_desc, #
c_desc, #
c_ptr,
M: tl.constexpr, #
N: tl.constexpr, #
K: tl.constexpr, #
L: tl.constexpr,
output_type: tl.constexpr, #
ELEM_PER_BYTE_A: tl.constexpr, #
ELEM_PER_BYTE_B: tl.constexpr, #
VEC_SIZE: tl.constexpr, #
BLOCK_M: tl.constexpr, #
BLOCK_N: tl.constexpr, #
BLOCK_K: tl.constexpr, #
SPLIT_K: tl.constexpr, #
rep_m: tl.constexpr, #
rep_n: tl.constexpr, #
rep_k: tl.constexpr, #
NUM_STAGES: tl.constexpr, #
): #
if output_type == 0:
output_dtype = tl.float32
elif output_type == 1:
output_dtype = tl.float16
elif output_type == 2:
output_dtype = tl.float8e4nv
lid = tl.program_id(axis=0)
pid = tl.program_id(axis=1)
# 解析 pid: 现在包含 split-K 维度
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
num_pid_mn = num_pid_m * num_pid_n
# Split-K: 提取 K 分片索引
pid_k = pid // num_pid_mn # K 方向的分片索引 [0, SPLIT_K)
pid_mn = pid % num_pid_mn # M-N 平面的索引
pid_m = pid_mn % num_pid_m
pid_n = pid_mn // num_pid_m
offs_am = pid_m * BLOCK_M
offs_bn = pid_n * BLOCK_N
# ========== 修复 1: 正确计算 K 维度分片边界 ==========
total_k_tiles = tl.cdiv(K, BLOCK_K)
# 计算每个 split 的 tile 范围(确保不重叠不遗漏)
k_per_split = total_k_tiles // SPLIT_K # 整除部分
k_remainder = total_k_tiles % SPLIT_K # 余数
# 前 k_remainder 个 split 多处理 1 个 tile
if pid_k < k_remainder:
k_start_tile = pid_k * (k_per_split + 1)
k_end_tile = k_start_tile + (k_per_split + 1)
else:
k_start_tile = k_remainder * (k_per_split + 1) + (pid_k - k_remainder) * k_per_split
k_end_tile = k_start_tile + k_per_split
# ========== 修复 2: 初始偏移 ==========
offs_k_a = k_start_tile * (BLOCK_K // ELEM_PER_BYTE_A)
offs_k_b = k_start_tile * (BLOCK_K // ELEM_PER_BYTE_B)
offs_scale_m = pid_m * rep_m
offs_scale_n = pid_n * rep_n
offs_scale_k = k_start_tile * rep_k
MIXED_PREC: tl.constexpr = ELEM_PER_BYTE_A == 1 and ELEM_PER_BYTE_B == 2
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
num_k_tiles = k_end_tile - k_start_tile
for k in tl.range(0, num_k_tiles, num_stages=NUM_STAGES):
a = a_desc.load([lid,offs_am, offs_k_a]).reshape(BLOCK_M,BLOCK_K//ELEM_PER_BYTE_A)
b = b_desc.load([lid,offs_bn, offs_k_b]).reshape(BLOCK_N,BLOCK_K//ELEM_PER_BYTE_B)
scale_a = a_scale_desc.load([lid, offs_scale_m, offs_scale_k, 0, 0])
scale_b = b_scale_desc.load([lid, offs_scale_n, offs_scale_k, 0, 0])
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_b = scale_b.reshape(rep_n, rep_k, 32, 4, 4).trans(0, 3, 2, 1, 4).reshape(BLOCK_N, BLOCK_K // VEC_SIZE)
if MIXED_PREC:
accumulator = tl.dot_scaled(a, scale_a, "e4m3", b.T, scale_b, "e2m1", accumulator)
elif ELEM_PER_BYTE_A == 2 and ELEM_PER_BYTE_B == 2:
accumulator = tl.dot_scaled(a, scale_a, "e2m1", b.T, scale_b, "e2m1", accumulator)
else:
accumulator = tl.dot_scaled(a, scale_a, "e4m3", b.T, scale_b, "e4m3", accumulator)
offs_k_a += BLOCK_K // ELEM_PER_BYTE_A
offs_k_b += BLOCK_K // ELEM_PER_BYTE_B
offs_scale_k += rep_k
if SPLIT_K == 1:
# 无 split-K: 直接写入
accumulator_reshaped = accumulator.reshape(1, BLOCK_M, BLOCK_N)
c_desc.store([lid, offs_am, offs_bn], accumulator_reshaped.to(output_dtype))
else:
accumulator_reshaped = accumulator.reshape(1, 1, BLOCK_M, BLOCK_N)
c_ptr.store([lid, pid_k, offs_am, offs_bn],accumulator_reshaped)
def _reduction_launch_metadata(grid, kernel, args):
ret = {}
M, N, split_k = args["M"], args["N"], args["SPLIT_K"]
kernel_name = kernel.name
ret["name"] = f"{kernel_name} [M={M}, N={N}, K={K}, SPLIT_K={split_k}]"
ret["flops"] = M * N * split_k
return ret
@triton.jit
def split_k_reduce_kernel(
acc_desc, # TensorDescriptor for [L, SPLIT_K, M, N]
out_desc, # TensorDescriptor for [L, M, N]
M: tl.constexpr,
N: tl.constexpr,
SPLIT_K: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""
极致优化的 split-K reduction
- SPLIT_K 必须是 2 的幂次 (1, 2, 4, 8)
- M % BLOCK_M == 0, N % BLOCK_N == 0
- 使用 TensorDescriptor 优化内存访问
- 树形归约最小化指令延迟
"""
lid = tl.program_id(0)
pid_mn = tl.program_id(1)
num_pid_m = tl.cdiv(M, BLOCK_M)
pid_m = pid_mn % num_pid_m
pid_n = pid_mn // num_pid_m
offs_m = pid_m * BLOCK_M
offs_n = pid_n * BLOCK_N
# ========== 编译期展开 + 树形归约 ==========
result = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in tl.static_range(SPLIT_K):
acc = acc_desc.load([lid, k, offs_m, offs_n])
result += acc
# 存储结果
result_reshaped = result.reshape(1, BLOCK_M, BLOCK_N)
out_desc.store([lid, offs_m, offs_n], result_reshaped.to(tl.float16))
def block_scaled_matmul(a_desc, a_scale_desc, b_desc, b_scale_desc,c,c_acc, dtype_dst, M, N, K,L, rep_m, rep_n, rep_k, configs):
# output = torch.empty((L,M, N), dtype=dtype_dst, device="cuda")
if dtype_dst == torch.float32:
dtype_dst = 0
elif dtype_dst == torch.float16:
dtype_dst = 1
elif dtype_dst == torch.float8_e4m3fn:
dtype_dst = 2
else:
raise ValueError(f"Unsupported dtype: {dtype_dst}")
BLOCK_M = configs["BLOCK_SIZE_M"]
BLOCK_N = configs["BLOCK_SIZE_N"]
c_desc = TensorDescriptor.from_tensor(c,[1,BLOCK_M,BLOCK_N])
grid = (L,triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N)*configs["SPLIT_K"])
if configs["SPLIT_K"] > 1:
c_acc_desc = TensorDescriptor.from_tensor(c_acc,block_shape=[1,1,BLOCK_M,BLOCK_N])
else:
c_acc_desc = c_desc
block_scaled_matmul_kernel[grid](
a_desc,
a_scale_desc,
b_desc,
b_scale_desc,
c_desc,
c_acc_desc,
M,
N,
K,
L,
dtype_dst,
configs["ELEM_PER_BYTE_A"],
configs["ELEM_PER_BYTE_B"],
configs["VEC_SIZE"],
configs["BLOCK_SIZE_M"],
configs["BLOCK_SIZE_N"],
configs["BLOCK_SIZE_K"],
configs["SPLIT_K"],
rep_m,
rep_n,
rep_k,
configs["num_stages"],
)
if configs["SPLIT_K"] > 1:
REDUCR_BLOCK_M = configs["REDUCR_BLOCK_M"]
REDUCR_BLOCK_N = configs["REDUCR_BLOCK_N"]
c_reduc_desc = TensorDescriptor.from_tensor(c,[1,REDUCR_BLOCK_M,REDUCR_BLOCK_N])
c_acc_reducr_desc = TensorDescriptor.from_tensor(c_acc,[1,1,REDUCR_BLOCK_M,REDUCR_BLOCK_N])
reducr_grid = (L,triton.cdiv(M, REDUCR_BLOCK_M) * triton.cdiv(N, REDUCR_BLOCK_N))
split_k_reduce_kernel[reducr_grid](
c_acc_reducr_desc,
c_reduc_desc,
M,
N,
configs["SPLIT_K"],
REDUCR_BLOCK_M,
REDUCR_BLOCK_N,
)
return
# Helper function for ceiling division
def ceil_div(a, b):
return (a + b - 1) // b
def get_block_config(M, N, K, L):
"""根据矩阵维度返回 block 配置"""
# BLOCK_M/N: 128 或 256 (必须是128的倍数)
BLOCK_M = 128
BLOCK_N = 128
# BLOCK_K: 256 或 512 (必须是64的倍数)
BLOCK_K = 256
# num_stages: 2-4
num_stages = 4
m_block_num = ceil_div(M, BLOCK_M)
n_block_num = ceil_div(N, BLOCK_N)
split_k_num = ceil_div(148,m_block_num*n_block_num*L)
if split_k_num <=1:
split_k_num = 1
if split_k_num>8:
split_k_num = 8
block_max_split = ceil_div(K,BLOCK_K)
split_k_num = min(split_k_num,block_max_split)
if split_k_num >= 8:
split_k_num = 8
elif split_k_num >= 4:
split_k_num = 4
elif split_k_num >= 2:
split_k_num = 2
else:
split_k_num = 1
REDUCR_BLOCK_M = 1
REDUCR_BLOCK_N = BLOCK_N
while REDUCR_BLOCK_N*2 < N:
REDUCR_BLOCK_N *= 2
return {
"BLOCK_SIZE_M": BLOCK_M,
"BLOCK_SIZE_N": BLOCK_N,
"BLOCK_SIZE_K": BLOCK_K,
"REDUCR_BLOCK_M": REDUCR_BLOCK_M,
"REDUCR_BLOCK_N": REDUCR_BLOCK_N,
"SPLIT_K": split_k_num,
"ELEM_PER_BYTE_A": 2,
"ELEM_PER_BYTE_B": 2,
"VEC_SIZE": 16,
"num_stages": num_stages,
}
def custom_kernel(data:input_t)->output_t:
a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, sfa_permuted, sfb_permuted, c_ref = data
# Get dimensions from MxNxL layout
_, _, l = c_ref.shape
# Get dimensions from the tensors
m, k_half, l = a_ref.shape # a_ref is [m, k//2, l] in float4_e2m1fn_x2
n,_,_ = b_ref.shape
k = k_half * 2 # Actual k dimension (float4 packs 2 elements per byte)
# Call torch._scaled_mm to compute the GEMM result
configs = get_block_config(m, n, k,l)
BLOCK_M = configs["BLOCK_SIZE_M"]
BLOCK_N = configs["BLOCK_SIZE_N"]
BLOCK_K = configs["BLOCK_SIZE_K"]
SPLIT_K = configs["SPLIT_K"]
ELEM_PER_BYTE_A = configs["ELEM_PER_BYTE_A"]
ELEM_PER_BYTE_B = configs["ELEM_PER_BYTE_B"]
VEC_SIZE=configs["VEC_SIZE"]
rep_m = BLOCK_M // 128
rep_n = BLOCK_N // 128
rep_k = BLOCK_K // VEC_SIZE // 4
a_per = a_ref.permute(2,0,1)
b_per = b_ref.permute(2,0,1)
a = a_per.view(torch.uint8)
b = b_per.view(torch.uint8)
c = c_ref.permute(2,0,1)
m_row = ceil_div(m,BLOCK_M)
n_row = ceil_div(n,BLOCK_N)
# Convert the scale factor tensor to blocked format
a_desc = TensorDescriptor.from_tensor(a,[1,BLOCK_M,BLOCK_K // ELEM_PER_BYTE_A])
b_desc = TensorDescriptor.from_tensor(b,[1,BLOCK_N,BLOCK_K // ELEM_PER_BYTE_B])
# c_desc = TensorDescriptor.from_tensor(c,[1,BLOCK_M,BLOCK_N])
# sfa_per_cpu = sfa_ref_cpu.permute(2,0,1)
# sfb_per_cpu = sfb_ref_cpu.permute(2,0,1)
# scale_a = to_blocked(sfa_per_cpu)
# scale_b = to_blocked(sfb_per_cpu)
_,_,m_row,_,k_row,_ = sfa_permuted.shape
_,_,n_row,_,_,_ = sfb_permuted.shape
sfa_per = sfa_permuted.permute(5,2,4,0,1,3).reshape(l,m_row,k_row,2,256)
sfb_per = sfb_permuted.permute(5,2,4,0,1,3).reshape(l,n_row,k_row,2,256)
a_scale_desc = TensorDescriptor.from_tensor(sfa_per,block_shape=[1,rep_m,rep_k,2,256])
b_scale_desc = TensorDescriptor.from_tensor(sfb_per,block_shape=[1,rep_n,rep_k,2,256])
# (m, k) @ (n, k).T -> (m, n)
if SPLIT_K > 1:
# float32 buffer 用于 split-K 累加
c_acc = torch.empty((l, SPLIT_K,m,n), dtype=torch.float32, device="cuda")
block_scaled_matmul(
a_desc, a_scale_desc, b_desc, b_scale_desc,
c, # 不用于 SPLIT_K>1
c_acc, # atomic_add 目标
torch.float16, # 累加精度
m, n, k, l, rep_m, rep_n, rep_k, configs
)
else:
block_scaled_matmul(
a_desc,
a_scale_desc,
b_desc,
b_scale_desc,
c,
c,
torch.float16,
m,
n,
k,
l,
rep_m,
rep_n,
rep_k,
configs
)
return c_refscrolls · 390 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