submission 747989
bigmodel_wuzhigang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 280 lines, June 9 Researcher Reciprocity License v1.0.
team_mla_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-747989?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:3f1a719e35761cd47176a7e1d5633e9af23d9231ffb5e2487a5eb514ba7f2eac
license declaredunknown
license concludedunknown
authorsbigmodel_wuzhigang
imported2026-08-15
Kernel source
team_mla_v3.py280 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
Team MLA trunk v3.
Single-knob exact-family A/B on top of team_mla_v1:
- keep the same `256/1024` pg2 + bf16Q route
- change only `num_kv_splits` for that route from 4 to 8
- keep all other shapes and policies unchanged
"""
import torch
import triton
import triton.language as tl
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import mla_decode_fwd
from task import input_t, output_t
FP8_DTYPE = aiter_dtypes.fp8
BF16 = torch.bfloat16
FP8_MAX = float(torch.finfo(FP8_DTYPE).max)
FIXED_AMAX = 16.0
PG2_256_NUM_KV_SPLITS = 8
PG2_256_KV_GRANULARITY = 8
_meta_cache = {}
_alloc_cache = {}
_bf16_np_cache = {}
@triton.jit
def _q_to_fp8_kernel(
q_ptr,
out_ptr,
scale_ptr,
amax_ptr,
FP8_MAX: tl.constexpr,
N,
BLOCK: tl.constexpr,
):
amax = tl.load(amax_ptr)
amax = tl.where(amax < 1e-12, 1e-12, amax)
scale = amax / FP8_MAX
if tl.program_id(0) == 0:
tl.store(scale_ptr, scale)
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < N
x = tl.load(q_ptr + offs, mask=mask, other=0.0).to(tl.float32)
x = x / scale
x = tl.clamp(x, -FP8_MAX, FP8_MAX)
tl.store(out_ptr + offs, x.to(out_ptr.dtype.element_ty), mask=mask)
def _build_meta(
batch_size,
kv_seq_len,
q_seq_len,
nq,
nkv,
num_kv_splits,
page_size,
dtype_q,
dtype_kv,
qo_indptr,
kv_indptr,
kv_granularity,
):
total_kv = batch_size * kv_seq_len
if page_size == 1:
num_pages = total_kv
kv_indptr_pages = kv_indptr
seq_lens = kv_indptr[1:] - kv_indptr[:-1]
kv_last_page_len = seq_lens.to(torch.int32)
else:
num_pages = total_kv // page_size
kv_indptr_pages = kv_indptr // page_size
seq_lens = kv_indptr[1:] - kv_indptr[:-1]
kv_last_page_len = (seq_lens % page_size).to(torch.int32)
kv_last_page_len = torch.where(kv_last_page_len == 0, page_size, kv_last_page_len)
info = get_mla_metadata_info_v1(
batch_size,
q_seq_len,
nq,
dtype_q,
dtype_kv,
is_sparse=False,
fast_mode=False,
num_kv_splits=num_kv_splits,
intra_batch_mode=True,
)
work = [torch.empty(shape, dtype=dtype, device="cuda") for shape, dtype in info]
wm, wi, wis, ri, rfm, rpm = work
get_mla_metadata_v1(
qo_indptr,
kv_indptr_pages,
kv_last_page_len,
nq // nkv,
nkv,
True,
wm,
wis,
wi,
ri,
rfm,
rpm,
page_size=page_size,
kv_granularity=kv_granularity,
max_seqlen_qo=q_seq_len,
uni_seqlen_qo=q_seq_len,
fast_mode=False,
max_split_per_batch=num_kv_splits,
intra_batch_mode=True,
dtype_q=dtype_q,
dtype_kv=dtype_kv,
)
kv_indices = torch.arange(num_pages, dtype=torch.int32, device="cuda")
return wm, wi, wis, ri, rfm, rpm, kv_indices, kv_last_page_len, kv_indptr_pages, page_size
def _choose_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
if batch_size == 256 and kv_seq_len == 1024:
return PG2_256_NUM_KV_SPLITS
if batch_size <= 4:
return 4
if batch_size <= 64:
return 8
return 16
def custom_kernel(data: input_t) -> output_t:
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"]
q_seq_len = config["q_seq_len"]
sm_scale = config["sm_scale"]
kv_seq_len = config["kv_seq_len"]
use_small_bf16_np = batch_size <= 32 and kv_seq_len <= 1024
use_bf16_pg2_256 = batch_size == 256 and kv_seq_len == 1024
if use_small_bf16_np:
cache_key = (batch_size, kv_seq_len)
if cache_key not in _bf16_np_cache:
total_kv = batch_size * kv_seq_len
kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
_bf16_np_cache[cache_key] = (kv_indices, kv_last_page_len)
kv_indices, kv_last_page_len = _bf16_np_cache[cache_key]
alloc_key = ("bf16_np", q.shape[0], nq, dv)
if alloc_key not in _alloc_cache:
_alloc_cache[alloc_key] = torch.empty((q.shape[0], nq, dv), dtype=BF16, device="cuda")
o = _alloc_cache[alloc_key]
kv_buffer_bf16 = kv_data["bf16"].view(-1, 1, nkv, dq)
mla_decode_fwd(
q,
kv_buffer_bf16,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
q_seq_len,
page_size=1,
nhead_kv=nkv,
sm_scale=sm_scale,
intra_batch_mode=False,
)
return o
if use_bf16_pg2_256:
page_size = 2
dtype_q = BF16
dtype_kv = FP8_DTYPE
kv_granularity = PG2_256_KV_GRANULARITY
use_fp8_q = False
else:
page_size = 1 if kv_seq_len <= 1024 else 8
dtype_q = FP8_DTYPE
dtype_kv = FP8_DTYPE
kv_granularity = max(1, 16 // page_size)
use_fp8_q = True
num_kv_splits = _choose_num_kv_splits(batch_size, kv_seq_len)
cache_key = (
batch_size,
kv_seq_len,
num_kv_splits,
page_size,
str(dtype_q),
str(dtype_kv),
kv_granularity,
)
if cache_key not in _meta_cache:
_meta_cache[cache_key] = _build_meta(
batch_size,
kv_seq_len,
q_seq_len,
nq,
nkv,
num_kv_splits,
page_size,
dtype_q,
dtype_kv,
qo_indptr,
kv_indptr,
kv_granularity,
)
wm, wi, wis, ri, rfm, rpm, kv_indices, kv_last_page_len, kv_indptr_pages, ps = _meta_cache[cache_key]
kv_buffer_fp8, kv_scale = kv_data["fp8"]
kv_buffer_4d = kv_buffer_fp8.view(-1, ps, nkv, kv_buffer_fp8.shape[-1])
if use_fp8_q:
alloc_key = ("fp8_structmix", q.shape[0], nq, dv, dq)
if alloc_key not in _alloc_cache:
amax_buf = torch.full((1,), FIXED_AMAX, dtype=torch.float32, device="cuda")
_alloc_cache[alloc_key] = (
torch.empty((q.shape[0], nq, dv), dtype=BF16, device="cuda"),
amax_buf,
torch.empty(1, dtype=torch.float32, device="cuda"),
torch.empty(q.shape[0] * nq * dq, dtype=FP8_DTYPE, device="cuda"),
)
o, amax_buf, scale_buf, q_fp8_flat = _alloc_cache[alloc_key]
n = q.numel()
block = 4096
grid = ((n + block - 1) // block,)
_q_to_fp8_kernel[grid](
q,
q_fp8_flat,
scale_buf,
amax_buf,
FP8_MAX=FP8_MAX,
N=n,
BLOCK=block,
)
q_input = q_fp8_flat.view(q.shape[0], nq, dq)
kwargs = {"q_scale": scale_buf, "kv_scale": kv_scale}
else:
alloc_key = ("bf16_pg2split8_256", q.shape[0], nq, dv)
if alloc_key not in _alloc_cache:
_alloc_cache[alloc_key] = torch.empty((q.shape[0], nq, dv), dtype=BF16, device="cuda")
o = _alloc_cache[alloc_key]
q_input = q
kwargs = {"kv_scale": kv_scale}
mla_decode_fwd(
q_input,
kv_buffer_4d,
o,
qo_indptr,
kv_indptr_pages,
kv_indices,
kv_last_page_len,
q_seq_len,
page_size=ps,
nhead_kv=nkv,
sm_scale=sm_scale,
logit_cap=0.0,
num_kv_splits=num_kv_splits,
intra_batch_mode=True,
work_meta_data=wm,
work_indptr=wi,
work_info_set=wis,
reduce_indptr=ri,
reduce_final_map=rfm,
reduce_partial_map=rpm,
**kwargs,
)
return o
scrolls · 280 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