Skip to content
KernelIndex
Search⌘K

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
NVFP4 GEMVsuite of 3 cases
NVIDIA B200
32.2µs
#171 of 678
2025-11-29

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.

fp8output_dtype = tl.float8e4nv
tile-k = 256BLOCK_K = 256
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128
vector-width = float4k = 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_ref
scrolls · 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