Skip to content
KernelIndex
Search⌘K

submission 102999

nrehiew · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

bestv2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-102999?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
26.9µs
#125 of 678
2025-11-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c33d1e04a75532f7f511695f9dd26a2cb7e2d538e3b30cff9474207be839a8ad
license declaredunknown
license concludedunknown
authorsnrehiew
imported2026-08-15

Techniques

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

fp4logical_k_size_scale = (K + sf_vec_size - 1) // sf_vec_size # sf_vec_size-wide FP4 groups
warp-specializationwarp_specialize: tl.constexpr,

Kernel source

bestv2.py359 lines
import os
import torch
import triton
import triton.language as tl
from typing import Optional
from task import input_t, output_t
from reference import generate_input
from triton.tools.tensor_descriptor import TensorDescriptor

# print(triton.__version__) # 3.5.0

# os.environ["TRITON_PRINT_AUTOTUNING"] = "1"
sf_vec_size = 16
elements_per_byte = 2

ACC_DTYPE = tl.float32

CONFIGS = {
    (7168, 16384, 1): {
        "block_size_k": 4096,
        "num_stages": 3,
        "num_warps": 2,
        "block_size_m": 2,
        # "block_size_n": 64,what about
    },
    # (4096, 7168, 8): {
    #     "block_size_k": 512,
    #     "num_stages": 4,
    #     "num_warps": 1,
    #     "block_size_m": 4,
    # },
    # (7168, 2048, 4): {
    #     "block_size_k": 256,
    #     "num_stages": 4,
    #     "num_warps": 2,
    #     "block_size_m": 8,
    # },
    (4096, 7168, 8): {
        "block_size_k": 128,
        "num_stages": 4,
        "num_warps": 4,
        "block_size_m": 128,
        "block_size_n": 64,
    },
    (7168, 2048, 4): {
        "block_size_k": 128,
        "num_stages": 4,
        "num_warps": 4,
        "block_size_m": 128,
        "block_size_n": 64,
    },
}


def get_config(m, k, l):
    key = (m, k, l)
    if key in CONFIGS:
        return CONFIGS[key]
    else:
        return {
            "block_size_k": 128,
            "num_stages": 4,
            "num_warps": 4,
            "block_size_m": 128,
            "block_size_n": 64,
        }  # need 128 because there is a really small sized test case


@triton.jit
def triton_kernel_tma(
    a_desc,
    b_desc,
    out_desc,
    sfa_desc,
    sfb_desc,
    M,
    K,
    BLOCK_SIZE_K: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K_SCALE: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    elements_per_byte: tl.constexpr = elements_per_byte,
    sf_vec_size: tl.constexpr = sf_vec_size,
):
    m_pid = tl.program_id(0)
    l_pid = tl.program_id(1)
    row_start = m_pid * BLOCK_SIZE_M

    packed_k = (K + elements_per_byte - 1) // elements_per_byte
    acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=ACC_DTYPE)

    for k_byte in tl.range(0, packed_k, BLOCK_SIZE_K, num_stages=num_stages):
        a_val_uint8 = a_desc.load([l_pid, row_start, k_byte])
        b_val_uint8 = b_desc.load([l_pid, 0, k_byte])
        scale_offset = (k_byte * elements_per_byte) // sf_vec_size
        a_scale_val = sfa_desc.load([l_pid, row_start, scale_offset])
        b_scale_val = sfb_desc.load([l_pid, 0, scale_offset])

        a_val_uint8 = tl.reshape(a_val_uint8, (BLOCK_SIZE_M, BLOCK_SIZE_K))
        b_val_uint8 = tl.reshape(b_val_uint8, (BLOCK_SIZE_N, BLOCK_SIZE_K))
        a_scale_val = tl.reshape(a_scale_val, (BLOCK_SIZE_M, BLOCK_SIZE_K_SCALE))
        b_scale_val = tl.reshape(b_scale_val, (BLOCK_SIZE_N, BLOCK_SIZE_K_SCALE))

        acc = tl.dot_scaled(a_val_uint8, a_scale_val, "e2m1", b_val_uint8.T, b_scale_val, "e2m1", acc)

    out_desc.store([l_pid, row_start, 0], tl.reshape(acc, (1, BLOCK_SIZE_M, BLOCK_SIZE_N)))


