submission 690509
GWinfinity · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 120 lines, June 9 Researcher Reciprocity License v1.0.
submission_4d.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-690509?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:77f37d3f3dde0704bfc4e9560140746b69b3d5a5d7c2f3ae0c5e033f14acc70b
license declaredunknown
license concludedunknown
authorsGWinfinity
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
elif "mxfp4" in kv_data:Kernel source
submission_4d.py120 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA Decode - 4D KV Buffer Fix for Updated Aiter
"""
import torch
from typing import Dict, Tuple, Any
import os
import sys
NUM_HEADS = 16
KV_LORA_RANK = 512
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
input_t = Tuple[torch.Tensor, Dict[str, Any], torch.Tensor, torch.Tensor, Dict[str, Any]]
output_t = torch.Tensor
# Import aiter
_aiter_available = False
try:
aiter_path = os.path.expanduser('~/aiter')
if aiter_path not in sys.path:
sys.path.insert(0, aiter_path)
import aiter
from aiter.mla import mla_decode_fwd
_aiter_available = True
except Exception:
pass
def custom_kernel(data: input_t) -> output_t:
"""MLA decode with 4D kv_buffer for updated Aiter."""
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
dtype = q.dtype
device = q.device
# Parse KV
kv_buffer = None
kv_scale = None
if "bf16" in kv_data:
kv_buffer = kv_data["bf16"]
elif "fp8" in kv_data:
kv_fp8, kv_scale = kv_data["fp8"]
kv_buffer = kv_fp8
elif "mxfp4" in kv_data:
total_kv = int(kv_indptr[-1].item())
kv_buffer = torch.zeros((total_kv, 1, QK_HEAD_DIM), dtype=dtype, device=device)
else:
raise ValueError("No valid KV data")
# Aiter path - with 4D kv_buffer
if _aiter_available and dtype == torch.bfloat16:
total_q = q.shape[0]
# Convert 3D [total_kv, 1, head_dim] to 4D [num_pages, page_size, nhead_kv, head_dim]
# For non-paged case: page_size=1, nhead_kv=1
if kv_buffer.dim() == 3:
# [total_kv, 1, head_dim] -> [total_kv, 1, 1, head_dim]
kv_buffer_4d = kv_buffer.unsqueeze(2)
else:
kv_buffer_4d = kv_buffer
output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM),
dtype=torch.bfloat16, device=device)
# Use token-level indices (original behavior)
total_kv_pages = kv_buffer_4d.shape[0] # num_pages = total_kv when page_size=1
kv_indices = torch.arange(total_kv_pages, dtype=torch.int32, device=device)
kv_last_page_lens = kv_indptr[1:batch_size+1] - kv_indptr[:batch_size]
kwargs = {}
if kv_scale is not None:
kwargs['kv_scale'] = kv_scale
kwargs['q_scale'] = torch.ones(1, dtype=torch.float32, device=device)
mla_decode_fwd(
q=q,
kv_buffer=kv_buffer_4d, # 4D: [num_pages, page_size, nhead_kv, head_dim]
o=output,
qo_indptr=qo_indptr,
kv_indptr=kv_indptr,
kv_indices=kv_indices,
kv_last_page_lens=kv_last_page_lens,
max_seqlen_q=1,
page_size=1, # page_size=1 means each "page" is one token
nhead_kv=1,
sm_scale=SM_SCALE,
**kwargs
)
return output
# Fallback
import torch.nn.functional as F
kv_len = int(kv_indptr[1].item()) - int(kv_indptr[0].item())
if "fp8" in kv_data and kv_scale is not None:
kv_buffer = kv_buffer.to(dtype) * kv_scale.view(1, 1, 1)
elif "fp8" in kv_data:
kv_buffer = kv_buffer.to(dtype)
q_view = q.view(batch_size, NUM_HEADS, QK_HEAD_DIM)
kv_view = kv_buffer.view(batch_size, kv_len, QK_HEAD_DIM)
q_c = q_view[:, :, :KV_LORA_RANK]
q_r = q_view[:, :, KV_LORA_RANK:]
k_c = kv_view[:, :, :KV_LORA_RANK]
k_r = kv_view[:, :, KV_LORA_RANK:]
scores = torch.matmul(q_c, k_c.transpose(-2, -1))
scores = scores + torch.matmul(q_r, k_r.transpose(-2, -1))
scores = scores * SM_SCALE
attn = torch.softmax(scores, dim=-1, dtype=torch.float32).to(dtype)
return torch.matmul(attn, k_c)
scrolls · 120 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