Skip to content
KernelIndex
Search⌘K

submission 607557

jpy794 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-607557?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD Instinct MI355X
61.9µs
#265 of 766
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1a0b8af895b084bb8696acdc9fd8661daa9987ed53c2e33af7d5773370863392
license declaredunknown
license concludedunknown
authorsjpy794
imported2026-08-15

Techniques

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

fp4if USE_MXFP4 and "mxfp4" in kv_data:
fp8q_nope = (q_nope_bf16 / q_scale[:, None]).to(tl.float8e4nv)
mmaqk = tl.dot(q_nope, kv_nope) + tl.dot(q_pe, kv_pe)
num-warps = 4num_warps = 4
split-kSPLITKV_BATCH_1 = 4
stages = 3num_stages = 3
tile-n = 64BLOCK_N = 64

Kernel source

submission.py679 lines
#!POPCORN leaderboard amd-mixed-mla
import torch
import triton
import triton.language as tl
import os

from task import input_t, output_t


NUM_HEADS = 16
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_DIM = 512
SM_SCALE = 1.0 / (QK_DIM ** 0.5)
MXFP4_BLOCK = 32
MXFP4_BLOCKS = QK_DIM // MXFP4_BLOCK
PACKED_QK_DIM = QK_DIM // 2
USE_MXFP4 = os.getenv("MLA_USE_MXFP4", "0") == "1"

BLOCK_H = 16
BLOCK_N = 64
BLOCK_C = 128
BLOCK_R = 64
SPLITKV_BATCH_1 = 4
SPLITKV_BATCH_2 = 32
SPLITKV_BATCH_3 = 64
SPLITKV_KV = 4096


@triton.jit
def _fp4_table(x):
    x = x.to(tl.int32)
    v = tl.where(x == 0, 0.0, 0.5)
    v = tl.where(x == 2, 1.0, v)
    v = tl.where(x == 3, 1.5, v)
    v = tl.where(x == 4, 2.0, v)
    v = tl.where(x == 5, 3.0, v)
    v = tl.where(x == 6, 4.0, v)
    v = tl.where(x == 7, 6.0, v)
    return v.to(tl.float32)


@triton.jit
def _decode_e8m0(scale_u8):
    return tl.exp2(scale_u8.to(tl.float32) - 127.0)


@triton.jit
def _load_mxfp4(
    packed_ptr,
    scale_ptr,
    token_idx,
    dim_offsets,
    mask_n,
    mask_d,
    PACKED_DIM_CONST: tl.constexpr,
    BLOCKS_CONST: tl.constexpr,
    BLOCK_SIZE_CONST: tl.constexpr,
):
    packed_idx = dim_offsets // 2
    packed_ptrs = packed_ptr + token_idx[None, :] * PACKED_DIM_CONST + packed_idx[:, None]
    packed = tl.load(
        packed_ptrs,
        mask=mask_d[:, None] & mask_n[None, :],
        other=0,
    ).to(tl.int32)

    lo = packed & 0xF
    hi = (packed >> 4) & 0xF
    nibbles = tl.where((dim_offsets[:, None] & 1) == 0, lo, hi)

    sign = tl.where((nibbles & 0x8) == 0, 1.0, -1.0)
    mag = _fp4_table(nibbles & 0x7)

    block_idx = dim_offsets // BLOCK_SIZE_CONST
    scale_ptrs = scale_ptr + token_idx[None, :] * BLOCKS_CONST + block_idx[:, None]
    block_scale = _decode_e8m0(
        tl.load(
            scale_ptrs,
            mask=mask_d[:, None] & mask_n[None, :],
            other=0,
        ).to(tl.int32)
    )
    return sign * mag * block_scale


@triton.jit
def _decode_mxfp4_tile_with_scale(
    packed_tile,
    scale_t,
    BLOCK_SIZE_CONST: tl.constexpr,
    CHUNK_DIM_CONST: tl.constexpr,
):
    offs_d = tl.arange(0, CHUNK_DIM_CONST)
    packed_idx = offs_d // 2
    packed = packed_tile[packed_idx, :].to(tl.int32)
    lo = packed & 0xF
    hi = (packed >> 4) & 0xF
    nibbles = tl.where((offs_d[:, None] & 1) == 0, lo, hi)

    sign = tl.where((nibbles & 0x8) == 0, 1.0, -1.0)
    mag = _fp4_table(nibbles & 0x7)

    block_idx = offs_d // BLOCK_SIZE_CONST
    block_scale = scale_t[block_idx, :]
    return sign * mag * block_scale