@triton.jit
def triton_kernel_naive(
    a_ptr,
    b_ptr,
    out_ptr,
    M,
    K,
    L,
    am_stride,
    ak_stride,
    al_stride,
    bm_stride,
    bk_stride,
    bl_stride,
    outm_stride,
    outk_stride,
    outl_stride,  # outk_stride = 1
    sfa_ptr,
    sfb_ptr,
    sfa_m_stride,
    sfa_k_stride,
    sfa_l_stride,  # M x (K // sf_vec_size) x L
    sfb_m_stride,
    sfb_k_stride,
    sfb_l_stride,  # 1 x (K // sf_vec_size) x L
    BLOCK_SIZE_K: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    warp_specialize: tl.constexpr,
    elements_per_byte: tl.constexpr = elements_per_byte,
    sf_vec_size: tl.constexpr = sf_vec_size,
):
    # each iteration we want to load BLOCK_SIZE_K bytes of a and b = 2 * BLOCK_SIZE_K FP4 values for a and b each
    # with 2 * BLOCK_SIZE_K FP4 values we need (2 * BLOCK_SIZE_K) // sf_vec_size scale factors (FP8) each load
    m_pid = tl.program_id(0)
    l_pid = tl.program_id(1)
    row_start = m_pid * BLOCK_SIZE_M
    row_offsets = row_start + tl.arange(0, BLOCK_SIZE_M)

    a_row_ptr = a_ptr + row_offsets * am_stride + l_pid * al_stride
    b_starting_ptr = b_ptr + bl_stride * l_pid
    a_scale_row_ptr = sfa_ptr + row_offsets * sfa_m_stride + l_pid * sfa_l_stride
    b_scale_starting_ptr = sfb_ptr + sfb_l_stride * l_pid
    out_row_ptr = out_ptr + row_offsets * outm_stride + l_pid * outl_stride

    logical_k_size_scale = (K + sf_vec_size - 1) // sf_vec_size  # sf_vec_size-wide FP4 groups
    NUM_LOGICAL_ITEMS_K_PER_ITERATION: tl.constexpr = BLOCK_SIZE_K * elements_per_byte
    BLOCK_SIZE_K_SCALE: tl.constexpr = (NUM_LOGICAL_ITEMS_K_PER_ITERATION + sf_vec_size - 1) // sf_vec_size

    scale_offsets = tl.arange(0, BLOCK_SIZE_K_SCALE)
    val_offsets = tl.arange(0, BLOCK_SIZE_K)

    val_idx = 0
    acc_vec = tl.zeros((BLOCK_SIZE_M,), dtype=ACC_DTYPE)
    ONES_MASK = 0xFFFF
    for scale_idx in tl.range(0, logical_k_size_scale, BLOCK_SIZE_K_SCALE, num_stages=num_stages, warp_specialize=warp_specialize):
        # each loop should process uint8 values so 2 values per iteration
        scale_idx_offsets = scale_offsets + scale_idx

        a_scale_ptr_curr = a_scale_row_ptr[:, None] + scale_idx_offsets[None, :] * sfa_k_stride
        b_scale_ptr_curr = b_scale_starting_ptr + scale_idx_offsets * sfb_k_stride

        # Load scales as f16 for SIMD optimization
        a_scale_val = tl.load(a_scale_ptr_curr, cache_modifier=".cv", eviction_policy="evict_first").to(tl.float16)
        b_scale_val = tl.load(b_scale_ptr_curr, cache_modifier=".cg").to(tl.float16)

        val_idx_offsets = val_offsets + val_idx
        a_ptr_curr = a_row_ptr[:, None] + val_idx_offsets[None, :] * ak_stride
        b_ptr_curr = b_starting_ptr + val_idx_offsets * bk_stride
        a_val_uint8 = tl.load(a_ptr_curr, cache_modifier=".cv", eviction_policy="evict_first")  # [BLOCK_SIZE_M, BLOCK_SIZE_K]
        b_val_uint8 = tl.load(b_ptr_curr, cache_modifier=".cg")  # [BLOCK_SIZE_K]

        # Source: https://github.com/triton-lang/triton/blob/main/python/triton_kernels/triton_kernels/numerics_details/mxfp_details/_upcast_from_mxfp.py
        # Convert FP4 to f16x2 packed format (2 f16 values per u32)
        a_packed_u32 = tl.inline_asm_elementwise(
            asm="""
            {
            .reg .b8 in_8;
            .reg .f16x2 out;
            cvt.u8.u32 in_8, $1;
            cvt.rn.f16x2.e2m1x2 out, in_8;
            mov.b32 $0, out;
            }
            """,
            constraints="=r,r",
            args=[a_val_uint8],
            dtype=tl.uint32,
            is_pure=True,
            pack=1,
        )
        b_packed_u32 = tl.inline_asm_elementwise(
            asm="""
            {
            .reg .b8 in_8;
            .reg .f16x2 out;
            cvt.u8.u32 in_8, $1;
            cvt.rn.f16x2.e2m1x2 out, in_8;
            mov.b32 $0, out;
            }
            """,
            constraints="=r,r",
            args=[b_val_uint8],
            dtype=tl.uint32,
            is_pure=True,
            pack=1,
        )

        # Unpack f16x2 to individual f16 values
        a_lo_f16 = (a_packed_u32 & ONES_MASK).to(tl.uint16).to(tl.float16, bitcast=True)
        a_hi_f16 = (a_packed_u32 >> 16).to(tl.uint16).to(tl.float16, bitcast=True)
        b_lo_f16 = (b_packed_u32 & ONES_MASK).to(tl.uint16).to(tl.float16, bitcast=True)
        b_hi_f16 = (b_packed_u32 >> 16).to(tl.uint16).to(tl.float16, bitcast=True)

        # Interleave to get [M, K*2] for both a and b
        a_val_f16 = tl.interleave(a_lo_f16, a_hi_f16)  # [BLOCK_SIZE_M, BLOCK_SIZE_K * 2]
        b_val_f16 = tl.interleave(b_lo_f16, b_hi_f16)  # [BLOCK_SIZE_K * 2]
        b_val_f16 = tl.broadcast_to(b_val_f16[None, :], (BLOCK_SIZE_M, NUM_LOGICAL_ITEMS_K_PER_ITERATION))

        # Broadcast scales to match [M, K*2]
        a_scale_broadcast = tl.broadcast_to(a_scale_val[:, :, None], (BLOCK_SIZE_M, BLOCK_SIZE_K_SCALE, sf_vec_size))
        b_scale_broadcast = tl.broadcast_to(b_scale_val[None, :, None], (BLOCK_SIZE_M, BLOCK_SIZE_K_SCALE, sf_vec_size))

        a_scale_f16 = a_scale_broadcast.reshape((BLOCK_SIZE_M, NUM_LOGICAL_ITEMS_K_PER_ITERATION))
        b_scale_f16 = b_scale_broadcast.reshape((BLOCK_SIZE_M, NUM_LOGICAL_ITEMS_K_PER_ITERATION))

        #     # Use FMA in f16 for the multiply chain, convert to f32 for accumulation
        #     # FMA: result = ((a * b) * scale_a) * scale_b
        temp1 = a_val_f16 * b_val_f16
        temp2 = temp1 * a_scale_f16
        result_f16 = temp2 * b_scale_f16

        acc_vec += tl.sum(result_f16, axis=1).to(tl.float32)

        val_idx += BLOCK_SIZE_K  # in bytes

    # # acc = tl.sum(acc_vec, axis=1)
    # # tl.store(out_row_ptr, acc.to(tl.float16))
    tl.store(out_row_ptr, acc_vec.to(tl.float16))


