submission 653808
lonk · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 392 lines, June 9 Researcher Reciprocity License v1.0.
amd-mixed-mla-hybrid-v57-legacy28.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-653808?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:be044a8b768ac8aada83434b39f25856d12eb9a2efd71ff847d2c62f18164022
license declaredunknown
license concludedunknown
authorslonk
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
scores = tl.dot(q_lat, tl.trans(kv_lat))num-warps = 1
num_warps=1,persistent-kernel
def _persistent_decode(data: input_t) -> output_t:tile-k = 32
BK = 32Kernel source
amd-mixed-mla-hybrid-v57-legacy28.py392 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# Hybrid v57 legacy-28 sweep: keep the validated Triton baseline for
# seven shapes and route the gating shape through the legacy 28-split metadata regime
# without any cross-call caches.
PAGE_SIZE = 1
MAX_KV_SPLITS = 32
TARGET_TILE_COUNT = 512
MIN_KV_TOKENS_PER_SPLIT = 256
def _aiter_symbols():
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
return aiter_dtypes.fp8, get_mla_metadata_info_v1, get_mla_metadata_v1, mla_decode_fwd
def _quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
fp8_dtype, _, _, _ = _aiter_symbols()
finfo = torch.finfo(fp8_dtype)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = (amax / finfo.max).to(torch.float32).reshape(1)
fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(fp8_dtype)
return fp8_tensor, scale
def _choose_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
del batch_size, kv_seq_len
return 28
def _make_metadata(
batch_size: int,
max_q_len: int,
nhead: int,
nhead_kv: int,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
kv_last_page_len: torch.Tensor,
num_kv_splits: int,
) -> dict[str, torch.Tensor]:
_, get_mla_metadata_info_v1, get_mla_metadata_v1, _ = _aiter_symbols()
info = get_mla_metadata_info_v1(
batch_size,
max_q_len,
nhead,
q_dtype,
kv_dtype,
is_sparse=False,
fast_mode=False,
num_kv_splits=num_kv_splits,
intra_batch_mode=True,
)
work = [torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info]
(
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
) = work
get_mla_metadata_v1(
qo_indptr,
kv_indptr,
kv_last_page_len,
nhead // nhead_kv,
nhead_kv,
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=num_kv_splits,
intra_batch_mode=True,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
return {
"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,
}
def _persistent_decode(data: input_t) -> output_t:
_, _, _, mla_decode_fwd = _aiter_symbols()
q, kv_data, qo_indptr, kv_indptr, config = data
q_fp8, q_scale = _quantize_fp8(q)
kv_buffer_fp8, kv_scale = kv_data["fp8"]
batch_size = config["batch_size"]
num_heads = config["num_heads"]
num_kv_heads = config["num_kv_heads"]
q_seq_len = config["q_seq_len"]
kv_seq_len = config["kv_seq_len"]
total_kv = kv_buffer_fp8.shape[0]
num_kv_splits = _choose_num_kv_splits(batch_size, kv_seq_len)
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=q.device)
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
out = torch.empty((q.shape[0], num_heads, config["v_head_dim"]), dtype=torch.bfloat16, device=q.device)
meta = _make_metadata(
batch_size,
q_seq_len,
num_heads,
num_kv_heads,
q_fp8.dtype,
kv_buffer_fp8.dtype,
qo_indptr,
kv_indptr,
kv_last_page_len,
num_kv_splits,
)
mla_decode_fwd(
q_fp8,
kv_buffer_fp8.view(total_kv, PAGE_SIZE, num_kv_heads, kv_buffer_fp8.shape[-1]),
out,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
q_seq_len,
page_size=PAGE_SIZE,
nhead_kv=num_kv_heads,
sm_scale=config["sm_scale"],
logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
**meta,
)
return out
@triton.jit
def _mla_s1(
Q,
KV,
KV_SCALE,
QO_ind,
KV_ind,
PO,
PM,
PL,
sm_scale,
sq0,
sq1,
skv,
spo0,
spo1,
spo2,
sml0,
sml1,
NS: tl.constexpr,
BK: tl.constexpr,
):
bid = tl.program_id(0)
sid = tl.program_id(1)
q_off = tl.load(QO_ind + bid)
kv_lo = tl.load(KV_ind + bid)
kv_hi = tl.load(KV_ind + bid + 1)
kv_n = kv_hi - kv_lo
kvs = tl.load(KV_SCALE)
csz = tl.cdiv(kv_n, NS)
lo = kv_lo + sid * csz
hi = tl.minimum(lo + csz, kv_hi)
cs = kvs * sm_scale
h16 = tl.arange(0, 16)
d512 = tl.arange(0, 512)
d64 = tl.arange(0, 64)
q_lat = tl.load(Q + q_off * sq0 + h16[:, None] * sq1 + d512[None, :])
q_rope = tl.load(Q + q_off * sq0 + h16[:, None] * sq1 + 512 + d64[None, :])
mi = tl.full([16], float("-inf"), dtype=tl.float32)
li = tl.zeros([16], dtype=tl.float32)
acc = tl.zeros([16, 512], dtype=tl.float32)
bk_range = tl.arange(0, BK)
pos = lo
while pos < hi:
remaining = hi - pos
mask = bk_range < remaining
kv_offs = pos + bk_range
kv_lat = tl.load(
KV + kv_offs[:, None] * skv + d512[None, :],
mask=mask[:, None],
other=0.0,
).to(tl.bfloat16)
kv_rope = tl.load(
KV + kv_offs[:, None] * skv + 512 + d64[None, :],
mask=mask[:, None],
other=0.0,
).to(tl.bfloat16)
scores = tl.dot(q_lat, tl.trans(kv_lat))
scores += tl.dot(q_rope, tl.trans(kv_rope))
scores *= cs
scores = tl.where(mask[None, :], scores, float("-inf"))
block_max = tl.max(scores, axis=1)
mn = tl.maximum(mi, block_max)
alpha = tl.exp(mi - mn)
exp_s = tl.exp(scores - mn[:, None])
exp_s = tl.where(mask[None, :], exp_s, 0.0)
li = li * alpha + tl.sum(exp_s, axis=1)
acc = acc * alpha[:, None]
acc += tl.dot(exp_s.to(tl.bfloat16), kv_lat)
mi = mn
pos += BK
po_result = (acc * kvs).to(tl.bfloat16)
tl.store(PO + q_off * spo0 + sid * spo1 + h16[:, None] * spo2 + d512[None, :], po_result)
tl.store(PM + q_off * sml0 + sid * sml1 + h16, mi)
tl.store(PL + q_off * sml0 + sid * sml1 + h16, li)
@triton.jit
def _mla_s2(
PO,
PM,
PL,
Out,
QO_ind,
spo0,
spo1,
spo2,
sml0,
sml1,
so0,
so1,
NS: tl.constexpr,
):
hid = tl.program_id(0)
bid = tl.program_id(1)
q_off = tl.load(QO_ind + bid)
mlb = q_off * sml0 + hid
gm = float("-inf")
for s in range(NS):
gm = tl.maximum(gm, tl.load(PM + mlb + s * sml1))
acc = tl.zeros([512], dtype=tl.float32)
lsum = 0.0
for s in range(NS):
ms = tl.load(PM + mlb + s * sml1)
ls = tl.load(PL + mlb + s * sml1)
w = tl.exp(ms - gm)
off = q_off * spo0 + s * spo1 + hid * spo2
po = tl.load(PO + off + tl.arange(0, 512)).to(tl.float32)
acc += w * po
lsum += w * ls
acc = acc / lsum
off_o = q_off * so0 + hid * so1
tl.store(Out + off_o + tl.arange(0, 512), acc.to(tl.bfloat16))
def _triton_decode(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_fp8, kv_scale = kv_data["fp8"]
bs = config["batch_size"]
nh = config["num_heads"]
vd = config["v_head_dim"]
qkd = config["qk_head_dim"]
sm = config["sm_scale"]
tq = q.shape[0]
total_kv = kv_fp8.shape[0]
kv_seq_len = total_kv // bs
out = torch.empty((tq, nh, vd), dtype=torch.bfloat16, device=q.device)
if bs >= 128 and kv_seq_len >= 4096:
BK = 32
ns = 8
nw = 4
elif bs >= 128:
BK = 32
ns = 2
nw = 2
elif bs == 64 and kv_seq_len >= 4096:
BK = 32
ns = 8
nw = 2
elif bs == 64:
BK = 16
ns = 8
nw = 4
elif bs == 32 and kv_seq_len >= 4096:
BK = 32
ns = 16
nw = 2
elif bs == 32:
BK = 16
ns = 8
nw = 4
elif bs <= 8 and kv_seq_len >= 4096:
BK = 32
ns = 32 if bs <= 4 else 16
nw = 4
else:
BK = 16
ns = 16 if bs <= 16 else 8
nw = 4
po = torch.empty((tq, ns, nh, vd), dtype=torch.bfloat16, device=q.device)
pm = torch.empty((tq, ns, nh), dtype=torch.float32, device=q.device)
pl = torch.empty((tq, ns, nh), dtype=torch.float32, device=q.device)
_mla_s1[(bs, ns)](
q,
kv_fp8,
kv_scale,
qo_indptr,
kv_indptr,
po,
pm,
pl,
sm,
q.stride(0),
q.stride(1),
qkd,
po.stride(0),
po.stride(1),
po.stride(2),
pm.stride(0),
pm.stride(1),
NS=ns,
BK=BK,
num_warps=nw,
)
_mla_s2[(nh, bs)](
po,
pm,
pl,
out,
qo_indptr,
po.stride(0),
po.stride(1),
po.stride(2),
pm.stride(0),
pm.stride(1),
out.stride(0),
out.stride(1),
NS=ns,
num_warps=1,
)
return out
def custom_kernel(data: input_t) -> output_t:
_, _, _, _, config = data
if config["batch_size"] == 256 and config["kv_seq_len"] == 8192:
return _persistent_decode(data)
return _triton_decode(data)
scrolls · 392 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