@triton.jit
def _mla_decode_kernel(
    q_ptr,
    kv_ptr,
    kv_indptr_ptr,
    o_ptr,
    kv_scale_ptr,
    NUM_HEADS_CONST: tl.constexpr,
    KV_LORA_RANK_CONST: tl.constexpr,
    QK_ROPE_HEAD_DIM_CONST: tl.constexpr,
    QK_DIM_CONST: tl.constexpr,
    V_DIM_CONST: tl.constexpr,
    SM_SCALE_CONST: tl.constexpr,
    MAX_KV_CONST: tl.constexpr,
    BLOCK_H_CONST: tl.constexpr,
    BLOCK_N_CONST: tl.constexpr,
    BLOCK_C_CONST: tl.constexpr,
    BLOCK_R_CONST: tl.constexpr,
):
    pid_b = tl.program_id(0)
    pid_h = tl.program_id(1)

    offs_h = pid_h * BLOCK_H_CONST + tl.arange(0, BLOCK_H_CONST)
    offs_c = tl.arange(0, BLOCK_C_CONST)
    offs_r = tl.arange(0, BLOCK_R_CONST)
    offs_v = tl.arange(0, V_DIM_CONST)

    mask_h = offs_h < NUM_HEADS_CONST
    mask_c = offs_c < KV_LORA_RANK_CONST
    mask_r = offs_r < QK_ROPE_HEAD_DIM_CONST
    mask_v = offs_v < V_DIM_CONST

    q_nope_ptrs = (
        q_ptr
        + ((pid_b * NUM_HEADS_CONST + offs_h[:, None]) * QK_DIM_CONST + offs_c[None, :])
    )
    q_pe_ptrs = (
        q_ptr
        + ((pid_b * NUM_HEADS_CONST + offs_h[:, None]) * QK_DIM_CONST + KV_LORA_RANK_CONST + offs_r[None, :])
    )
    q_nope_bf16 = tl.load(q_nope_ptrs, mask=mask_h[:, None] & mask_c[None, :], other=0.0)
    q_pe_bf16 = tl.load(q_pe_ptrs, mask=mask_h[:, None] & mask_r[None, :], other=0.0)
    q_amax = tl.maximum(tl.max(tl.abs(q_nope_bf16), axis=1), tl.max(tl.abs(q_pe_bf16), axis=1))
    q_scale = tl.maximum(q_amax / 448.0, 1e-12)
    q_nope = (q_nope_bf16 / q_scale[:, None]).to(tl.float8e4nv)
    q_pe = (q_pe_bf16 / q_scale[:, None]).to(tl.float8e4nv)

    kv_start = tl.load(kv_indptr_ptr + pid_b)
    kv_end = tl.load(kv_indptr_ptr + pid_b + 1)
    logits_scale = q_scale[:, None] * tl.load(kv_scale_ptr) * SM_SCALE_CONST
    v_scale = tl.load(kv_scale_ptr)

    e_max = tl.full((BLOCK_H_CONST,), float("-inf"), tl.float32)
    e_sum = tl.zeros((BLOCK_H_CONST,), tl.float32)
    acc = tl.zeros((BLOCK_H_CONST, V_DIM_CONST), tl.float32)

    for start_n in range(0, MAX_KV_CONST, BLOCK_N_CONST):
        offs_n = start_n + tl.arange(0, BLOCK_N_CONST)
        token_idx = kv_start + offs_n
        mask_n = token_idx < kv_end

        kv_nope_ptrs = kv_ptr + token_idx[None, :] * QK_DIM_CONST + offs_c[:, None]
        kv_pe_ptrs = (
            kv_ptr
            + token_idx[None, :] * QK_DIM_CONST
            + KV_LORA_RANK_CONST
            + offs_r[:, None]
        )
        v_ptrs = kv_ptr + token_idx[:, None] * QK_DIM_CONST + offs_v[None, :]

        kv_nope = tl.load(
            kv_nope_ptrs,
            mask=mask_c[:, None] & mask_n[None, :],
            other=0.0,
        )
        kv_pe = tl.load(
            kv_pe_ptrs,
            mask=mask_r[:, None] & mask_n[None, :],
            other=0.0,
        )
        v = tl.load(
            v_ptrs,
            mask=mask_n[:, None] & mask_v[None, :],
            other=0.0,
        )

        qk = tl.dot(q_nope, kv_nope) + tl.dot(q_pe, kv_pe)
        qk = qk * logits_scale
        qk = tl.where(mask_h[:, None] & mask_n[None, :], qk, float("-inf"))

        n_e_max = tl.maximum(tl.max(qk, axis=1), e_max)
        re_scale = tl.math.exp2((e_max - n_e_max) * 1.4426950408889634)
        p = tl.math.exp2((qk - n_e_max[:, None]) * 1.4426950408889634)

        acc *= re_scale[:, None]
        acc += tl.dot(p.to(v.dtype), v)
        e_sum = e_sum * re_scale + tl.sum(p, axis=1)
        e_max = n_e_max

    out_ptrs = (
        o_ptr
        + ((pid_b * NUM_HEADS_CONST + offs_h[:, None]) * V_DIM_CONST + offs_v[None, :])
    )
    tl.store(
        out_ptrs,
        (acc * v_scale) / e_sum[:, None],
        mask=mask_h[:, None] & mask_v[None, :],
    )


