submission 745751
farhan-navas · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 470 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-745751?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:e4118cf31163c5401986acf571a05065b9019d9a8e752213e9f79b883e70eb1b
license declaredunknown
license concludedunknown
authorsfarhan-navas
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
q_nope = q_nope_raw.to(tl.float8e4nv) # [H, 512] fp8mma
s = tl.dot(q_nope, tl.trans(k_nope))num-warps = 4
num_warps=4,online-softmax
m_new = tl.maximum(m_i, tl.max(s_scaled, axis=1))stages = 2
num_stages=2,Kernel source
submission.py470 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA Decode Submission — AMD MI355X (gfx950)
Hybrid: AITER (a16w8/a8w8) + Triton MLA decode kernel.
Based on SGLang's production Triton MLA decode patterns.
"""
from task import input_t, output_t
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1") # -2-3μs per kernel launch
import torch
import triton
import triton.language as tl
# Eval server: Triton 3.6.0, GPUTarget(backend='hip', arch='gfx950', warp_size=64)
PAGE_SIZE = 1
NUM_KV_SPLITS = 32
SM_SCALE = 1.0 / (576 ** 0.5)
_cache = {}
# ============================================================
# AITER path (existing optimized baseline)
# ============================================================
def _quantize_fp8(tensor):
from aiter import dtypes as aiter_dtypes
from aiter.ops.quant import per_tensor_quant_hip
return per_tensor_quant_hip(tensor, scale=None, quant_dtype=aiter_dtypes.fp8)
_metadata_done = set()
def _get_cached_buffers(batch_size, nq, nkv, dv, total_q, total_kv_len, q_dtype, kv_dtype, kv_splits=32):
from aiter import get_mla_metadata_info_v1
key = (batch_size, nq, total_q, total_kv_len, q_dtype, kv_splits)
if key not in _cache:
max_q_len = 1
info = get_mla_metadata_info_v1(
batch_size, max_q_len, nq, q_dtype, kv_dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=kv_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
o = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
_cache[key] = (work, o, kv_indices)
return _cache[key]
_aiter_mla = None
_aiter_meta = None
def _aiter_path(data, use_bf16_q=False, num_splits=None):
global _aiter_mla, _aiter_meta
if _aiter_mla is None:
from aiter.mla import mla_decode_fwd
from aiter import get_mla_metadata_v1
_aiter_mla = mla_decode_fwd
_aiter_meta = get_mla_metadata_v1
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
nq = config["num_heads"]
nkv = config["num_kv_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
max_q_len = 1
if use_bf16_q:
q_input = q.view(-1, nq, dq)
q_scale = None
else:
q_fp8, q_scale = _quantize_fp8(q)
q_input = q_fp8.view(-1, nq, dq)
kv_buffer_fp8, kv_scale = kv_data["fp8"]
# Avoid .item() CPU-GPU sync — compute from config
total_kv_len = batch_size * config["kv_seq_len"]
total_q = q.shape[0]
# Cache kv_last_page_len (same for same shape)
kv_lpl_key = ("kv_lpl", batch_size, config["kv_seq_len"])
if kv_lpl_key not in _cache:
_cache[kv_lpl_key] = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
kv_last_page_len = _cache[kv_lpl_key]
kv_splits = num_splits if num_splits is not None else NUM_KV_SPLITS
work, o, kv_indices = _get_cached_buffers(
batch_size, nq, nkv, dv, total_q, total_kv_len,
q_input.dtype, kv_buffer_fp8.dtype, kv_splits=kv_splits,
)
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = work
kv_buffer_4d = kv_buffer_fp8.view(kv_buffer_fp8.shape[0], PAGE_SIZE, nkv, kv_buffer_fp8.shape[-1])
meta_key = (batch_size, nq, total_q, total_kv_len, q_input.dtype, kv_splits)
if meta_key not in _metadata_done:
_aiter_meta(
qo_indptr, kv_indptr, kv_last_page_len,
nq // nkv, nkv, True,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=PAGE_SIZE,
kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=max_q_len,
uni_seqlen_qo=max_q_len,
fast_mode=False,
max_split_per_batch=kv_splits,
intra_batch_mode=True,
dtype_q=q_input.dtype,
dtype_kv=kv_buffer_fp8.dtype,
)
_metadata_done.add(meta_key)
_aiter_mla(
q_input,
kv_buffer_4d,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
max_q_len,
page_size=PAGE_SIZE,
nhead_kv=nkv,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=kv_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
work_meta_data=work_metadata,
work_indptr=work_indptr,
work_info_set=work_info_set,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
)
return o
# ============================================================
# Triton MLA Decode — Stage 1 (split-K attention)
# ============================================================
# Approach 1: Native fp8×fp8 tl.dot — no cast overhead
# Q cast to fp8 once before tile loop, KV loaded as raw fp8
# tl.dot(fp8, fp8) → native MFMA fp8 on gfx950
# bf16 fallback when USE_FP8_DOT=False
@triton.jit
def _mla_stage1(
Q, # [batch_size, num_heads, 576] bf16 — contiguous
KV, # [total_kv, 576] fp8 or bf16
kv_scale, # scalar float (1.0 for bf16, actual scale for fp8)
kv_indptr, # [batch_size + 1] int32
partial_out, # [batch_size * num_kv_splits * num_heads * DV] f32
partial_lse, # [batch_size * num_kv_splits * num_heads] f32
stride_qb, # Q stride: batch
stride_qh, # Q stride: head
stride_kvt, # KV stride: token
NUM_KV_SPLITS: tl.constexpr,
NUM_HEADS: tl.constexpr,
SM_SCALE: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_DMODEL: tl.constexpr, # 512
BLOCK_DPE: tl.constexpr, # 64
BLOCK_DV: tl.constexpr, # 512
USE_FP8_DOT: tl.constexpr, # True = native fp8 dot, False = bf16 dot
):
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
kv_start = tl.load(kv_indptr + batch_id)
kv_end = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_end - kv_start
split_size = tl.cdiv(kv_len, NUM_KV_SPLITS)
my_start = kv_start + split_id * split_size
my_end = tl.minimum(my_start + split_size, kv_end)
my_len = my_end - my_start
partial_base = batch_id * NUM_KV_SPLITS * NUM_HEADS * BLOCK_DV + split_id * NUM_HEADS * BLOCK_DV
lse_base = batch_id * NUM_KV_SPLITS * NUM_HEADS + split_id * NUM_HEADS
offs_h = tl.arange(0, NUM_HEADS)
if my_len <= 0:
tl.store(partial_lse + lse_base + offs_h, tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32))
return
# Load Q and cast once — fp8 for native dot, bf16 for fallback
offs_d = tl.arange(0, BLOCK_DMODEL)
offs_pe = tl.arange(0, BLOCK_DPE)
q_base = batch_id * stride_qb
q_nope_raw = tl.load(Q + q_base + offs_h[:, None] * stride_qh + offs_d[None, :])
q_pe_raw = tl.load(Q + q_base + offs_h[:, None] * stride_qh + (BLOCK_DMODEL + offs_pe[None, :]))
if USE_FP8_DOT:
# Cast Q bf16 → fp8 ONCE (amortized over all KV tiles)
q_nope = q_nope_raw.to(tl.float8e4nv) # [H, 512] fp8
q_pe = q_pe_raw.to(tl.float8e4nv) # [H, 64] fp8
else:
q_nope = q_nope_raw # [H, 512] bf16
q_pe = q_pe_raw # [H, 64] bf16
# Online softmax state
m_i = tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NUM_HEADS], dtype=tl.float32)
acc = tl.zeros([NUM_HEADS, BLOCK_DV], dtype=tl.float32)
# Absorb kv_scale into QK prescale
scale = SM_SCALE * kv_scale * 1.4426950408889634
offs_n = tl.arange(0, BLOCK_N)
for start in range(0, my_len, BLOCK_N):
n_valid = tl.minimum(BLOCK_N, my_len - start)
kv_idx = my_start + start + offs_n
mask_n = offs_n < n_valid
# Load K_nope and K_pe — no cast for fp8 dot, bf16 cast for fallback
k_ptrs_nope = KV + kv_idx[:, None] * stride_kvt + offs_d[None, :]
k_ptrs_pe = KV + kv_idx[:, None] * stride_kvt + (BLOCK_DMODEL + offs_pe[None, :])
if USE_FP8_DOT:
k_nope = tl.load(k_ptrs_nope, mask=mask_n[:, None], other=0.0) # raw fp8
k_pe = tl.load(k_ptrs_pe, mask=mask_n[:, None], other=0.0) # raw fp8
else:
k_nope = tl.load(k_ptrs_nope, mask=mask_n[:, None], other=0.0).to(tl.bfloat16)
k_pe = tl.load(k_ptrs_pe, mask=mask_n[:, None], other=0.0).to(tl.bfloat16)
# QK: native fp8×fp8 dot or bf16 dot → FP32 accumulator
s = tl.dot(q_nope, tl.trans(k_nope))
s += tl.dot(q_pe, tl.trans(k_pe))
# Mask + scale
s = tl.where(mask_n[None, :], s, float("-inf"))
s_scaled = s * scale
# Online softmax
m_new = tl.maximum(m_i, tl.max(s_scaled, axis=1))
alpha = tl.math.exp2(m_i - m_new)
p = tl.math.exp2(s_scaled - m_new[:, None])
l_i = l_i * alpha + tl.sum(p, axis=1)
acc = acc * alpha[:, None]
# Load V — no cast for fp8 dot, bf16 cast for fallback
offs_v = tl.arange(0, BLOCK_DV)
v_ptrs = KV + kv_idx[:, None] * stride_kvt + offs_v[None, :]
if USE_FP8_DOT:
v = tl.load(v_ptrs, mask=mask_n[:, None], other=0.0) # raw fp8
# Cast P to fp8 for native fp8×fp8 PV dot
acc = tl.dot(p.to(tl.float8e4nv), v, acc)
else:
v = tl.load(v_ptrs, mask=mask_n[:, None], other=0.0).to(tl.bfloat16)
acc += tl.dot(p.to(tl.bfloat16), v)
m_i = m_new
# Normalize and apply deferred V scale
acc = acc * (kv_scale / l_i[:, None])
# Store LSE in log2 domain: lse = log2(l_i) + m_i
lse = tl.math.log2(l_i) + m_i
tl.store(partial_lse + lse_base + offs_h, lse)
# Store normalized partial output [NUM_HEADS, BLOCK_DV] as 2D block
offs_v = tl.arange(0, BLOCK_DV)
out_offs = offs_h[:, None] * BLOCK_DV + offs_v[None, :]
tl.store(partial_out + partial_base + out_offs, acc)
# ============================================================
# Triton MLA Decode — Stage 2 (reduction)
# ============================================================
@triton.jit
def _mla_stage2(
partial_out, # [batch_size * num_kv_splits * num_heads * DV] f32
partial_lse, # [batch_size * num_kv_splits * num_heads] f32
output, # [batch_size, num_heads, DV] bf16
stride_ob, # output stride batch
stride_oh, # output stride head
NUM_KV_SPLITS: tl.constexpr,
NUM_HEADS: tl.constexpr,
BLOCK_DV: tl.constexpr,
):
batch_id = tl.program_id(0)
head_id = tl.program_id(1)
# Load all LSEs for this batch x head, find global max
lse_base = batch_id * NUM_KV_SPLITS * NUM_HEADS + head_id
out_base = batch_id * NUM_KV_SPLITS * NUM_HEADS * BLOCK_DV + head_id * BLOCK_DV
offs_v = tl.arange(0, BLOCK_DV)
acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
weight_sum = tl.zeros([1], dtype=tl.float32)
# First pass: find global max LSE
m_global = tl.full([1], float("-inf"), dtype=tl.float32)
for s in range(NUM_KV_SPLITS):
lse_s = tl.load(partial_lse + lse_base + s * NUM_HEADS)
m_global = tl.maximum(m_global, lse_s)
# Second pass: weighted sum with rescaling
for s in range(NUM_KV_SPLITS):
lse_s = tl.load(partial_lse + lse_base + s * NUM_HEADS)
w = tl.math.exp2(lse_s - m_global)
partial = tl.load(partial_out + out_base + s * NUM_HEADS * BLOCK_DV + offs_v)
acc += w * partial
weight_sum += w
acc = acc / weight_sum
# Store
out_ptr = output + batch_id * stride_ob + head_id * stride_oh + offs_v
tl.store(out_ptr, acc.to(tl.bfloat16))
# ============================================================
# Triton path wrapper
# ============================================================
_triton_cache = {}
def _triton_path(data, use_fp8=True, block_n=None, splits=None, stages=None, s1_waves=None, s2_waves=None, warps=4):
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
num_heads = config["num_heads"]
dv = config["v_head_dim"]
total_q = q.shape[0]
# fp8 KV: 2x less bandwidth than bf16 (576 vs 1152 bytes/token)
if use_fp8:
kv_tensor, kv_scale_val = kv_data["fp8"] # [total_kv, 1, 576] fp8 + scalar
kv_scale_f = float(kv_scale_val)
else:
kv_tensor = kv_data["bf16"] # [total_kv, 1, 576] bf16
kv_scale_f = 1.0
# Reshape KV: [total_kv, 1, 576] → [total_kv, 576]
kv_flat = kv_tensor.view(-1, 576)
# Q is [total_q, num_heads, 576] bf16, already contiguous
q_3d = q.view(total_q, num_heads, 576)
BLOCK_N = block_n if block_n is not None else (32 if use_fp8 else 16)
BLOCK_DMODEL = 512
BLOCK_DPE = 64
BLOCK_DV = 512
num_kv_splits = splits if splits is not None else NUM_KV_SPLITS
# Allocate/reuse partial buffers
cache_key = ("triton", batch_size, num_heads, num_kv_splits, BLOCK_N)
if cache_key not in _triton_cache:
partial_out = torch.empty(
batch_size * num_kv_splits * num_heads * BLOCK_DV,
dtype=torch.float32, device="cuda"
)
partial_lse = torch.empty(
batch_size * num_kv_splits * num_heads,
dtype=torch.float32, device="cuda"
)
out = torch.empty(
(total_q, num_heads, BLOCK_DV),
dtype=torch.bfloat16, device="cuda"
)
_triton_cache[cache_key] = (partial_out, partial_lse, out)
partial_out, partial_lse, out = _triton_cache[cache_key]
if out.shape[0] != total_q:
out = torch.empty((total_q, num_heads, BLOCK_DV), dtype=torch.bfloat16, device="cuda")
_triton_cache[cache_key] = (partial_out, partial_lse, out)
# Stage 1: split-K attention
grid1 = (batch_size, num_kv_splits)
s1_kwargs = dict(
NUM_KV_SPLITS=num_kv_splits,
NUM_HEADS=num_heads,
SM_SCALE=SM_SCALE,
BLOCK_N=BLOCK_N,
BLOCK_DMODEL=BLOCK_DMODEL,
BLOCK_DPE=BLOCK_DPE,
BLOCK_DV=BLOCK_DV,
USE_FP8_DOT=use_fp8,
num_warps=warps,
num_stages=stages if stages is not None else 1,
)
if s1_waves is not None:
s1_kwargs["waves_per_eu"] = s1_waves
_mla_stage1[grid1](
q_3d, kv_flat, kv_scale_f, kv_indptr,
partial_out, partial_lse,
q_3d.stride(0), q_3d.stride(1),
kv_flat.stride(0),
**s1_kwargs,
)
# Stage 2: reduction
grid2 = (batch_size, num_heads)
s2_kwargs = dict(
NUM_KV_SPLITS=num_kv_splits,
NUM_HEADS=num_heads,
BLOCK_DV=BLOCK_DV,
num_warps=4,
num_stages=2,
)
if s2_waves is not None:
s2_kwargs["waves_per_eu"] = s2_waves
_mla_stage2[grid2](
partial_out, partial_lse, out,
out.stride(0), out.stride(1),
**s2_kwargs,
)
return out
# ============================================================
# Entry point
# ============================================================
# EVOLVE-BLOCK-START mla_kernel
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
total_kv = batch_size * kv_seq_len
if batch_size >= 256:
# AITER for BS=256 — per-shape splits
if kv_seq_len >= 8192:
return _aiter_path(data, use_bf16_q=False, num_splits=64)
else:
return _aiter_path(data, use_bf16_q=True, num_splits=16)
elif batch_size >= 64:
# AITER for BS≥64 — per-shape splits
if kv_seq_len >= 8192:
return _aiter_path(data, use_bf16_q=False, num_splits=32)
else:
return _aiter_path(data, use_bf16_q=True, num_splits=16)
elif batch_size >= 32 and kv_seq_len >= 8192:
# AITER a8w8 for BS=32 KV=8192
return _aiter_path(data, use_bf16_q=False, num_splits=16)
elif batch_size <= 4:
# Triton bf16 N=32, warps=8, stages=2, waves_per_eu=1
return _triton_path(data, use_fp8=False, block_n=32,
splits=16 if kv_seq_len <= 1024 else 32,
stages=2, s1_waves=1, s2_waves=4, warps=8)
else:
# BS=32 KV=1024: Triton bf16 N=32, warps=8, splits=8
return _triton_path(data, use_fp8=False, block_n=32,
splits=8, stages=2, s2_waves=4, warps=8)
# EVOLVE-BLOCK-END mla_kernel
scrolls · 470 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