submission 609348
Danishlynx · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 202 lines, June 9 Researcher Reciprocity License v1.0.
submission_v11.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-609348?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:d904f4b342fce6da8f79acdd4b2c777a23b9583826a99d876249858375a3bafc
license declaredunknown
license concludedunknown
authorsDanishlynx
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
fp4, fsc = kv_data["mxfp4"]online-softmax
float pv_acc[4][32]; float running_max[4], running_sum[4];Kernel source
submission_v11.py202 lines
# /// script
# leaderboard = "amd-mixed-mla"
# ///
"""XK89: Clean separate-buffer chunking for (64,8192).
Each chunk gets its own split_out/split_lse tensors (pre-allocated).
After all chunks, copy into combined buffer and reduce.
Fixes the deterministic corruption from offset pointer writes.
"""
import sys
import torch
import aiter
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_D = aiter_dtypes.fp8
_S = 1.0 / (576 ** 0.5)
_c = {}
def P(*a): print(*a, file=sys.stderr, flush=True)
_NS = {
(4, 1024): 16, (4, 8192): 16,
(32, 1024): 8, (32, 8192): 8,
(64, 1024): 4, (64, 8192): 4,
(256, 1024): 1, (256, 8192): 1,
}
_MFMA_NS = {
(4, 8192): 64,
(64, 8192): 64,
}
_CHUNK = 16
_ATTN_SRC = r"""
#define NHEADS 16
#define HEAD_DIM 576
#define V_DIM 512
#define QK_TILE 16
#define PV_BLOCK 128
#define SM_SCALE 0.041666667f
#define LOG2E 1.4426950408889634f
#define MYEXP(x) __builtin_exp2f((x) * LOG2E)
#define NEG_INF (-1e30f)
typedef unsigned char u8; typedef unsigned short u16; typedef unsigned int u32;
typedef u8 __attribute__((ext_vector_type(32))) reg256_t;
typedef float __attribute__((ext_vector_type(4))) f32x4_t;
#define READLANE_F32(val, lane_id) ({ int _rl_i; float _rl_f = (val); __builtin_memcpy(&_rl_i, &_rl_f, 4); int _rl_r = __builtin_amdgcn_readlane(_rl_i, (int)(lane_id)); float _rl_out; __builtin_memcpy(&_rl_out, &_rl_r, 4); _rl_out; })
#define BF16_TO_FP8(bf16_bits) ({ u16 _bb = (bf16_bits); unsigned _sign = (_bb >> 15) & 1u; int _exp = (int)((_bb >> 7) & 0xFFu) - 127; unsigned _mant = _bb & 0x7Fu; int _fp8_exp = _exp + 7; u8 _result; if (_fp8_exp <= 0) { _result = (u8)(_sign << 7); } else if (_fp8_exp >= 15) { _result = (u8)((_sign << 7) | 0x7Eu); } else { _result = (u8)((_sign << 7) | (_fp8_exp << 3) | (_mant >> 4)); } _result; })
#define F32_TO_FP8(fval) ({ float _fv = (fval); u32 _fb; __builtin_memcpy(&_fb, &_fv, 4); unsigned _sign = (_fb >> 31) & 1u; int _exp = (int)((_fb >> 23) & 0xFFu) - 127; unsigned _mant = (_fb >> 20) & 0x7u; int _fp8_exp = _exp + 7; u8 _result; if (_fv == 0.0f || _fp8_exp <= 0) { _result = (u8)(_sign << 7); } else if (_fp8_exp >= 15) { _result = (u8)((_sign << 7) | 0x7Eu); } else { _result = (u8)((_sign << 7) | (_fp8_exp << 3) | _mant); } _result; })
__attribute__((shared)) float lds_scores[16 * 128];
__attribute__((shared)) u8 lds_p[16 * 128];
extern "C" __attribute__((global)) void __attribute__((amdgpu_flat_work_group_size(64, 64)))
mla_attn_fp4(const u16* __restrict__ Q_bf16, const u8* __restrict__ KV_fp4, const u8* __restrict__ KV_scale, const int* __restrict__ kv_indptr, float* __restrict__ split_out, float* __restrict__ split_lse, long ns, long bs) {
unsigned wg_id = __builtin_amdgcn_workgroup_id_x(); unsigned tid = __builtin_amdgcn_workitem_id_x(); unsigned group = tid >> 4; unsigned lane = tid & 15u;
int ins = (int)ns, ibs = (int)bs; int batch = wg_id / ins, split = wg_id % ins; if (batch >= ibs) return;
int kv_start = kv_indptr[batch]; int total_tokens = kv_indptr[batch + 1] - kv_start;
int tps = (total_tokens + ins - 1) / ins; int my_start = split * tps; int my_end = my_start + tps;
if (my_end > total_tokens) my_end = total_tokens; int my_tokens = my_end - my_start;
if (my_tokens <= 0) { for (int s = 0; s < 4; s++) { int head = group * 4 + s; split_lse[wg_id * 16 + head] = NEG_INF; float* so = split_out + wg_id * 16 * V_DIM + head * V_DIM; for (int fc = 0; fc < 32; fc++) { int vf = lane + fc * 16; if (vf < V_DIM) so[vf] = 0.0f; } } return; }
int kv_off = kv_start + my_start;
reg256_t q_a[5]; const u16* qh = Q_bf16 + batch * NHEADS * HEAD_DIM;
for (int call = 0; call < 5; call++) { int kb = call * 128 + group * 32; for (int i = 0; i < 32; i++) q_a[call][i] = 0; if (kb < HEAD_DIM) { int valid = HEAD_DIM - kb; if (valid > 32) valid = 32; const u16* qp = qh + lane * HEAD_DIM + kb; for (int i = 0; i < valid; i++) q_a[call][i] = BF16_TO_FP8(qp[i]); } }
float pv_acc[4][32]; float running_max[4], running_sum[4];
for (int s = 0; s < 4; s++) { for (int fc = 0; fc < 32; fc++) pv_acc[s][fc] = 0.0f; running_max[s] = NEG_INF; running_sum[s] = 0.0f; }
for (int blk = 0; blk < my_tokens; blk += PV_BLOCK) {
int bt = my_tokens - blk; if (bt > PV_BLOCK) bt = PV_BLOCK; int nt = (bt + QK_TILE - 1) / QK_TILE;
for (int tile = 0; tile < nt; tile++) { int ts = tile * QK_TILE; int tc = bt - ts; if (tc > QK_TILE) tc = QK_TILE; int gtok = kv_off + blk + ts + lane; int valid = ((int)lane < tc) ? 1 : 0; f32x4_t c_qk; for (int i = 0; i < 4; i++) c_qk[i] = 0.0f; for (int call = 0; call < 5; call++) { int kb = call * 128 + group * 32; reg256_t b_reg; for (int i = 0; i < 32; i++) b_reg[i] = 0; u8 kv_sc = 127; if (valid && kb < HEAD_DIM) { int bo = kb / 2, vb = (HEAD_DIM - kb + 1) / 2; if (vb > 16) vb = 16; const u8* kp = KV_fp4 + gtok * 288 + bo; for (int i = 0; i < vb; i++) b_reg[i] = kp[i]; int sg = kb / 32; if (sg >= 18) sg = 17; kv_sc = KV_scale[gtok * 24 + sg]; } c_qk = (f32x4_t)__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(q_a[call], b_reg, c_qk, 0, 4, 0, 127, 0, kv_sc); } for (int s = 0; s < 4; s++) { float sc = valid ? (c_qk[s] * SM_SCALE) : NEG_INF; lds_scores[(group * 4 + s) * PV_BLOCK + ts + lane] = sc; } }
__builtin_amdgcn_s_barrier();
float block_max[4], block_sum[4];
for (int s = 0; s < 4; s++) { int head = group * 4 + s; float lmax = NEG_INF; float msc[8]; for (int j = 0; j < 8; j++) { int tok = lane + j * 16; msc[j] = (tok < bt) ? lds_scores[head * PV_BLOCK + tok] : NEG_INF; if (msc[j] > lmax) lmax = msc[j]; } float mx = lmax; int bl = group * 16; for (int j = 0; j < 16; j++) { float v = READLANE_F32(lmax, bl+j); if (v > mx) mx = v; } block_max[s] = mx; float lsum = 0.0f; for (int j = 0; j < 8; j++) { int tok = lane + j * 16; float p = (tok < bt) ? MYEXP(msc[j] - mx) : 0.0f; lsum += p; lds_p[head * PV_BLOCK + tok] = F32_TO_FP8(p * 256.0f); } float tsum = 0.0f; for (int j = 0; j < 16; j++) tsum += READLANE_F32(lsum, bl+j); block_sum[s] = tsum; }
__builtin_amdgcn_s_barrier();
for (int s = 0; s < 4; s++) { float om = running_max[s]; float nm = (block_max[s] > om) ? block_max[s] : om; float eo = MYEXP(om - nm), eb = MYEXP(block_max[s] - nm); for (int fc = 0; fc < 32; fc++) pv_acc[s][fc] *= eo; running_sum[s] = running_sum[s] * eo + block_sum[s] * eb; running_max[s] = nm; }
reg256_t p_a; for (int b = 0; b < 32; b++) { int tok = group * 32 + b; p_a[b] = (tok < bt) ? lds_p[lane * PV_BLOCK + tok] : 0; }
for (int fc = 0; fc < 32; fc++) { f32x4_t c_pv; for (int i = 0; i < 4; i++) c_pv[i] = 0.0f; int vf = lane + fc * 16; if (vf < V_DIM) { int fb = vf / 2, ns2 = vf & 1, sg = vf / 32; if (sg >= 16) sg = 15; reg256_t v_b; for (int i = 0; i < 16; i++) { int t0 = group*32+i*2, t1 = t0+1; u8 n0=0,n1=0; if (t0 < bt) { int g0=kv_off+blk+t0; u8 b0=KV_fp4[g0*288+fb]; n0=ns2?((b0>>4)&0xFu):(b0&0xFu); } if (t1 < bt) { int g1=kv_off+blk+t1; u8 b1=KV_fp4[g1*288+fb]; n1=ns2?((b1>>4)&0xFu):(b1&0xFu); } v_b[i] = n0 | (n1 << 4); } int rt = group*32+16; if (rt >= bt) rt = bt > 0 ? bt-1 : 0; u8 vsc = KV_scale[(kv_off+blk+rt)*24+sg]; c_pv = (f32x4_t)__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(p_a, v_b, c_pv, 0, 4, 0, 119, 0, vsc); } for (int s = 0; s < 4; s++) { float corr = MYEXP(block_max[s] - running_max[s]); pv_acc[s][fc] += c_pv[s] * corr; } }
__builtin_amdgcn_s_barrier();
}
for (int s = 0; s < 4; s++) { int head = group * 4 + s; split_lse[wg_id * 16 + head] = running_max[s] + __builtin_log2f(running_sum[s] + 1e-30f) / LOG2E; float inv = (running_sum[s] > 1e-9f) ? (1.0f / running_sum[s]) : 0.0f; float* so = split_out + wg_id * 16 * V_DIM + head * V_DIM; for (int fc = 0; fc < 32; fc++) { int vf = lane + fc * 16; if (vf < V_DIM) so[vf] = pv_acc[s][fc] * inv; } }
}
"""
_REDUCE_SRC = r"""
#define V_DIM 512
#define LOG2E 1.4426950408889634f
#define NEG_INF (-1e30f)
typedef unsigned char u8; typedef unsigned short u16; typedef unsigned int u32;
extern "C" __attribute__((global)) void __attribute__((amdgpu_flat_work_group_size(64, 64)))
mla_reduce(const float* __restrict__ split_out, const float* __restrict__ split_lse, u16* __restrict__ final_out, long ns, long bs) {
unsigned wg_id = __builtin_amdgcn_workgroup_id_x(); unsigned tid = __builtin_amdgcn_workitem_id_x();
int ins = (int)ns, batch = wg_id / 4, hg = wg_id % 4; if (batch >= (int)bs) return;
for (int hs = 0; hs < 4; hs++) { int head = hg * 4 + hs; float max_lse = NEG_INF; for (int s = 0; s < ins; s++) { float lse = split_lse[(batch*ins+s)*16+head]; if (lse > max_lse) max_lse = lse; } for (int fi = 0; fi < 8; fi++) { int vf = tid + fi * 64; if (vf >= V_DIM) break; float ws = 0.0f, wt = 0.0f; for (int s = 0; s < ins; s++) { int wi = batch*ins+s; float lse = split_lse[wi*16+head]; float w = __builtin_exp2f((lse-max_lse)*LOG2E); ws += w * split_out[wi*16*V_DIM+head*V_DIM+vf]; wt += w; } u32 fb; float r = ws/(wt+1e-30f); __builtin_memcpy(&fb,&r,4); final_out[(batch*16+head)*V_DIM+vf] = (u16)(fb>>16); } }
}
"""
_attn_prog = None
_reduce_prog = None
_dev = None
_compiled = False
def _compile():
global _attn_prog, _reduce_prog, _dev, _compiled
if _compiled: return
_compiled = True
try:
from tinygrad.runtime.ops_hip import HIPDevice, HIPProgram
_dev = HIPDevice("")
b1 = _dev.compiler.compile(_ATTN_SRC); _attn_prog = HIPProgram(_dev, "mla_attn_fp4", b1)
b2 = _dev.compiler.compile(_REDUCE_SRC); _reduce_prog = HIPProgram(_dev, "mla_reduce", b2)
P(f"[XK89] OK")
except Exception as e:
P(f"[XK89] FAIL: {e}"); _attn_prog = None; _reduce_prog = None
_compile()
_mc = {}
_dd = {}
def _run(q, kv_data, kv_indptr, bs, kvsl):
if _attn_prog is None or _reduce_prog is None: return None
key = (bs, kvsl)
ns = _MFMA_NS.get(key)
if ns is None: return None
tw = bs * ns
if key not in _mc:
n_chunks = max(1, (bs + _CHUNK - 1) // _CHUNK)
# Separate buffers per chunk
chunk_sos = [torch.empty((_CHUNK * ns, 16, 512), dtype=torch.float32, device="cuda") for _ in range(n_chunks)]
chunk_sls = [torch.empty((_CHUNK * ns, 16), dtype=torch.float32, device="cuda") for _ in range(n_chunks)]
so_full = torch.empty((tw, 16, 512), dtype=torch.float32, device="cuda")
sl_full = torch.empty((tw, 16), dtype=torch.float32, device="cuda")
ob = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")
# Pre-create ki_chunks
_mc[key] = (chunk_sos, chunk_sls, so_full, sl_full, ob, n_chunks, ns)
chunk_sos, chunk_sls, so_full, sl_full, ob, n_chunks, ns = _mc[key]
fp4, fsc = kv_data["mxfp4"]
ki = kv_indptr.to(torch.int32).contiguous()
q16 = q.view(bs, 16, 576).contiguous().view(torch.uint16)
try:
if bs <= _CHUNK:
# Direct write — no copy needed for single chunk
_attn_prog(q16.data_ptr(), fp4.data_ptr(), fsc.data_ptr(), ki.data_ptr(),
so_full.data_ptr(), sl_full.data_ptr(), ns, bs,
global_size=(tw * 64, 1, 1), local_size=(64, 1, 1))
else:
ki_chunks = [ki[ci*_CHUNK:ci*_CHUNK+min(_CHUNK,bs-ci*_CHUNK)+1].contiguous() for ci in range(n_chunks)]
# Zero chunk buffers to prevent stale data
for ci in range(n_chunks):
chunk_sos[ci].zero_()
chunk_sls[ci].zero_()
for ci in range(n_chunks):
cs = ci * _CHUNK
cb = min(_CHUNK, bs - cs)
ctw = cb * ns
_attn_prog(q16.data_ptr() + cs * 16 * 576 * 2, fp4.data_ptr(), fsc.data_ptr(),
ki_chunks[ci].data_ptr(), chunk_sos[ci].data_ptr(), chunk_sls[ci].data_ptr(),
ns, cb, global_size=(ctw * 64, 1, 1), local_size=(64, 1, 1))
_dev.synchronize()
# Copy all chunks into combined buffer
offset = 0
for ci in range(n_chunks):
cb = min(_CHUNK, bs - ci * _CHUNK)
ctw = cb * ns
so_full[offset:offset+ctw].copy_(chunk_sos[ci][:ctw])
sl_full[offset:offset+ctw].copy_(chunk_sls[ci][:ctw])
offset += ctw
_reduce_prog(so_full.data_ptr(), sl_full.data_ptr(), ob.view(torch.uint16).data_ptr(),
ns, bs, global_size=(bs * 4 * 64, 1, 1), local_size=(64, 1, 1))
_dev.synchronize()
if key not in _dd:
_dd[key] = True
P(f"[XK89] OK bs={bs} kvsl={kvsl} ns={ns} chunks={n_chunks}")
return ob
except Exception as e:
P(f"[XK89] FAIL: {e}")
return None
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs, kvsl = config["batch_size"], config["kv_seq_len"]
key = (bs, kvsl)
if _attn_prog is not None and key in _MFMA_NS:
r = _run(q, kv_data, kv_indptr, bs, kvsl)
if r is not None: return r
if key not in _c:
ns = _NS.get(key, max(1, 256 // bs))
ki = torch.arange(bs * kvsl, dtype=torch.int32, device="cuda")
kl = torch.full((bs,), kvsl, dtype=torch.int32, device="cuda")
info = get_mla_metadata_info_v1(bs, 1, 16, _D, _D, is_sparse=False, fast_mode=True, num_kv_splits=ns, intra_batch_mode=True)
wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
get_mla_metadata_v1(qo_indptr, kv_indptr, kl, 16, 1, True, wk[0], wk[2], wk[1], wk[3], wk[4], wk[5],
page_size=1, kv_granularity=32, max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True, max_split_per_batch=ns, intra_batch_mode=True, dtype_q=_D, dtype_kv=_D)
np_ = wk[5].size(0)
sd = torch.empty((np_, 1, 16, 576), dtype=torch.float32, device="cuda")
sl = torch.empty((np_, 1, 16, 1), dtype=torch.float32, device="cuda")
qs = torch.ones(1, dtype=torch.float32, device="cuda")
qf = torch.empty((bs, 16, 576), dtype=_D, device="cuda")
ob = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")
_c[key] = (ki, kl, wk, sd, sl, qs, qf, ob, ns)
torch.cuda.synchronize()
ki, kl, wk, sd, sl, qs, qf, ob, ns2 = _c[key]
kf, k_scale = kv_data["fp8"]
qf.copy_(q.view(-1, 16, 576))
aiter.mla_decode_stage1_asm_fwd(qf, kf.view(-1, 1, 1, 576), qo_indptr, kv_indptr, ki, kl,
None, wk[0], wk[1], wk[2], 1, 1, 1, _S, sd, sl, ob, q_scale=qs, kv_scale=k_scale)
if ns2 > 1:
aiter.mla_reduce_v1(sd, sl, wk[3], wk[4], wk[5], 1, ob, None)
return ob
scrolls · 202 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 588791.
# /// script# leaderboard = "amd-mixed-mla"# ///- """XK72: Custom FP4 for kvsl=8192 only (passes rtol=0.1), AITER for kvsl=1024.- kvsl=8192 shapes pass custom kernel tolerance. kvsl=1024 shapes fail (3x error).- Hybrid: custom 4.3μs for 8192 shapes, AITER for 1024 shapes.- Expected geomean: ~25-30μs (huge drop from 56μs).+ """XK89: Clean separate-buffer chunking for (64,8192).+ Each chunk gets its own split_out/split_lse tensors (pre-allocated).+ After all chunks, copy into combined buffer and reduce.+ Fixes the deterministic corruption from offset pointer writes."""import sys- import structimport torchimport aiterfrom task import input_t, output_t⋯ 11 unchanged lines(64, 1024): 4, (64, 8192): 4,(256, 1024): 1, (256, 8192): 1,}- # Custom ONLY for kvsl=8192 shapes (pass tolerance)_MFMA_NS = {(4, 8192): 64,+ (64, 8192): 64,}+ _CHUNK = 16_ATTN_SRC = r"""#define NHEADS 16⋯ 5 unchanged lines#define LOG2E 1.4426950408889634f#define MYEXP(x) __builtin_exp2f((x) * LOG2E)#define NEG_INF (-1e30f)-- typedef unsigned char u8;- typedef unsigned short u16;- typedef unsigned int u32;+ typedef unsigned char u8; typedef unsigned short u16; typedef unsigned int u32;typedef u8 __attribute__((ext_vector_type(32))) reg256_t;typedef float __attribute__((ext_vector_type(4))) f32x4_t;-- #define READLANE_F32(val, lane_id) ({ \- int _rl_i; \- float _rl_f = (val); \- __builtin_memcpy(&_rl_i, &_rl_f, 4); \- int _rl_r = __builtin_amdgcn_readlane(_rl_i, (int)(lane_id)); \- float _rl_out; \- __builtin_memcpy(&_rl_out, &_rl_r, 4); \- _rl_out; \- })-- #define BF16_TO_FP8(bf16_bits) ({ \- u16 _bb = (bf16_bits); \- unsigned _sign = (_bb >> 15) & 1u; \- int _exp = (int)((_bb >> 7) & 0xFFu) - 127; \- unsigned _mant = _bb & 0x7Fu; \- int _fp8_exp = _exp + 7; \- u8 _result; \- if (_fp8_exp <= 0) { _result = (u8)(_sign << 7); } \- else if (_fp8_exp >= 15) { _result = (u8)((_sign << 7) | 0x7Eu); } \- else { _result = (u8)((_sign << 7) | (_fp8_exp << 3) | (_mant >> 4)); } \- _result; \- })-- #define F32_TO_FP8(fval) ({ \- float _fv = (fval); \- u32 _fb; \- __builtin_memcpy(&_fb, &_fv, 4); \- unsigned _sign = (_fb >> 31) & 1u; \- int _exp = (int)((_fb >> 23) & 0xFFu) - 127; \- unsigned _mant = (_fb >> 20) & 0x7u; \- int _fp8_exp = _exp + 7; \- u8 _result; \- if (_fv == 0.0f || _fp8_exp <= 0) { _result = (u8)(_sign << 7); } \- else if (_fp8_exp >= 15) { _result = (u8)((_sign << 7) | 0x7Eu); } \- else { _result = (u8)((_sign << 7) | (_fp8_exp << 3) | _mant); } \- _result; \- })-+ #define READLANE_F32(val, lane_id) ({ int _rl_i; float _rl_f = (val); __builtin_memcpy(&_rl_i, &_rl_f, 4); int _rl_r = __builtin_amdgcn_readlane(_rl_i, (int)(lane_id)); float _rl_out; __builtin_memcpy(&_rl_out, &_rl_r, 4); _rl_out; })+ #define BF16_TO_FP8(bf16_bits) ({ u16 _bb = (bf16_bits); unsigned _sign = (_bb >> 15) & 1u; int _exp = (int)((_bb >> 7) & 0xFFu) - 127; unsigned _mant = _bb & 0x7Fu; int _fp8_exp = _exp + 7; u8 _result; if (_fp8_exp <= 0) { _result = (u8)(_sign << 7); } else if (_fp8_exp >= 15) { _result = (u8)((_sign << 7) | 0x7Eu); } else { _result = (u8)((_sign << 7) | (_fp8_exp << 3) | (_mant >> 4)); } _result; })+ #define F32_TO_FP8(fval) ({ float _fv = (fval); u32 _fb; __builtin_memcpy(&_fb, &_fv, 4); unsigned _sign = (_fb >> 31) & 1u; int _exp = (int)((_fb >> 23) & 0xFFu) - 127; unsigned _mant = (_fb >> 20) & 0x7u; int _fp8_exp = _exp + 7; u8 _result; if (_fv == 0.0f || _fp8_exp <= 0) { _result = (u8)(_sign << 7); } else if (_fp8_exp >= 15) { _result = (u8)((_sign << 7) | 0x7Eu); } else { _result = (u8)((_sign << 7) | (_fp8_exp << 3) | _mant); } _result; })__attribute__((shared)) float lds_scores[16 * 128];__attribute__((shared)) u8 lds_p[16 * 128];-extern "C" __attribute__((global)) void __attribute__((amdgpu_flat_work_group_size(64, 64)))- mla_attn_fp4(- const u16* __restrict__ Q_bf16,- const u8* __restrict__ KV_fp4,- const u8* __restrict__ KV_scale,- const int* __restrict__ kv_indptr,- float* __restrict__ split_out,- float* __restrict__ split_lse,- long ns, long bs- ) {- unsigned wg_id = __builtin_amdgcn_workgroup_id_x();- unsigned tid = __builtin_amdgcn_workitem_id_x();- unsigned group = tid >> 4;- unsigned lane = tid & 15u;- int ins = (int)ns, ibs = (int)bs;- int batch = wg_id / ins, split = wg_id % ins;- if (batch >= ibs) return;- int kv_start = kv_indptr[batch];- int total_tokens = kv_indptr[batch + 1] - kv_start;- int tps = (total_tokens + ins - 1) / ins;- int my_start = split * tps;- int my_end = my_start + tps;- if (my_end > total_tokens) my_end = total_tokens;- int my_tokens = my_end - my_start;- if (my_tokens <= 0) {- for (int s = 0; s < 4; s++) {- int head = group * 4 + s;- split_lse[wg_id * 16 + head] = NEG_INF;- float* so = split_out + wg_id * 16 * V_DIM + head * V_DIM;- for (int fc = 0; fc < 32; fc++) { int vf = lane + fc * 16; if (vf < V_DIM) so[vf] = 0.0f; }- }- return;- }+ mla_attn_fp4(const u16* __restrict__ Q_bf16, const u8* __restrict__ KV_fp4, const u8* __restrict__ KV_scale, const int* __restrict__ kv_indptr, float* __restrict__ split_out, float* __restrict__ split_lse, long ns, long bs) {+ unsigned wg_id = __builtin_amdgcn_workgroup_id_x(); unsigned tid = __builtin_amdgcn_workitem_id_x(); unsigned group = tid >> 4; unsigned lane = tid & 15u;+ int ins = (int)ns, ibs = (int)bs; int batch = wg_id / ins, split = wg_id % ins; if (batch >= ibs) return;+ int kv_start = kv_indptr[batch]; int total_tokens = kv_indptr[batch + 1] - kv_start;+ int tps = (total_tokens + ins - 1) / ins; int my_start = split * tps; int my_end = my_start + tps;+ if (my_end > total_tokens) my_end = total_tokens; int my_tokens = my_end - my_start;+ if (my_tokens <= 0) { for (int s = 0; s < 4; s++) { int head = group * 4 + s; split_lse[wg_id * 16 + head] = NEG_INF; float* so = split_out + wg_id * 16 * V_DIM + head * V_DIM; for (int fc = 0; fc < 32; fc++) { int vf = lane + fc * 16; if (vf < V_DIM) so[vf] = 0.0f; } } return; }int kv_off = kv_start + my_start;- reg256_t q_a[5];- const u16* qh = Q_bf16 + batch * NHEADS * HEAD_DIM;- for (int call = 0; call < 5; call++) {- int kb = call * 128 + group * 32;- for (int i = 0; i < 32; i++) q_a[call][i] = 0;- if (kb < HEAD_DIM) {- int valid = HEAD_DIM - kb; if (valid > 32) valid = 32;- const u16* qp = qh + lane * HEAD_DIM + kb;- for (int i = 0; i < valid; i++) q_a[call][i] = BF16_TO_FP8(qp[i]);- }- }- float pv_acc[4][32];- float running_max[4], running_sum[4];- for (int s = 0; s < 4; s++) {- for (int fc = 0; fc < 32; fc++) pv_acc[s][fc] = 0.0f;- running_max[s] = NEG_INF; running_sum[s] = 0.0f;- }+ reg256_t q_a[5]; const u16* qh = Q_bf16 + batch * NHEADS * HEAD_DIM;+ for (int call = 0; call < 5; call++) { int kb = call * 128 + group * 32; for (int i = 0; i < 32; i++) q_a[call][i] = 0; if (kb < HEAD_DIM) { int valid = HEAD_DIM - kb; if (valid > 32) valid = 32; const u16* qp = qh + lane * HEAD_DIM + kb; for (int i = 0; i < valid; i++) q_a[call][i] = BF16_TO_FP8(qp[i]); } }+ float pv_acc[4][32]; float running_max[4], running_sum[4];+ for (int s = 0; s < 4; s++) { for (int fc = 0; fc < 32; fc++) pv_acc[s][fc] = 0.0f; running_max[s] = NEG_INF; running_sum[s] = 0.0f; }for (int blk = 0; blk < my_tokens; blk += PV_BLOCK) {- int bt = my_tokens - blk; if (bt > PV_BLOCK) bt = PV_BLOCK;- int nt = (bt + QK_TILE - 1) / QK_TILE;- for (int tile = 0; tile < nt; tile++) {- int ts = tile * QK_TILE;- int tc = bt - ts; if (tc > QK_TILE) tc = QK_TILE;- int gtok = kv_off + blk + ts + lane;- int valid = ((int)lane < tc) ? 1 : 0;- f32x4_t c_qk; for (int i = 0; i < 4; i++) c_qk[i] = 0.0f;- for (int call = 0; call < 5; call++) {- int kb = call * 128 + group * 32;- reg256_t b_reg; for (int i = 0; i < 32; i++) b_reg[i] = 0;- u8 kv_sc = 127;- if (valid && kb < HEAD_DIM) {- int bo = kb / 2, vb = (HEAD_DIM - kb + 1) / 2; if (vb > 16) vb = 16;- const u8* kp = KV_fp4 + gtok * 288 + bo;- for (int i = 0; i < vb; i++) b_reg[i] = kp[i];- int sg = kb / 32; if (sg >= 18) sg = 17;- kv_sc = KV_scale[gtok * 24 + sg];- }- c_qk = (f32x4_t)__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(- q_a[call], b_reg, c_qk, 0, 4, 0, 127, 0, kv_sc);- }- for (int s = 0; s < 4; s++) {- float sc = valid ? (c_qk[s] * SM_SCALE) : NEG_INF;- lds_scores[(group * 4 + s) * PV_BLOCK + ts + lane] = sc;- }- }+ int bt = my_tokens - blk; if (bt > PV_BLOCK) bt = PV_BLOCK; int nt = (bt + QK_TILE - 1) / QK_TILE;+ for (int tile = 0; tile < nt; tile++) { int ts = tile * QK_TILE; int tc = bt - ts; if (tc > QK_TILE) tc = QK_TILE; int gtok = kv_off + blk + ts + lane; int valid = ((int)lane < tc) ? 1 : 0; f32x4_t c_qk; for (int i = 0; i < 4; i++) c_qk[i] = 0.0f; for (int call = 0; call < 5; call++) { int kb = call * 128 + group * 32; reg256_t b_reg; for (int i = 0; i < 32; i++) b_reg[i] = 0; u8 kv_sc = 127; if (valid && kb < HEAD_DIM) { int bo = kb / 2, vb = (HEAD_DIM - kb + 1) / 2; if (vb > 16) vb = 16; const u8* kp = KV_fp4 + gtok * 288 + bo; for (int i = 0; i < vb; i++) b_reg[i] = kp[i]; int sg = kb / 32; if (sg >= 18) sg = 17; kv_sc = KV_scale[gtok * 24 + sg]; } c_qk = (f32x4_t)__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(q_a[call], b_reg, c_qk, 0, 4, 0, 127, 0, kv_sc); } for (int s = 0; s < 4; s++) { float sc = valid ? (c_qk[s] * SM_SCALE) : NEG_INF; lds_scores[(group * 4 + s) * PV_BLOCK + ts + lane] = sc; } }__builtin_amdgcn_s_barrier();float block_max[4], block_sum[4];- for (int s = 0; s < 4; s++) {- int head = group * 4 + s;- float lmax = NEG_INF; float msc[8];- for (int j = 0; j < 8; j++) { int tok = lane + j * 16; msc[j] = (tok < bt) ? lds_scores[head * PV_BLOCK + tok] : NEG_INF; if (msc[j] > lmax) lmax = msc[j]; }- float mx = lmax; int bl = group * 16;- for (int j = 0; j < 16; j++) { float v = READLANE_F32(lmax, bl+j); if (v > mx) mx = v; }- block_max[s] = mx;- float lsum = 0.0f;- for (int j = 0; j < 8; j++) { int tok = lane + j * 16; float p = (tok < bt) ? MYEXP(msc[j] - mx) : 0.0f; lsum += p; lds_p[head * PV_BLOCK + tok] = F32_TO_FP8(p * 256.0f); }- float tsum = 0.0f;- for (int j = 0; j < 16; j++) tsum += READLANE_F32(lsum, bl+j);- block_sum[s] = tsum;- }+ for (int s = 0; s < 4; s++) { int head = group * 4 + s; float lmax = NEG_INF; float msc[8]; for (int j = 0; j < 8; j++) { int tok = lane + j * 16; msc[j] = (tok < bt) ? lds_scores[head * PV_BLOCK + tok] : NEG_INF; if (msc[j] > lmax) lmax = msc[j]; } float mx = lmax; int bl = group * 16; for (int j = 0; j < 16; j++) { float v = READLANE_F32(lmax, bl+j); if (v > mx) mx = v; } block_max[s] = mx; float lsum = 0.0f; for (int j = 0; j < 8; j++) { int tok = lane + j * 16; float p = (tok < bt) ? MYEXP(msc[j] - mx) : 0.0f; lsum += p; lds_p[head * PV_BLOCK + tok] = F32_TO_FP8(p * 256.0f); } float tsum = 0.0f; for (int j = 0; j < 16; j++) tsum += READLANE_F32(lsum, bl+j); block_sum[s] = tsum; }__builtin_amdgcn_s_barrier();- for (int s = 0; s < 4; s++) {- float om = running_max[s]; float nm = (block_max[s] > om) ? block_max[s] : om;- float eo = MYEXP(om - nm), eb = MYEXP(block_max[s] - nm);- for (int fc = 0; fc < 32; fc++) pv_acc[s][fc] *= eo;- running_sum[s] = running_sum[s] * eo + block_sum[s] * eb; running_max[s] = nm;- }- reg256_t p_a;- for (int b = 0; b < 32; b++) { int tok = group * 32 + b; p_a[b] = (tok < bt) ? lds_p[lane * PV_BLOCK + tok] : 0; }- for (int fc = 0; fc < 32; fc++) {- f32x4_t c_pv; for (int i = 0; i < 4; i++) c_pv[i] = 0.0f;- int vf = lane + fc * 16;- if (vf < V_DIM) {- int fb = vf / 2, ns2 = vf & 1, sg = vf / 32; if (sg >= 16) sg = 15;- reg256_t v_b;- for (int i = 0; i < 16; i++) {- int t0 = group*32+i*2, t1 = t0+1; u8 n0=0,n1=0;- if (t0 < bt) { int g0=kv_off+blk+t0; u8 b0=KV_fp4[g0*288+fb]; n0=ns2?((b0>>4)&0xFu):(b0&0xFu); }- if (t1 < bt) { int g1=kv_off+blk+t1; u8 b1=KV_fp4[g1*288+fb]; n1=ns2?((b1>>4)&0xFu):(b1&0xFu); }- v_b[i] = n0 | (n1 << 4);- }- int rt = group*32+16; if (rt >= bt) rt = bt > 0 ? bt-1 : 0;- u8 vsc = KV_scale[(kv_off+blk+rt)*24+sg];- c_pv = (f32x4_t)__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(p_a, v_b, c_pv, 0, 4, 0, 119, 0, vsc);- }- for (int s = 0; s < 4; s++) { float corr = MYEXP(block_max[s] - running_max[s]); pv_acc[s][fc] += c_pv[s] * corr; }- }+ for (int s = 0; s < 4; s++) { float om = running_max[s]; float nm = (block_max[s] > om) ? block_max[s] : om; float eo = MYEXP(om - nm), eb = MYEXP(block_max[s] - nm); for (int fc = 0; fc < 32; fc++) pv_acc[s][fc] *= eo; running_sum[s] = running_sum[s] * eo + block_sum[s] * eb; running_max[s] = nm; }+ reg256_t p_a; for (int b = 0; b < 32; b++) { int tok = group * 32 + b; p_a[b] = (tok < bt) ? lds_p[lane * PV_BLOCK + tok] : 0; }+ for (int fc = 0; fc < 32; fc++) { f32x4_t c_pv; for (int i = 0; i < 4; i++) c_pv[i] = 0.0f; int vf = lane + fc * 16; if (vf < V_DIM) { int fb = vf / 2, ns2 = vf & 1, sg = vf / 32; if (sg >= 16) sg = 15; reg256_t v_b; for (int i = 0; i < 16; i++) { int t0 = group*32+i*2, t1 = t0+1; u8 n0=0,n1=0; if (t0 < bt) { int g0=kv_off+blk+t0; u8 b0=KV_fp4[g0*288+fb]; n0=ns2?((b0>>4)&0xFu):(b0&0xFu); } if (t1 < bt) { int g1=kv_off+blk+t1; u8 b1=KV_fp4[g1*288+fb]; n1=ns2?((b1>>4)&0xFu):(b1&0xFu); } v_b[i] = n0 | (n1 << 4); } int rt = group*32+16; if (rt >= bt) rt = bt > 0 ? bt-1 : 0; u8 vsc = KV_scale[(kv_off+blk+rt)*24+sg]; c_pv = (f32x4_t)__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(p_a, v_b, c_pv, 0, 4, 0, 119, 0, vsc); } for (int s = 0; s < 4; s++) { float corr = MYEXP(block_max[s] - running_max[s]); pv_acc[s][fc] += c_pv[s] * corr; } }__builtin_amdgcn_s_barrier();}- for (int s = 0; s < 4; s++) {- int head = group * 4 + s;- split_lse[wg_id * 16 + head] = running_max[s] + __builtin_log2f(running_sum[s] + 1e-30f) / LOG2E;- float inv = (running_sum[s] > 1e-9f) ? (1.0f / running_sum[s]) : 0.0f;- float* so = split_out + wg_id * 16 * V_DIM + head * V_DIM;- for (int fc = 0; fc < 32; fc++) { int vf = lane + fc * 16; if (vf < V_DIM) so[vf] = pv_acc[s][fc] * inv; }- }+ for (int s = 0; s < 4; s++) { int head = group * 4 + s; split_lse[wg_id * 16 + head] = running_max[s] + __builtin_log2f(running_sum[s] + 1e-30f) / LOG2E; float inv = (running_sum[s] > 1e-9f) ? (1.0f / running_sum[s]) : 0.0f; float* so = split_out + wg_id * 16 * V_DIM + head * V_DIM; for (int fc = 0; fc < 32; fc++) { int vf = lane + fc * 16; if (vf < V_DIM) so[vf] = pv_acc[s][fc] * inv; } }}"""-_REDUCE_SRC = r"""#define V_DIM 512#define LOG2E 1.4426950408889634f⋯ 2 unchanged linesextern "C" __attribute__((global)) void __attribute__((amdgpu_flat_work_group_size(64, 64)))mla_reduce(const float* __restrict__ split_out, const float* __restrict__ split_lse, u16* __restrict__ final_out, long ns, long bs) {unsigned wg_id = __builtin_amdgcn_workgroup_id_x(); unsigned tid = __builtin_amdgcn_workitem_id_x();- int ins = (int)ns, batch = wg_id / 4, hg = wg_id % 4;- if (batch >= (int)bs) return;- for (int hs = 0; hs < 4; hs++) {- int head = hg * 4 + hs; float max_lse = NEG_INF;- for (int s = 0; s < ins; s++) { float lse = split_lse[(batch*ins+s)*16+head]; if (lse > max_lse) max_lse = lse; }- for (int fi = 0; fi < 8; fi++) {- int vf = tid + fi * 64; if (vf >= V_DIM) break;- float ws = 0.0f, wt = 0.0f;- for (int s = 0; s < ins; s++) { int wi = batch*ins+s; float lse = split_lse[wi*16+head]; float w = __builtin_exp2f((lse-max_lse)*LOG2E); ws += w * split_out[wi*16*V_DIM+head*V_DIM+vf]; wt += w; }- u32 fb; float r = ws/(wt+1e-30f); __builtin_memcpy(&fb,&r,4); final_out[(batch*16+head)*V_DIM+vf] = (u16)(fb>>16);- }- }+ int ins = (int)ns, batch = wg_id / 4, hg = wg_id % 4; if (batch >= (int)bs) return;+ for (int hs = 0; hs < 4; hs++) { int head = hg * 4 + hs; float max_lse = NEG_INF; for (int s = 0; s < ins; s++) { float lse = split_lse[(batch*ins+s)*16+head]; if (lse > max_lse) max_lse = lse; } for (int fi = 0; fi < 8; fi++) { int vf = tid + fi * 64; if (vf >= V_DIM) break; float ws = 0.0f, wt = 0.0f; for (int s = 0; s < ins; s++) { int wi = batch*ins+s; float lse = split_lse[wi*16+head]; float w = __builtin_exp2f((lse-max_lse)*LOG2E); ws += w * split_out[wi*16*V_DIM+head*V_DIM+vf]; wt += w; } u32 fb; float r = ws/(wt+1e-30f); __builtin_memcpy(&fb,&r,4); final_out[(batch*16+head)*V_DIM+vf] = (u16)(fb>>16); } }}"""⋯ 8 unchanged linestry:from tinygrad.runtime.ops_hip import HIPDevice, HIPProgram_dev = HIPDevice("")- b1 = _dev.compiler.compile(_ATTN_SRC)- P(f"[XK72] Attn: {len(b1)} bytes")- _attn_prog = HIPProgram(_dev, "mla_attn_fp4", b1)- b2 = _dev.compiler.compile(_REDUCE_SRC)- _reduce_prog = HIPProgram(_dev, "mla_reduce", b2)- P("[XK72] OK")+ b1 = _dev.compiler.compile(_ATTN_SRC); _attn_prog = HIPProgram(_dev, "mla_attn_fp4", b1)+ b2 = _dev.compiler.compile(_REDUCE_SRC); _reduce_prog = HIPProgram(_dev, "mla_reduce", b2)+ P(f"[XK89] OK")except Exception as e:- P(f"[XK72] FAIL: {e}")- _attn_prog = None; _reduce_prog = None+ P(f"[XK89] FAIL: {e}"); _attn_prog = None; _reduce_prog = None_compile()_mc = {}⋯ 5 unchanged linesif ns is None: return Nonetw = bs * nsif key not in _mc:- _mc[key] = (torch.empty((tw, 16, 512), dtype=torch.float32, device="cuda"), torch.empty((tw, 16), dtype=torch.float32, device="cuda"), torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda"))- so, sl, ob = _mc[key]+ n_chunks = max(1, (bs + _CHUNK - 1) // _CHUNK)+ # Separate buffers per chunk+ chunk_sos = [torch.empty((_CHUNK * ns, 16, 512), dtype=torch.float32, device="cuda") for _ in range(n_chunks)]+ chunk_sls = [torch.empty((_CHUNK * ns, 16), dtype=torch.float32, device="cuda") for _ in range(n_chunks)]+ so_full = torch.empty((tw, 16, 512), dtype=torch.float32, device="cuda")+ sl_full = torch.empty((tw, 16), dtype=torch.float32, device="cuda")+ ob = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")+ # Pre-create ki_chunks+ _mc[key] = (chunk_sos, chunk_sls, so_full, sl_full, ob, n_chunks, ns)+ chunk_sos, chunk_sls, so_full, sl_full, ob, n_chunks, ns = _mc[key]fp4, fsc = kv_data["mxfp4"]ki = kv_indptr.to(torch.int32).contiguous()q16 = q.view(bs, 16, 576).contiguous().view(torch.uint16)try:- _attn_prog(q16.data_ptr(), fp4.data_ptr(), fsc.data_ptr(), ki.data_ptr(), so.data_ptr(), sl.data_ptr(), ns, bs, global_size=(tw * 64, 1, 1), local_size=(64, 1, 1))- _reduce_prog(so.data_ptr(), sl.data_ptr(), ob.view(torch.uint16).data_ptr(), ns, bs, global_size=(bs * 4 * 64, 1, 1), local_size=(64, 1, 1))+ if bs <= _CHUNK:+ # Direct write — no copy needed for single chunk+ _attn_prog(q16.data_ptr(), fp4.data_ptr(), fsc.data_ptr(), ki.data_ptr(),+ so_full.data_ptr(), sl_full.data_ptr(), ns, bs,+ global_size=(tw * 64, 1, 1), local_size=(64, 1, 1))+ else:+ ki_chunks = [ki[ci*_CHUNK:ci*_CHUNK+min(_CHUNK,bs-ci*_CHUNK)+1].contiguous() for ci in range(n_chunks)]+ # Zero chunk buffers to prevent stale data+ for ci in range(n_chunks):+ chunk_sos[ci].zero_()+ chunk_sls[ci].zero_()+ for ci in range(n_chunks):+ cs = ci * _CHUNK+ cb = min(_CHUNK, bs - cs)+ ctw = cb * ns+ _attn_prog(q16.data_ptr() + cs * 16 * 576 * 2, fp4.data_ptr(), fsc.data_ptr(),+ ki_chunks[ci].data_ptr(), chunk_sos[ci].data_ptr(), chunk_sls[ci].data_ptr(),+ ns, cb, global_size=(ctw * 64, 1, 1), local_size=(64, 1, 1))+ _dev.synchronize()+ # Copy all chunks into combined buffer+ offset = 0+ for ci in range(n_chunks):+ cb = min(_CHUNK, bs - ci * _CHUNK)+ ctw = cb * ns+ so_full[offset:offset+ctw].copy_(chunk_sos[ci][:ctw])+ sl_full[offset:offset+ctw].copy_(chunk_sls[ci][:ctw])+ offset += ctw+ _reduce_prog(so_full.data_ptr(), sl_full.data_ptr(), ob.view(torch.uint16).data_ptr(),+ ns, bs, global_size=(bs * 4 * 64, 1, 1), local_size=(64, 1, 1))_dev.synchronize()if key not in _dd:_dd[key] = True- P(f"[XK72] OK bs={bs} kvsl={kvsl} ns={ns}")+ P(f"[XK89] OK bs={bs} kvsl={kvsl} ns={ns} chunks={n_chunks}")return obexcept Exception as e:- P(f"[XK72] FAIL: {e}")+ P(f"[XK89] FAIL: {e}")return Nonedef custom_kernel(data: input_t) -> output_t:
scrolls · 344 diff lines total
Best evidence level for this revision: reported
JSON