submission 113120
gilsaia · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 220 lines, June 9 Researcher Reciprocity License v1.0.
triton_4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-113120?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:1fa3f02468cfd6e228bbcbfa2b1b3311f645ccbd45a279338fb769a979fa5f81
license declaredunknown
license concludedunknown
authorsgilsaia
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
output_dtype = tl.float8e4nvtile-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_4.py220 lines
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 = args["M"], args["N"], args["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}]"
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, #
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, #
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)
num_pid_m = tl.cdiv(M, BLOCK_M)
pid_m = pid % num_pid_m
pid_n = pid // num_pid_m
offs_am = pid_m * BLOCK_M
offs_bn = pid_n * BLOCK_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
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)
for k in tl.range(0, tl.cdiv(K, BLOCK_K), 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
# 创建一个mask,只保留第一列
offs_n = tl.arange(0, BLOCK_N)
mask_n = offs_n == 0 # [BLOCK_N], 只有第一个是 True
# 使用mask提取第一列
# 方法: 将其他列置零,然后沿N维度求和
masked_acc = tl.where(mask_n[None, :], accumulator, 0.0) # [BLOCK_M, BLOCK_N]
result = tl.sum(masked_acc, axis=1) # [BLOCK_M]
result_r = result.reshape(1, BLOCK_M)
c_desc.store([lid,offs_am], result_r.to(output_dtype))
def block_scaled_matmul(a_desc, a_scale_desc, b_desc, b_scale_desc,c_out_desc, 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(output, [1,BLOCK_M, BLOCK_N])
grid = (L,triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N))
block_scaled_matmul_kernel[grid](
a_desc,
a_scale_desc,
b_desc,
b_scale_desc,
c_out_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"],
rep_m,
rep_n,
rep_k,
configs["num_stages"],
)
return
def ceil_div(a, b):
return (a + b - 1) // b
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 = 1 # GEMV operation: N dimension is 1/ with pad
n_padded_128 = 128
k = k_half * 2 # Actual k dimension (float4 packs 2 elements per byte)
# Call torch._scaled_mm to compute the GEMV result
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 256
ELEM_PER_BYTE_A = 2
ELEM_PER_BYTE_B = 2
VEC_SIZE=16
rep_m = BLOCK_M // 128
rep_n = BLOCK_N // 128
rep_k = BLOCK_K // VEC_SIZE // 4
configs = {
"BLOCK_SIZE_M": BLOCK_M,
"BLOCK_SIZE_N": BLOCK_N,
"BLOCK_SIZE_K": BLOCK_K,
"ELEM_PER_BYTE_A": ELEM_PER_BYTE_A,
"ELEM_PER_BYTE_B": ELEM_PER_BYTE_B,
"VEC_SIZE": 16,
"num_stages": 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)
# 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])
_,_,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])
c = c_ref.permute(2,0,1).reshape(l,m)
c_out_desc = TensorDescriptor.from_tensor(c,[1,BLOCK_M])
# (m, k) @ (n, k).T -> (m, n)
block_scaled_matmul(
a_desc,
a_scale_desc,
b_desc,
b_scale_desc,
c_out_desc,
torch.float16,
m,
n_padded_128,
k,
l,
rep_m,
rep_n,
rep_k,
configs
)
return c_refscrolls · 220 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