Skip to content
KernelIndex
Search⌘K

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
NVFP4 GEMMsuite of 3 cases
NVIDIA B200
17.2µs
#169 of 369
2025-12-07

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.

fp8output_dtype = tl.float8e4nv
split-kM, N, K, split_k = args["M"], args["N"], args["K"], args["SPLIT_K"]
stages = 4num_stages = 4
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_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_ref
scrolls · 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