Skip to content
KernelIndex
Search⌘K

submission 115186

sahanp · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 161 lines, June 9 Researcher Reciprocity License v1.0.

gemv.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-115186?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
27.9µs
#134 of 678
2025-11-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f53974eafc6054077ed5ffa6d76758df55f8bb3949d85176bfbd45928dd5f591
license declaredunknown
license concludedunknown
authorssahanp
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

autotunetriton.Config(
split-kdef batched_block_scaled_matmul_kernel_splitk(
tile-k = 256BLOCK_K = 256
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128

Kernel source

gemv.py161 lines
import torch
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor

# Fixed constants
NVFP4_VEC_SIZE = 16
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 256
ELEM_PER_BYTE = 2
ROWS_PER_SCALE_CHUNK = 128
K_PACKED = BLOCK_K // ELEM_PER_BYTE


def get_tma_configs():
    configs = []
    for num_stages in [2, 3, 4, 5]:
        for num_warps in [4, 8]:
            configs.append(
                triton.Config(
                    {'NUM_STAGES': num_stages},
                    num_warps=num_warps,
                    num_stages=num_stages,
                )
            )
    return configs


@triton.autotune(
    configs=get_tma_configs(),
    key=['M_per_L', 'K'],
)
@triton.jit
def batched_block_scaled_matmul_kernel_splitk(
    a_desc, a_scale_desc, b_desc, b_scale_desc,
    c_ptr,
    stride_c_m, stride_c_l,
    M_per_L: tl.constexpr, K: tl.constexpr,
    D_chunks_per_L: 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,
    SPLIT_K: tl.constexpr,
    num_m_blocks_per_L: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    
    pid_k = pid % SPLIT_K
    pid_tile = pid // SPLIT_K
    
    pid_m = pid_tile % num_m_blocks_per_L
    pid_l = pid_tile // num_m_blocks_per_L
    
    offs_am = pid_l * M_per_L + pid_m * BLOCK_M
    offs_bn = pid_l * BLOCK_N
    offs_scale_m = pid_l * D_chunks_per_L + pid_m * rep_m
    offs_scale_n = pid_l * rep_n

    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    K_STEP: tl.constexpr = BLOCK_K // 2
    
    k_iters_total = tl.cdiv(K, BLOCK_K)
    k_iters_per_split = tl.cdiv(k_iters_total, SPLIT_K)
    k_iter_start = pid_k * k_iters_per_split
    k_iter_end = tl.minimum(k_iter_start + k_iters_per_split, k_iters_total)
    
    offs_k = k_iter_start * K_STEP
    offs_scale_k = k_iter_start * rep_k

    for _ in tl.range(0, k_iter_end - k_iter_start, num_stages=NUM_STAGES):
        a = a_desc.load([offs_am, offs_k])
        b = b_desc.load([offs_bn, offs_k])
        scale_a = a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
        scale_b = b_scale_desc.load([0, 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)

        accumulator = tl.dot_scaled(a, scale_a, "e2m1", b.T, scale_b, "e2m1", accumulator)

        offs_k += K_STEP
        offs_scale_k += rep_k

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    m_mask = offs_m < M_per_L
    
    col0_selector = (tl.arange(0, BLOCK_N) == 0).to(tl.float32)
    result = tl.sum(accumulator * col0_selector[None, :], axis=1)
    
    c_ptrs = c_ptr + offs_m * stride_c_m + pid_l * stride_c_l
    
    if SPLIT_K > 1:
        tl.atomic_add(c_ptrs, result.to(tl.float16), mask=m_mask)
    else:
        tl.store(c_ptrs, result.to(tl.float16), mask=m_mask)


def custom_kernel(data):
    a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, sfa_permuted, sfb_permuted, c_ref = data
    
    device = a_ref.device
    M, K_packed_dim, L = a_ref.shape
    N_ext = b_ref.shape[0]
    _, K_blocks, _ = sfa_ref_cpu.shape
    K_real = K_blocks * NVFP4_VEC_SIZE
    
    D_chunks_m = M // ROWS_PER_SCALE_CHUNK
    D_chunks_n = N_ext // ROWS_PER_SCALE_CHUNK
    K_chunks = K_blocks // 4
    rep_m, rep_n, rep_k = 1, 1, 4
    
    a_u8 = a_ref.view(torch.uint8)
    a_stacked = a_u8.permute(2, 0, 1).reshape(L * M, K_packed_dim)
    
    b_u8 = b_ref.view(torch.uint8)
    b_stacked = b_u8.permute(2, 0, 1).reshape(L * N_ext, K_packed_dim)
    
    sfa_5d = (sfa_permuted
                .permute(5, 2, 4, 0, 1, 3)  # (L, D_chunks_m, K_chunks, 32, 4, 4)
                .reshape(1, L * D_chunks_m, K_chunks, 2, 256))
    sfb_5d = (sfb_permuted
                .permute(5, 2, 4, 0, 1, 3)
                .reshape(1, L * D_chunks_n, K_chunks, 2, 256))
    
    num_m_blocks = triton.cdiv(M, BLOCK_M)
    total_tiles = num_m_blocks * L
    
    # TMA descriptors
    a_desc = TensorDescriptor.from_tensor(a_stacked, [BLOCK_M, K_PACKED])
    b_desc = TensorDescriptor.from_tensor(b_stacked, [BLOCK_N, K_PACKED])
    a_scale_desc = TensorDescriptor.from_tensor(sfa_5d, [1, rep_m, rep_k, 2, 256])
    b_scale_desc = TensorDescriptor.from_tensor(sfb_5d, [1, rep_n, rep_k, 2, 256])
    
    c_slice = c_ref[:, 0, :]
    
    # Split-K heuristic
    # if K_packed_dim == 16384:
    #     SPLIT_K = 4
    # else:
    SPLIT_K = 1
    
    # if SPLIT_K > 1:
    #     c_slice.zero_()
    
    grid = (total_tiles * SPLIT_K,)
    
    batched_block_scaled_matmul_kernel_splitk[grid](
        a_desc, a_scale_desc, b_desc, b_scale_desc,
        c_slice, c_slice.stride(0), c_slice.stride(1),
        M, K_real,
        D_chunks_m,
        NVFP4_VEC_SIZE,
        BLOCK_M, BLOCK_N, BLOCK_K,
        rep_m, rep_n, rep_k,
        SPLIT_K=SPLIT_K,
        num_m_blocks_per_L=num_m_blocks,
    )
    
    return c_ref
scrolls · 161 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