submission 587798
Mihir Shah · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 115 lines, June 9 Researcher Reciprocity License v1.0.
final.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-587798?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:1d12a279dababe01e5a75e94dc34abd20616351da9d2c7e7ebfcb8c07651daf3
license declaredunknown
license concludedunknown
authorsMihir Shah
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_p, kv_s = kv_data["mxfp4"]mma
scores = tl.dot(q_le, tl.trans(vke)) + tl.dot(q_lo, tl.trans(vko))num-warps = 8
BLOCK_N=16, TILE_SIZE=1024, num_warps=8, num_stages=1stages = 1
BLOCK_N=16, TILE_SIZE=1024, num_warps=8, num_stages=1tile-n = 16
BLOCK_N=16, TILE_SIZE=1024, num_warps=8, num_stages=1Kernel source
final.py115 lines
import torch
import triton
import triton.language as tl
NUM_HEADS = 16
V_HEAD_DIM = 512
NUM_KV_SPLITS = 8
@triton.jit
def mla_mxfp4_kernel_seq_split(
Q, KV_packed, KV_scales, Out, LUT,
qo_indptr, kv_indptr, sm_scale,
stride_qb, stride_qh, stride_kvs, stride_kvs_s,
stride_ob, stride_oh,
BLOCK_N: tl.constexpr, TILE_SIZE: tl.constexpr,
):
batch_id = tl.program_id(0)
tile_id = tl.program_id(1)
offs_h = tl.arange(0, 16)
kv_start = tl.load(kv_indptr + batch_id) + (tile_id * TILE_SIZE)
kv_end = tl.minimum(tl.load(kv_indptr + batch_id + 1), kv_start + TILE_SIZE)
if kv_start >= kv_end:
return
# Latent dimensions (512)
idx_lat_e = tl.arange(0, 256) * 2
idx_lat_o = idx_lat_e + 1
# RoPE dimensions (64)
idx_rop_e = tl.arange(0, 32) * 2 + 512
idx_rop_o = idx_rop_e + 1
q_base = Q + (batch_id * stride_qb) + (offs_h[:, None] * stride_qh)
q_le = tl.load(q_base + idx_lat_e[None, :]).to(tl.bfloat16)
q_lo = tl.load(q_base + idx_lat_o[None, :]).to(tl.bfloat16)
q_re = tl.load(q_base + idx_rop_e[None, :]).to(tl.bfloat16)
q_ro = tl.load(q_base + idx_rop_o[None, :]).to(tl.bfloat16)
acc_e = tl.zeros([16, 256], dtype=tl.float32)
acc_o = tl.zeros([16, 256], dtype=tl.float32)
m_i = tl.zeros([16], dtype=tl.float32) - 1.0e20
l_i = tl.zeros([16], dtype=tl.float32)
for start_n in range(kv_start, kv_end, BLOCK_N):
offs_n = start_n + tl.arange(0, BLOCK_N)
mask_n = offs_n < kv_end
# Load 576 dims (256 bytes latent + 32 bytes RoPE)
kv_ptr = KV_packed + (offs_n[:, None] * stride_kvs)
pkd_l = tl.load(kv_ptr + tl.arange(0, 256)[None, :], mask=mask_n[:, None], other=0)
pkd_r = tl.load(kv_ptr + 256 + tl.arange(0, 32)[None, :], mask=mask_n[:, None], other=0)
# Load 18 scales (16 latent + 2 RoPE)
s_ptr = KV_scales + (offs_n[:, None] * stride_kvs_s)
sl = tl.exp2(tl.load(s_ptr + tl.arange(0, 16)[None, :], mask=mask_n[:, None], other=0).to(tl.float32) - 127.0).to(tl.bfloat16)
sr = tl.exp2(tl.load(s_ptr + 16 + tl.arange(0, 2)[None, :], mask=mask_n[:, None], other=0).to(tl.float32) - 127.0).to(tl.bfloat16)
sl_b = tl.reshape(tl.broadcast_to(sl[:, :, None], [BLOCK_N, 16, 16]), [BLOCK_N, 256])
sr_b = tl.reshape(tl.broadcast_to(sr[:, :, None], [BLOCK_N, 2, 16]), [BLOCK_N, 32])
# Dequantize via LUT
vke = tl.load(LUT + (pkd_l & 0x0F).to(tl.int32)).to(tl.bfloat16) * sl_b
vko = tl.load(LUT + ((pkd_l >> 4) & 0x0F).to(tl.int32)).to(tl.bfloat16) * sl_b
rke = tl.load(LUT + (pkd_r & 0x0F).to(tl.int32)).to(tl.bfloat16) * sr_b
rko = tl.load(LUT + ((pkd_r >> 4) & 0x0F).to(tl.int32)).to(tl.bfloat16) * sr_b
# Full 576-dim dot product
scores = tl.dot(q_le, tl.trans(vke)) + tl.dot(q_lo, tl.trans(vko))
scores += tl.dot(q_re, tl.trans(rke)) + tl.dot(q_ro, tl.trans(rko))
scores *= sm_scale
scores = tl.where(mask_n[None, :], scores, -1.0e20)
m_ij = tl.max(scores, axis=1)
p = tl.exp(scores - m_ij[:, None])
l_ij = tl.sum(p, axis=1)
m_next = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_next)
beta = tl.exp(m_ij - m_next)
# Accumulate latent values only (512 dims)
p_bf = p.to(tl.bfloat16)
acc_e = acc_e * alpha[:, None] + tl.dot(p_bf, vke) * beta[:, None]
acc_o = acc_o * alpha[:, None] + tl.dot(p_bf, vko) * beta[:, None]
l_i = l_i * alpha + l_ij * beta
m_i = m_next
out_base = Out + (batch_id * 16 * 512) + (offs_h[:, None] * 512)
tl.store(out_base + idx_lat_e[None, :], (acc_e / l_i[:, None]).to(Out.dtype.element_ty))
tl.store(out_base + idx_lat_o[None, :], (acc_o / l_i[:, None]).to(Out.dtype.element_ty))
def custom_kernel(data):
q, kv_data, qo_indptr, kv_indptr, config = data
kv_p, kv_s = kv_data["mxfp4"]
# E2M1 Standard LUT
lut = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
-0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0],
dtype=torch.float32, device="cuda")
out = torch.empty((q.shape[0], 16, 512), dtype=torch.bfloat16, device="cuda")
grid = (config["batch_size"], NUM_KV_SPLITS)
mla_mxfp4_kernel_seq_split[grid](
q, kv_p.view(torch.uint8), kv_s.view(torch.uint8), out, lut,
qo_indptr, kv_indptr, config["sm_scale"],
q.stride(0), q.stride(1), kv_p.stride(0), kv_s.stride(0),
512 * 16, 512,
BLOCK_N=16, TILE_SIZE=1024, num_warps=8, num_stages=1
)
return outscrolls · 115 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