@triton.jit
def _mla_decode_splitkv_kernel(
    q_ptr,
    kv_ptr,
    kv_indptr_ptr,
    partial_acc_ptr,
    partial_max_ptr,
    partial_sum_ptr,
    kv_scale_ptr,
    NUM_SPLITS_CONST: tl.constexpr,
    NUM_HEADS_CONST: tl.constexpr,
    KV_LORA_RANK_CONST: tl.constexpr,
    QK_ROPE_HEAD_DIM_CONST: tl.constexpr,
    QK_DIM_CONST: tl.constexpr,
    V_DIM_CONST: tl.constexpr,
    SM_SCALE_CONST: tl.constexpr,
    MAX_SPLIT_KV_CONST: tl.constexpr,
    BLOCK_H_CONST: tl.constexpr,
    BLOCK_N_CONST: tl.constexpr,
    BLOCK_C_CONST: tl.constexpr,
    BLOCK_R_CONST: tl.constexpr,
):
    pid_b = tl.program_id(0)
    pid_h = tl.program_id(1)
    pid_s = tl.program_id(2)

    offs_h = pid_h * BLOCK_H_CONST + tl.arange(0, BLOCK_H_CONST)
    offs_c = tl.arange(0, BLOCK_C_CONST)
    offs_r = tl.arange(0, BLOCK_R_CONST)
    offs_v = tl.arange(0, V_DIM_CONST)

    mask_h = offs_h < NUM_HEADS_CONST
    mask_c = offs_c < KV_LORA_RANK_CONST
    mask_r = offs_r < QK_ROPE_HEAD_DIM_CONST
    mask_v = offs_v < V_DIM_CONST

    q_nope_ptrs = (
        q_ptr
        + ((pid_b * NUM_HEADS_CONST + offs_h[:, None]) * QK_DIM_CONST + offs_c[None, :])
    )
    q_pe_ptrs = (
        q_ptr
        + ((pid_b * NUM_HEADS_CONST + offs_h[:, None]) * QK_DIM_CONST + KV_LORA_RANK_CONST + offs_r[None, :])
    )
    q_nope_bf16 = tl.load(q_nope_ptrs, mask=mask_h[:, None] & mask_c[None, :], other=0.0)
    q_pe_bf16 = tl.load(q_pe_ptrs, mask=mask_h[:, None] & mask_r[None, :], other=0.0)
    q_amax = tl.maximum(tl.max(tl.abs(q_nope_bf16), axis=1), tl.max(tl.abs(q_pe_bf16), axis=1))
    q_scale = tl.maximum(q_amax / 448.0, 1e-12)
    q_nope = (q_nope_bf16 / q_scale[:, None]).to(tl.float8e4nv)
    q_pe = (q_pe_bf16 / q_scale[:, None]).to(tl.float8e4nv)

    kv_start = tl.load(kv_indptr_ptr + pid_b)
    kv_end = tl.load(kv_indptr_ptr + pid_b + 1)
    kv_len = kv_end - kv_start
    split_start = kv_start + (kv_len * pid_s) // NUM_SPLITS_CONST
    split_end = kv_start + (kv_len * (pid_s + 1)) // NUM_SPLITS_CONST

    logits_scale = q_scale[:, None] * tl.load(kv_scale_ptr) * SM_SCALE_CONST
    e_max = tl.full((BLOCK_H_CONST,), float("-inf"), tl.float32)
    e_sum = tl.zeros((BLOCK_H_CONST,), tl.float32)
    acc = tl.zeros((BLOCK_H_CONST, V_DIM_CONST), tl.float32)

    for start_n in range(0, MAX_SPLIT_KV_CONST, BLOCK_N_CONST):
        offs_n = start_n + tl.arange(0, BLOCK_N_CONST)
        token_idx = split_start + offs_n
        mask_n = token_idx < split_end

        kv_nope_ptrs = kv_ptr + token_idx[None, :] * QK_DIM_CONST + offs_c[:, None]
        kv_pe_ptrs = (
            kv_ptr
            + token_idx[None, :] * QK_DIM_CONST
            + KV_LORA_RANK_CONST
            + offs_r[:, None]
        )
        v_ptrs = kv_ptr + token_idx[:, None] * QK_DIM_CONST + offs_v[None, :]

        kv_nope = tl.load(
            kv_nope_ptrs,
            mask=mask_c[:, None] & mask_n[None, :],
            other=0.0,
        )
        kv_pe = tl.load(
            kv_pe_ptrs,
            mask=mask_r[:, None] & mask_n[None, :],
            other=0.0,
        )
        v = tl.load(
            v_ptrs,
            mask=mask_n[:, None] & mask_v[None, :],
            other=0.0,
        )

        qk = tl.dot(q_nope, kv_nope) + tl.dot(q_pe, kv_pe)
        qk = qk * logits_scale
        qk = tl.where(mask_h[:, None] & mask_n[None, :], qk, float("-inf"))

        n_e_max = tl.maximum(tl.max(qk, axis=1), e_max)
        re_scale = tl.math.exp2((e_max - n_e_max) * 1.4426950408889634)
        p = tl.math.exp2((qk - n_e_max[:, None]) * 1.4426950408889634)

        acc *= re_scale[:, None]
        acc += tl.dot(p.to(v.dtype), v)
        e_sum = e_sum * re_scale + tl.sum(p, axis=1)
        e_max = n_e_max

    acc_ptrs = (
        partial_acc_ptr
        + ((((pid_b * NUM_SPLITS_CONST + pid_s) * NUM_HEADS_CONST) + offs_h[:, None]) * V_DIM_CONST + offs_v[None, :])
    )
    max_ptrs = partial_max_ptr + (((pid_b * NUM_SPLITS_CONST + pid_s) * NUM_HEADS_CONST) + offs_h)
    sum_ptrs = partial_sum_ptr + (((pid_b * NUM_SPLITS_CONST + pid_s) * NUM_HEADS_CONST) + offs_h)
    tl.store(acc_ptrs, acc, mask=mask_h[:, None] & mask_v[None, :])
    tl.store(max_ptrs, e_max, mask=mask_h)
    tl.store(sum_ptrs, e_sum, mask=mask_h)


