Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
53.9µs
#179 of 766
2026-04-02

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.

fp4fp4_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 out
scrolls · 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