def custom_kernel(
    data: input_t,
) -> output_t:
    a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_ref = data

    m, k_packed, l = a_ref.shape
    k = k_packed * elements_per_byte
    _, _, l = c_ref.shape

    a_ref = a_ref.view(torch.uint8)
    b_ref = b_ref.view(torch.uint8)

    config = get_config(m, k, l)
    if l == 1:
        grid = lambda meta: (triton.cdiv(m, meta["BLOCK_SIZE_M"]), l)

        triton_kernel_naive[grid](
            a_ref,
            b_ref,
            c_ref,
            m,
            k_packed * elements_per_byte,
            l,
            a_ref.stride(0),
            a_ref.stride(1),
            a_ref.stride(2),
            b_ref.stride(0),
            b_ref.stride(1),
            b_ref.stride(2),
            c_ref.stride(0),
            c_ref.stride(1),
            c_ref.stride(2),
            sfa_ref,
            sfb_ref,
            sfa_ref.stride(0),
            sfa_ref.stride(1),
            sfa_ref.stride(2),
            sfb_ref.stride(0),
            sfb_ref.stride(1),
            sfb_ref.stride(2),
            BLOCK_SIZE_K=config["block_size_k"],
            BLOCK_SIZE_M=config["block_size_m"],
            num_warps=config["num_warps"],
            num_stages=config["num_stages"],
            warp_specialize=False,
        )
        return c_ref
    else:
        block_m = config["block_size_m"]
        block_k = config["block_size_k"]
        block_n = config["block_size_n"]
        block_k_scale = (block_k * elements_per_byte + sf_vec_size - 1) // sf_vec_size

        a_uint8 = a_ref.view(torch.uint8).permute(2, 0, 1)
        b_uint8 = b_ref.view(torch.uint8).permute(2, 0, 1)
        sfa_perm = sfa_ref.permute(2, 0, 1)
        sfb_perm = sfb_ref.permute(2, 0, 1)

        a_desc = TensorDescriptor.from_tensor(a_uint8, [1, block_m, block_k])
        b_desc = TensorDescriptor.from_tensor(b_uint8, [1, block_n, block_k])
        sfa_desc = TensorDescriptor.from_tensor(sfa_perm, [1, block_m, block_k_scale])
        sfb_desc = TensorDescriptor.from_tensor(sfb_perm, [1, block_n, block_k_scale])

        c_buf = torch.empty((l, m, block_n), dtype=torch.float16, device="cuda")
        out_desc = TensorDescriptor.from_tensor(c_buf, [1, block_m, block_n])

        grid = (triton.cdiv(m, block_m), l)
        triton_kernel_tma[grid](
            a_desc,
            b_desc,
            out_desc,
            sfa_desc,
            sfb_desc,
            m,
            k,
            BLOCK_SIZE_K=block_k,
            BLOCK_SIZE_M=block_m,
            BLOCK_SIZE_N=block_n,
            BLOCK_SIZE_K_SCALE=block_k_scale,
            num_warps=config["num_warps"],
            num_stages=config["num_stages"],
        )

        out_first_col = c_buf[:, :, 0].permute(1, 0)
        return out_first_col.unsqueeze(1)


shapes = [
    {"m": 7168, "k": 16384, "l": 1, "seed": 1111},
    {"m": 4096, "k": 7168, "l": 8, "seed": 1111},
    {"m": 7168, "k": 2048, "l": 4, "seed": 1111},
    # {"m": 128, "k": 256, "l": 1, "seed": 1111},
    # {"m": 128, "k": 1536, "l": 1, "seed": 1111},
    # {"m": 128, "k": 3072, "l": 1, "seed": 1111},
    # {"m": 256, "k": 7168, "l": 1, "seed": 1111},
    # {"m": 256, "k": 7168, "l": 1, "seed": 1111},
    # {"m": 2432, "k": 4608, "l": 2, "seed": 1111},
    # {"m": 512, "k": 1536, "l": 2, "seed": 1111},
]

for shape in shapes:
    key_ = (shape["m"], shape["k"], shape["l"])
    data_ = generate_input(**shape)
    for _ in range(5):
        out = custom_kernel(data_)
    torch.cuda.synchronize()
    torch.cuda.empty_cache()
scrolls · 359 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