@triton.jit
def _mla_reduce_splitkv_kernel(
    partial_acc_ptr,
    partial_max_ptr,
    partial_sum_ptr,
    o_ptr,
    kv_scale_ptr,
    NUM_SPLITS_CONST: tl.constexpr,
    NUM_HEADS_CONST: tl.constexpr,
    V_DIM_CONST: tl.constexpr,
    BLOCK_H_CONST: tl.constexpr,
):
    pid_b = tl.program_id(0)
    pid_h = tl.program_id(1)

    offs_h = pid_h * BLOCK_H_CONST + tl.arange(0, BLOCK_H_CONST)
    offs_v = tl.arange(0, V_DIM_CONST)
    mask_h = offs_h < NUM_HEADS_CONST
    mask_v = offs_v < V_DIM_CONST

    global_max = tl.full((BLOCK_H_CONST,), float("-inf"), tl.float32)
    global_sum = tl.zeros((BLOCK_H_CONST,), tl.float32)

    for split_idx in range(0, NUM_SPLITS_CONST):
        max_ptrs = partial_max_ptr + (((pid_b * NUM_SPLITS_CONST + split_idx) * NUM_HEADS_CONST) + offs_h)
        sum_ptrs = partial_sum_ptr + (((pid_b * NUM_SPLITS_CONST + split_idx) * NUM_HEADS_CONST) + offs_h)
        split_max = tl.load(max_ptrs, mask=mask_h, other=float("-inf"))
        split_sum = tl.load(sum_ptrs, mask=mask_h, other=0.0)
        next_max = tl.maximum(global_max, split_max)
        global_sum = global_sum * tl.math.exp2((global_max - next_max) * 1.4426950408889634)
        global_sum += split_sum * tl.math.exp2((split_max - next_max) * 1.4426950408889634)
        global_max = next_max

    acc = tl.zeros((BLOCK_H_CONST, V_DIM_CONST), tl.float32)
    for split_idx in range(0, NUM_SPLITS_CONST):
        max_ptrs = partial_max_ptr + (((pid_b * NUM_SPLITS_CONST + split_idx) * NUM_HEADS_CONST) + offs_h)
        acc_ptrs = (
            partial_acc_ptr
            + ((((pid_b * NUM_SPLITS_CONST + split_idx) * NUM_HEADS_CONST) + offs_h[:, None]) * V_DIM_CONST + offs_v[None, :])
        )
        split_max = tl.load(max_ptrs, mask=mask_h, other=float("-inf"))
        scale = tl.math.exp2((split_max - global_max) * 1.4426950408889634)
        split_acc = tl.load(acc_ptrs, mask=mask_h[:, None] & mask_v[None, :], other=0.0)
        acc += split_acc * scale[:, None]

    v_scale = tl.load(kv_scale_ptr)
    out_ptrs = (
        o_ptr
        + ((pid_b * NUM_HEADS_CONST + offs_h[:, None]) * V_DIM_CONST + offs_v[None, :])
    )
    tl.store(
        out_ptrs,
        (acc * v_scale) / global_sum[:, None],
        mask=mask_h[:, None] & mask_v[None, :],
    )


