submission 694789
DiegoCao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 207 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-694789?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:0a2a04ce0ef704e140f6184e02948c120b9455c741db5d6697f15b0e5500a8d5
license declaredunknown
license concludedunknown
authorsDiegoCao
imported2026-08-26
Kernel source
submission.py207 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA Decode v10: Single-launch FP8 quant for small tensors.
For small batch sizes (bs=4, Q has only 36K elements), the ENTIRE Q tensor
fits in a single Triton thread block. This enables fusing amax reduction +
quantization into 1 launch instead of 3 (zero + amax + quant).
For large tensors: fall back to 2-launch approach (atomic amax + quant).
Direct ASM kernel calls with pre-allocated buffers for attention.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_FP8 = aiter_dtypes.fp8
_FP8_MAX = float(torch.finfo(_FP8).max)
_SM = 1.0 / (576.0 ** 0.5)
# ---------------------------------------------------------------------------
# Single-block fused amax+quant kernel (1 launch for small tensors)
# ---------------------------------------------------------------------------
@triton.jit
def _fused_quant_small_kernel(
X_ptr, Out_ptr, Scale_ptr,
N: tl.constexpr, FP8_MAX: tl.constexpr, BLOCK: tl.constexpr,
):
"""Single-block: compute amax + quantize in one pass."""
off = tl.arange(0, BLOCK)
mask = off < N
x = tl.load(X_ptr + off, mask=mask, other=0.0).to(tl.float32)
# Global amax across entire block
amax = tl.max(tl.abs(x))
amax = tl.maximum(amax, 1e-12)
inv_scale = FP8_MAX / amax
tl.store(Scale_ptr, amax / FP8_MAX)
# Quantize
q = x * inv_scale
q = tl.minimum(tl.maximum(q, -FP8_MAX), FP8_MAX)
tl.store(Out_ptr + off, q.to(Out_ptr.dtype.element_ty), mask=mask)
# ---------------------------------------------------------------------------
# Multi-block atomic amax + quant (2 launches for large tensors)
# ---------------------------------------------------------------------------
_BLOCK_LARGE = 4096
@triton.jit
def _amax_kernel(X_ptr, Amax_ptr, N: tl.constexpr, BLOCK: tl.constexpr):
pid = tl.program_id(0)
off = pid * BLOCK + tl.arange(0, BLOCK)
mask = off < N
x = tl.load(X_ptr + off, mask=mask, other=0.0).to(tl.float32)
local_max = tl.max(tl.abs(x))
tl.atomic_max(Amax_ptr, local_max)
@triton.jit
def _quant_kernel(
X_ptr, Out_ptr, Amax_ptr, Scale_ptr, N: tl.constexpr,
FP8_MAX: tl.constexpr, BLOCK: tl.constexpr,
):
pid = tl.program_id(0)
off = pid * BLOCK + tl.arange(0, BLOCK)
mask = off < N
amax = tl.load(Amax_ptr).to(tl.float32)
amax = tl.maximum(amax, 1e-12)
inv_scale = FP8_MAX / amax
if pid == 0:
tl.store(Scale_ptr, amax / FP8_MAX)
x = tl.load(X_ptr + off, mask=mask, other=0.0).to(tl.float32)
q = x * inv_scale
q = tl.minimum(tl.maximum(q, -FP8_MAX), FP8_MAX)
tl.store(Out_ptr + off, q.to(Out_ptr.dtype.element_ty), mask=mask)
# ---------------------------------------------------------------------------
# Quantization dispatch
# ---------------------------------------------------------------------------
# Threshold for single-block kernel: N must fit in one block
# Max block size ~65536 elements (16 warps, safe for AMD CDNA4)
_SINGLE_BLOCK_MAX = 65536
_quant_bufs: dict = {}
def _quant_q(q):
"""FP8 quantize Q: 1 launch for small tensors, 3 launches for large."""
N = q.numel()
c = _quant_bufs.get(N)
if c is None:
c = {
"out": torch.empty(N, dtype=_FP8, device="cuda"),
"amax": torch.zeros(1, dtype=torch.float32, device="cuda"),
"scale": torch.empty(1, dtype=torch.float32, device="cuda"),
}
_quant_bufs[N] = c
if N <= _SINGLE_BLOCK_MAX:
# Single-block fused kernel: 1 launch!
BLOCK = triton.next_power_of_2(N)
_fused_quant_small_kernel[(1,)](
q, c["out"], c["scale"],
N=N, FP8_MAX=_FP8_MAX, BLOCK=BLOCK,
num_warps=min(16, max(1, BLOCK // 256)),
)
else:
# Multi-block: 3 launches (zero + amax + quant)
c["amax"].zero_()
grid = ((N + _BLOCK_LARGE - 1) // _BLOCK_LARGE,)
_amax_kernel[grid](q, c["amax"], N=N, BLOCK=_BLOCK_LARGE)
_quant_kernel[grid](q, c["out"], c["amax"], c["scale"], N=N, FP8_MAX=_FP8_MAX, BLOCK=_BLOCK_LARGE)
return c["out"].view(q.shape), c["scale"]
# ---------------------------------------------------------------------------
# Attention state cache
# ---------------------------------------------------------------------------
_SPLITS = {
(4, 1024): 8, (4, 8192): 16,
(32, 1024): 16, (32, 8192): 32,
(64, 1024): 16, (64, 8192): 32,
(256, 1024): 16, (256, 8192): 32,
}
_cache: dict = {}
class _State:
__slots__ = (
'nks', 'kv_indices', 'kv_lpl', 'output',
'wm', 'wi', 'wis', 'ri', 'rfm', 'rpm',
'logits', 'attn_lse', 'final_lse',
)
def __init__(self, bs, kvlen, q_dtype, kv_dtype, qo_indptr, kv_indptr):
total_kv = bs * kvlen
self.nks = _SPLITS.get((bs, kvlen), 32)
self.kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
self.kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device="cuda")
self.output = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")
info = get_mla_metadata_info_v1(
bs, 1, 16, q_dtype, kv_dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=self.nks, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
self.wm, self.wi, self.wis, self.ri, self.rfm, self.rpm = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, self.kv_lpl,
16, 1, True,
self.wm, self.wis, self.wi, self.ri, self.rfm, self.rpm,
page_size=1, kv_granularity=16,
max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=False, max_split_per_batch=self.nks,
intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
rpm_size = self.rpm.size(0)
self.logits = torch.empty((rpm_size, 1, 16, 512), dtype=torch.float32, device="cuda")
self.attn_lse = torch.empty((rpm_size, 1, 16, 1), dtype=torch.float32, device="cuda")
self.final_lse = torch.empty((bs, 16), dtype=torch.float32, device="cuda")
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kvlen = config["kv_seq_len"]
kv_buf, kv_sc = kv_data["fp8"]
q_fp8, q_sc = _quant_q(q)
key = (bs, kvlen)
s = _cache.get(key)
if s is None:
s = _State(bs, kvlen, q_fp8.dtype, kv_buf.dtype, qo_indptr, kv_indptr)
_cache[key] = s
# Direct ASM kernel calls with pre-allocated buffers
aiter.mla_decode_stage1_asm_fwd(
q_fp8.view(-1, 16, 576),
kv_buf.view(-1, 1, 1, 576),
qo_indptr, kv_indptr,
s.kv_indices, s.kv_lpl,
None, s.wm, s.wi, s.wis,
1, 1, 1, _SM,
s.logits, s.attn_lse, s.output,
q_sc, kv_sc,
)
aiter.mla_reduce_v1(
s.logits, s.attn_lse,
s.ri, s.rfm, s.rpm,
1, s.output, s.final_lse,
)
return s.output
scrolls · 207 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