submission 755102
Law1912 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 238 lines, June 9 Researcher Reciprocity License v1.0.
svm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-755102?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:7ce405b628a401a480ca83f433a839b28572c1f9d5beafc3ef7287ef6d7fb52d
license declaredunknown
license concludedunknown
authorsLaw1912
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
FP8 = tl.float8e4nvmma
logits = tl.dot(q_c, tl.trans(kc)) + tl.dot(q_r, tl.trans(kr))num-warps = 4
num_warps=4, num_stages=2, waves_per_eu=2,stages = 2
num_warps=4, num_stages=2, waves_per_eu=2,tile-k = 64
NH=_H, DC=_DL, DR=64, BK=64,Kernel source
svm.py238 lines
"""Two-pass grouped FP8 attention for decode with log-sum-exp merging."""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
_H, _DQ, _DL, _DV = 16, 576, 512, 512
_PARTITION_MAP = [
(4, 1024, 16), (4, 8192, 32),
(32, 1024, 4), (32, 8192, 8),
(64, 1024, 4), (64, 8192, 8),
(256, 1024, 1), (256, 8192, 2),
]
def _select_partitions(batch, seqlen):
for b, s, p in _PARTITION_MAP:
if b == batch and s == seqlen:
return p
limit = max(1, seqlen // 64)
if batch >= 256:
return min(2, limit)
if batch >= 64:
return min(max(1, 512 // batch), limit)
p = min(max(1, 768 // batch), limit)
while p > 1 and batch * p > 912:
p -= 1
return p
@triton.jit
def _compute_chunks(
q_ptr, kv_ptr, acc_buf, lse_buf, final_buf,
token_map, kv_spans,
scale_f, scale_kv_ptr,
qs0: tl.constexpr, qs1: tl.constexpr,
kvs0: tl.constexpr,
os0: tl.constexpr, os1: tl.constexpr,
NH: tl.constexpr, DC: tl.constexpr, DR: tl.constexpr, BK: tl.constexpr,
NP: tl.constexpr, NITERS: tl.constexpr, FUSE_OUT: tl.constexpr,
NWGS: tl.constexpr,
NUM_XCD: tl.constexpr = 8,
):
raw = tl.program_id(0)
wgs_per = (NWGS + NUM_XCD - 1) // NUM_XCD
n_full = NWGS % NUM_XCD
n_full = NUM_XCD if n_full == 0 else n_full
xid = raw % NUM_XCD
lid = raw // NUM_XCD
if xid < n_full:
gid = xid * wgs_per + lid
else:
gid = n_full * wgs_per + (xid - n_full) * (wgs_per - 1) + lid
sample = gid // NP
part = gid % NP
ROPE_BASE: tl.constexpr = 512
L2E: tl.constexpr = 1.4426950408889634
LN2_VAL: tl.constexpr = 0.6931471805599453
FP8 = tl.float8e4nv
ks = tl.load(scale_kv_ptr).to(tl.float32)
combined = scale_f * ks * L2E
tok = tl.load(token_map + sample)
kv0 = tl.load(kv_spans + sample)
kv1 = tl.load(kv_spans + sample + 1)
total = kv1 - kv0
per_p = tl.cdiv(total, NP)
origin = per_p * part
hh = tl.arange(0, NH)
cc = tl.arange(0, DC)
rr = tl.arange(0, DR)
qoff = tok * qs0
q_c = tl.load(q_ptr + qoff + hh[:, None] * qs1 + cc[None, :]).to(FP8)
q_r = tl.load(q_ptr + qoff + hh[:, None] * qs1 + (ROPE_BASE + rr[None, :])).to(FP8)
mx = tl.full([NH], value=float("-inf"), dtype=tl.float32)
sm = tl.zeros([NH], dtype=tl.float32)
ov = tl.zeros([NH, DC], dtype=tl.float32)
loop_n = NITERS if NITERS > 0 else (tl.minimum(origin + per_p, total) - origin) // BK
for i in range(loop_n):
idx = kv0 + origin + i * BK + tl.arange(0, BK)
kc = tl.load(kv_ptr + idx[:, None] * kvs0 + cc[None, :], cache_modifier=".cg")
kr = tl.load(kv_ptr + idx[:, None] * kvs0 + (ROPE_BASE + rr[None, :]), cache_modifier=".cg")
logits = tl.dot(q_c, tl.trans(kc)) + tl.dot(q_r, tl.trans(kr))
logits *= combined
new_mx = tl.maximum(tl.max(logits, 1), mx)
alpha = tl.math.exp2(mx - new_mx)
beta = tl.math.exp2(logits - new_mx[:, None])
ov = ov * alpha[:, None] + tl.dot(beta.to(FP8), kc)
sm = sm * alpha + tl.sum(beta, 1)
mx = new_mx
guard = tl.where(sm > 0, sm, 1.0)
rcp = tl.inline_asm_elementwise("v_rcp_f32_e32 $0, $1", "=v, v",
[guard], dtype=tl.float32,
is_pure=True, pack=1)
normalized = ov * rcp[:, None]
if FUSE_OUT:
dst = tok * os0
tl.store(final_buf + dst + hh[:, None] * os1 + cc[None, :],
(normalized * ks).to(tl.bfloat16))
else:
BW: tl.constexpr = 512
s_part = BW
s_head = NP * s_part
s_batch = NH * s_head
a = sample * s_batch + hh * s_head + part * s_part
tl.store(acc_buf + a[:, None] + cc[None, :], normalized)
lse_val = tl.where(sm > 0,
(mx + tl.math.log2(sm)) * LN2_VAL,
float("-inf"))
la = sample * (NH * NP) + hh * NP + part
tl.store(lse_buf + la, lse_val)
@triton.jit
def _merge_chunks(
acc_buf, lse_buf, final_buf, token_map, scale_kv_ptr,
NP: tl.constexpr, DV: tl.constexpr,
NB: tl.constexpr, NWGS: tl.constexpr,
os0: tl.constexpr, os1: tl.constexpr,
NH: tl.constexpr,
NUM_XCD: tl.constexpr = 8,
):
L2E: tl.constexpr = 1.4426950408889634
raw = tl.program_id(0)
wgs_per = (NWGS + NUM_XCD - 1) // NUM_XCD
n_full = NWGS % NUM_XCD
n_full = NUM_XCD if n_full == 0 else n_full
xid = raw % NUM_XCD
lid = raw // NUM_XCD
if xid < n_full:
gid = xid * wgs_per + lid
else:
gid = n_full * wgs_per + (xid - n_full) * (wgs_per - 1) + lid
b = gid % NB
h = gid // NB
ks = tl.load(scale_kv_ptr).to(tl.float32)
tok = tl.load(token_map + b)
dd = tl.arange(0, DV)
BW: tl.constexpr = 512
s_part = BW
s_head = NP * s_part
s_batch = NH * s_head
data_off = b * s_batch + h * s_head
lse_off = b * (NH * NP) + h * NP
best = float("-inf")
wsum = 0.0
result = tl.zeros([DV], dtype=tl.float32)
for k in range(NP):
chunk_v = tl.load(acc_buf + data_off + k * s_part + dd)
chunk_l = tl.load(lse_buf + lse_off + k)
nb = tl.maximum(chunk_l, best)
wa = tl.math.exp2((best - nb) * L2E)
wb = tl.math.exp2((chunk_l - nb) * L2E)
result = result * wa + wb * chunk_v
wsum = wsum * wa + wb
best = nb
clamped = tl.maximum(wsum, 1e-12)
inv = tl.inline_asm_elementwise("v_rcp_f32_e32 $0, $1", "=v, v",
[clamped], dtype=tl.float32,
is_pure=True, pack=1)
base = tok * os0 + h * os1
tl.store(final_buf + base + dd, (result * inv * ks).to(tl.bfloat16))
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_raw, kv_sc = kv_data["fp8"]
kv = kv_raw.view(torch.float8_e4m3fn).view(-1, _DQ)
B = config["batch_size"]
sm = config["sm_scale"]
T = q.shape[0]
S = kv.shape[0] // B if B > 0 else 0
P = _select_partitions(B, S)
fused = P == 1
chunk_len = (S + P - 1) // P
fixed = chunk_len // 64
if chunk_len % 64 != 0 or fixed > 4:
fixed = 0
output = torch.empty((T, _H, _DV), dtype=torch.bfloat16, device=q.device)
if fused:
ab = torch.empty(1, dtype=torch.float32, device=q.device)
lb = torch.empty(1, dtype=torch.float32, device=q.device)
else:
ab = torch.empty((B, _H, P, _DL), dtype=torch.float32, device=q.device)
lb = torch.empty((B, _H, P), dtype=torch.float32, device=q.device)
g1 = B * P
_compute_chunks[(g1,)](
q, kv, ab, lb, output,
qo_indptr, kv_indptr,
sm, kv_sc,
qs0=_H * _DQ, qs1=_DQ,
kvs0=_DQ,
os0=_H * _DV, os1=_DV,
NH=_H, DC=_DL, DR=64, BK=64,
NP=P, NITERS=fixed, FUSE_OUT=fused,
NWGS=g1,
num_warps=4, num_stages=2, waves_per_eu=2,
)
if not fused:
g2 = _H * B
_merge_chunks[(g2,)](
ab, lb, output, qo_indptr, kv_sc,
NP=P, DV=_DL,
NB=B, NWGS=g2,
os0=_H * _DV, os1=_DV,
NH=_H,
num_warps=4, num_stages=2,
)
return output
scrolls · 238 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