submission 668995
Rakesh Jarupula · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 364 lines, June 9 Researcher Reciprocity License v1.0.
amd_mixed_mla.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-668995?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:2afeadc4bcc73a5887f7523cf91f018abf3788dd8818dd758b7c8167f08e31f4
license declaredunknown
license concludedunknown
authorsRakesh Jarupula
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
online-softmax
m_new = tl.maximum(m_i, score)split-k
- Triton split-K: grid=(B, H, S), each CTA handles a KV slicetile-n = 8
BN = 8 # inner unroll factorKernel source
amd_mixed_mla.py364 lines
"""
Custom MLA (Multi-head Latent Attention) decode kernel optimized for MI355X (CDNA3).
Key design:
- Triton split-K: grid=(B, H, S), each CTA handles a KV slice
- Separate K (576-dim) and V (512-dim) loads per KV token — avoids shape mismatch
- FP8 KV: scalar dequant inline (multiply by kv_scale)
- Online softmax (running m, l) — Flash Attention style
- Reduction kernel merges splits via LSE (numerically stable)
- PyTorch bmm fallback for safety
DeepSeek R1 forward_absorb MLA:
num_heads=16, num_kv_heads=1 (MQA), qk_head_dim=576, v_head_dim=512
Decode: q_seq_len=1, kv_seq_len up to 8k
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from utils import make_match_reference
# -----------------------------------------------------------------------
# Constants
# -----------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576 # kv_lora_rank(512) + qk_rope_head_dim(64)
V_HEAD_DIM = 512 # = kv_lora_rank
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
_NUM_KV_SPLITS = 32
# -----------------------------------------------------------------------
# Triton kernel 1/2: per-split attention with online softmax
#
# Grid = (batch_size, H, S)
#
# Notes on K vs V:
# KV buffer layout: (total_kv, 1, D) where D=576
# K uses all D=576 dims for score computation
# V uses first Dv=512 dims for output accumulation
# We load K with tile BD=1024 (masked to D=576)
# We load V with tile BDv=512 (masked to Dv=512) from same ptr
# -----------------------------------------------------------------------
@triton.jit
def _mla_split_fwd(
Q_ptr, # (total_q, H, D) bf16
KV_ptr, # (total_kv, 1, D) fp8
kv_scale_ptr, # () scalar f32
O_part_ptr, # (B, H, S, Dv) f32
LSE_part_ptr, # (B, H, S) f32
QO_Indptr, # (B+1,) i32
KV_Indptr, # (B+1,) i32
# Q strides
sq0, sq1, sq2,
# KV strides
skv0, skv1, skv2,
# O_part strides
so0, so1, so2, so3,
# LSE_part strides
sl0, sl1, sl2,
# Compile-time constants
D: tl.constexpr, # 576
Dv: tl.constexpr, # 512
BD: tl.constexpr, # >= D, power of 2 (1024)
BDv: tl.constexpr, # >= Dv, power of 2 (512)
BN: tl.constexpr, # KV tokens per inner unroll
S: tl.constexpr, # num_kv_splits
SM: tl.constexpr, # sm_scale (float literal)
):
b = tl.program_id(0)
h = tl.program_id(1)
s = tl.program_id(2)
# Decode: 1 query token per batch element
qt = tl.load(QO_Indptr + b)
kv0 = tl.load(KV_Indptr + b)
kv1 = tl.load(KV_Indptr + b + 1)
kv_len = kv1 - kv0
per_split = tl.cdiv(kv_len, S)
s_kv0 = kv0 + s * per_split
s_kv1 = tl.minimum(s_kv0 + per_split, kv1)
lse_ptr = LSE_part_ptr + b * sl0 + h * sl1 + s * sl2
if s_kv0 >= s_kv1:
tl.store(lse_ptr, float("-inf"))
return
# Load Q (D=576 dims) as f32 — tile BD=1024 with mask
d_idx = tl.arange(0, BD)
q = tl.load(Q_ptr + qt * sq0 + h * sq1 + d_idx * sq2,
mask=d_idx < D, other=0.0).to(tl.float32)
# V index for separate V load (BDv=512 tile)
v_idx = tl.arange(0, BDv)
kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
# Online softmax state
m_i = float("-inf")
l_i = 0.0
acc = tl.zeros([BDv], dtype=tl.float32)
kv_tok = s_kv0
while kv_tok < s_kv1:
for i in tl.static_range(0, BN):
tok = kv_tok + i
if tok < s_kv1:
kv_base = KV_ptr + tok * skv0 + 0 * skv1
# Load K: full D=576 dims (tile BD=1024, mask D)
k = tl.load(kv_base + d_idx * skv2,
mask=d_idx < D, other=0.0).to(tl.float32) * kv_scale
# Score: dot(q, k) * sm_scale
score = tl.sum(q * k) * SM
# Online softmax update
m_new = tl.maximum(m_i, score)
e = tl.exp(score - m_new)
r = tl.exp(m_i - m_new)
l_i = l_i * r + e
acc = acc * r
# Load V: first Dv=512 dims (tile BDv=512, mask Dv)
# Same kv_base pointer, different index tile
v = tl.load(kv_base + v_idx * skv2,
mask=v_idx < Dv, other=0.0).to(tl.float32) * kv_scale
acc = acc + e * v
m_i = m_new
kv_tok += BN
acc = acc / tl.maximum(l_i, 1e-8)
o_base = O_part_ptr + b * so0 + h * so1 + s * so2
tl.store(o_base + v_idx * so3, acc, mask=v_idx < Dv)
tl.store(lse_ptr, m_i + tl.log(tl.maximum(l_i, 1e-8)))
# -----------------------------------------------------------------------
# Triton kernel 2/2: reduction across splits
# Grid = (B, H)
# -----------------------------------------------------------------------
@triton.jit
def _mla_reduce(
O_part_ptr, # (B, H, S, Dv) f32
LSE_part_ptr, # (B, H, S) f32
OUT_ptr, # (total_q, H, Dv) bf16
QO_Indptr, # (B+1,) i32
so0, so1, so2, so3,
sl0, sl1, sl2,
ot0, ot1, ot2,
Dv: tl.constexpr,
BDv: tl.constexpr,
S: tl.constexpr,
):
b = tl.program_id(0)
h = tl.program_id(1)
qt = tl.load(QO_Indptr + b)
v_idx = tl.arange(0, BDv)
# Find global max LSE
m_g = float("-inf")
for s in tl.static_range(0, S):
m_g = tl.maximum(m_g, tl.load(LSE_part_ptr + b * sl0 + h * sl1 + s * sl2))
# Weighted sum
acc = tl.zeros([BDv], dtype=tl.float32)
l_tot = 0.0
for s in tl.static_range(0, S):
lse_s = tl.load(LSE_part_ptr + b * sl0 + h * sl1 + s * sl2)
w = tl.exp(lse_s - m_g)
l_tot = l_tot + w
o_s = tl.load(O_part_ptr + b * so0 + h * so1 + s * so2 + v_idx * so3,
mask=v_idx < Dv, other=0.0)
acc = acc + w * o_s
acc = acc / tl.maximum(l_tot, 1e-8)
tl.store(OUT_ptr + qt * ot0 + h * ot1 + v_idx * ot2,
acc.to(tl.bfloat16), mask=v_idx < Dv)
# -----------------------------------------------------------------------
# Triton dispatch
# -----------------------------------------------------------------------
def _triton_decode_fp8(
q: torch.Tensor,
kv_fp8: torch.Tensor,
kv_scale: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
S: int,
) -> torch.Tensor:
B = config["batch_size"]
H = config["num_heads"]
D = config["qk_head_dim"] # 576
Dv = config["v_head_dim"] # 512
SM = float(config["sm_scale"])
BD = triton.next_power_of_2(D) # 1024
BDv = triton.next_power_of_2(Dv) # 512
BN = 8 # inner unroll factor
O_part = torch.empty((B, H, S, Dv), dtype=torch.float32, device=q.device)
LSE_part = torch.full( (B, H, S), float("-inf"), dtype=torch.float32, device=q.device)
_mla_split_fwd[(B, H, S)](
q, kv_fp8, kv_scale,
O_part, LSE_part,
qo_indptr, kv_indptr,
q.stride(0), q.stride(1), q.stride(2),
kv_fp8.stride(0), kv_fp8.stride(1), kv_fp8.stride(2),
O_part.stride(0), O_part.stride(1), O_part.stride(2), O_part.stride(3),
LSE_part.stride(0), LSE_part.stride(1), LSE_part.stride(2),
D=D, Dv=Dv, BD=BD, BDv=BDv, BN=BN, S=S, SM=SM,
)
output = torch.empty((q.shape[0], H, Dv), dtype=torch.bfloat16, device=q.device)
_mla_reduce[(B, H)](
O_part, LSE_part, output, qo_indptr,
O_part.stride(0), O_part.stride(1), O_part.stride(2), O_part.stride(3),
LSE_part.stride(0), LSE_part.stride(1), LSE_part.stride(2),
output.stride(0), output.stride(1), output.stride(2),
Dv=Dv, BDv=BDv, S=S,
)
return output
# -----------------------------------------------------------------------
# PyTorch fallback: FP8 KV + batch-parallel bmm
# -----------------------------------------------------------------------
def _torch_decode_fp8(
q: torch.Tensor,
kv_fp8: torch.Tensor,
kv_scale: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
) -> torch.Tensor:
B = config["batch_size"]
H = config["num_heads"]
D = config["qk_head_dim"]
Dv = config["v_head_dim"]
sm = float(config["sm_scale"])
sc = float(kv_scale.item())
# Dequantize all KV at once
kv = kv_fp8.to(torch.float32).mul_(sc).to(torch.bfloat16) # (total_kv, 1, D)
total_q = q.shape[0]
out = torch.zeros((total_q, H, Dv), dtype=torch.bfloat16, device=q.device)
for b in range(B):
qs = int(qo_indptr[b]); qe = int(qo_indptr[b + 1])
ks = int(kv_indptr[b]); ke = int(kv_indptr[b + 1])
q_b = q[qs:qe] # (Lq, H, D) bf16
k_b = kv[ks:ke, 0, :] # (Lkv, D)
v_b = k_b[:, :Dv] # (Lkv, Dv)
qt = q_b.permute(1, 0, 2).float() # (H, Lq, D)
kt = k_b.unsqueeze(0).expand(H, -1, -1).float() # (H, Lkv, D)
sc_mat = torch.bmm(qt, kt.transpose(1, 2)).mul_(sm) # (H, Lq, Lkv)
attn = torch.softmax(sc_mat, dim=-1).to(torch.bfloat16)
vt = v_b.unsqueeze(0).expand(H, -1, -1) # (H, Lkv, Dv)
out[qs:qe] = torch.bmm(attn, vt).permute(1, 0, 2)
return out
# -----------------------------------------------------------------------
# PyTorch fallback: MXFP4 KV (dequant + bmm)
# -----------------------------------------------------------------------
def _torch_decode_mxfp4(
q: torch.Tensor,
kv_fp4: torch.Tensor,
kv_scale_e8m0: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
) -> torch.Tensor:
from aiter.utility.fp4_utils import mxfp4_to_f32, e8m0_to_f32
B = config["batch_size"]
H = config["num_heads"]
D = config["qk_head_dim"]
Dv = config["v_head_dim"]
sm = float(config["sm_scale"])
total_kv = kv_fp4.shape[0]
nb = D // 32 # 576/32 = 18 scale blocks
kv_f32 = mxfp4_to_f32(kv_fp4.reshape(total_kv, D // 2)) # (total_kv, D)
sc = e8m0_to_f32(kv_scale_e8m0)[:total_kv, :nb] # (total_kv, 18)
kv_f32 = (kv_f32.view(total_kv, nb, 32) * sc.unsqueeze(-1)).view(total_kv, D)
kv_bf = kv_f32.to(torch.bfloat16)
total_q = q.shape[0]
out = torch.zeros((total_q, H, Dv), dtype=torch.bfloat16, device=q.device)
for b in range(B):
qs = int(qo_indptr[b]); qe = int(qo_indptr[b + 1])
ks = int(kv_indptr[b]); ke = int(kv_indptr[b + 1])
q_b = q[qs:qe]
k_b = kv_bf[ks:ke]
v_b = k_b[:, :Dv]
qt = q_b.permute(1, 0, 2).float()
kt = k_b.unsqueeze(0).expand(H, -1, -1).float()
scores = torch.bmm(qt, kt.transpose(1, 2)).mul_(sm)
attn = torch.softmax(scores, dim=-1).to(torch.bfloat16)
vt = v_b.unsqueeze(0).expand(H, -1, -1)
out[qs:qe] = torch.bmm(attn, vt).permute(1, 0, 2)
return out
# -----------------------------------------------------------------------
# Entry point
# -----------------------------------------------------------------------
_use_triton: bool = True
def custom_kernel(data: input_t) -> output_t:
"""
MLA decode kernel — DeepSeek R1 forward_absorb path.
1. Triton split-K + FP8 KV (primary — maximizes GPU occupancy)
2. PyTorch bmm + FP8 KV (fallback)
"""
global _use_triton
q, kv_data, qo_indptr, kv_indptr, config = data
kv_fp8, kv_scale = kv_data["fp8"]
if _use_triton:
try:
return _triton_decode_fp8(
q, kv_fp8, kv_scale,
qo_indptr, kv_indptr, config,
S=_NUM_KV_SPLITS,
)
except Exception:
_use_triton = False
return _torch_decode_fp8(q, kv_fp8, kv_scale, qo_indptr, kv_indptr, config)
check_implementation = make_match_reference(custom_kernel, rtol=1e-01, atol=1e-01)
scrolls · 364 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