@triton.jit
def _mla_decode_mxfp4_kernel(
    q_ptr,
    kv_packed_ptr,
    kv_scale_ptr,
    kv_indptr_ptr,
    o_ptr,
    NUM_HEADS_CONST: tl.constexpr,
    KV_LORA_RANK_CONST: tl.constexpr,
    QK_ROPE_HEAD_DIM_CONST: tl.constexpr,
    QK_DIM_CONST: tl.constexpr,
    V_DIM_CONST: tl.constexpr,
    SM_SCALE_CONST: tl.constexpr,
    MAX_KV_CONST: tl.constexpr,
    BLOCK_H_CONST: tl.constexpr,
    BLOCK_N_CONST: tl.constexpr,
    BLOCK_C_CONST: tl.constexpr,
    BLOCK_R_CONST: tl.constexpr,
    PACKED_DIM_CONST: tl.constexpr,
    BLOCKS_CONST: tl.constexpr,
    BLOCK_SIZE_CONST: tl.constexpr,
):
    pid_b = tl.program_id(0)
    pid_h = tl.program_id(1)

    offs_h = pid_h * BLOCK_H_CONST + tl.arange(0, BLOCK_H_CONST)
    offs_v = tl.arange(0, V_DIM_CONST)

    PACKED_V_CONST: tl.constexpr = KV_LORA_RANK_CONST // 2
    SCALE_V_CONST: tl.constexpr = KV_LORA_RANK_CONST // BLOCK_SIZE_CONST
    PACKED_R_CONST: tl.constexpr = BLOCK_R_CONST // 2
    SCALE_R_CONST: tl.constexpr = BLOCK_R_CONST // BLOCK_SIZE_CONST

    mask_h = offs_h < NUM_HEADS_CONST
    mask_v = offs_v < V_DIM_CONST

    offs_r = tl.arange(0, BLOCK_R_CONST) + KV_LORA_RANK_CONST
    mask_r = offs_r < QK_DIM_CONST
    offs_c = tl.arange(0, KV_LORA_RANK_CONST)
    q_nope_ptrs = q_ptr + ((pid_b * NUM_HEADS_CONST + offs_h[:, None]) * QK_DIM_CONST + offs_c[None, :])
    q_nope_bf16 = tl.load(q_nope_ptrs, mask=mask_h[:, None] & (offs_c[None, :] < KV_LORA_RANK_CONST), other=0.0)
    q_nope_descale = tl.full((BLOCK_H_CONST, SCALE_V_CONST), 127, dtype=tl.uint8)
    q_pe_ptrs = q_ptr + ((pid_b * NUM_HEADS_CONST + offs_h[:, None]) * QK_DIM_CONST + offs_r[None, :])
    q_pe_bf16 = tl.load(q_pe_ptrs, mask=mask_h[:, None] & mask_r[None, :], other=0.0)
    q_amax = tl.maximum(tl.max(tl.abs(q_nope_bf16), axis=1), tl.max(tl.abs(q_pe_bf16), axis=1))
    q_scale = tl.maximum(q_amax / 448.0, 1e-12)
    q_nope = (q_nope_bf16 / q_scale[:, None]).to(tl.float8e4nv)
    q_pe = (q_pe_bf16 / q_scale[:, None]).to(tl.float8e4nv)
    q_pe_descale = tl.full((BLOCK_H_CONST, BLOCK_R_CONST // BLOCK_SIZE_CONST), 127, dtype=tl.uint8)

    kv_start = tl.load(kv_indptr_ptr + pid_b)
    kv_end = tl.load(kv_indptr_ptr + pid_b + 1)
    logits_scale = q_scale[:, None] * (SM_SCALE_CONST * 1.4426950408889634)

    e_max = tl.full((BLOCK_H_CONST,), float("-inf"), tl.float32)
    e_sum = tl.zeros((BLOCK_H_CONST,), tl.float32)
    acc = tl.zeros((BLOCK_H_CONST, V_DIM_CONST), tl.float32)

    for start_n in range(0, MAX_KV_CONST, BLOCK_N_CONST):
        offs_n = start_n + tl.arange(0, BLOCK_N_CONST)
        token_idx = kv_start + offs_n
        mask_n = token_idx < kv_end

        offs_lp = tl.arange(0, PACKED_V_CONST)
        latent_ptrs = kv_packed_ptr + token_idx[None, :] * PACKED_DIM_CONST + offs_lp[:, None]
        latent_packed = tl.load(
            latent_ptrs,
            mask=(offs_lp[:, None] < PACKED_V_CONST) & mask_n[None, :],
            other=0,
        )

        offs_ls = tl.arange(0, SCALE_V_CONST)
        latent_scale_ptrs = kv_scale_ptr + token_idx[:, None] * BLOCKS_CONST + offs_ls[None, :]
        latent_scale = tl.load(
            latent_scale_ptrs,
            mask=mask_n[:, None] & (offs_ls[None, :] < SCALE_V_CONST),
            other=0,
        )

        offs_rp = tl.arange(0, PACKED_R_CONST)
        rope_ptrs = kv_packed_ptr + token_idx[None, :] * PACKED_DIM_CONST + ((KV_LORA_RANK_CONST // 2) + offs_rp)[:, None]
        kv_pe = tl.load(
            rope_ptrs,
            mask=(offs_rp[:, None] < PACKED_R_CONST) & mask_n[None, :],
            other=0,
        )
        offs_rs = (KV_LORA_RANK_CONST // BLOCK_SIZE_CONST) + tl.arange(0, SCALE_R_CONST)
        rope_scale_ptrs = kv_scale_ptr + token_idx[:, None] * BLOCKS_CONST + offs_rs[None, :]
        kv_pe_scale = tl.load(
            rope_scale_ptrs,
            mask=mask_n[:, None] & (offs_rs[None, :] < BLOCKS_CONST),
            other=0,
        )

        qk = tl.zeros((BLOCK_H_CONST, BLOCK_N_CONST), dtype=tl.float32)
        qk = tl.dot_scaled(
            q_nope,
            q_nope_descale,
            "e4m3",
            latent_packed,
            latent_scale,
            "e2m1",
            fast_math=True,
            acc=qk,
        )
        qk = tl.dot_scaled(
            q_pe,
            q_pe_descale,
            "e4m3",
            kv_pe,
            kv_pe_scale,
            "e2m1",
            fast_math=True,
            acc=qk,
        )
        qk = qk * logits_scale
        qk = tl.where(mask_h[:, None] & mask_n[None, :], qk, float("-inf"))

        n_e_max = tl.maximum(tl.max(qk, axis=1), e_max)
        re_scale = tl.math.exp2(e_max - n_e_max)
        p = tl.math.exp2(qk - n_e_max[:, None])

        acc *= re_scale[:, None]
        latent_scale_t = _decode_e8m0(tl.trans(latent_scale).to(tl.int32))
        v = _decode_mxfp4_tile_with_scale(
            latent_packed,
            latent_scale_t,
            BLOCK_SIZE_CONST=BLOCK_SIZE_CONST,
            CHUNK_DIM_CONST=KV_LORA_RANK_CONST,
        ).to(q_pe.dtype)
        acc += tl.dot(p.to(v.dtype), tl.trans(v))
        e_sum = e_sum * re_scale + tl.sum(p, axis=1)
        e_max = n_e_max

    out_ptrs = o_ptr + ((pid_b * NUM_HEADS_CONST + offs_h[:, None]) * V_DIM_CONST + offs_v[None, :])
    tl.store(
        out_ptrs,
        acc / e_sum[:, None],
        mask=mask_h[:, None] & mask_v[None, :],
    )


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, _, kv_indptr, config = data

    if int(config["q_seq_len"]) != 1:
        raise RuntimeError("custom_kernel only supports q_seq_len=1")
    if int(config["num_heads"]) != NUM_HEADS:
        raise RuntimeError(f"custom_kernel expects num_heads={NUM_HEADS}")
    if abs(float(config["sm_scale"]) - SM_SCALE) > 1e-6:
        raise RuntimeError(f"custom_kernel expects sm_scale={SM_SCALE}")

    batch_size = int(config["batch_size"])
    kv_seq_len = int(config["kv_seq_len"])
    q = q.contiguous().view(batch_size, NUM_HEADS, QK_DIM)
    o = torch.empty((batch_size, NUM_HEADS, V_DIM), device=q.device, dtype=torch.bfloat16)

    num_warps = 4
    num_stages = 3
    waves_per_eu = 0
    split_kv = 0
    if kv_seq_len >= SPLITKV_KV:
        if batch_size <= SPLITKV_BATCH_1:
            split_kv = 8
        elif batch_size <= SPLITKV_BATCH_2:
            split_kv = 4
        elif batch_size <= SPLITKV_BATCH_3:
            split_kv = 4

    grid = (batch_size, triton.cdiv(NUM_HEADS, BLOCK_H))
    if USE_MXFP4 and "mxfp4" in kv_data:
        kv_packed, kv_scale = kv_data["mxfp4"]
        kv_packed = kv_packed.contiguous().view(-1, PACKED_QK_DIM).view(torch.uint8)
        kv_scale = kv_scale[:, :MXFP4_BLOCKS].contiguous().view(-1, MXFP4_BLOCKS).view(torch.uint8)
        _mla_decode_mxfp4_kernel[grid](
            q,
            kv_packed,
            kv_scale,
            kv_indptr,
            o,
            NUM_HEADS_CONST=NUM_HEADS,
            KV_LORA_RANK_CONST=KV_LORA_RANK,
            QK_ROPE_HEAD_DIM_CONST=QK_ROPE_HEAD_DIM,
            QK_DIM_CONST=QK_DIM,
            V_DIM_CONST=V_DIM,
            SM_SCALE_CONST=SM_SCALE,
            MAX_KV_CONST=kv_seq_len,
            BLOCK_H_CONST=BLOCK_H,
            BLOCK_N_CONST=BLOCK_N,
            BLOCK_C_CONST=BLOCK_C,
            BLOCK_R_CONST=BLOCK_R,
            PACKED_DIM_CONST=PACKED_QK_DIM,
            BLOCKS_CONST=MXFP4_BLOCKS,
            BLOCK_SIZE_CONST=MXFP4_BLOCK,
            num_warps=num_warps,
            num_stages=num_stages,
            waves_per_eu=waves_per_eu,
            matrix_instr_nonkdim=16,
        )
    else:
        if "fp8" not in kv_data:
            raise RuntimeError("custom_kernel expects kv_data['fp8'] or kv_data['mxfp4']")
        kv, kv_scale = kv_data["fp8"]
        kv = kv.contiguous().view(-1, QK_DIM)
        if split_kv > 1:
            partial_acc = torch.empty(
                (batch_size, split_kv, NUM_HEADS, V_DIM),
                device=q.device,
                dtype=torch.float32,
            )
            partial_max = torch.empty(
                (batch_size, split_kv, NUM_HEADS),
                device=q.device,
                dtype=torch.float32,
            )
            partial_sum = torch.empty(
                (batch_size, split_kv, NUM_HEADS),
                device=q.device,
                dtype=torch.float32,
            )
            split_grid = (batch_size, triton.cdiv(NUM_HEADS, BLOCK_H), split_kv)
            max_split_kv = triton.cdiv(kv_seq_len, split_kv * BLOCK_N) * BLOCK_N
            _mla_decode_splitkv_kernel[split_grid](
                q,
                kv,
                kv_indptr,
                partial_acc,
                partial_max,
                partial_sum,
                kv_scale,
                NUM_SPLITS_CONST=split_kv,
                NUM_HEADS_CONST=NUM_HEADS,
                KV_LORA_RANK_CONST=KV_LORA_RANK,
                QK_ROPE_HEAD_DIM_CONST=QK_ROPE_HEAD_DIM,
                QK_DIM_CONST=QK_DIM,
                V_DIM_CONST=V_DIM,
                SM_SCALE_CONST=SM_SCALE,
                MAX_SPLIT_KV_CONST=max_split_kv,
                BLOCK_H_CONST=BLOCK_H,
                BLOCK_N_CONST=BLOCK_N,
                BLOCK_C_CONST=BLOCK_C,
                BLOCK_R_CONST=BLOCK_R,
                num_warps=num_warps,
                num_stages=2,
                waves_per_eu=waves_per_eu,
                matrix_instr_nonkdim=16,
            )
            _mla_reduce_splitkv_kernel[grid](
                partial_acc,
                partial_max,
                partial_sum,
                o,
                kv_scale,
                NUM_SPLITS_CONST=split_kv,
                NUM_HEADS_CONST=NUM_HEADS,
                V_DIM_CONST=V_DIM,
                BLOCK_H_CONST=BLOCK_H,
                num_warps=4,
                num_stages=2,
                waves_per_eu=0,
                matrix_instr_nonkdim=16,
            )
        else:
            _mla_decode_kernel[grid](
                q,
                kv,
                kv_indptr,
                o,
                kv_scale,
                NUM_HEADS_CONST=NUM_HEADS,
                KV_LORA_RANK_CONST=KV_LORA_RANK,
                QK_ROPE_HEAD_DIM_CONST=QK_ROPE_HEAD_DIM,
                QK_DIM_CONST=QK_DIM,
                V_DIM_CONST=V_DIM,
                SM_SCALE_CONST=SM_SCALE,
                MAX_KV_CONST=kv_seq_len,
                BLOCK_H_CONST=BLOCK_H,
                BLOCK_N_CONST=BLOCK_N,
                BLOCK_C_CONST=BLOCK_C,
                BLOCK_R_CONST=BLOCK_R,
                num_warps=num_warps,
                num_stages=num_stages,
                waves_per_eu=waves_per_eu,
                matrix_instr_nonkdim=16,
            )
    return o.contiguous()
scrolls · 679 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