submission 694148
sky · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 112 lines, June 9 Researcher Reciprocity License v1.0.
submission3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-694148?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:9cb458629cfc09d68b99cb59c6c68d80a4e9a8a7ab80b790be782b02cd4f5113
license declaredunknown
license concludedunknown
authorssky
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
if "mxfp4" in kv_data:Kernel source
submission3.py112 lines
import torch
import math
# -----------------------------
# Constants
# -----------------------------
NUM_HEADS = 16
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM)
# -----------------------------
# Dequant
# -----------------------------
def dequant_fp8(x, scale):
return x.to(torch.bfloat16) * scale
def dequant_mxfp4(fp4_data, scale_e8m0):
from aiter.utility.fp4_utils import mxfp4_to_f32, e8m0_to_f32
B, M, N2 = fp4_data.shape
N = N2 * 2
rows = B * M
x = mxfp4_to_f32(fp4_data.reshape(rows, N2))
scales = e8m0_to_f32(scale_e8m0)
block = 32
num_blocks = N // block
scales = scales[:rows, :num_blocks]
x = x.view(rows, num_blocks, block)
x = x * scales.unsqueeze(-1)
return x.view(B, M, N).to(torch.bfloat16)
# -----------------------------
# FAST batched attention
# -----------------------------
def fast_attention(q, k, v):
"""
q: (B, 16, 576)
k: (B, L, 576)
v: (B, L, 512)
"""
# (B, 16, L)
scores = torch.matmul(q.to(torch.float32), k.transpose(1, 2).to(torch.float32))
scores *= SM_SCALE
probs = torch.softmax(scores, dim=-1)
# (B, 16, 512)
out = torch.matmul(probs, v.to(torch.float32))
return out.to(torch.bfloat16)
# -----------------------------
# Main kernel (vectorized)
# -----------------------------
def mla_decode_kernel(data):
q, kv_data, qo_indptr, kv_indptr, config = data
device = q.device
batch_size = qo_indptr.shape[0] - 1
q = q.view(batch_size, NUM_HEADS, QK_HEAD_DIM)
# -------------------------
# Resolve KV
# -------------------------
if "mxfp4" in kv_data:
kv_buffer, scale = kv_data["mxfp4"]
kv_full = dequant_mxfp4(kv_buffer, scale)
elif "fp8" in kv_data:
kv_buffer, scale = kv_data["fp8"]
kv_full = dequant_fp8(kv_buffer, scale)
else:
kv_full = kv_data["bf16"]
kv_full = kv_full.squeeze(1) # (total_kv, 576)
outputs = []
# -------------------------
# Batch (still segmented, but vectorized inside)
# -------------------------
for b in range(batch_size):
kv_start = kv_indptr[b].item()
kv_end = kv_indptr[b + 1].item()
k = kv_full[kv_start:kv_end, :QK_HEAD_DIM].unsqueeze(0) # (1, L, 576)
v = kv_full[kv_start:kv_end, :V_HEAD_DIM].unsqueeze(0) # (1, L, 512)
qb = q[b].unsqueeze(0) # (1, 16, 576)
out = fast_attention(qb, k, v) # (1, 16, 512)
outputs.append(out)
return torch.cat(outputs, dim=0)
# -----------------------------
# Entry point
# -----------------------------
def custom_kernel(data):
return mla_decode_kernel(data)scrolls · 112 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