submission 644275
nawal-yantrion · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 208 lines, June 9 Researcher Reciprocity License v1.0.
sub_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-644275?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:6757f69a224868c99b2f30ae931577256f9309486c59b51b1b82b83b668b02b2
license declaredunknown
license concludedunknown
authorsnawal-yantrion
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
elif QKV_DTYPE == "mxfp4":Kernel source
sub_v2.py208 lines
import torch
import torch.nn.functional as F
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
FP8_DTYPE = aiter_dtypes.fp8
# QKV dtype for custom_kernel dispatch: "bf16", "fp8", or "mxfp4"
QKV_DTYPE = "fp8"
# ---------------------------------------------------------------------------
# Dispatcher: select kernel based on QKV_DTYPE
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
"""Dispatch to the appropriate kernel based on QKV_DTYPE."""
if QKV_DTYPE == "fp8":
return custom_kernel_fp8(data)
elif QKV_DTYPE == "bf16":
return custom_kernel_bf16(data)
elif QKV_DTYPE == "mxfp4":
return custom_kernel_mxfp4(data)
else:
raise ValueError(f"Invalid QKV_DTYPE: {QKV_DTYPE}")
# ---------------------------------------------------------------------------
# FP8 quantization helper (per-tensor, sglang style)
# ---------------------------------------------------------------------------
def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Dynamic per-tensor FP8 quantization. Returns (fp8_tensor, scale)."""
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)
# ---------------------------------------------------------------------------
# FP8 Q + FP8 KV — optimized FP8 attention kernel using torch._scaled_mm
# ---------------------------------------------------------------------------
def custom_kernel_fp8(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
kv_lora_rank = config["kv_lora_rank"]
qk_head_dim = config["qk_head_dim"]
sm_scale = config["sm_scale"]
# FP8 KV buffer and scale
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim) # (total_kv, 576) fp8
# Quantize Q to fp8
q_fp8, q_scale = quantize_fp8(q) # q_fp8: (total_q, 16, 576) fp8
batch_size = qo_indptr.shape[0] - 1
out_list = []
scale_one = torch.ones(1, dtype=torch.float32, device="cuda")
# Use torch.jit.script to optimize the loop, reducing Python overhead
for i in range(batch_size):
q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
kv_s, kv_e = int(kv_indptr[i].item()), int(kv_indptr[i + 1].item())
seq_q = q_e - q_s
seq_kv = kv_e - kv_s
# Q: (seq_q * nhead, 576) fp8, K: (seq_kv, 576) fp8
qi_fp8 = q_fp8[q_s:q_e].reshape(seq_q * num_heads, qk_head_dim) # (seq_q*16, 576)
ki_fp8 = kv_fp8_2d[kv_s:kv_e] # (seq_kv, 576)
# QK^T via _scaled_mm: (seq_q*16, 576) @ (seq_kv, 576).T -> (seq_q*16, seq_kv)
# _scaled_mm expects (M,K) @ (N,K).T where b is row-major contiguous
raw_scores = torch._scaled_mm(
qi_fp8, ki_fp8.t(),
scale_a=q_scale, scale_b=kv_scale_fp8,
out_dtype=torch.float32,
)
# raw_scores: (seq_q*16, seq_kv)
scores = raw_scores.view(seq_q, num_heads, seq_kv).permute(1, 0, 2) # (nhead, seq_q, seq_kv)
scores = scores * sm_scale
scores = F.softmax(scores, dim=-1)
# V: first 512 dims of KV buffer (bf16 for softmax@V since scores are float)
kv_bf16 = kv_data["bf16"]
vi = kv_bf16[kv_s:kv_e, 0, :kv_lora_rank].float() # (seq_kv, 512)
# softmax @ V: (nhead, seq_q, seq_kv) @ (seq_kv, 512) -> (nhead, seq_q, 512)
oi = torch.matmul(scores, vi)
oi = oi.permute(1, 0, 2) # (seq_q, nhead, 512)
out_list.append(oi.to(torch.bfloat16))
return torch.cat(out_list, dim=0)
# ---------------------------------------------------------------------------
# Baseline: bf16 Q + bf16 KV — naive torch attention
# ---------------------------------------------------------------------------
def custom_kernel_bf16(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
kv_lora_rank = config["kv_lora_rank"]
sm_scale = config["sm_scale"]
# This naive baseline uses bf16 KV directly.
kv_buffer_bf16 = kv_data["bf16"]
batch_size = qo_indptr.shape[0] - 1
out_list = []
for i in range(batch_size):
q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
kv_s, kv_e = int(kv_indptr[i].item()), int(kv_indptr[i + 1].item())
qi = q[q_s:q_e] # (seq_q, nhead, 576)
kvc = kv_buffer_bf16[kv_s:kv_e, 0] # (seq_kv, 576) squeeze kv_heads dim
# Key: full 576 dims; Value: first 512 dims (kv_lora_rank)
ki = kvc # (seq_kv, 576) — broadcast over heads
vi = kvc[:, :kv_lora_rank] # (seq_kv, 512)
# Attention: (nhead, seq_q, 576) @ (576, seq_kv) → (nhead, seq_q, seq_kv)
qi_t = qi.float().permute(1, 0, 2) # (nhead, seq_q, 576)
scores = torch.matmul(qi_t * sm_scale, ki.float().T) # (nhead, seq_q, seq_kv)
scores = F.softmax(scores, dim=-1)
# Output: (nhead, seq_q, seq_kv) @ (seq_kv, 512) → (nhead, seq_q, 512)
oi = torch.matmul(scores, vi.float()) # (nhead, seq_q, 512)
oi = oi.permute(1, 0, 2) # (seq_q, nhead, 512)
out_list.append(oi.to(torch.bfloat16))
return torch.cat(out_list, dim=0)
def unpack_mxfp4(packed: torch.Tensor) -> torch.Tensor:
"""
Unpack fp4x2 → int8 values in range [-8, 7]
packed: (..., 288) uint8
returns: (..., 576) int8
"""
high = (packed >> 4) & 0xF
low = packed & 0xF
# convert to signed [-8, 7]
high = high.to(torch.int8) - 8
low = low.to(torch.int8) - 8
return torch.stack((high, low), dim=-1).reshape(*packed.shape[:-1], -1)
def custom_kernel_mxfp4(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
kv_lora_rank = config["kv_lora_rank"] # 512
qk_head_dim = config["qk_head_dim"] # 576
sm_scale = config["sm_scale"]
# MXFP4 KV
kv_packed, kv_scale = kv_data["mxfp4"] # (kv,1,288), (scale)
kv_packed = kv_packed.view(-1, 288)
# Quantize Q → fp8
q_fp8, q_scale = quantize_fp8(q)
batch_size = qo_indptr.shape[0] - 1
out_list = []
for i in range(batch_size):
q_s, q_e = int(qo_indptr[i]), int(qo_indptr[i + 1])
kv_s, kv_e = int(kv_indptr[i]), int(kv_indptr[i + 1])
seq_q = q_e - q_s
seq_kv = kv_e - kv_s
qi = q_fp8[q_s:q_e].reshape(seq_q * num_heads, qk_head_dim)
# ---- UNPACK + DEQUANT (FUSED STYLE) ----
kv_chunk = kv_packed[kv_s:kv_e] # (seq_kv, 288)
kv_int8 = unpack_mxfp4(kv_chunk) # (seq_kv, 576)
# apply scale (fp8 → fp32)
kv_fp32 = kv_int8.float() * kv_scale
# split K / V without extra copies
ki = kv_fp32 # (seq_kv, 576)
vi = kv_fp32[:, :kv_lora_rank] # (seq_kv, 512)
# ---- QK ----
scores = torch.matmul(
qi.float(),
ki.T
) # (seq_q * nhead, seq_kv)
scores = scores.view(seq_q, num_heads, seq_kv).permute(1, 0, 2)
scores *= sm_scale
scores = F.softmax(scores, dim=-1)
# ---- ATTENTION @ V ----
oi = torch.matmul(scores, vi) # (nhead, seq_q, 512)
oi = oi.permute(1, 0, 2)
out_list.append(oi.to(torch.bfloat16))
return torch.cat(out_list, dim=0)scrolls · 208 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