submission 700955
Maxwell Cipher · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1293 lines, June 9 Researcher Reciprocity License v1.0.
mla_cute_05.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-700955?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:3e5bf5097948db9c5c5e11a7f61e4e069f56fb3eb243aea7fc8329316bc29cc7
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
fp4_packed, scale_e8m0 = kv_data["mxfp4"]persistent-kernel
- Small shapes: bf16 non-persistent (page_size=1)shared-memory
__shared__ unsigned char lkv[LDS_KV];Kernel source
mla_cute_05.py1293 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
Locked baseline dispatch (v93-style):
- Small shapes: bf16 non-persistent (page_size=1)
- Larger shapes: a16w8 persistent (bf16 Q + fp8 KV, page_size=2)
- Per-shape fast_mode overrides
This intentionally avoids fp8 Q quantization overhead/rounding drift and is the
stable baseline before custom-kernel experimentation.
"""
from __future__ import annotations
import hashlib
import os
import sys
import time
import torch
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
DEFAULT_PERSIST_PAGE_SIZE = 2
NUM_KV_SPLITS = 32
# Best-known fast_mode cherry-picks from prior tuning.
FAST_MODE_OVERRIDES = {
(4, 8192): True,
(64, 1024): True,
}
# Experimental per-shape page-size overrides (off by default).
# Enable with MLA_POLICY_MODE=pg8_probe.
PAGE_SIZE_OVERRIDES_PG8_PROBE = {
(4, 8192): 8,
(32, 8192): 8,
(64, 8192): 8,
(256, 8192): 8,
}
FP8_DTYPE = aiter_dtypes.fp8
TRACE_ROUTE = os.getenv("MLA_TRACE_ROUTE", "0") == "1"
FORCE_FAST_MODE_ENV = os.getenv("MLA_FORCE_FAST_MODE")
POLICY_MODE = os.getenv("MLA_POLICY_MODE", "v93")
VALID_POLICY_MODES = {"v93", "pg8_probe"}
ENABLE_MFMA_V71 = os.getenv("MLA_ENABLE_MFMA_V71", "0") == "1"
VALIDATE_MFMA_V71 = os.getenv("MLA_VALIDATE_MFMA_V71", "0") == "1"
VALIDATE_MFMA_V71_MAX_ABS = float(os.getenv("MLA_VALIDATE_MFMA_V71_MAX_ABS", "0"))
VALIDATE_MFMA_V71_RTOL = 1e-1
VALIDATE_MFMA_V71_ATOL = 1e-1
CUSTOM_V71_OFFLOAD_ARCH = os.getenv("MLA_MFMA_V71_ARCH", "gfx950")
CUSTOM_V71_BUILD_TAG = os.getenv("MLA_MFMA_V71_BUILD_TAG", "").strip()
CUSTOM_V71_VARIANT = os.getenv("MLA_MFMA_V71_VARIANT", "v71").strip().lower()
VALID_CUSTOM_V71_VARIANTS = {"v71", "v74"}
DEFAULT_CUSTOM_LONG_KV_SHAPES = {
(32, 8192),
(64, 8192),
(256, 8192),
}
DEFAULT_CUSTOM_NSPLIT_OVERRIDES: dict[tuple[int, int], int] = {}
CUSTOM_V71_PACKED_DIM = 288
# Submission-local guided split override probe for v74 long-KV shapes.
# This file is an immutable experiment snapshot and intentionally overrides only
# the split map (shape/config derived, no input-data dependency).
_LOCAL_CUSTOM_NSPLIT_OVERRIDES = {
(32, 8192): 8,
(64, 8192): 8,
(256, 8192): 6,
}
_cache: dict[tuple, tuple] = {}
_hip_custom = None
_hip_custom_tried = False
_hip_custom_ok = False
_custom_variant_resolved: str | None = None
_custom_entry_name: str | None = None
_custom_v71_scratch_cache: dict[tuple[int, int, int], tuple[torch.Tensor, torch.Tensor]] = {}
_custom_v71_module_name: str | None = None
_baseline_compare_out_cache: dict[tuple[int, int], torch.Tensor] = {}
def _next_custom_v71_module_name(cpp_src: str, hip_src: str, resolved_variant: str) -> str:
"""
Build a deterministic extension module name keyed by source+arch(+tag),
avoiding stale binary reuse while still allowing cache reuse.
"""
key = "\n".join(
[
cpp_src,
hip_src,
CUSTOM_V71_OFFLOAD_ARCH,
CUSTOM_V71_BUILD_TAG,
resolved_variant,
]
).encode("utf-8")
digest = hashlib.sha1(key).hexdigest()[:16]
return f"mla_custom_v71_{digest}"
def _log_route(msg: str) -> None:
if TRACE_ROUTE:
print(f"[MLA route] {msg}", file=sys.stderr, flush=True)
def _parse_custom_shapes_env(raw: str | None) -> set[tuple[int, int]]:
"""
Parse MLA_CUSTOM_SHAPES like: "32x8192,64x8192,256x8192"
"""
if not raw:
return set(DEFAULT_CUSTOM_LONG_KV_SHAPES)
shapes: set[tuple[int, int]] = set()
for token in raw.split(","):
item = token.strip().lower()
if not item:
continue
if "x" not in item:
_log_route(f"ignoring malformed shape token '{token}'")
continue
bs_raw, kv_raw = item.split("x", 1)
try:
shapes.add((int(bs_raw), int(kv_raw)))
except ValueError:
_log_route(f"ignoring malformed shape token '{token}'")
return shapes if shapes else set(DEFAULT_CUSTOM_LONG_KV_SHAPES)
def _parse_custom_nsplits_env(raw: str | None) -> dict[tuple[int, int], int]:
"""
Parse MLA_CUSTOM_NSPLITS like: "32x8192:4,64x8192:8,256x8192:16"
"""
if not raw:
return dict(DEFAULT_CUSTOM_NSPLIT_OVERRIDES)
overrides: dict[tuple[int, int], int] = {}
for token in raw.split(","):
item = token.strip().lower()
if not item:
continue
if ":" not in item:
_log_route(f"ignoring malformed nsplit token '{token}'")
continue
shape_raw, nsplit_raw = item.split(":", 1)
if "x" not in shape_raw:
_log_route(f"ignoring malformed nsplit token '{token}'")
continue
bs_raw, kv_raw = shape_raw.split("x", 1)
try:
bs = int(bs_raw)
kv = int(kv_raw)
ns = int(nsplit_raw)
except ValueError:
_log_route(f"ignoring malformed nsplit token '{token}'")
continue
if ns <= 0:
_log_route(f"ignoring non-positive nsplit token '{token}'")
continue
overrides[(bs, kv)] = ns
return overrides if overrides else dict(DEFAULT_CUSTOM_NSPLIT_OVERRIDES)
if POLICY_MODE not in VALID_POLICY_MODES:
POLICY_MODE = "v93"
_log_route("invalid MLA_POLICY_MODE, falling back to v93")
if CUSTOM_V71_VARIANT not in VALID_CUSTOM_V71_VARIANTS:
_log_route(
f"invalid MLA_MFMA_V71_VARIANT='{CUSTOM_V71_VARIANT}', "
"falling back to v71"
)
CUSTOM_V71_VARIANT = "v71"
CUSTOM_LONG_KV_SHAPES = _parse_custom_shapes_env(os.getenv("MLA_CUSTOM_SHAPES"))
CUSTOM_NSPLIT_OVERRIDES = _parse_custom_nsplits_env(os.getenv("MLA_CUSTOM_NSPLITS"))
def _device_key(device: torch.device) -> int:
return -1 if device.index is None else int(device.index)
def _legal_page_size(kv_seq_len: int, desired: int) -> int:
ps = max(1, int(desired))
while ps > 1 and kv_seq_len % ps != 0:
ps //= 2
return max(1, ps)
def _get_pg1_state(batch_size: int, kv_seq_len: int, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
key = ("pg1", _device_key(device), batch_size, kv_seq_len)
if key not in _cache:
_cache[key] = (
torch.arange(batch_size * kv_seq_len, dtype=torch.int32, device=device),
torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device),
)
return _cache[key]
def _get_a16w8_persistent_state(
batch_size: int,
kv_seq_len: int,
page_size: int,
fast_mode: bool,
qo_indptr: torch.Tensor,
device: torch.device,
) -> tuple[dict[str, torch.Tensor], torch.Tensor, torch.Tensor, torch.Tensor]:
key = ("a16w8_persist", _device_key(device), batch_size, kv_seq_len, page_size, fast_mode)
if key not in _cache:
num_pages = (batch_size * kv_seq_len) // page_size
kv_indices = torch.arange(num_pages, dtype=torch.int32, device=device)
kv_indptr_pages = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * (kv_seq_len // page_size)
kv_last_page_len = torch.full((batch_size,), page_size, dtype=torch.int32, device=device)
info = get_mla_metadata_info_v1(
batch_size,
1,
NUM_HEADS,
torch.bfloat16,
FP8_DTYPE,
is_sparse=False,
fast_mode=fast_mode,
num_kv_splits=NUM_KV_SPLITS,
intra_batch_mode=True,
)
work = [torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info]
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = work
get_mla_metadata_v1(
qo_indptr,
kv_indptr_pages,
kv_last_page_len,
NUM_HEADS // NUM_KV_HEADS,
NUM_KV_HEADS,
True,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
page_size=page_size,
kv_granularity=max(page_size, 16),
max_seqlen_qo=1,
uni_seqlen_qo=1,
fast_mode=fast_mode,
max_split_per_batch=NUM_KV_SPLITS,
intra_batch_mode=True,
dtype_q=torch.bfloat16,
dtype_kv=FP8_DTYPE,
)
meta = {
"work_meta_data": work_metadata,
"work_indptr": work_indptr,
"work_info_set": work_info_set,
"reduce_indptr": reduce_indptr,
"reduce_final_map": reduce_final_map,
"reduce_partial_map": reduce_partial_map,
}
_cache[key] = (meta, kv_indices, kv_indptr_pages, kv_last_page_len)
return _cache[key]
def _run_bf16_non_persistent(
q: torch.Tensor,
kv_bf16: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
batch_size: int,
kv_seq_len: int,
out: torch.Tensor,
) -> None:
q_view = q.view(-1, NUM_HEADS, QK_HEAD_DIM)
kv_4d = kv_bf16.view(-1, 1, NUM_KV_HEADS, kv_bf16.shape[-1])
page_ids, last_page_len = _get_pg1_state(batch_size, kv_seq_len, q.device)
mla_decode_fwd(
q_view,
kv_4d,
out,
qo_indptr,
kv_indptr,
page_ids,
last_page_len,
1,
page_size=1,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
intra_batch_mode=False,
)
def _run_a16w8_persistent(
q: torch.Tensor,
kv_fp8: torch.Tensor,
kv_scale: torch.Tensor,
qo_indptr: torch.Tensor,
batch_size: int,
kv_seq_len: int,
out: torch.Tensor,
) -> None:
forced_ps = int(os.getenv("MLA_FORCE_PAGE_SIZE", "0"))
if forced_ps > 0:
desired_ps = forced_ps
elif POLICY_MODE == "pg8_probe":
desired_ps = PAGE_SIZE_OVERRIDES_PG8_PROBE.get((batch_size, kv_seq_len), DEFAULT_PERSIST_PAGE_SIZE)
else:
desired_ps = DEFAULT_PERSIST_PAGE_SIZE
page_size = _legal_page_size(kv_seq_len, desired_ps)
if FORCE_FAST_MODE_ENV is None:
fast_mode = FAST_MODE_OVERRIDES.get((batch_size, kv_seq_len), False)
else:
fast_mode = FORCE_FAST_MODE_ENV == "1"
q_view = q.view(-1, NUM_HEADS, QK_HEAD_DIM)
kv_4d = kv_fp8.view(-1, page_size, NUM_KV_HEADS, kv_fp8.shape[-1])
meta, kv_indices, kv_indptr_pages, kv_last_page_len = _get_a16w8_persistent_state(
batch_size, kv_seq_len, page_size, fast_mode, qo_indptr, q.device
)
mla_decode_fwd(
q_view,
kv_4d,
out,
qo_indptr,
kv_indptr_pages,
kv_indices,
kv_last_page_len,
1,
page_size=page_size,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=NUM_KV_SPLITS,
kv_scale=kv_scale,
intra_batch_mode=True,
**meta,
)
def _should_use_custom_v71(batch_size: int, kv_seq_len: int) -> bool:
return ENABLE_MFMA_V71 and (batch_size, kv_seq_len) in CUSTOM_LONG_KV_SHAPES
_HIP_SRC_V71 = r"""
#include <hip/hip_runtime.h>
#include <torch/extension.h>
#include <cmath>
typedef float v4f32 __attribute__((ext_vector_type(4)));
typedef short v4i16 __attribute__((ext_vector_type(4)));
__device__ __forceinline__ float e8m0f(unsigned char e){
if(e==0) return 0.0f;
return __uint_as_float(((unsigned int)e) << 23);
}
typedef __bf16 bf16x2_t __attribute__((ext_vector_type(2)));
__device__ __forceinline__ unsigned int hw_fp4_to_bf16_pair(unsigned int packed, float scale, int byte_sel) {
bf16x2_t r;
switch(byte_sel) {
case 0: r = __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(packed, scale, 0); break;
case 1: r = __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(packed, scale, 1); break;
case 2: r = __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(packed, scale, 2); break;
default: r = __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(packed, scale, 3); break;
}
unsigned int result;
memcpy(&result, &r, 4);
return result;
}
__device__ __forceinline__ float bf16_lo_to_f32(unsigned int pair) {
return __uint_as_float((pair & 0xFFFFu) << 16);
}
__device__ __forceinline__ float bf16_hi_to_f32(unsigned int pair) {
return __uint_as_float((pair >> 16) << 16);
}
#define NH_C 16
#define VD_C 512
#define BN 16
#define PACKED_C 288
#define BSZ 256
#define HPW 4
#define N_K_ITER 36
#define LDS_KV (BN*PACKED_C)
#define LDS_SC (BN*32)
__global__ __attribute__((amdgpu_flat_work_group_size(BSZ, BSZ)))
void mla_attn_v71(
const short* __restrict__ Q,
const unsigned char* __restrict__ KVf,
const unsigned char* __restrict__ KVs,
float* __restrict__ po,
float* __restrict__ pl,
const int* __restrict__ ki,
float sms, int tq, int ns, int nsc,
int q_stride_tok, int q_stride_h)
{
__shared__ unsigned char lkv[LDS_KV];
__shared__ unsigned char lsc[LDS_SC];
__shared__ float lqk[64 * 4];
int tid = threadIdx.x;
int wid = tid / 64;
int ln = tid % 64;
int lmn = ln % 16;
int lk = ln / 16;
int hs = wid * HPW;
int sid = blockIdx.x;
int bid = blockIdx.y;
int kvs_ = ki[bid], kve = ki[bid+1], kvl = kve - kvs_;
int per_split = (kvl > 0) ? (kvl + ns - 1) / ns : 0;
int ms = kvs_ + sid * per_split, me = min(ms + per_split, kve);
if (kvl <= 0 || ms >= kve) {
int base = sid * tq * NH_C + bid * NH_C;
for (int h = tid; h < NH_C; h += BSZ) pl[base + h] = -1e30f;
for (int i = tid; i < NH_C * VD_C; i += BSZ) po[(base + i / VD_C) * VD_C + i % VD_C] = 0.0f;
return;
}
v4i16 q_regs[N_K_ITER];
if (wid == 0) {
for (int k = 0; k < N_K_ITER; k++) {
int d_base = k * 16 + lk * 4;
const short* qp = Q + bid * q_stride_tok + lmn * q_stride_h + d_base;
if (d_base + 3 < 576) {
q_regs[k][0] = qp[0];
q_regs[k][1] = qp[1];
q_regs[k][2] = qp[2];
q_regs[k][3] = qp[3];
} else {
q_regs[k] = {0, 0, 0, 0};
}
}
}
float v_acc[HPW * 8];
#pragma unroll
for (int i = 0; i < HPW * 8; i++) v_acc[i] = 0.0f;
float row_max[HPW], row_sum[HPW];
#pragma unroll
for (int i = 0; i < HPW; i++) { row_max[i] = -1e30f; row_sum[i] = 0.0f; }
float exp_s[HPW];
for (int kp = ms; kp < me; kp += BN) {
int ch = min(BN, me - kp);
{
const unsigned int* src_kv = (const unsigned int*)(KVf + kp * PACKED_C);
unsigned int* dst_kv = (unsigned int*)lkv;
int total_u32 = (BN * PACKED_C) / 4;
for (int i = tid; i < total_u32; i += BSZ) {
dst_kv[i] = src_kv[i];
}
if (ch < BN) {
int valid_u32 = (ch * PACKED_C) / 4;
for (int i = valid_u32 + tid; i < total_u32; i += BSZ)
dst_kv[i] = 0;
}
}
{
const unsigned char* src_sc = KVs + kp * nsc;
unsigned int* sc_u32 = (unsigned int*)lsc;
for (int i = tid; i < (BN * 32) / 4; i += BSZ)
sc_u32[i] = 0x7F7F7F7F;
for (int t = tid; t < BN; t += BSZ) {
if (t < ch) {
for (int c = 0; c < nsc; c++)
lsc[t * 32 + c] = src_sc[t * nsc + c];
}
}
}
__syncthreads();
if (wid == 0) {
v4f32 qk0 = {0, 0, 0, 0};
int k_byte_base = lmn * PACKED_C + (lk * 4) / 2;
int k_sc_base = lmn * 32;
for (int k = 0; k < N_K_ITER; k++) {
int d_base = k * 16;
int d_start = d_base + lk * 4;
v4i16 k0_reg;
{
if (lmn < ch && d_start + 3 < 576) {
int byte_off = k_byte_base + d_base / 2;
unsigned int packed = *(const unsigned short*)&lkv[byte_off];
float sc = e8m0f(lsc[k_sc_base + d_start / 32]);
unsigned int pair0 = hw_fp4_to_bf16_pair(packed, sc, 0);
unsigned int pair1 = hw_fp4_to_bf16_pair(packed, sc, 1);
k0_reg[0] = (short)(pair0 & 0xFFFF);
k0_reg[1] = (short)((pair0 >> 16) & 0xFFFF);
k0_reg[2] = (short)(pair1 & 0xFFFF);
k0_reg[3] = (short)((pair1 >> 16) & 0xFFFF);
} else {
k0_reg = {0, 0, 0, 0};
}
}
qk0 = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(q_regs[k], k0_reg, qk0, 0, 0, 0);
}
lqk[ln * 4 + 0] = qk0[0];
lqk[ln * 4 + 1] = qk0[1];
lqk[ln * 4 + 2] = qk0[2];
lqk[ln * 4 + 3] = qk0[3];
}
__syncthreads();
for (int hi = 0; hi < HPW; hi++) {
int h = hs + hi;
int hp = h % 4;
int hg = h / 4;
float ls = lqk[(hg * 16 + lmn) * 4 + hp] * sms;
if (lmn >= ch) ls = -1e30f;
float bm = ls;
for (int o = 8; o >= 1; o >>= 1) bm = fmaxf(bm, __shfl_xor(bm, o));
float new_max = fmaxf(row_max[hi], bm);
float rescale = expf(row_max[hi] - new_max);
row_sum[hi] *= rescale;
#pragma unroll
for (int d = 0; d < 8; d++) v_acc[hi * 8 + d] *= rescale;
float es = expf(ls - new_max);
if (lmn >= ch) es = 0.0f;
float su = es;
for (int o = 8; o >= 1; o >>= 1) su += __shfl_xor(su, o);
row_sum[hi] += su;
row_max[hi] = new_max;
exp_s[hi] = es;
}
{
int vd_base = ln * 8;
int v_byte_base = vd_base / 2;
int v_sc_idx = vd_base / 32;
int v_off = v_byte_base;
int v_sc_off = v_sc_idx;
for (int t = 0; t < ch; t++, v_off += PACKED_C, v_sc_off += 32) {
unsigned int packed = *(const unsigned int*)&lkv[v_off];
float sc = e8m0f(lsc[v_sc_off]);
unsigned int p0 = hw_fp4_to_bf16_pair(packed, sc, 0);
unsigned int p1 = hw_fp4_to_bf16_pair(packed, sc, 1);
unsigned int p2 = hw_fp4_to_bf16_pair(packed, sc, 2);
unsigned int p3 = hw_fp4_to_bf16_pair(packed, sc, 3);
float v0 = bf16_lo_to_f32(p0), v1 = bf16_hi_to_f32(p0);
float v2 = bf16_lo_to_f32(p1), v3 = bf16_hi_to_f32(p1);
float v4 = bf16_lo_to_f32(p2), v5 = bf16_hi_to_f32(p2);
float v6 = bf16_lo_to_f32(p3), v7 = bf16_hi_to_f32(p3);
#pragma unroll
for (int hi = 0; hi < HPW; hi++) {
float w = __shfl(exp_s[hi], t);
v_acc[hi*8+0] += w * v0; v_acc[hi*8+1] += w * v1;
v_acc[hi*8+2] += w * v2; v_acc[hi*8+3] += w * v3;
v_acc[hi*8+4] += w * v4; v_acc[hi*8+5] += w * v5;
v_acc[hi*8+6] += w * v6; v_acc[hi*8+7] += w * v7;
}
}
}
__syncthreads();
}
int base = sid * tq * NH_C + bid * NH_C;
for (int hi = 0; hi < HPW; hi++) {
int h = hs + hi;
float inv = (row_sum[hi] > 0.0f) ? 1.0f / row_sum[hi] : 0.0f;
if (ln == 0) {
pl[base + h] = (row_sum[hi] > 0.0f) ? row_max[hi] + logf(row_sum[hi]) : -1e30f;
}
for (int d = 0; d < 8; d++) {
int vd = ln * 8 + d;
if (vd < VD_C) {
po[(base + h) * VD_C + vd] = v_acc[hi * 8 + d] * inv;
}
}
}
}
__global__ void reduce_lse_v71(
const float* __restrict__ po,
const float* __restrict__ pl,
short* __restrict__ out,
int ns, int tbh)
{
int h = blockIdx.x, v = threadIdx.x;
if (h >= tbh || v >= 512) return;
__shared__ float smx;
__shared__ float sw[64];
if (v == 0) {
float m = -1e30f;
for (int s = 0; s < ns; s++) m = fmaxf(m, pl[s * tbh + h]);
smx = m;
}
__syncthreads();
float m = smx;
if (v < ns) sw[v] = expf(pl[v * tbh + h] - m);
__syncthreads();
if (v == 0) {
float s = 0;
for (int i = 0; i < ns; i++) s += sw[i];
smx = (s > 0) ? 1.0f / s : 0.0f;
}
__syncthreads();
float inv = smx;
float acc = 0;
for (int s = 0; s < ns; s++) {
acc += sw[s] * po[(s * tbh + h) * 512 + v];
}
float final_val = acc * inv;
unsigned int fi = __float_as_uint(final_val);
unsigned int rnd = ((fi >> 16) & 1) + 0x7FFF;
out[h * 512 + v] = (short)((fi + rnd) >> 16);
}
void hip_mla_v71(
torch::Tensor Q, torch::Tensor KVf, torch::Tensor KVs, torch::Tensor ki,
torch::Tensor po, torch::Tensor pl, torch::Tensor out,
float sms, int bs, int ns, int tq, int nsc,
int q_stride_tok, int q_stride_h)
{
mla_attn_v71<<<dim3(ns, bs), BSZ, 0, 0>>>(
(const short*)Q.data_ptr(),
(const unsigned char*)KVf.data_ptr(),
(const unsigned char*)KVs.data_ptr(),
po.data_ptr<float>(),
pl.data_ptr<float>(),
ki.data_ptr<int>(),
sms, tq, ns, nsc, q_stride_tok, q_stride_h);
int tbh = tq * NH_C;
reduce_lse_v71<<<tbh, 512>>>(
po.data_ptr<float>(),
pl.data_ptr<float>(),
(short*)out.data_ptr(),
ns, tbh);
}
"""
_CPP_SRC_V71 = (
"void hip_mla_v71("
"torch::Tensor Q,torch::Tensor KVf,torch::Tensor KVs,torch::Tensor ki,"
"torch::Tensor po,torch::Tensor pl,torch::Tensor out,"
"float sms,int bs,int ns,int tq,int nsc,"
"int q_stride_tok,int q_stride_h);"
)
# v74 variant: all-heads-per-lane PV MFMA kernel (ported from mla_brrr_v74).
_HIP_SRC_V74 = r"""
#include <hip/hip_runtime.h>
#include <torch/extension.h>
#include <cmath>
typedef float v4f32 __attribute__((ext_vector_type(4)));
typedef short v4i16 __attribute__((ext_vector_type(4)));
__device__ __forceinline__ float e8m0f(unsigned char e){
if(e==0) return 0.0f;
return __uint_as_float(((unsigned int)e) << 23);
}
typedef __bf16 bf16x2_t __attribute__((ext_vector_type(2)));
__device__ __forceinline__ unsigned int hw_fp4_to_bf16_pair(unsigned int packed, float scale, int byte_sel) {
bf16x2_t r;
switch(byte_sel) {
case 0: r = __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(packed, scale, 0); break;
case 1: r = __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(packed, scale, 1); break;
case 2: r = __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(packed, scale, 2); break;
default: r = __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(packed, scale, 3); break;
}
unsigned int result;
memcpy(&result, &r, 4);
return result;
}
__device__ __forceinline__ float bf16_lo_to_f32(unsigned int pair) {
return __uint_as_float((pair & 0xFFFFu) << 16);
}
__device__ __forceinline__ float bf16_hi_to_f32(unsigned int pair) {
return __uint_as_float((pair >> 16) << 16);
}
__device__ __forceinline__ short f32_to_bf16(float f) {
unsigned int fi = __float_as_uint(f);
unsigned int rounding = ((fi >> 16) & 1) + 0x7FFF;
return (short)((fi + rounding) >> 16);
}
#define NH_C 16
#define VD_C 512
#define BN 16
#define PACKED_C 288
#define BSZ 64
#define N_K_ITER 36
#define N_V_TILES 32
#define LDS_KV (BN*PACKED_C)
#define LDS_SC (BN*32)
#define LDS_EXP_OFF (LDS_KV + LDS_SC)
#define LDS_EXP_SIZE (16*16*4)
#define LDS_TOTAL (LDS_EXP_OFF + LDS_EXP_SIZE)
__global__ __attribute__((amdgpu_flat_work_group_size(BSZ, BSZ)))
void mla_attn_v74(
const short* __restrict__ Q,
const unsigned char* __restrict__ KVf,
const unsigned char* __restrict__ KVs,
float* __restrict__ po,
float* __restrict__ pl,
const int* __restrict__ ki,
float sms, int tq, int ns, int nsc,
int q_stride_tok, int q_stride_h)
{
__shared__ unsigned char lds_all[LDS_TOTAL];
unsigned char* lkv = lds_all;
unsigned char* lsc_buf = lds_all + LDS_KV;
float* lds_exp = (float*)(lds_all + LDS_EXP_OFF); // [16][16] floats
int tid = threadIdx.x;
int ln = tid;
int lmn = ln % 16;
int lk = ln / 16;
int sid = blockIdx.x;
int bid = blockIdx.y;
int kvs_ = ki[bid], kve = ki[bid+1], kvl = kve - kvs_;
int per_split = (kvl > 0) ? (kvl + ns - 1) / ns : 0;
int ms = kvs_ + sid * per_split, me = min(ms + per_split, kve);
if (kvl <= 0 || ms >= kve) {
int base = sid * tq * NH_C + bid * NH_C;
for (int h = tid; h < NH_C; h += BSZ) pl[base + h] = -1e30f;
for (int i = tid; i < NH_C * VD_C; i += BSZ) po[(base + i / VD_C) * VD_C + i % VD_C] = 0.0f;
return;
}
v4f32 pv_acc[N_V_TILES];
#pragma unroll
for (int t = 0; t < N_V_TILES; t++) pv_acc[t] = {0, 0, 0, 0};
float row_max[NH_C], row_sum[NH_C];
#pragma unroll
for (int i = 0; i < NH_C; i++) { row_max[i] = -1e30f; row_sum[i] = 0.0f; }
float all_exp_s[NH_C];
int num_tiles = (me - ms + BN - 1) / BN;
if (num_tiles <= 0) {
int base = sid * tq * NH_C + bid * NH_C;
for (int h = tid; h < NH_C; h += BSZ) pl[base + h] = -1e30f;
for (int i = tid; i < NH_C * VD_C; i += BSZ) po[(base + i / VD_C) * VD_C + i % VD_C] = 0.0f;
return;
}
const short* q_base = Q + bid * q_stride_tok;
for (int tile = 0; tile < num_tiles; tile++) {
int kp = ms + tile * BN;
int ch = min(BN, me - kp);
{
const unsigned int* src_kv = (const unsigned int*)(KVf + kp * PACKED_C);
unsigned int* dst_kv = (unsigned int*)lkv;
int total_u32 = (BN * PACKED_C) / 4;
for (int i = tid; i < total_u32; i += BSZ) dst_kv[i] = src_kv[i];
if (ch < BN) {
int valid_u32 = (ch * PACKED_C) / 4;
for (int i = valid_u32 + tid; i < total_u32; i += BSZ) dst_kv[i] = 0;
}
}
{
const unsigned char* src_sc = KVs + kp * nsc;
unsigned int* sc_u32 = (unsigned int*)lsc_buf;
for (int i = tid; i < (BN * 32) / 4; i += BSZ) sc_u32[i] = 0x7F7F7F7F;
for (int t = tid; t < BN && t < ch; t += BSZ) {
for (int c = 0; c < nsc; c++) lsc_buf[t * 32 + c] = src_sc[t * nsc + c];
}
}
__syncthreads();
v4f32 qk0 = {0, 0, 0, 0};
int k_byte_base = lmn * PACKED_C + (lk * 4) / 2;
int k_sc_base = lmn * 32;
for (int k = 0; k < N_K_ITER; k++) {
int d_base = k * 16;
int d_start = d_base + lk * 4;
v4i16 q_reg;
if (d_start + 3 < 576) {
const short* qp = q_base + lmn * q_stride_h + d_start;
q_reg[0] = qp[0];
q_reg[1] = qp[1];
q_reg[2] = qp[2];
q_reg[3] = qp[3];
} else {
q_reg = {0, 0, 0, 0};
}
v4i16 k_reg;
if (lmn < ch && d_start + 3 < 576) {
int byte_off = k_byte_base + d_base / 2;
unsigned int packed = *(const unsigned short*)&lkv[byte_off];
float sc = e8m0f(lsc_buf[k_sc_base + d_start / 32]);
unsigned int pair0 = hw_fp4_to_bf16_pair(packed, sc, 0);
unsigned int pair1 = hw_fp4_to_bf16_pair(packed, sc, 1);
k_reg[0] = (short)(pair0 & 0xFFFF);
k_reg[1] = (short)((pair0 >> 16) & 0xFFFF);
k_reg[2] = (short)(pair1 & 0xFFFF);
k_reg[3] = (short)((pair1 >> 16) & 0xFFFF);
} else {
k_reg = {0, 0, 0, 0};
}
qk0 = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(q_reg, k_reg, qk0, 0, 0, 0);
}
for (int h = 0; h < NH_C; h++) {
int hp = h % 4;
int hg = h / 4;
float rs = qk0[hp];
float ls = __shfl(rs, hg * 16 + lmn) * sms;
if (lmn >= ch) ls = -1e30f;
float bm = ls;
for (int o = 8; o >= 1; o >>= 1) bm = fmaxf(bm, __shfl_xor(bm, o));
float new_max = fmaxf(row_max[h], bm);
float rescale = expf(row_max[h] - new_max);
if (h / 4 == lk) {
int v = h % 4;
row_sum[h] *= rescale;
#pragma unroll
for (int t = 0; t < N_V_TILES; t++) pv_acc[t][v] *= rescale;
} else {
row_sum[h] *= rescale;
}
float es = expf(ls - new_max);
if (lmn >= ch) es = 0.0f;
float su = es;
for (int o = 8; o >= 1; o >>= 1) su += __shfl_xor(su, o);
row_sum[h] += su;
row_max[h] = new_max;
all_exp_s[h] = es;
}
if (lk == 0) {
for (int h = 0; h < NH_C; h++) lds_exp[lmn * NH_C + h] = all_exp_s[h];
}
int v_tok_off[4];
v_tok_off[0] = (lk * 4) * PACKED_C;
v_tok_off[1] = v_tok_off[0] + PACKED_C;
v_tok_off[2] = v_tok_off[1] + PACKED_C;
v_tok_off[3] = v_tok_off[2] + PACKED_C;
int v_sc_tok_off[4];
v_sc_tok_off[0] = (lk * 4) * 32;
v_sc_tok_off[1] = v_sc_tok_off[0] + 32;
v_sc_tok_off[2] = v_sc_tok_off[1] + 32;
v_sc_tok_off[3] = v_sc_tok_off[2] + 32;
for (int vt = 0; vt < N_V_TILES; vt++) {
int v_dim = lmn + vt * 16;
int v_byte_idx = v_dim / 2;
int v_nibble_hi = v_dim & 1;
int v_sc_idx = v_dim / 32;
v4i16 p_reg;
#pragma unroll
for (int i = 0; i < 4; i++) {
int tok = lk * 4 + i;
float w = lds_exp[tok * NH_C + lmn];
p_reg[i] = f32_to_bf16(w);
}
v4i16 v_reg;
#pragma unroll
for (int i = 0; i < 4; i++) {
int tok = lk * 4 + i;
if (tok < ch && v_dim < VD_C) {
unsigned int pk_u32 = (unsigned int)lkv[v_tok_off[i] + v_byte_idx];
float sc = e8m0f(lsc_buf[v_sc_tok_off[i] + v_sc_idx]);
unsigned int pair = hw_fp4_to_bf16_pair(pk_u32, sc, 0);
v_reg[i] = v_nibble_hi ? (short)((pair >> 16) & 0xFFFF) : (short)(pair & 0xFFFF);
} else {
v_reg[i] = 0;
}
}
pv_acc[vt] = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(p_reg, v_reg, pv_acc[vt], 0, 0, 0);
}
__syncthreads();
}
int base = sid * tq * NH_C + bid * NH_C;
for (int v = 0; v < 4; v++) {
int h = lk * 4 + v;
float inv = (row_sum[h] > 0.0f) ? 1.0f / row_sum[h] : 0.0f;
for (int vt = 0; vt < N_V_TILES; vt++) {
int vd = vt * 16 + lmn;
if (vd < VD_C) po[(base + h) * VD_C + vd] = pv_acc[vt][v] * inv;
}
}
if (lmn == 0) {
for (int v = 0; v < 4; v++) {
int h = lk * 4 + v;
pl[base + h] = (row_sum[h] > 0.0f) ? row_max[h] + logf(row_sum[h]) : -1e30f;
}
}
}
__global__ void reduce_lse_v74(
const float* __restrict__ po,
const float* __restrict__ pl,
short* __restrict__ out,
int ns, int tbh)
{
int h = blockIdx.x, v = threadIdx.x;
if (h >= tbh || v >= 512) return;
__shared__ float smx;
__shared__ float sw[64];
if (v == 0) {
float m = -1e30f;
for (int s = 0; s < ns; s++) m = fmaxf(m, pl[s * tbh + h]);
smx = m;
}
__syncthreads();
float m = smx;
if (v < ns) sw[v] = expf(pl[v * tbh + h] - m);
__syncthreads();
if (v == 0) {
float s = 0;
for (int i = 0; i < ns; i++) s += sw[i];
smx = (s > 0) ? 1.0f / s : 0.0f;
}
__syncthreads();
float inv = smx;
float acc = 0;
for (int s = 0; s < ns; s++) acc += sw[s] * po[(s * tbh + h) * 512 + v];
float final_val = acc * inv;
unsigned int fi = __float_as_uint(final_val);
unsigned int rnd = ((fi >> 16) & 1) + 0x7FFF;
out[h * 512 + v] = (short)((fi + rnd) >> 16);
}
void hip_mla_v74(
torch::Tensor Q, torch::Tensor KVf, torch::Tensor KVs, torch::Tensor ki,
torch::Tensor po, torch::Tensor pl, torch::Tensor out,
float sms, int bs, int ns, int tq, int nsc,
int q_stride_tok, int q_stride_h)
{
mla_attn_v74<<<dim3(ns, bs), BSZ, 0, 0>>>(
(const short*)Q.data_ptr(),
(const unsigned char*)KVf.data_ptr(),
(const unsigned char*)KVs.data_ptr(),
po.data_ptr<float>(),
pl.data_ptr<float>(),
ki.data_ptr<int>(),
sms, tq, ns, nsc, q_stride_tok, q_stride_h);
int tbh = tq * NH_C;
reduce_lse_v74<<<tbh, 512>>>(
po.data_ptr<float>(),
pl.data_ptr<float>(),
(short*)out.data_ptr(),
ns, tbh);
}
"""
_CPP_SRC_V74 = (
"void hip_mla_v74("
"torch::Tensor Q,torch::Tensor KVf,torch::Tensor KVs,torch::Tensor ki,"
"torch::Tensor po,torch::Tensor pl,torch::Tensor out,"
"float sms,int bs,int ns,int tq,int nsc,"
"int q_stride_tok,int q_stride_h);"
)
def _select_custom_variant_sources() -> tuple[str, str, str, str]:
"""
Resolve custom backend variant to (cpp_src, hip_src, entry_name, variant_tag).
"""
if CUSTOM_V71_VARIANT == "v74":
return _CPP_SRC_V74, _HIP_SRC_V74, "hip_mla_v74", "v74"
return _CPP_SRC_V71, _HIP_SRC_V71, "hip_mla_v71", "v71"
def _load_custom_v71_impl() -> bool:
"""
MFMA probe harness loader.
Compile custom HIP v71 long-KV path on demand.
"""
global _hip_custom, _hip_custom_tried, _hip_custom_ok
global _custom_variant_resolved, _custom_entry_name, _custom_v71_module_name
if _hip_custom_tried:
return _hip_custom_ok
_hip_custom_tried = True
try:
from torch.utils.cpp_extension import load_inline
cpp_src, hip_src, entry_name, variant_tag = _select_custom_variant_sources()
t0 = time.time()
_custom_v71_module_name = _next_custom_v71_module_name(
cpp_src=cpp_src,
hip_src=hip_src,
resolved_variant=variant_tag,
)
_hip_custom = load_inline(
name=_custom_v71_module_name,
cpp_sources=[cpp_src],
cuda_sources=[hip_src],
functions=[entry_name],
extra_cuda_cflags=[
f"--offload-arch={CUSTOM_V71_OFFLOAD_ARCH}",
"-std=c++20",
"-U__HIP_NO_HALF_OPERATORS__",
"-U__HIP_NO_HALF_CONVERSIONS__",
"-O3",
],
verbose=False,
)
_custom_variant_resolved = variant_tag
_custom_entry_name = entry_name
_hip_custom_ok = True
_log_route(
"custom-backend compile ok "
f"name={_custom_v71_module_name} "
f"variant={_custom_variant_resolved} "
f"entry={_custom_entry_name} "
f"arch={CUSTOM_V71_OFFLOAD_ARCH} "
f"in {time.time()-t0:.2f}s"
)
except Exception as exc:
_log_route(f"custom-backend compile failed: {exc}")
_hip_custom_ok = False
return _hip_custom_ok
def _compute_custom_splits(
batch_size: int,
kv_seq_len: int,
variant_hint: str | None = None,
cu: int = 256,
) -> int:
avg = kv_seq_len
variant_key = (variant_hint or "").lower()
if variant_key.startswith("v74"):
# v74 uses BSZ=64 (1 wavefront), increasing per-CU concurrent block capacity.
# Model this with higher effective block capacity and lower per-split overhead.
effective_cu = cu * 4
overhead = 72.0
max_splits = 24
else:
effective_cu = cu
overhead = 84.1
max_splits = 16
best_score = -1.0
best_splits = 1
for i in range(1, max_splits + 1):
total = batch_size * i
waves = (total + effective_cu - 1) // effective_cu
util = total / (waves * effective_cu)
eff = avg / (avg + overhead * i)
score = util * eff
if score > best_score:
best_score = score
best_splits = i
return best_splits
def _resolve_custom_splits(batch_size: int, kv_seq_len: int, variant_hint: str | None = None) -> int:
key = (batch_size, kv_seq_len)
if key in _LOCAL_CUSTOM_NSPLIT_OVERRIDES:
return max(1, min(32, int(_LOCAL_CUSTOM_NSPLIT_OVERRIDES[key])))
if key in CUSTOM_NSPLIT_OVERRIDES:
return max(1, min(32, int(CUSTOM_NSPLIT_OVERRIDES[key])))
return _compute_custom_splits(batch_size, kv_seq_len, variant_hint=variant_hint)
def _get_custom_v71_scratch(ns: int, tbh: int, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
key = (_device_key(device), ns, tbh)
if key not in _custom_v71_scratch_cache:
po = torch.empty((ns, tbh, V_HEAD_DIM), dtype=torch.float32, device=device)
pl = torch.empty((ns, tbh), dtype=torch.float32, device=device)
_custom_v71_scratch_cache[key] = (po, pl)
return _custom_v71_scratch_cache[key]
def _get_baseline_compare_out(num_q: int, device: torch.device) -> torch.Tensor:
key = (_device_key(device), num_q)
if key not in _baseline_compare_out_cache:
_baseline_compare_out_cache[key] = torch.empty(
(num_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device
)
return _baseline_compare_out_cache[key]
def _run_custom_v71(
q: torch.Tensor,
kv_data: dict,
kv_indptr: torch.Tensor,
batch_size: int,
kv_seq_len: int,
) -> torch.Tensor | None:
if not _load_custom_v71_impl():
return None
try:
variant_label = _custom_variant_resolved or CUSTOM_V71_VARIANT or "unknown"
fp4_packed, scale_e8m0 = kv_data["mxfp4"]
kvf = fp4_packed.reshape(-1, CUSTOM_V71_PACKED_DIM).contiguous().view(torch.uint8)
kvs = scale_e8m0.contiguous().view(torch.uint8)
if kvs.dim() != 2:
kvs = kvs.view(-1, kvs.numel() // kvf.shape[0])
q_view = q.view(q.shape[0], NUM_HEADS, QK_HEAD_DIM)
ns = _resolve_custom_splits(batch_size, kv_seq_len, variant_hint=variant_label)
tbh = q_view.shape[0] * NUM_HEADS
po, pl = _get_custom_v71_scratch(ns, tbh, q.device)
po.zero_()
pl.fill_(-1e30)
out = torch.empty((q_view.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)
if _hip_custom is None or _custom_entry_name is None:
return None
run_fn = getattr(_hip_custom, _custom_entry_name)
run_fn(
q_view,
kvf,
kvs,
kv_indptr,
po,
pl,
out,
SM_SCALE,
batch_size,
ns,
q_view.shape[0],
kvs.shape[1],
q_view.stride(0),
q_view.stride(1),
)
_log_route(
"custom-backend active "
f"variant={variant_label} "
f"bs={batch_size} kv={kv_seq_len} ns={ns}"
)
return out
except Exception as exc:
_log_route(
"custom-backend runtime failed "
f"variant={_custom_variant_resolved or CUSTOM_V71_VARIANT or 'unknown'} "
f"err={exc}"
)
return None
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = int(config["batch_size"])
kv_seq_len = int(config["kv_seq_len"])
# Best-known routing boundary from v93-family.
if batch_size <= 32 and kv_seq_len <= 1024:
out = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)
_log_route(f"bf16-np bs={batch_size} kv={kv_seq_len}")
_run_bf16_non_persistent(
q=q,
kv_bf16=kv_data["bf16"],
qo_indptr=qo_indptr,
kv_indptr=kv_indptr,
batch_size=batch_size,
kv_seq_len=kv_seq_len,
out=out,
)
return out
if _should_use_custom_v71(batch_size, kv_seq_len):
custom_out = _run_custom_v71(
q=q,
kv_data=kv_data,
kv_indptr=kv_indptr,
batch_size=batch_size,
kv_seq_len=kv_seq_len,
)
if custom_out is not None:
need_baseline_compare = VALIDATE_MFMA_V71 or VALIDATE_MFMA_V71_MAX_ABS > 0
if need_baseline_compare:
baseline_out = _get_baseline_compare_out(q.shape[0], q.device)
kv_fp8, kv_scale = kv_data["fp8"]
_run_a16w8_persistent(
q=q,
kv_fp8=kv_fp8,
kv_scale=kv_scale,
qo_indptr=qo_indptr,
batch_size=batch_size,
kv_seq_len=kv_seq_len,
out=baseline_out,
)
diff = (custom_out.float() - baseline_out.float()).abs()
max_abs = float(diff.max().item())
if VALIDATE_MFMA_V71:
mean_abs = float(diff.mean().item())
flat = diff.reshape(-1)
p95_abs = float(torch.quantile(flat, 0.95).item()) if flat.numel() > 0 else 0.0
p99_abs = float(torch.quantile(flat, 0.99).item()) if flat.numel() > 0 else 0.0
variant_label = _custom_variant_resolved or CUSTOM_V71_VARIANT or "unknown"
pass_tol = bool(
torch.allclose(
custom_out,
baseline_out,
rtol=VALIDATE_MFMA_V71_RTOL,
atol=VALIDATE_MFMA_V71_ATOL,
)
)
_log_route(
"custom-backend-vs-baseline "
f"variant={variant_label} "
f"max_abs={max_abs:.6f} "
f"p95_abs={p95_abs:.6f} "
f"p99_abs={p99_abs:.6f} "
f"mean_abs={mean_abs:.6f} "
f"tol_pass={pass_tol}"
)
if VALIDATE_MFMA_V71_MAX_ABS > 0 and max_abs > VALIDATE_MFMA_V71_MAX_ABS:
variant_label = _custom_variant_resolved or CUSTOM_V71_VARIANT or "unknown"
_log_route(
"custom-backend rejected by max_abs gate "
f"variant={variant_label} "
f"threshold={VALIDATE_MFMA_V71_MAX_ABS:.6f}"
)
return baseline_out
return custom_out
out = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)
_log_route(f"a16w8-persist bs={batch_size} kv={kv_seq_len}")
kv_fp8, kv_scale = kv_data["fp8"]
_run_a16w8_persistent(
q=q,
kv_fp8=kv_fp8,
kv_scale=kv_scale,
qo_indptr=qo_indptr,
batch_size=batch_size,
kv_seq_len=kv_seq_len,
out=out,
)
return outscrolls · 1293 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 682562.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON