submission 619884
ftyghome · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 684 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-619884?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:7bdfc270d5faf11788f1d3156a8261c0095ea96937029a519e1d29ebc073eda3
license declaredunknown
license concludedunknown
authorsftyghome
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
if USE_MXFP4 and "mxfp4" in kv_data:fp8
q_nope = (q_nope_bf16 / q_scale[:, None]).to(tl.float8e4nv)mma
qk = tl.dot(q_nope, kv_nope) + tl.dot(q_pe, kv_pe)num-warps = 4
num_warps = 4split-k
SPLITKV_BATCH_1 = 4stages = 3
num_stages = 3tile-n = 64
BLOCK_N = 64Kernel source
submission.py684 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 = False # 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_fp8_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)
v_scale = tl.load(kv_fp8_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
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,
) / v_scale
v = v.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 * v_scale) / 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_fp8_scale = kv_data["fp8"]
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_fp8_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 · 684 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