submission 748147
s.k.9151 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 703 lines, June 9 Researcher Reciprocity License v1.0.
submission_hybrid_dot_scaled_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-748147?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:27a27efae49aa2e17f0599667fbdd4625998a20147c88aefdb5b8dcfde81d641
license declaredunknown
license concludedunknown
authorss.k.9151
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
batch_size * kv_seq_len <= 65536 -> Triton dot_scaled (native MXFP4 MFMA)mma
c = tl.dot(pb, v_block.to(tl.bfloat16))persistent-kernel
otherwise -> aiter a8w8 (FP8 persistent kernel)split-k
Triton path: two-stage split-K with tl.dot_scaled for score computationtile-n = 64
BN = 64Kernel source
submission_hybrid_dot_scaled_v1.py703 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
Hybrid MLA decode kernel: Triton dot_scaled for small batch, aiter a8w8 for large batch.
Dispatch logic:
batch_size * kv_seq_len <= 65536 -> Triton dot_scaled (native MXFP4 MFMA)
otherwise -> aiter a8w8 (FP8 persistent kernel)
Triton path: two-stage split-K with tl.dot_scaled for score computation
and manual FP4 dequant for V accumulation.
Aiter path: cached metadata, cached kv_indices, pre-allocated output,
per-config NUM_KV_SPLITS tuning.
"""
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from task import input_t, output_t
from utils import make_match_reference
# ---------------------------------------------------------------------------
# Aiter imports (triggers ~222s JIT build)
# ---------------------------------------------------------------------------
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.utility.fp4_utils import (
dynamic_mxfp4_quant,
mxfp4_to_f32,
e8m0_to_f32,
)
# ---------------------------------------------------------------------------
# Shared constants
# ---------------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM # 576
V_HEAD_DIM = KV_LORA_RANK # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
# ---------------------------------------------------------------------------
# Triton path constants
# ---------------------------------------------------------------------------
NSPLITS_SMALL = 8 # for kv_seq_len <= 1024 (grid=batch*8, optimal for bs=32)
NSPLITS_MED = 4 # for large batch + kv>1024 (grid=batch*4, less overhead)
NSPLITS_LARGE = 64 # for small batch + kv>1024 (grid=batch*64, max parallelism)
PADDED_DIM = 768 # 576 padded to 768 = 3 * 256
PADDED_BYTES = 384 # 768 / 2 = 384 packed fp4x2 bytes
PADDED_SCALES = 24 # 768 / 32 = 24 scale blocks
# ---------------------------------------------------------------------------
# Aiter path constants
# ---------------------------------------------------------------------------
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
Q_DTYPE = "fp8"
KV_DTYPE = "fp8"
# ---------------------------------------------------------------------------
# Hybrid dispatch threshold
# ---------------------------------------------------------------------------
TRITON_THRESHOLD = 65536 # batch_size * kv_seq_len <= this -> use Triton
# ===========================================================================
# TRITON KERNELS
# ===========================================================================
# ---------------------------------------------------------------------------
# FP4 E2M1 dequantization helper (used only in V phase)
# ---------------------------------------------------------------------------
@triton.jit
def _dq4(x):
"""Dequantize FP4 E2M1 nibble (0-15) to float32."""
n = x.to(tl.int32)
e = (n >> 1) & 3
m = (n & 1).to(tl.float32)
val = tl.where(e == 0,
m * 0.5,
(1.0 + m * 0.5) * tl.math.exp2((e - 1).to(tl.float32)))
sg = 1.0 - ((n >> 3) & 1).to(tl.float32) * 2.0
return sg * val
# ---------------------------------------------------------------------------
# Stage 1: Split-K partial attention with native MXFP4 dot_scaled for scores
# ---------------------------------------------------------------------------
@triton.jit
def mla_s1(
Q, # query tensor: (total_q, 16, 576) bf16, row-major (ORIGINAL)
KV, # kv_fp4x2: (total_kv, 288) uint8 (packed fp4x2, ORIGINAL)
Sc, # kv_scales: (total_kv, 18+) uint8 (e8m0, ORIGINAL)
Ind, # kv_indptr: (batch+1,) int32
PO, # partial_out: (batch*nsplits*16, 512) float32
PM, # partial_max: (batch*nsplits*16,) float32
PS, # partial_sum: (batch*nsplits*16,) float32
Q_pad, # padded Q: (total_q, 16, 768) bf16
KV_pad, # padded KV: (total_kv, 384) uint8
Sc_pad, # padded scales: (total_kv, 24) uint8
sm_scale, # softmax scale factor
nsplits: tl.constexpr, # number of splits
sq_b: tl.constexpr, # Q stride for batch dim (= 16 * 576)
sq_h: tl.constexpr, # Q stride for head dim (= 576)
skv: tl.constexpr, # KV stride for token dim (= 288)
ssc_t: tl.constexpr, # Scale stride for token dim
ssc_b: tl.constexpr, # Scale stride (unused, kept for compat)
sq_b_pad: tl.constexpr, # Padded Q stride for batch dim (= 16 * 768)
sq_h_pad: tl.constexpr, # Padded Q stride for head dim (= 768)
skv_pad: tl.constexpr, # Padded KV stride for token dim (= 384)
ssc_pad: tl.constexpr, # Padded Scale stride for token dim (= 24)
BN: tl.constexpr, # tile size for KV tokens (64)
):
bid = tl.program_id(0) # batch index
sid = tl.program_id(1) # split index
# Compute token range for this split
ks = tl.load(Ind + bid)
ke = tl.load(Ind + bid + 1)
kl = ke - ks
cs = tl.cdiv(kl, nsplits)
ms = ks + sid * cs
me = tl.minimum(ms + cs, ke)
si = bid * nsplits + sid # flat split index
hr = tl.arange(0, 16) # head range [0..15]
nr = tl.arange(0, BN) # token range within tile [0..BN-1]
# Early exit for empty splits
if ms >= me:
tl.store(PM + si * 16 + hr, tl.full([16], float('-inf'), tl.float32))
tl.store(PS + si * 16 + hr, tl.zeros([16], tl.float32))
return
# Running softmax state
mi = tl.full([16], float('-inf'), tl.float32) # max logit per head
li = tl.zeros([16], tl.float32) # sum of exp per head
# 16 V accumulators: each (16, 32) for 16 heads x 32 dims = one scale block
a0 = tl.zeros([16, 32], tl.float32)
a1 = tl.zeros([16, 32], tl.float32)
a2 = tl.zeros([16, 32], tl.float32)
a3 = tl.zeros([16, 32], tl.float32)
a4 = tl.zeros([16, 32], tl.float32)
a5 = tl.zeros([16, 32], tl.float32)
a6 = tl.zeros([16, 32], tl.float32)
a7 = tl.zeros([16, 32], tl.float32)
a8 = tl.zeros([16, 32], tl.float32)
a9 = tl.zeros([16, 32], tl.float32)
a10 = tl.zeros([16, 32], tl.float32)
a11 = tl.zeros([16, 32], tl.float32)
a12 = tl.zeros([16, 32], tl.float32)
a13 = tl.zeros([16, 32], tl.float32)
a14 = tl.zeros([16, 32], tl.float32)
a15 = tl.zeros([16, 32], tl.float32)
# Iterate over KV tiles within this split's range
for ts in range(ms, me, BN):
tid = ts + nr # absolute token indices for this tile
vm = tid < me # validity mask
# ---- Score computation using tl.dot_scaled for native MXFP4 MFMA ----
scores = tl.zeros([BN, 16], tl.float32)
for ki in tl.static_range(3):
# K tile: (BN, 128) packed fp4x2 from padded KV
k_tile = tl.load(KV_pad + tid[:, None] * skv_pad + ki * 128 + tl.arange(0, 128)[None, :],
mask=vm[:, None], other=0)
# K scales: (BN, 8) raw e8m0 for this BK=256 chunk
k_sc = tl.load(Sc_pad + tid[:, None] * ssc_pad + ki * 8 + tl.arange(0, 8)[None, :],
mask=vm[:, None], other=127)
# Q^T tile: (256, 16) bf16
d_offs = ki * 256 + tl.arange(0, 256)
h_offs = tl.arange(0, 16)
q_t = tl.load(Q_pad + bid * sq_b_pad + h_offs[None, :] * sq_h_pad + d_offs[:, None],
mask=True, other=0.0).to(tl.bfloat16)
# Native MXFP4 x BF16 dot product via hardware MFMA
scores += tl.dot_scaled(k_tile, k_sc, "e2m1", q_t, None, "bf16")
# Transpose scores: (BN, 16) -> (16, BN) for softmax per head
sc = tl.trans(scores) * sm_scale
# ---- Online softmax ----
sc = tl.where(vm[None, :], sc, float('-inf'))
# Per-head max for this tile
tile_max = tl.max(sc, axis=1)
new_mi = tl.maximum(mi, tile_max)
# Correction factor for previous accumulators
alpha = tl.math.exp2((mi - new_mi) * 1.4426950408889634)
# Softmax weights for this tile
p = tl.math.exp2((sc - new_mi[:, None]) * 1.4426950408889634)
tile_sum = tl.sum(p, axis=1)
# Update running state
li = li * alpha + tile_sum
mi = new_mi
# Rescale previous accumulators
a0 = a0 * alpha[:, None]
a1 = a1 * alpha[:, None]
a2 = a2 * alpha[:, None]
a3 = a3 * alpha[:, None]
a4 = a4 * alpha[:, None]
a5 = a5 * alpha[:, None]
a6 = a6 * alpha[:, None]
a7 = a7 * alpha[:, None]
a8 = a8 * alpha[:, None]
a9 = a9 * alpha[:, None]
a10 = a10 * alpha[:, None]
a11 = a11 * alpha[:, None]
a12 = a12 * alpha[:, None]
a13 = a13 * alpha[:, None]
a14 = a14 * alpha[:, None]
a15 = a15 * alpha[:, None]
# ---- V accumulation: p @ V for first 512 dims = 16 blocks of 32 ----
pb = p.to(tl.bfloat16)
for vb in tl.static_range(16):
# Load packed fp4x2 for V block vb: (BN, 16) uint8
v_raw = tl.load(KV + (tid[:, None] * skv + vb * 16 + tl.arange(0, 16)[None, :]),
mask=vm[:, None], other=0)
# Unpack
vlo = (v_raw & 0xF).to(tl.uint8)
vhi = (v_raw >> 4).to(tl.uint8)
# Dequant
fvlo = _dq4(vlo)
fvhi = _dq4(vhi)
# Load e8m0 scale
vsc_raw = tl.load(Sc + tid * ssc_t + vb, mask=vm, other=0)
vsc_exp = (vsc_raw.to(tl.int32) - 127).to(tl.float32)
vsc_f = tl.math.exp2(vsc_exp)
fvlo = fvlo * vsc_f[:, None]
fvhi = fvhi * vsc_f[:, None]
# Interleave even/odd to reconstruct 32 contiguous values per token
v_joined = tl.join(fvlo, fvhi)
v_block = tl.reshape(v_joined, [BN, 32])
# p @ V_block: (16, BN) @ (BN, 32) -> (16, 32)
c = tl.dot(pb, v_block.to(tl.bfloat16))
if vb == 0: a0 += c.to(tl.float32)
if vb == 1: a1 += c.to(tl.float32)
if vb == 2: a2 += c.to(tl.float32)
if vb == 3: a3 += c.to(tl.float32)
if vb == 4: a4 += c.to(tl.float32)
if vb == 5: a5 += c.to(tl.float32)
if vb == 6: a6 += c.to(tl.float32)
if vb == 7: a7 += c.to(tl.float32)
if vb == 8: a8 += c.to(tl.float32)
if vb == 9: a9 += c.to(tl.float32)
if vb == 10: a10 += c.to(tl.float32)
if vb == 11: a11 += c.to(tl.float32)
if vb == 12: a12 += c.to(tl.float32)
if vb == 13: a13 += c.to(tl.float32)
if vb == 14: a14 += c.to(tl.float32)
if vb == 15: a15 += c.to(tl.float32)
# ---- Store partial results ----
tl.store(PM + si * 16 + hr, mi)
tl.store(PS + si * 16 + hr, li)
# Store partial output: (16 heads, 512 dims) as 16 blocks of 32
po_rows = (si * 16 + hr)[:, None] * 512
dr = tl.arange(0, 32)
tl.store(PO + po_rows + 0 * 32 + dr[None, :], a0)
tl.store(PO + po_rows + 1 * 32 + dr[None, :], a1)
tl.store(PO + po_rows + 2 * 32 + dr[None, :], a2)
tl.store(PO + po_rows + 3 * 32 + dr[None, :], a3)
tl.store(PO + po_rows + 4 * 32 + dr[None, :], a4)
tl.store(PO + po_rows + 5 * 32 + dr[None, :], a5)
tl.store(PO + po_rows + 6 * 32 + dr[None, :], a6)
tl.store(PO + po_rows + 7 * 32 + dr[None, :], a7)
tl.store(PO + po_rows + 8 * 32 + dr[None, :], a8)
tl.store(PO + po_rows + 9 * 32 + dr[None, :], a9)
tl.store(PO + po_rows + 10 * 32 + dr[None, :], a10)
tl.store(PO + po_rows + 11 * 32 + dr[None, :], a11)
tl.store(PO + po_rows + 12 * 32 + dr[None, :], a12)
tl.store(PO + po_rows + 13 * 32 + dr[None, :], a13)
tl.store(PO + po_rows + 14 * 32 + dr[None, :], a14)
tl.store(PO + po_rows + 15 * 32 + dr[None, :], a15)
# ---------------------------------------------------------------------------
# Stage 2: Reduce partial results across splits
# ---------------------------------------------------------------------------
@triton.jit
def mla_s2(
PO, # partial_out: (batch*nsplits*16, 512) float32
PM, # partial_max: (batch*nsplits*16,) float32
PS, # partial_sum: (batch*nsplits*16,) float32
Out, # output: (batch, 16, 512) bf16
nsplits: tl.constexpr,
so_b: tl.constexpr, # output stride for batch (= 16 * 512)
so_h: tl.constexpr, # output stride for head (= 512)
):
bid = tl.program_id(0) # batch index
hid = tl.program_id(1) # head index [0..15]
# Find global max across all splits for this (batch, head)
global_max = tl.full([], float('-inf'), tl.float32)
for s in range(nsplits):
si = bid * nsplits + s
m = tl.load(PM + si * 16 + hid)
global_max = tl.maximum(global_max, m)
# Compute weighted sum with log-sum-exp correction
global_sum = tl.zeros([], tl.float32)
acc = tl.zeros([512], tl.float32)
dr = tl.arange(0, 512)
for s in range(nsplits):
si = bid * nsplits + s
m = tl.load(PM + si * 16 + hid)
l = tl.load(PS + si * 16 + hid)
# Correction weight: exp(m_split - m_global)
alpha = tl.math.exp2((m - global_max) * 1.4426950408889634)
w = alpha * l
# Load partial output row
po_base = (si * 16 + hid) * 512
pv = tl.load(PO + po_base + dr)
acc += alpha * pv
global_sum += w
# Normalize and store
acc = acc / global_sum
out_base = bid * so_b + hid * so_h
tl.store(Out + out_base + dr, acc.to(tl.bfloat16))
# ===========================================================================
# TRITON PATH: buffer caches and entry point
# ===========================================================================
_triton_buf_cache: dict = {}
def _get_triton_buffers(batch_size: int, nsplits: int):
"""Return (or allocate) partial output/max/sum buffers for Triton path."""
key = (batch_size, nsplits)
if key not in _triton_buf_cache:
total_splits = batch_size * nsplits
total_rows = total_splits * NUM_HEADS
po = torch.empty((total_rows, V_HEAD_DIM), dtype=torch.float32, device="cuda")
pm = torch.empty((total_rows,), dtype=torch.float32, device="cuda")
ps = torch.empty((total_rows,), dtype=torch.float32, device="cuda")
_triton_buf_cache[key] = (po, pm, ps)
return _triton_buf_cache[key]
_triton_out_cache: dict = {}
def _get_triton_output(batch_size: int):
"""Return (or allocate) output tensor for Triton path."""
if batch_size not in _triton_out_cache:
_triton_out_cache[batch_size] = torch.empty(
(batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"
)
return _triton_out_cache[batch_size]
_triton_pad_cache: dict = {}
def _get_padded_tensors(q, kv_flat, sc_flat, batch_size, total_kv):
"""Pad Q (576->768), KV (288->384), scales (18->24) for dot_scaled."""
key = (batch_size, total_kv)
if key not in _triton_pad_cache:
# Pad KV: (total_kv, 288) -> (total_kv, 384) uint8, zeros for padding
kv_padded = torch.zeros((total_kv, PADDED_BYTES), dtype=torch.uint8, device="cuda")
kv_padded[:, :288] = kv_flat
# Pad scales: (total_kv, 18) -> (total_kv, 24) uint8, 127 = neutral e8m0 (2^0)
sc_padded = torch.full((total_kv, PADDED_SCALES), 127, dtype=torch.uint8, device="cuda")
sc_cols = sc_flat.shape[1]
sc_padded[:, :sc_cols] = sc_flat[:, :sc_cols]
# Pad Q: (batch, 16, 576) -> (batch, 16, 768) bf16, zeros for padding
q_padded = torch.zeros((batch_size, NUM_HEADS, PADDED_DIM), dtype=torch.bfloat16, device="cuda")
q_padded[:, :, :QK_HEAD_DIM] = q[:batch_size]
_triton_pad_cache[key] = (q_padded, kv_padded, sc_padded)
else:
q_padded, kv_padded, sc_padded = _triton_pad_cache[key]
# Update with current data (cache is for allocation reuse)
kv_padded[:, :288] = kv_flat
sc_cols = sc_flat.shape[1]
sc_padded[:, :sc_cols] = sc_flat[:, :sc_cols]
q_padded[:, :, :QK_HEAD_DIM] = q[:batch_size]
return q_padded, kv_padded, sc_padded
def _triton_dot_scaled_path(q, kv_data, kv_indptr, config):
"""Triton dot_scaled MLA decode for small batch * kv_seq_len cases."""
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
# Unpack MXFP4 KV cache
kv_fp4x2, kv_scales = kv_data["mxfp4"]
# View as uint8 for Triton pointer arithmetic
kv_u8 = kv_fp4x2.view(torch.uint8)
sc_u8 = kv_scales.view(torch.uint8)
# Flatten KV to 2D: (total_kv, 288)
kv_flat = kv_u8.reshape(kv_u8.shape[0], -1)
sc_flat = sc_u8
total_kv = kv_flat.shape[0]
# Strides for original tensors
skv = kv_flat.stride(0)
ssc_t = sc_flat.stride(0)
# Q strides
sq_b = q.stride(0) * 1
sq_h = q.stride(1)
# Pad tensors for dot_scaled
q_padded, kv_padded, sc_padded = _get_padded_tensors(q, kv_flat, sc_flat, batch_size, total_kv)
# Padded strides
sq_b_pad = q_padded.stride(0)
sq_h_pad = q_padded.stride(1)
skv_pad = kv_padded.stride(0)
ssc_pad = sc_padded.stride(0)
# Select nsplits: 3-way dispatch for optimal grid size
if kv_seq_len <= 1024:
nsplits = NSPLITS_SMALL
elif batch_size >= 64:
nsplits = NSPLITS_MED
else:
nsplits = NSPLITS_LARGE
# Allocate/reuse buffers
po, pm, ps = _get_triton_buffers(batch_size, nsplits)
out = _get_triton_output(batch_size)
BN = 64
# Stage 1: compute partial attention
grid_s1 = (batch_size, nsplits)
mla_s1[grid_s1](
q, kv_flat, sc_flat, kv_indptr,
po, pm, ps,
q_padded, kv_padded, sc_padded,
SM_SCALE,
nsplits=nsplits,
sq_b=sq_b,
sq_h=sq_h,
skv=skv,
ssc_t=ssc_t,
ssc_b=0,
sq_b_pad=sq_b_pad,
sq_h_pad=sq_h_pad,
skv_pad=skv_pad,
ssc_pad=ssc_pad,
BN=BN,
)
# Stage 2: reduce across splits
grid_s2 = (batch_size, NUM_HEADS)
mla_s2[grid_s2](
po, pm, ps, out,
nsplits=nsplits,
so_b=out.stride(0),
so_h=out.stride(1),
)
return out
# ===========================================================================
# AITER PATH: caches, helpers, and entry point
# ===========================================================================
# Per-config NUM_KV_SPLITS tuning table
_SPLIT_TABLE = {
(4, 1024): 4,
(4, 8192): 16,
(32, 1024): 4,
(32, 8192): 16,
(64, 1024): 4,
(64, 8192): 16,
(256, 1024): 8,
(256, 8192): 16,
}
_DEFAULT_NUM_KV_SPLITS = 32
# Metadata buffer cache
_aiter_metadata_cache: dict = {}
# kv_indices cache
_aiter_kv_indices_cache: dict = {}
# Output tensor cache
_aiter_output_cache: dict = {}
def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Dynamic per-tensor FP8 quantization."""
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
return fp8_tensor, scale.to(torch.float32).reshape(1)
def _make_mla_decode_metadata(
batch_size, max_q_len, nhead, nhead_kv,
q_dtype, kv_dtype,
qo_indptr, kv_indptr, kv_last_page_len,
num_kv_splits=_DEFAULT_NUM_KV_SPLITS,
):
"""Allocate and populate work buffers for persistent mla_decode_fwd."""
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(s, dtype=t, device="cuda") for s, t 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 _get_num_kv_splits(batch_size: int, kv_len_per_seq: int) -> int:
return _SPLIT_TABLE.get((batch_size, kv_len_per_seq), _DEFAULT_NUM_KV_SPLITS)
def _get_cached_kv_indices(total_kv_len: int) -> torch.Tensor:
if total_kv_len not in _aiter_kv_indices_cache:
_aiter_kv_indices_cache[total_kv_len] = torch.arange(
total_kv_len, dtype=torch.int32, device="cuda"
)
return _aiter_kv_indices_cache[total_kv_len]
def _get_cached_output(total_q: int, nq: int, dv: int) -> torch.Tensor:
key = (total_q, nq, dv)
cached = _aiter_output_cache.get(key)
if cached is None:
cached = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
_aiter_output_cache[key] = cached
return cached
def _get_cached_metadata(
batch_size, max_q_len, total_kv_len, nhead, nhead_kv,
q_dtype, kv_dtype,
qo_indptr, kv_indptr, kv_last_page_len,
num_kv_splits,
):
cache_key = (batch_size, max_q_len, total_kv_len, q_dtype, kv_dtype, num_kv_splits)
if cache_key not in _aiter_metadata_cache:
_aiter_metadata_cache[cache_key] = _make_mla_decode_metadata(
batch_size, max_q_len, nhead, nhead_kv,
q_dtype, kv_dtype,
qo_indptr, kv_indptr, kv_last_page_len,
num_kv_splits=num_kv_splits,
)
return _aiter_metadata_cache[cache_key]
def _aiter_mla_decode(
q, kv_buffer, qo_indptr, kv_indptr, config,
q_scale=None, kv_scale=None,
):
"""MLA decode attention using aiter persistent-mode kernel."""
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"]
total_kv_len = int(kv_indptr[-1].item())
kv_buffer_4d = kv_buffer.view(kv_buffer.shape[0], PAGE_SIZE, nkv, kv_buffer.shape[-1])
max_q_len = q_seq_len
kv_len_per_seq = total_kv_len // batch_size
num_kv_splits = _get_num_kv_splits(batch_size, kv_len_per_seq)
kv_indices = _get_cached_kv_indices(total_kv_len)
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
meta = _get_cached_metadata(
batch_size, max_q_len, total_kv_len, nq, nkv,
q.dtype, kv_buffer.dtype,
qo_indptr, kv_indptr, kv_last_page_len,
num_kv_splits=num_kv_splits,
)
o = _get_cached_output(q.shape[0], nq, dv)
mla_decode_fwd(
q.view(-1, nq, dq),
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=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
**meta,
)
return o
def _aiter_path(q, kv_data, qo_indptr, kv_indptr, config):
"""Aiter a8w8 FP8 path for large batch cases."""
# Quantize Q to FP8
q_input, q_scale = quantize_fp8(q)
# Use FP8 KV cache
kv_buffer_fp8, kv_scale = kv_data["fp8"]
return _aiter_mla_decode(
q_input, kv_buffer_fp8, qo_indptr, kv_indptr, config,
q_scale=q_scale, kv_scale=kv_scale,
)
# ===========================================================================
# HYBRID DISPATCH
# ===========================================================================
def custom_kernel(data: input_t) -> output_t:
"""
Hybrid MLA decode: Triton dot_scaled for small workloads, aiter a8w8 for large.
Dispatch threshold: batch_size * kv_seq_len <= 65536 -> Triton
This covers: bs=4/kv=1k, bs=4/kv=8k, bs=32/kv=1k, bs=64/kv=1k
Aiter handles: bs=32/kv=8k, bs=64/kv=8k, bs=256/kv=1k, bs=256/kv=8k
"""
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
workload = batch_size * kv_seq_len
if workload <= TRITON_THRESHOLD:
return _triton_dot_scaled_path(q, kv_data, kv_indptr, config)
else:
return _aiter_path(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 703 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