submission 754244
Navdeep Singh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 195 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-754244?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:1a9e5b0b7690d189f28ca4528129e11b71beea6636b86870a101a083c634f892
license declaredunknown
license concludedunknown
authorsNavdeep Singh
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
scores = tl.dot(q_lat_fp8, tl.trans(p_lat)) + tl.dot(q_rop_fp8, tl.trans(p_rop))num-warps = 4
num_warps=4,persistent-kernel
num_splits = tl.num_programs(1)stages = 3
for start_n in tl.range(kv_start, kv_end, BLOCK_N, num_stages=3):Kernel source
submission.py195 lines
import torch
import triton
import triton.language as tl
from typing import TypeVar
input_t = TypeVar("input_t")
output_t = TypeVar("output_t")
@triton.jit
def mla_stage1_fp8(
Q, KV_fp8, KV_scale,
Workspace_V, Workspace_M, Workspace_L, Out,
kv_indptr, sm_scale,
stride_qb, stride_qh,
stride_obs, stride_oh,
BLOCK_N: tl.constexpr,
TILE_SIZE: tl.constexpr,
WRITE_OUT: tl.constexpr,
):
LOG2E = 1.4426950408889634
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
num_splits = tl.num_programs(1)
kv_seq_start = tl.load(kv_indptr + batch_id)
kv_seq_end = tl.load(kv_indptr + batch_id + 1)
kv_start = kv_seq_start + split_id * TILE_SIZE
kv_end = tl.minimum(kv_start + TILE_SIZE, kv_seq_end)
offs_h = tl.arange(0, 16)
if kv_start >= kv_end:
if not WRITE_OUT:
ws_ml_off = (batch_id * num_splits + split_id) * 16 + offs_h
tl.store(Workspace_M + ws_ml_off, -1.0e20)
tl.store(Workspace_L + ws_ml_off, 0.0)
return
global_scale = tl.load(KV_scale).to(tl.float32)
scale_factor = 200.0
outer_scale = (sm_scale * LOG2E * global_scale) / scale_factor
v_scale = 350.0
q_base = Q + batch_id * stride_qb
q_lat_raw = tl.load(q_base + offs_h[:, None] * stride_qh + tl.arange(0, 512)[None, :]).to(tl.float32)
q_rop_raw = tl.load(q_base + offs_h[:, None] * stride_qh + 512 + tl.arange(0, 64)[None, :]).to(tl.float32)
fp8_ty = KV_fp8.dtype.element_ty
q_lat_fp8 = (q_lat_raw * scale_factor).to(fp8_ty)
q_rop_fp8 = (q_rop_raw * scale_factor).to(fp8_ty)
m_i = tl.full([16], -1.0e20, dtype=tl.float32)
l_i = tl.zeros([16], dtype=tl.float32)
acc = tl.zeros([16, 512], dtype=tl.float32)
for start_n in tl.range(kv_start, kv_end, BLOCK_N, num_stages=3):
offs_n = start_n + tl.arange(0, BLOCK_N)
mask_n = offs_n < kv_end
kv_ptr = KV_fp8 + offs_n[:, None] * 576
p_lat = tl.load(kv_ptr + tl.arange(0, 512)[None, :], mask=mask_n[:, None], other=0.0)
p_rop = tl.load(kv_ptr + 512 + tl.arange(0, 64)[None, :], mask=mask_n[:, None], other=0.0)
scores = tl.dot(q_lat_fp8, tl.trans(p_lat)) + tl.dot(q_rop_fp8, tl.trans(p_rop))
scores = scores * outer_scale
scores = tl.where(mask_n[None, :], scores, -1.0e20)
m_ij = tl.max(scores, axis=1)
p = tl.exp2(scores - m_ij[:, None])
l_ij = tl.sum(p, axis=1)
m_next = tl.maximum(m_i, m_ij)
alpha = tl.exp2(m_i - m_next)
beta = tl.exp2(m_ij - m_next)
# Scale applied *before* accumulation to prevent precision blowout
p_beta_fp8 = (p * beta[:, None] * v_scale).to(fp8_ty)
acc = acc * alpha[:, None] + tl.dot(p_beta_fp8, p_lat)
l_i = l_i * alpha + l_ij * beta
m_i = m_next
acc = (acc * global_scale) / v_scale
if WRITE_OUT:
out_ptr = Out + batch_id * stride_obs + offs_h[:, None] * stride_oh + tl.arange(0, 512)[None, :]
tl.store(out_ptr, (acc / l_i[:, None]).to(tl.bfloat16))
else:
ws_ml_off = (batch_id * num_splits + split_id) * 16 + offs_h
ws_v_off = ((batch_id * num_splits + split_id) * 16 + offs_h[:, None]) * 512
tl.store(Workspace_V + ws_v_off + tl.arange(0, 512)[None, :], acc.to(tl.bfloat16))
tl.store(Workspace_M + ws_ml_off, m_i)
tl.store(Workspace_L + ws_ml_off, l_i)
@triton.jit
def mla_stage2_reduce(
Workspace_V, Workspace_M, Workspace_L, Out,
stride_obs, stride_oh,
num_splits: tl.constexpr, BLOCK_SPLITS: tl.constexpr,
):
pid = tl.program_id(0)
batch_id = pid // 16
head_id = pid % 16
offs_s = tl.arange(0, BLOCK_SPLITS)
mask_s = offs_s < num_splits
base_ml = (batch_id * num_splits + offs_s) * 16 + head_id
m_s = tl.load(Workspace_M + base_ml, mask=mask_s, other=-1.0e20)
l_s = tl.load(Workspace_L + base_ml, mask=mask_s, other=0.0)
m_g = tl.max(m_s, axis=0)
alpha = tl.exp2(m_s - m_g)
l_g = tl.sum(l_s * alpha, axis=0)
off_v = ((batch_id * num_splits + offs_s) * 16 + head_id) * 512
v_all = tl.load(Workspace_V + off_v[:, None] + tl.arange(0, 512)[None, :], mask=mask_s[:, None], other=0.0).to(tl.float32)
v_acc = tl.sum(v_all * alpha[:, None], axis=0)
tl.store(Out + batch_id * stride_obs + head_id * stride_oh + tl.arange(0, 512), (v_acc / l_g).to(tl.bfloat16))
def _heuristics(bs: int, kv_len: int):
# The "God Grid" mapping: Target ~256 total blocks globally.
# We strictly enforce `splits=1` (WRITE_OUT) for large batch sizes.
if bs >= 256:
splits = 1 # 256 total blocks
elif bs >= 64:
splits = 4 # 256 total blocks
elif bs >= 32:
splits = 8 # 256 total blocks
elif bs >= 16:
splits = 16 # 256 total blocks
else:
splits = 32 # 128 total blocks (leaves room for scheduler)
block_n = 128
if kv_len // splits < 128:
block_n = 64
if kv_len // splits < 64:
block_n = 32
tile_size = ((kv_len + splits - 1) // splits)
# Ensure memory alignment
tile_size = ((tile_size + block_n - 1) // block_n) * block_n
return int(splits), int(tile_size), int(block_n), 8, 3
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv_len = config.get("kv_seq_len", 8192)
kv_p_fp8, kv_s_fp8 = kv_data["fp8"]
kv_s_fp8 = kv_s_fp8.view(1)
splits, tile_size, block_n, warps, stages = _heuristics(bs, kv_len)
out = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device=q.device)
write_out = (splits == 1)
ws_v, ws_m, ws_l = None, None, None
if not write_out:
ws_v = torch.empty((bs, splits, 16, 512), dtype=torch.bfloat16, device=q.device)
ws_m = torch.empty((bs, splits, 16), dtype=torch.float32, device=q.device)
ws_l = torch.empty((bs, splits, 16), dtype=torch.float32, device=q.device)
mla_stage1_fp8[(bs, splits)](
q, kv_p_fp8, kv_s_fp8,
ws_v, ws_m, ws_l, out,
kv_indptr, config["sm_scale"],
q.stride(0), q.stride(1),
out.stride(0), out.stride(1),
BLOCK_N=block_n, TILE_SIZE=tile_size,
WRITE_OUT=write_out,
num_warps=warps, num_stages=stages,
)
if not write_out:
block_splits = triton.next_power_of_2(splits)
mla_stage2_reduce[(bs * 16,)](
ws_v, ws_m, ws_l, out,
out.stride(0), out.stride(1),
num_splits=splits, BLOCK_SPLITS=block_splits,
num_warps=4,
)
return outscrolls · 195 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