submission 739756
divc13 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 233 lines, June 9 Researcher Reciprocity License v1.0.
submission_v94.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-739756?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:7186d9ddb65c333c84f7bef41bd0cf2d2e00b9f324593c16154e7edcbe9318ee
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
- bs32/64 kv1k: low splits + ns=3 (v92 autotune), less stage2 overheadfp8
FP8_DTYPE = tl.float8e4nvmma
qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)num-warps = 4
num_warps=4, num_stages=ns, waves_per_eu=2,stages = 2
num_warps=4, num_stages=2,tile-n = 64
NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,Kernel source
submission_v94.py233 lines
"""MLA decode Triton v94: cherry-picked optimal (splits, ns) per shape.
Best configs from v84/v92/v93 experiments:
- bs4: high splits (v84), GPU underutilized
- bs32/64 kv1k: low splits + ns=3 (v92 autotune), less stage2 overhead
- bs32/64 kv8k: moderate splits (v93 autotune)
- bs256: force_single kv1k, 2-split kv8k (v93 autotune)
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
NUM_HEADS = 16
KV_LORA_RANK = 512
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
FP8_DTYPE = tl.float8e4nv
_SHAPE_CONFIGS = {
(4, 1024): (16, 2),
(4, 8192): (32, 3),
(32, 1024): (4, 3),
(32, 8192): (8, 3),
(64, 1024): (4, 3),
(64, 8192): (8, 2),
(256, 1024): (1, 2),
(256, 8192): (2, 2),
}
@triton.jit
def _remap_xcd(pid, GRID_SIZE, NUM_XCDS: tl.constexpr = 8):
pids_per_xcd = (GRID_SIZE + NUM_XCDS - 1) // NUM_XCDS
tall_xcds = GRID_SIZE % NUM_XCDS
tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
xcd = pid % NUM_XCDS
local_pid = pid // NUM_XCDS
if xcd < tall_xcds:
return xcd * pids_per_xcd + local_pid
else:
return tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid
@triton.jit
def _stage1_kernel(
Q, KV, Mid_O, Out,
qo_indptr, kv_indptr,
sm_scale, kv_scale_ptr,
stride_qt, stride_qh, stride_kvt,
stride_mb, stride_mh, stride_ms,
stride_ot, stride_oh,
NUM_KV_SPLITS: tl.constexpr, BLOCK_N: tl.constexpr,
BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr, BLOCK_R: tl.constexpr,
DIRECT_OUT: tl.constexpr, GRID_SIZE: tl.constexpr,
):
raw_pid = tl.program_id(0)
pid = _remap_xcd(raw_pid, GRID_SIZE)
cur_batch = pid // NUM_KV_SPLITS
split_id = pid % NUM_KV_SPLITS
heads = tl.arange(0, BLOCK_H)
kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
score_scale = sm_scale * kv_scale
q_tok = tl.load(qo_indptr + cur_batch)
kv_s = tl.load(kv_indptr + cur_batch)
kv_e = tl.load(kv_indptr + cur_batch + 1)
kv_len = kv_e - kv_s
kv_per_split = tl.cdiv(kv_len, NUM_KV_SPLITS)
sp_start = kv_per_split * split_id
sp_end = tl.minimum(sp_start + kv_per_split, kv_len)
offs_c = tl.arange(0, BLOCK_C)
offs_r = tl.arange(0, BLOCK_R)
q_base = q_tok * stride_qt
q_nope = tl.load(Q + q_base + heads[:, None] * stride_qh + offs_c[None, :]).to(FP8_DTYPE)
q_rope = tl.load(Q + q_base + heads[:, None] * stride_qh + (512 + offs_r[None, :])).to(FP8_DTYPE)
e_max = tl.full([BLOCK_H], value=float("-inf"), dtype=tl.float32)
e_sum = tl.zeros([BLOCK_H], dtype=tl.float32)
acc = tl.zeros([BLOCK_H, BLOCK_C], dtype=tl.float32)
num_tokens = sp_end - sp_start
num_full = num_tokens // BLOCK_N
full_end = sp_start + num_full * BLOCK_N
for start_n in range(sp_start, full_end, BLOCK_N):
kv_idx = kv_s + start_n + tl.arange(0, BLOCK_N)
kv_nope = tl.load(KV + kv_idx[:, None] * stride_kvt + offs_c[None, :],
cache_modifier=".cg")
kv_rope = tl.load(KV + kv_idx[:, None] * stride_kvt + (512 + offs_r[None, :]),
cache_modifier=".cg")
k_nope = tl.trans(kv_nope)
k_rope = tl.trans(kv_rope)
qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)
qk = qk * score_scale
new_max = tl.maximum(tl.max(qk, 1), e_max)
rescale = tl.exp(e_max - new_max)
p = tl.exp(qk - new_max[:, None])
acc = acc * rescale[:, None]
acc = acc + tl.dot(p.to(FP8_DTYPE), kv_nope)
e_sum = e_sum * rescale + tl.sum(p, 1)
e_max = new_max
if full_end < sp_end:
offs_n = full_end + tl.arange(0, BLOCK_N)
mask_n = offs_n < sp_end
kv_idx = kv_s + offs_n
kv_nope = tl.load(KV + kv_idx[:, None] * stride_kvt + offs_c[None, :],
mask=mask_n[:, None], other=0.0, cache_modifier=".cg")
kv_rope = tl.load(KV + kv_idx[:, None] * stride_kvt + (512 + offs_r[None, :]),
mask=mask_n[:, None], other=0.0, cache_modifier=".cg")
k_nope = tl.trans(kv_nope)
k_rope = tl.trans(kv_rope)
qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)
qk = qk * score_scale
qk = tl.where(mask_n[None, :], qk, float("-inf"))
new_max = tl.maximum(tl.max(qk, 1), e_max)
rescale = tl.exp(e_max - new_max)
p = tl.exp(qk - new_max[:, None])
acc = acc * rescale[:, None]
acc = acc + tl.dot(p.to(FP8_DTYPE), kv_nope)
e_sum = e_sum * rescale + tl.sum(p, 1)
e_max = new_max
safe_esum = tl.where(e_sum > 0, e_sum, 1.0)
result = acc / safe_esum[:, None]
if DIRECT_OUT:
o_base = q_tok * stride_ot
tl.store(Out + o_base + heads[:, None] * stride_oh + offs_c[None, :],
(result * kv_scale).to(tl.bfloat16))
else:
mid_base = cur_batch * stride_mb + heads * stride_mh + split_id * stride_ms
tl.store(Mid_O + mid_base[:, None] + offs_c[None, :], result)
lse = tl.where(e_sum > 0, e_max + tl.log(e_sum), float("-inf"))
tl.store(Mid_O + mid_base + 512, lse)
@triton.jit
def _stage2_kernel(
Mid_O, Out, qo_indptr, kv_scale_ptr,
stride_mb, stride_mh, stride_ms, stride_ot, stride_oh,
NUM_KV_SPLITS: tl.constexpr, BLOCK_DV: tl.constexpr,
batch: tl.constexpr, GRID_SIZE: tl.constexpr,
):
raw_pid = tl.program_id(0)
pid = _remap_xcd(raw_pid, GRID_SIZE)
cur_batch = pid % batch
cur_head = pid // batch
kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
q_tok = tl.load(qo_indptr + cur_batch)
offs_d = tl.arange(0, BLOCK_DV)
e_max = float("-inf")
e_sum = 0.0
acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
mid_base = cur_batch * stride_mb + cur_head * stride_mh
for s in range(NUM_KV_SPLITS):
tv = tl.load(Mid_O + mid_base + s * stride_ms + offs_d)
lse = tl.load(Mid_O + mid_base + s * stride_ms + 512)
new_max = tl.maximum(lse, e_max)
old_scale = tl.exp(e_max - new_max)
exp_lse = tl.exp(lse - new_max)
acc = acc * old_scale + exp_lse * tv
e_sum = e_sum * old_scale + exp_lse
e_max = new_max
result = acc / tl.maximum(e_sum, 1e-12) * kv_scale
o_base = q_tok * stride_ot + cur_head * stride_oh
tl.store(Out + o_base + offs_d, result.to(tl.bfloat16))
def _get_config(batch_size, kv_len):
key = (batch_size, kv_len)
if key in _SHAPE_CONFIGS:
return _SHAPE_CONFIGS[key]
max_tiles = max(1, kv_len // 64)
if batch_size >= 256:
splits = min(2, max_tiles)
elif batch_size >= 64:
splits = min(max(1, 512 // batch_size), max_tiles)
else:
splits = min(max(1, 768 // batch_size), max_tiles)
while splits > 1 and batch_size * splits > 912:
splits -= 1
grid_size = batch_size * splits
ns = 3 if grid_size < 512 else 2
return splits, ns
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_fp8, kv_scale_t = kv_data["fp8"]
kv = kv_fp8.view(torch.float8_e4m3fn).view(-1, QK_HEAD_DIM)
batch_size = config["batch_size"]
sm_scale = config["sm_scale"]
total_q = q.shape[0]
kv_len = kv.shape[0] // batch_size if batch_size > 0 else 0
num_splits, ns = _get_config(batch_size, kv_len)
dev = q.device
direct_out = (num_splits == 1)
mid_o = torch.empty(
(batch_size, NUM_HEADS, num_splits, KV_LORA_RANK + 1),
dtype=torch.float32, device=dev,
) if not direct_out else torch.empty(1, dtype=torch.float32, device=dev)
out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
grid_size1 = batch_size * num_splits
_stage1_kernel[(grid_size1,)](
q, kv, mid_o, out, qo_indptr, kv_indptr,
sm_scale, kv_scale_t,
q.stride(0), q.stride(1), kv.stride(0),
mid_o.stride(0) if not direct_out else 0,
mid_o.stride(1) if not direct_out else 0,
mid_o.stride(2) if not direct_out else 0,
out.stride(0), out.stride(1),
NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,
BLOCK_C=512, BLOCK_R=64, DIRECT_OUT=direct_out,
GRID_SIZE=grid_size1,
num_warps=4, num_stages=ns, waves_per_eu=2,
)
if not direct_out:
grid_size2 = NUM_HEADS * batch_size
_stage2_kernel[(grid_size2,)](
mid_o, out, qo_indptr, kv_scale_t,
mid_o.stride(0), mid_o.stride(1), mid_o.stride(2),
out.stride(0), out.stride(1),
NUM_KV_SPLITS=num_splits, BLOCK_DV=512,
batch=batch_size, GRID_SIZE=grid_size2,
num_warps=4, num_stages=2,
)
return out
scrolls · 233 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 710716.
+ """MLA decode Triton v94: cherry-picked optimal (splits, ns) per shape.++ Best configs from v84/v92/v93 experiments:+ - bs4: high splits (v84), GPU underutilized+ - bs32/64 kv1k: low splits + ns=3 (v92 autotune), less stage2 overhead+ - bs32/64 kv8k: moderate splits (v93 autotune)+ - bs256: force_single kv1k, 2-split kv8k (v93 autotune)+ """+import torch- import os- from torch.utils.cpp_extension import load_inline+ import triton+ import triton.language as tlfrom task import input_t, output_t- # ---------------------------------------------------------------------------- # MLA decode v249: Tuned split strategy + parallel pack- # - Target 912 grid for large workloads, 768+ for medium- # - Parallel packing: 2-step lshl_or chain instead of 3-step- # - Base: v245 (restored v232 barriers + ubfe/lshl_or)- # ---------------------------------------------------------------------------+ NUM_HEADS = 16+ KV_LORA_RANK = 512+ QK_HEAD_DIM = 576+ V_HEAD_DIM = 512+ FP8_DTYPE = tl.float8e4nv- os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"-- HIP_SRC = r"""- #include <torch/extension.h>- #include <hip/hip_runtime.h>-- static constexpr int WARP_SIZE = 64;- static constexpr int NUM_WARPS = 4;- static constexpr int BLOCK_SIZE = WARP_SIZE * NUM_WARPS;-- static constexpr int QK_DIM = 576;- static constexpr int V_DIM = 512;- static constexpr int NUM_HEADS = 16;- static constexpr int NUM_K_CHUNKS = QK_DIM / 32; // 18- static constexpr int NUM_K128_CHUNKS = (QK_DIM + 127) / 128; // 5- static constexpr int SUPER_TILE = 32;- static constexpr int SV_CHUNKS = 8;- static constexpr int KV_TILE_BYTES = SUPER_TILE * QK_DIM; // 18432-- // 18432 bytes / 16 bytes per uint4 / 256 threads = 4.5 -> 5 rounds- static constexpr int PF_UINT4S = KV_TILE_BYTES / 16; // 1152- static constexpr int PF_ROUNDS = (PF_UINT4S + BLOCK_SIZE - 1) / BLOCK_SIZE; // 5-- typedef float __attribute__((ext_vector_type(4))) v4f32;- typedef unsigned int __attribute__((ext_vector_type(4))) u32x4;- typedef int __attribute__((ext_vector_type(4))) i32x4;- typedef int __attribute__((ext_vector_type(8))) i32x8;- typedef unsigned int __attribute__((address_space(3)))* lds_ptr_t;-- extern "C" __device__ void __llvm_amdgcn_raw_buffer_load_lds(- i32x4 rsrc, lds_ptr_t lds_ptr, int size,- int voffset, int soffset, int offset, int aux)- __asm("llvm.amdgcn.raw.buffer.load.lds");-- struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };-- __device__ __forceinline__ i32x4 make_buffer_rsrc(const void* p, uint32_t bytes) {- buffer_resource r = {reinterpret_cast<uint64_t>(p), bytes, 0x110000};- return *reinterpret_cast<i32x4*>(&r);+ _SHAPE_CONFIGS = {+ (4, 1024): (16, 2),+ (4, 8192): (32, 3),+ (32, 1024): (4, 3),+ (32, 8192): (8, 3),+ (64, 1024): (4, 3),+ (64, 8192): (8, 2),+ (256, 1024): (1, 2),+ (256, 8192): (2, 2),}- __device__ __forceinline__ float bf16_to_f32(unsigned short v) {- return __uint_as_float(static_cast<unsigned int>(v) << 16);- }- __device__ __forceinline__ unsigned short f32_to_bf16(float v) {- unsigned int bits = __float_as_uint(v);- bits += 0x7FFF + ((bits >> 16) & 1);- return static_cast<unsigned short>(bits >> 16);- }+ @triton.jit+ def _remap_xcd(pid, GRID_SIZE, NUM_XCDS: tl.constexpr = 8):+ pids_per_xcd = (GRID_SIZE + NUM_XCDS - 1) // NUM_XCDS+ tall_xcds = GRID_SIZE % NUM_XCDS+ tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds+ xcd = pid % NUM_XCDS+ local_pid = pid // NUM_XCDS+ if xcd < tall_xcds:+ return xcd * pids_per_xcd + local_pid+ else:+ return tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid- __device__ __forceinline__ float fp8_to_f32(unsigned char b) {- return __builtin_amdgcn_cvt_f32_fp8(static_cast<int>(b), 0);- }- __device__ __forceinline__ v4f32 mfma_f32_16x16x128_fp8(- i32x8 A, i32x8 B, v4f32 C)- {- v4f32 D;- asm(- "v_mfma_f32_16x16x128_f8f6f4 %0, %1, %2, %3 cbsz:0 blgp:0"- : "=v"(D) : "v"(A), "v"(B), "v"(C));- return D;- }+ @triton.jit+ def _stage1_kernel(+ Q, KV, Mid_O, Out,+ qo_indptr, kv_indptr,+ sm_scale, kv_scale_ptr,+ stride_qt, stride_qh, stride_kvt,+ stride_mb, stride_mh, stride_ms,+ stride_ot, stride_oh,+ NUM_KV_SPLITS: tl.constexpr, BLOCK_N: tl.constexpr,+ BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr, BLOCK_R: tl.constexpr,+ DIRECT_OUT: tl.constexpr, GRID_SIZE: tl.constexpr,+ ):+ raw_pid = tl.program_id(0)+ pid = _remap_xcd(raw_pid, GRID_SIZE)+ cur_batch = pid // NUM_KV_SPLITS+ split_id = pid % NUM_KV_SPLITS+ heads = tl.arange(0, BLOCK_H)+ kv_scale = tl.load(kv_scale_ptr).to(tl.float32)+ score_scale = sm_scale * kv_scale+ q_tok = tl.load(qo_indptr + cur_batch)+ kv_s = tl.load(kv_indptr + cur_batch)+ kv_e = tl.load(kv_indptr + cur_batch + 1)+ kv_len = kv_e - kv_s+ kv_per_split = tl.cdiv(kv_len, NUM_KV_SPLITS)+ sp_start = kv_per_split * split_id+ sp_end = tl.minimum(sp_start + kv_per_split, kv_len)+ offs_c = tl.arange(0, BLOCK_C)+ offs_r = tl.arange(0, BLOCK_R)+ q_base = q_tok * stride_qt+ q_nope = tl.load(Q + q_base + heads[:, None] * stride_qh + offs_c[None, :]).to(FP8_DTYPE)+ q_rope = tl.load(Q + q_base + heads[:, None] * stride_qh + (512 + offs_r[None, :])).to(FP8_DTYPE)+ e_max = tl.full([BLOCK_H], value=float("-inf"), dtype=tl.float32)+ e_sum = tl.zeros([BLOCK_H], dtype=tl.float32)+ acc = tl.zeros([BLOCK_H, BLOCK_C], dtype=tl.float32)+ num_tokens = sp_end - sp_start+ num_full = num_tokens // BLOCK_N+ full_end = sp_start + num_full * BLOCK_N- __device__ __forceinline__ unsigned int lshl_or(unsigned int src, unsigned int shift, unsigned int base) {- unsigned int r;- asm("v_lshl_or_b32 %0, %1, %2, %3" : "=v"(r) : "v"(src), "v"(shift), "v"(base));- return r;- }+ for start_n in range(sp_start, full_end, BLOCK_N):+ kv_idx = kv_s + start_n + tl.arange(0, BLOCK_N)+ kv_nope = tl.load(KV + kv_idx[:, None] * stride_kvt + offs_c[None, :],+ cache_modifier=".cg")+ kv_rope = tl.load(KV + kv_idx[:, None] * stride_kvt + (512 + offs_r[None, :]),+ cache_modifier=".cg")+ k_nope = tl.trans(kv_nope)+ k_rope = tl.trans(kv_rope)+ qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)+ qk = qk * score_scale+ new_max = tl.maximum(tl.max(qk, 1), e_max)+ rescale = tl.exp(e_max - new_max)+ p = tl.exp(qk - new_max[:, None])+ acc = acc * rescale[:, None]+ acc = acc + tl.dot(p.to(FP8_DTYPE), kv_nope)+ e_sum = e_sum * rescale + tl.sum(p, 1)+ e_max = new_max- __device__ __forceinline__ unsigned int pack4_fp8(unsigned int b0, unsigned int b1,- unsigned int b2, unsigned int b3) {- unsigned int lo = lshl_or(b1, 8u, b0);- unsigned int hi = lshl_or(b3, 8u, b2);- return lshl_or(hi, 16u, lo);- }+ if full_end < sp_end:+ offs_n = full_end + tl.arange(0, BLOCK_N)+ mask_n = offs_n < sp_end+ kv_idx = kv_s + offs_n+ kv_nope = tl.load(KV + kv_idx[:, None] * stride_kvt + offs_c[None, :],+ mask=mask_n[:, None], other=0.0, cache_modifier=".cg")+ kv_rope = tl.load(KV + kv_idx[:, None] * stride_kvt + (512 + offs_r[None, :]),+ mask=mask_n[:, None], other=0.0, cache_modifier=".cg")+ k_nope = tl.trans(kv_nope)+ k_rope = tl.trans(kv_rope)+ qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)+ qk = qk * score_scale+ qk = tl.where(mask_n[None, :], qk, float("-inf"))+ new_max = tl.maximum(tl.max(qk, 1), e_max)+ rescale = tl.exp(e_max - new_max)+ p = tl.exp(qk - new_max[:, None])+ acc = acc * rescale[:, None]+ acc = acc + tl.dot(p.to(FP8_DTYPE), kv_nope)+ e_sum = e_sum * rescale + tl.sum(p, 1)+ e_max = new_max- // =========================================================================- // Tile processing: template specialization for full vs partial tiles- // FULL_TILE=true: stcnt=32 (compile-time), all bounds checks eliminated- // FULL_TILE=false: runtime stcnt with full bounds checks- // =========================================================================+ safe_esum = tl.where(e_sum > 0, e_sum, 1.0)+ result = acc / safe_esum[:, None]+ if DIRECT_OUT:+ o_base = q_tok * stride_ot+ tl.store(Out + o_base + heads[:, None] * stride_oh + offs_c[None, :],+ (result * kv_scale).to(tl.bfloat16))+ else:+ mid_base = cur_batch * stride_mb + heads * stride_mh + split_id * stride_ms+ tl.store(Mid_O + mid_base[:, None] + offs_c[None, :], result)+ lse = tl.where(e_sum > 0, e_max + tl.log(e_sum), float("-inf"))+ tl.store(Mid_O + mid_base + 512, lse)- template<bool FULL_TILE>- __device__ __forceinline__ void process_tile(- const unsigned char* __restrict__ kv_ptr,- unsigned char kv_lds[][KV_TILE_BYTES],- float s_W[][16][33],- const i32x8* q_128,- const float score_scale,- float mv[4], float lv[4],- float vacc[][4],- int& cur_buf,- const int mr, const int kg, const int tid, const int warp_id,- const int split_kv_start, const int split_kv_end,- const int total_tokens, const int num_st,- const int st_idx, const int stcnt_arg)- {- const int stcnt = FULL_TILE ? SUPER_TILE : stcnt_arg;- const int ta = FULL_TILE ? 16 : min(16, stcnt);- const int tb = FULL_TILE ? 16 : max(0, stcnt - 16);- const unsigned char* kv_cur = kv_lds[cur_buf];- // ---- INTERLEAVED DMA + QK: 1 DMA round per QK chunk ----- // PF_ROUNDS == NUM_K128_CHUNKS == 5, so natural 1:1 interleaving.- // DMA writes to kv_lds[nxt_buf] via VMEM, QK reads kv_lds[cur_buf] via LDS.- __builtin_amdgcn_s_setprio(3);- const int nxt_buf = cur_buf ^ 1;- const bool has_next = (st_idx + 1 < num_st);+ @triton.jit+ def _stage2_kernel(+ Mid_O, Out, qo_indptr, kv_scale_ptr,+ stride_mb, stride_mh, stride_ms, stride_ot, stride_oh,+ NUM_KV_SPLITS: tl.constexpr, BLOCK_DV: tl.constexpr,+ batch: tl.constexpr, GRID_SIZE: tl.constexpr,+ ):+ raw_pid = tl.program_id(0)+ pid = _remap_xcd(raw_pid, GRID_SIZE)+ cur_batch = pid % batch+ cur_head = pid // batch+ kv_scale = tl.load(kv_scale_ptr).to(tl.float32)+ q_tok = tl.load(qo_indptr + cur_batch)+ offs_d = tl.arange(0, BLOCK_DV)+ e_max = float("-inf")+ e_sum = 0.0+ acc = tl.zeros([BLOCK_DV], dtype=tl.float32)+ mid_base = cur_batch * stride_mb + cur_head * stride_mh+ for s in range(NUM_KV_SPLITS):+ tv = tl.load(Mid_O + mid_base + s * stride_ms + offs_d)+ lse = tl.load(Mid_O + mid_base + s * stride_ms + 512)+ new_max = tl.maximum(lse, e_max)+ old_scale = tl.exp(e_max - new_max)+ exp_lse = tl.exp(lse - new_max)+ acc = acc * old_scale + exp_lse * tv+ e_sum = e_sum * old_scale + exp_lse+ e_max = new_max+ result = acc / tl.maximum(e_sum, 1e-12) * kv_scale+ o_base = q_tok * stride_ot + cur_head * stride_oh+ tl.store(Out + o_base + offs_d, result.to(tl.bfloat16))- i32x4 srsrc = {};- if (has_next) {- const int nxt_start = split_kv_start + (st_idx + 1) * SUPER_TILE;- const int nxt_bytes = min(SUPER_TILE, split_kv_end - nxt_start) * QK_DIM;- const unsigned char* __restrict__ nsrc = kv_ptr +- static_cast<long long>(nxt_start) * QK_DIM;- srsrc = make_buffer_rsrc(nsrc, nxt_bytes);- }- v4f32 ca = {0, 0, 0, 0};- v4f32 cb = {0, 0, 0, 0};- #pragma unroll- for (int c = 0; c < NUM_K128_CHUNKS; c++) {- if (has_next) {- int dma_off = tid * 16 + c * BLOCK_SIZE * 16;- lds_ptr_t ldp = (lds_ptr_t)(reinterpret_cast<uintptr_t>(kv_lds[nxt_buf]) + dma_off);- __llvm_amdgcn_raw_buffer_load_lds(srsrc, ldp, 16, dma_off, 0, 0, 2);- }- i32x8 ba = {};- if (FULL_TILE || mr < ta) {- int base1 = mr * QK_DIM + c * 128 + 16 * kg;- #pragma unroll- for (int i = 0; i < 4; i++) {- int off = base1 + i * 4;- if (c * 128 + 16 * kg + i * 4 + 4 <= QK_DIM)- ba[i] = *reinterpret_cast<const int*>(&kv_cur[off]);- }- int base2 = mr * QK_DIM + c * 128 + 64 + 16 * kg;- #pragma unroll- for (int i = 0; i < 4; i++) {- int off = base2 + i * 4;- if (c * 128 + 64 + 16 * kg + i * 4 + 4 <= QK_DIM)- ba[4 + i] = *reinterpret_cast<const int*>(&kv_cur[off]);- }- }- i32x8 bb = {};- if (FULL_TILE || mr < tb) {- int base1 = (16 + mr) * QK_DIM + c * 128 + 16 * kg;- #pragma unroll- for (int i = 0; i < 4; i++) {- int off = base1 + i * 4;- if (c * 128 + 16 * kg + i * 4 + 4 <= QK_DIM)- bb[i] = *reinterpret_cast<const int*>(&kv_cur[off]);- }- int base2 = (16 + mr) * QK_DIM + c * 128 + 64 + 16 * kg;- #pragma unroll- for (int i = 0; i < 4; i++) {- int off = base2 + i * 4;- if (c * 128 + 64 + 16 * kg + i * 4 + 4 <= QK_DIM)- bb[4 + i] = *reinterpret_cast<const int*>(&kv_cur[off]);- }- }- ca = mfma_f32_16x16x128_fp8(q_128[c], ba, ca);- cb = mfma_f32_16x16x128_fp8(q_128[c], bb, cb);- }- __builtin_amdgcn_s_setprio(0);-- // ---- Online softmax (16-lane reduce only: offsets 8,4,2,1) ----- float sa0 = ca[0] * score_scale, sa1 = ca[1] * score_scale;- float sa2 = ca[2] * score_scale, sa3 = ca[3] * score_scale;- float sb0 = cb[0] * score_scale, sb1 = cb[1] * score_scale;- float sb2 = cb[2] * score_scale, sb3 = cb[3] * score_scale;-- if (!FULL_TILE && mr >= ta) { sa0 = sa1 = sa2 = sa3 = -1e30f; }- if (!FULL_TILE && mr >= tb) { sb0 = sb1 = sb2 = sb3 = -1e30f; }-- float tm0 = fmaxf(sa0, sb0), tm1 = fmaxf(sa1, sb1);- float tm2 = fmaxf(sa2, sb2), tm3 = fmaxf(sa3, sb3);- #pragma unroll- for (int off = 8; off >= 1; off >>= 1) {- tm0 = fmaxf(tm0, __shfl_xor(tm0, off));- tm1 = fmaxf(tm1, __shfl_xor(tm1, off));- tm2 = fmaxf(tm2, __shfl_xor(tm2, off));- tm3 = fmaxf(tm3, __shfl_xor(tm3, off));- }-- float nm0 = fmaxf(mv[0], tm0), nm1 = fmaxf(mv[1], tm1);- float nm2 = fmaxf(mv[2], tm2), nm3 = fmaxf(mv[3], tm3);- float rc0 = __expf(mv[0] - nm0), rc1 = __expf(mv[1] - nm1);- float rc2 = __expf(mv[2] - nm2), rc3 = __expf(mv[3] - nm3);- mv[0] = nm0; mv[1] = nm1; mv[2] = nm2; mv[3] = nm3;-- float wa0 = (FULL_TILE || mr < ta) ? __expf(sa0 - nm0) : 0.f;- float wa1 = (FULL_TILE || mr < ta) ? __expf(sa1 - nm1) : 0.f;- float wa2 = (FULL_TILE || mr < ta) ? __expf(sa2 - nm2) : 0.f;- float wa3 = (FULL_TILE || mr < ta) ? __expf(sa3 - nm3) : 0.f;- float wb0 = (FULL_TILE || mr < tb) ? __expf(sb0 - nm0) : 0.f;- float wb1 = (FULL_TILE || mr < tb) ? __expf(sb1 - nm1) : 0.f;- float wb2 = (FULL_TILE || mr < tb) ? __expf(sb2 - nm2) : 0.f;- float wb3 = (FULL_TILE || mr < tb) ? __expf(sb3 - nm3) : 0.f;-- float dl0 = wa0 + wb0, dl1 = wa1 + wb1;- float dl2 = wa2 + wb2, dl3 = wa3 + wb3;- #pragma unroll- for (int off = 8; off >= 1; off >>= 1) {- dl0 += __shfl_xor(dl0, off); dl1 += __shfl_xor(dl1, off);- dl2 += __shfl_xor(dl2, off); dl3 += __shfl_xor(dl3, off);- }- lv[0] = lv[0] * rc0 + dl0; lv[1] = lv[1] * rc1 + dl1;- lv[2] = lv[2] * rc2 + dl2; lv[3] = lv[3] * rc3 + dl3;-- // ---- W to per-warp LDS ----- s_W[warp_id][kg * 4 ][mr] = wa0;- s_W[warp_id][kg * 4 + 1][mr] = wa1;- s_W[warp_id][kg * 4 + 2][mr] = wa2;- s_W[warp_id][kg * 4 + 3][mr] = wa3;- s_W[warp_id][kg * 4 ][16 + mr] = wb0;- s_W[warp_id][kg * 4 + 1][16 + mr] = wb1;- s_W[warp_id][kg * 4 + 2][16 + mr] = wb2;- s_W[warp_id][kg * 4 + 3][16 + mr] = wb3;-- asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");-- float wvals[8];- #pragma unroll- for (int i = 0; i < 8; i++)- wvals[i] = s_W[warp_id][mr][i * 4 + kg];-- unsigned int wlo = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[0], wvals[1], 0, false);- wlo = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[2], wvals[3], wlo, true);- unsigned int whi = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[4], wvals[5], 0, false);- whi = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[6], wvals[7], whi, true);- long w_a = static_cast<long>(wlo) | (static_cast<long>(whi) << 32);-- // ---- SV MFMA from current LDS (each warp: 128 V dims) ----- __builtin_amdgcn_s_setprio(1);- const int v_warp_base = warp_id * SV_CHUNKS * 16;-- #pragma unroll- for (int vc = 0; vc < SV_CHUNKS; vc += 2) {- const int vd0 = v_warp_base + vc * 16 + mr;- const int vd0_align = vd0 & ~3;- const int vd0_shift = (vd0 & 3) * 8;- const int vd1 = v_warp_base + (vc + 1) * 16 + mr;- const int vd1_align = vd1 & ~3;- const int vd1_shift = (vd1 & 3) * 8;-- unsigned int blo0 = 0, bhi0 = 0;- unsigned int blo1 = 0, bhi1 = 0;- if (vd0 < V_DIM) {- unsigned int d0[8], d1[8];- #pragma unroll- for (int i = 0; i < 8; i++) {- int tok = i * 4 + kg;- if (FULL_TILE || tok < stcnt) {- const unsigned int* base = reinterpret_cast<const unsigned int*>(- &kv_cur[tok * QK_DIM + vd0_align]);- d0[i] = base[0];- d1[i] = base[4];- } else {- d0[i] = 0u;- d1[i] = 0u;- }- }- __builtin_amdgcn_sched_group_barrier(0x0040, 16, 0);- __builtin_amdgcn_sched_group_barrier(0x0002, 64, 0);- unsigned int e0[8], e1[8];- #pragma unroll- for (int i = 0; i < 8; i++) {- e0[i] = __builtin_amdgcn_ubfe(d0[i], vd0_shift, 8);- e1[i] = __builtin_amdgcn_ubfe(d1[i], vd1_shift, 8);- }- blo0 = pack4_fp8(e0[0], e0[1], e0[2], e0[3]);- bhi0 = pack4_fp8(e0[4], e0[5], e0[6], e0[7]);- blo1 = pack4_fp8(e1[0], e1[1], e1[2], e1[3]);- bhi1 = pack4_fp8(e1[4], e1[5], e1[6], e1[7]);- }- long v_b0 = static_cast<long>(blo0) | (static_cast<long>(bhi0) << 32);- long v_b1 = static_cast<long>(blo1) | (static_cast<long>(bhi1) << 32);-- v4f32 sc0 = {vacc[vc][0]*rc0, vacc[vc][1]*rc1, vacc[vc][2]*rc2, vacc[vc][3]*rc3};- v4f32 sc1 = {vacc[vc+1][0]*rc0, vacc[vc+1][1]*rc1, vacc[vc+1][2]*rc2, vacc[vc+1][3]*rc3};- sc0 = __builtin_amdgcn_mfma_f32_16x16x32_fp8_fp8(w_a, v_b0, sc0, 0, 0, 0);- sc1 = __builtin_amdgcn_mfma_f32_16x16x32_fp8_fp8(w_a, v_b1, sc1, 0, 0, 0);- vacc[vc][0] = sc0[0]; vacc[vc][1] = sc0[1];- vacc[vc][2] = sc0[2]; vacc[vc][3] = sc0[3];- vacc[vc+1][0] = sc1[0]; vacc[vc+1][1] = sc1[1];- vacc[vc+1][2] = sc1[2]; vacc[vc+1][3] = sc1[3];- }-- // ---- Wait for GLOBAL_LOAD_LDS and flip ----- __builtin_amdgcn_s_setprio(3);- if (has_next) {- asm volatile("s_waitcnt vmcnt(0)" ::: "memory");- __builtin_amdgcn_sched_barrier(0);- __builtin_amdgcn_s_barrier();- cur_buf = nxt_buf;- }- }-- // =========================================================================- // Full MFMA pipeline kernel with double-buffered LDS + K=128 QK MFMA- // =========================================================================-- __global__ __launch_bounds__(256, 3)- void mla_mfma_pipeline_kernel(- const unsigned short* __restrict__ q_ptr,- const unsigned char* __restrict__ kv_ptr,- float* __restrict__ partial_m,- float* __restrict__ partial_l,- float* __restrict__ partial_acc,- unsigned short* __restrict__ out_ptr,- const int* __restrict__ qo_indptr,- const int* __restrict__ kv_indptr,- const float* __restrict__ kv_scale_ptr,- const int num_splits,- const float sm_scale)- {- const int split_idx = blockIdx.x;- const int batch_idx = blockIdx.y;- const int warp_id = threadIdx.x / WARP_SIZE;- const int lane_id = threadIdx.x % WARP_SIZE;- const int tid = threadIdx.x;-- const float score_scale = sm_scale * (*kv_scale_ptr);-- const int q_start = qo_indptr[batch_idx];- const int kv_start = kv_indptr[batch_idx];- const int kv_end = kv_indptr[batch_idx + 1];- const int kv_len = kv_end - kv_start;-- const int tps = (kv_len + num_splits - 1) / num_splits;- const int split_kv_start = kv_start + split_idx * tps;- const int split_kv_end = min(split_kv_start + tps, kv_end);-- const int mr = lane_id & 0xF;- const int kg = lane_id >> 4;-- if (split_kv_start >= kv_end) {- if (lane_id < 16 && warp_id == 0) {- int head = lane_id;- int off = (batch_idx * NUM_HEADS + head) * num_splits + split_idx;- partial_m[off] = -1e30f;- partial_l[off] = 0.0f;- }- return;- }-- // ===== LDS: double-buffered KV + per-warp W =====- __shared__ __align__(16) unsigned char kv_lds[2][KV_TILE_BYTES];- __shared__ float s_W[NUM_WARPS][16][33];-- // ===== Super-tile iteration setup =====- const int total_tokens = split_kv_end - split_kv_start;- const int num_st = (total_tokens + SUPER_TILE - 1) / SUPER_TILE;-- // ===== PROLOGUE: issue DMA FIRST, then Q prep overlaps with DMA =====- {- const int first_bytes = min(SUPER_TILE, total_tokens) * QK_DIM;- const unsigned char* __restrict__ src0 = kv_ptr +- static_cast<long long>(split_kv_start) * QK_DIM;- i32x4 srsrc = make_buffer_rsrc(src0, first_bytes);- #pragma unroll- for (int r = 0; r < PF_ROUNDS; r++) {- int off = tid * 16 + r * BLOCK_SIZE * 16;- lds_ptr_t ldp = (lds_ptr_t)(reinterpret_cast<uintptr_t>(kv_lds[0]) + off);- __llvm_amdgcn_raw_buffer_load_lds(srsrc, ldp, 16, off, 0, 0, 2);- }- }-- // ===== Q preload as FP8 (overlapped with DMA in flight) =====- const unsigned short* qh = q_ptr +- (static_cast<long long>(q_start) * NUM_HEADS + mr) * QK_DIM;-- i32x8 q_128[NUM_K128_CHUNKS];- #pragma unroll- for (int c = 0; c < NUM_K128_CHUNKS; c++) {- unsigned int w[8];- int base1 = c * 128 + 16 * kg;- #pragma unroll- for (int i = 0; i < 4; i++) {- int d = base1 + i * 4;- float f0 = (d < QK_DIM) ? bf16_to_f32(qh[d]) : 0.f;- float f1 = (d + 1 < QK_DIM) ? bf16_to_f32(qh[d + 1]) : 0.f;- float f2 = (d + 2 < QK_DIM) ? bf16_to_f32(qh[d + 2]) : 0.f;- float f3 = (d + 3 < QK_DIM) ? bf16_to_f32(qh[d + 3]) : 0.f;- unsigned int pk = __builtin_amdgcn_cvt_pk_fp8_f32(f0, f1, 0, false);- pk = __builtin_amdgcn_cvt_pk_fp8_f32(f2, f3, pk, true);- w[i] = pk;- }- int base2 = c * 128 + 64 + 16 * kg;- #pragma unroll- for (int i = 0; i < 4; i++) {- int d = base2 + i * 4;- float f0 = (d < QK_DIM) ? bf16_to_f32(qh[d]) : 0.f;- float f1 = (d + 1 < QK_DIM) ? bf16_to_f32(qh[d + 1]) : 0.f;- float f2 = (d + 2 < QK_DIM) ? bf16_to_f32(qh[d + 2]) : 0.f;- float f3 = (d + 3 < QK_DIM) ? bf16_to_f32(qh[d + 3]) : 0.f;- unsigned int pk = __builtin_amdgcn_cvt_pk_fp8_f32(f0, f1, 0, false);- pk = __builtin_amdgcn_cvt_pk_fp8_f32(f2, f3, pk, true);- w[4 + i] = pk;- }- q_128[c] = *reinterpret_cast<i32x8*>(w);- }-- // ===== V accumulators + softmax state =====- float vacc[SV_CHUNKS][4];- #pragma unroll- for (int i = 0; i < SV_CHUNKS; i++)- vacc[i][0] = vacc[i][1] = vacc[i][2] = vacc[i][3] = 0.0f;-- float mv[4] = {-1e30f, -1e30f, -1e30f, -1e30f};- float lv[4] = {0.0f, 0.0f, 0.0f, 0.0f};-- // ===== Wait for DMA (Q prep ran while DMA was in flight) =====- asm volatile("s_waitcnt vmcnt(0)" ::: "memory");- __builtin_amdgcn_sched_barrier(0);- __builtin_amdgcn_s_barrier();-- int cur_buf = 0;-- // ===== MAIN LOOP: Two-phase for compile-time optimization =====- const int num_full = total_tokens / SUPER_TILE;- const int has_partial = (total_tokens % SUPER_TILE) != 0;-- // Phase 1: Full tiles — stcnt=32 is compile-time constant- for (int st_idx = 0; st_idx < num_full; st_idx++) {- process_tile<true>(kv_ptr, kv_lds, s_W, q_128, score_scale,- mv, lv, vacc, cur_buf, mr, kg, tid, warp_id,- split_kv_start, split_kv_end, total_tokens, num_st,- st_idx, SUPER_TILE);- }-- // Phase 2: Last tile with runtime stcnt (if partial)- if (has_partial) {- int last_stcnt = total_tokens - num_full * SUPER_TILE;- process_tile<false>(kv_ptr, kv_lds, s_W, q_128, score_scale,- mv, lv, vacc, cur_buf, mr, kg, tid, warp_id,- split_kv_start, split_kv_end, total_tokens, num_st,- num_full, last_stcnt);- }-- if (num_splits == 1) {- const float kv_scale = *kv_scale_ptr;- const int vwb = warp_id * SV_CHUNKS * 16;- #pragma unroll- for (int vc = 0; vc < SV_CHUNKS; vc++) {- int vd = vwb + vc * 16 + mr;- if (vd < V_DIM) {- #pragma unroll- for (int r = 0; r < 4; r++) {- int head = kg * 4 + r;- float inv_l = (lv[r] > 0.f) ? (kv_scale / lv[r]) : 0.f;- long long idx = (static_cast<long long>(q_start) * NUM_HEADS + head) * V_DIM + vd;- out_ptr[idx] = f32_to_bf16(vacc[vc][r] * inv_l);- }- }- }- } else {- if (warp_id == 0 && mr == 0) {- #pragma unroll- for (int r = 0; r < 4; r++) {- int head = kg * 4 + r;- int off = (batch_idx * NUM_HEADS + head) * num_splits + split_idx;- partial_m[off] = mv[r];- partial_l[off] = lv[r];- }- }- const int vwb = warp_id * SV_CHUNKS * 16;- #pragma unroll- for (int vc = 0; vc < SV_CHUNKS; vc++) {- int vd = vwb + vc * 16 + mr;- if (vd < V_DIM) {- #pragma unroll- for (int r = 0; r < 4; r++) {- int head = kg * 4 + r;- int off = (batch_idx * NUM_HEADS + head) * num_splits + split_idx;- partial_acc[static_cast<long long>(off) * V_DIM + vd] = vacc[vc][r];- }- }- }- }- }-- // =========================================================================- // Reduce kernel (template-specialized for compile-time loop unrolling)- // =========================================================================-- __global__ __launch_bounds__(512)- void mla_reduce_kernel_generic(- const float* __restrict__ partial_m,- const float* __restrict__ partial_l,- const float* __restrict__ partial_acc,- unsigned short* __restrict__ out_ptr,- const float* __restrict__ kv_scale_ptr,- const int num_splits)- {- const int item_idx = blockIdx.x;- const int tid = threadIdx.x;- const float kv_scale = *kv_scale_ptr;-- __shared__ float s_corr[128];- __shared__ float s_inv_l;-- const int base = item_idx * num_splits;-- float my_m = -1e30f;- float my_l = 0.0f;- if (tid < num_splits) {- my_m = partial_m[base + tid];- my_l = partial_l[base + tid];- }-- float merged_m = my_m;- #pragma unroll- for (int off = 32; off >= 1; off >>= 1)- merged_m = fmaxf(merged_m, __shfl_xor(merged_m, off));-- float my_c = 0.0f;- if (tid < num_splits && my_l > 0.f)- my_c = __expf(my_m - merged_m);- float weighted_l = my_l * my_c;-- if (tid < num_splits)- s_corr[tid] = my_c;-- float merged_l = weighted_l;- #pragma unroll- for (int off = 32; off >= 1; off >>= 1)- merged_l += __shfl_xor(merged_l, off);-- if (tid == 0)- s_inv_l = (merged_l > 0.f) ? (kv_scale / merged_l) : 0.f;- __syncthreads();-- if (tid < V_DIM) {- float val = 0.0f;- for (int s = 0; s < num_splits; ++s) {- val += partial_acc[(static_cast<long long>(base + s)) * V_DIM + tid]- * s_corr[s];- }- out_ptr[static_cast<long long>(item_idx) * V_DIM + tid] = f32_to_bf16(val * s_inv_l);- }- }-- template<int NUM_SPLITS>- __global__ __launch_bounds__(512)- void mla_reduce_kernel(- const float* __restrict__ partial_m,- const float* __restrict__ partial_l,- const float* __restrict__ partial_acc,- unsigned short* __restrict__ out_ptr,- const float* __restrict__ kv_scale_ptr)- {- const int item_idx = blockIdx.x;- const int tid = threadIdx.x;- const float kv_scale = *kv_scale_ptr;-- __shared__ float s_corr[NUM_SPLITS < 128 ? 128 : NUM_SPLITS];- __shared__ float s_inv_l;-- const int base = item_idx * NUM_SPLITS;-- float my_m = -1e30f;- float my_l = 0.0f;- if (tid < NUM_SPLITS) {- my_m = partial_m[base + tid];- my_l = partial_l[base + tid];- }-- float merged_m = my_m;- #pragma unroll- for (int off = 32; off >= 1; off >>= 1)- merged_m = fmaxf(merged_m, __shfl_xor(merged_m, off));-- float my_c = 0.0f;- if (tid < NUM_SPLITS && my_l > 0.f)- my_c = __expf(my_m - merged_m);- float weighted_l = my_l * my_c;-- if (tid < NUM_SPLITS)- s_corr[tid] = my_c;-- float merged_l = weighted_l;- #pragma unroll- for (int off = 32; off >= 1; off >>= 1)- merged_l += __shfl_xor(merged_l, off);-- if (tid == 0)- s_inv_l = (merged_l > 0.f) ? (kv_scale / merged_l) : 0.f;- __syncthreads();-- if (tid < V_DIM) {- float val = 0.0f;- #pragma unroll- for (int s = 0; s < NUM_SPLITS; ++s) {- val += partial_acc[(static_cast<long long>(base + s)) * V_DIM + tid]- * s_corr[s];- }- out_ptr[static_cast<long long>(item_idx) * V_DIM + tid] = f32_to_bf16(val * s_inv_l);- }- }-- torch::Tensor mla_decode(- torch::Tensor q, torch::Tensor kv_buffer,- torch::Tensor qo_indptr, torch::Tensor kv_indptr,- torch::Tensor kv_scale_tensor,- int64_t num_heads, int64_t num_splits,- float sm_scale,- torch::Tensor partial_m, torch::Tensor partial_l,- torch::Tensor partial_acc,- torch::Tensor output)- {- const int batch_size = qo_indptr.size(0) - 1;- const int num_items = batch_size * static_cast<int>(num_heads);-- dim3 grid1(static_cast<int>(num_splits), batch_size);- dim3 block1(BLOCK_SIZE);- mla_mfma_pipeline_kernel<<<grid1, block1>>>(- reinterpret_cast<const unsigned short*>(q.data_ptr()),- reinterpret_cast<const unsigned char*>(kv_buffer.data_ptr()),- partial_m.data_ptr<float>(), partial_l.data_ptr<float>(),- partial_acc.data_ptr<float>(),- reinterpret_cast<unsigned short*>(output.data_ptr()),- qo_indptr.data_ptr<int>(), kv_indptr.data_ptr<int>(),- kv_scale_tensor.data_ptr<float>(),- static_cast<int>(num_splits), sm_scale);- if (num_splits == 1) return output;-- dim3 grid2(num_items);- dim3 block2(512);-- #define REDUCE_DISPATCH(N) \- mla_reduce_kernel<N><<<grid2, block2>>>( \- partial_m.data_ptr<float>(), partial_l.data_ptr<float>(), \- partial_acc.data_ptr<float>(), \- reinterpret_cast<unsigned short*>(output.data_ptr()), \- kv_scale_tensor.data_ptr<float>())-- switch (static_cast<int>(num_splits)) {- case 3: REDUCE_DISPATCH(3); break;- case 8: REDUCE_DISPATCH(8); break;- case 12: REDUCE_DISPATCH(12); break;- case 16: REDUCE_DISPATCH(16); break;- case 24: REDUCE_DISPATCH(24); break;- case 64: REDUCE_DISPATCH(64); break;- default:- mla_reduce_kernel_generic<<<grid2, block2>>>(- partial_m.data_ptr<float>(), partial_l.data_ptr<float>(),- partial_acc.data_ptr<float>(),- reinterpret_cast<unsigned short*>(output.data_ptr()),- kv_scale_tensor.data_ptr<float>(),- static_cast<int>(num_splits));- break;- }- #undef REDUCE_DISPATCH-- return output;- }- """-- CPP_DECL = """- torch::Tensor mla_decode(- torch::Tensor q, torch::Tensor kv_buffer,- torch::Tensor qo_indptr, torch::Tensor kv_indptr,- torch::Tensor kv_scale_tensor,- int64_t num_heads, int64_t num_splits,- float sm_scale,- torch::Tensor partial_m, torch::Tensor partial_l,- torch::Tensor partial_acc,- torch::Tensor output);- """-- _module = load_inline(- name="mla_hip_v249_tuned_splits",- cpp_sources=CPP_DECL,- cuda_sources=HIP_SRC,- functions=["mla_decode"],- extra_cuda_cflags=[- "-O3", "-std=c++17",- "-ffast-math", "-funsafe-math-optimizations", "-ffp-contract=fast",- "-fno-gpu-rdc",- "-mllvm", "-amdgpu-early-inline-all=true",- "-mllvm", "-amdgpu-function-calls=false",- "-mllvm", "-amdgpu-max-memory-clause=64",- "-mllvm", "-amdgpu-load-store-vectorizer",- "-mllvm", "-amdgpu-early-ifcvt",- "-mllvm", "-amdgpu-internalize-symbols",- "-mllvm", "-amdgpu-scalarize-global-loads",- "-mllvm", "-amdgpu-dpp-combine",- "-mllvm", "-amdgpu-enable-pre-ra-optimizations",- "-mllvm", "-amdgpu-promote-alloca-to-vector-limit=256",- ],- verbose=False,- )---- def _choose_splits(batch_size, kv_len):- tiles = max(1, kv_len // 32)- max_useful = max(1, tiles // 2)-- if tiles > 64:- ideal = max(1, -(-768 // batch_size))+ def _get_config(batch_size, kv_len):+ key = (batch_size, kv_len)+ if key in _SHAPE_CONFIGS:+ return _SHAPE_CONFIGS[key]+ max_tiles = max(1, kv_len // 64)+ if batch_size >= 256:+ splits = min(2, max_tiles)+ elif batch_size >= 64:+ splits = min(max(1, 512 // batch_size), max_tiles)else:- target_wgs = max(512, batch_size * 8)- ideal = max(1, target_wgs // batch_size)-- splits = max(1, min(ideal, max_useful, 64))-+ splits = min(max(1, 768 // batch_size), max_tiles)while splits > 1 and batch_size * splits > 912:splits -= 1+ grid_size = batch_size * splits+ ns = 3 if grid_size < 512 else 2+ return splits, ns- return splits-def custom_kernel(data: input_t) -> output_t:q, kv_data, qo_indptr, kv_indptr, config = data-- kv_buffer_fp8, kv_scale = kv_data["fp8"]- kv_buffer = kv_buffer_fp8.view(-1, 576)-+ kv_fp8, kv_scale_t = kv_data["fp8"]+ kv = kv_fp8.view(torch.float8_e4m3fn).view(-1, QK_HEAD_DIM)batch_size = config["batch_size"]- num_heads = config["num_heads"]sm_scale = config["sm_scale"]- total_q = q.size(0)- num_items = batch_size * num_heads+ total_q = q.shape[0]+ kv_len = kv.shape[0] // batch_size if batch_size > 0 else 0- total_kv = kv_buffer.shape[0]- kv_len = total_kv // batch_size+ num_splits, ns = _get_config(batch_size, kv_len)- num_splits = _choose_splits(batch_size, kv_len)-dev = q.device- pm = torch.empty((num_items * num_splits,), dtype=torch.float32, device=dev)- pl = torch.empty((num_items * num_splits,), dtype=torch.float32, device=dev)- pa = torch.empty((num_items * num_splits, 512), dtype=torch.float32, device=dev)- out = torch.empty((total_q, num_heads, 512), dtype=torch.bfloat16, device=dev)+ direct_out = (num_splits == 1)+ mid_o = torch.empty(+ (batch_size, NUM_HEADS, num_splits, KV_LORA_RANK + 1),+ dtype=torch.float32, device=dev,+ ) if not direct_out else torch.empty(1, dtype=torch.float32, device=dev)- return _module.mla_decode(- q, kv_buffer, qo_indptr, kv_indptr,- kv_scale,- num_heads, num_splits,- sm_scale,- pm, pl, pa, out)+ out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)++ grid_size1 = batch_size * num_splits+ _stage1_kernel[(grid_size1,)](+ q, kv, mid_o, out, qo_indptr, kv_indptr,+ sm_scale, kv_scale_t,+ q.stride(0), q.stride(1), kv.stride(0),+ mid_o.stride(0) if not direct_out else 0,+ mid_o.stride(1) if not direct_out else 0,+ mid_o.stride(2) if not direct_out else 0,+ out.stride(0), out.stride(1),+ NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,+ BLOCK_C=512, BLOCK_R=64, DIRECT_OUT=direct_out,+ GRID_SIZE=grid_size1,+ num_warps=4, num_stages=ns, waves_per_eu=2,+ )++ if not direct_out:+ grid_size2 = NUM_HEADS * batch_size+ _stage2_kernel[(grid_size2,)](+ mid_o, out, qo_indptr, kv_scale_t,+ mid_o.stride(0), mid_o.stride(1), mid_o.stride(2),+ out.stride(0), out.stride(1),+ NUM_KV_SPLITS=num_splits, BLOCK_DV=512,+ batch=batch_size, GRID_SIZE=grid_size2,+ num_warps=4, num_stages=2,+ )++ return out
scrolls · 961 diff lines total
Best evidence level for this revision: reported
JSON