Skip to content
KernelIndex
Search⌘K

submission 747369

nataliakokoromyti · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 2110 lines, June 9 Researcher Reciprocity License v1.0.

baseline_743863.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-747369?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
48.2µs
#155 of 766
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:95834955d042e0c3d165a1f47d25c8d3ae33b36dd9646a4227cac7ad788db538
license declaredunknown
license concludedunknown
authorsnataliakokoromyti
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

shared-memory__shared__ uint8_t klds_buf[1][K_TILE_BYTES];
vector-width = uint4const uint4 lo4 = *reinterpret_cast<const uint4*>(pv);
warp-specializationstatic constexpr int DUET_PRODUCER_WAVES = 2;

Kernel source

baseline_743863.py2110 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
# Trimmed live-path variant of amd_mla_hip_occ2_fused_noblob.py.
# Keeps the active fused IQ2 kernel path and dispatch logic while removing
# disabled variants, unreachable wrappers, and no-blob vestigial scaffolding.
# Original file is preserved unchanged.


import os
from functools import lru_cache

import torch
from torch.utils.cpp_extension import load_inline

os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

# Diagnostic: print GPU info at module load (helps tune splits for target GPU)
try:
    _gpu_props = torch.cuda.get_device_properties(0)
    print(f"[mla] GPU: {_gpu_props.name}, CUs: {_gpu_props.multi_processor_count}", flush=True)
except Exception:
    pass

NUM_HEADS = 16
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK

HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <cstdint>
#include <cmath>
#include <algorithm>

namespace mla {

#define NEG_INF __int_as_float(0xff800000)

static constexpr int NUM_HEADS = 16;
static constexpr int HEADS_PER_WAVE = 4;
static constexpr int D_QK = 576;
static constexpr int D_V = 512;
static constexpr float SM_SCALE = 1.0f / 24.0f;
static constexpr int BLOCK_THREADS = 256;
static constexpr int WAVE_SIZE = 64;
static constexpr int V_ELEMS = D_V / WAVE_SIZE;
static constexpr int TILE_ROWS = 32;
static constexpr int FP8_MAX = 448;
static constexpr int K_COL_BLOCK = 64;
static constexpr int K_NUM_BLOCKS = D_QK / K_COL_BLOCK;
static constexpr int K_NUM_ROWS_PER_SUBBLOCK = 4;
static constexpr int K_NUM_PADDING_DW = 2;
static constexpr int K_NUM_BYTES_PER_ROW = K_COL_BLOCK;
static constexpr int K_NUM_BYTES_PER_SUBBLOCK =
    K_NUM_ROWS_PER_SUBBLOCK * K_NUM_BYTES_PER_ROW + K_NUM_PADDING_DW * 4;
static constexpr int K_NUM_BYTES_PER_BLOCK =
    K_NUM_BYTES_PER_SUBBLOCK * (TILE_ROWS / K_NUM_ROWS_PER_SUBBLOCK);
static constexpr int K_TILE_BYTES = K_NUM_BYTES_PER_BLOCK * K_NUM_BLOCKS;
static constexpr int V_TILE_BYTES = TILE_ROWS * D_V;
static constexpr int VT_DV_SLICE = 16;
static constexpr int VT_NUM_SLICES = D_V / VT_DV_SLICE;
// Column-major VT layout: 8 bytes per (group, col) contiguous for ds_read_b64
// Padding between groups avoids LDS bank conflicts (128+8=136, not multiple of 128)
static constexpr int VT_GROUP_ROWS = 8;
static constexpr int VT_GROUP_PAD = 8;
static constexpr int VT_GROUP_STRIDE = VT_DV_SLICE * VT_GROUP_ROWS + VT_GROUP_PAD; // 136
static constexpr int VT_NUM_GROUPS = TILE_ROWS / VT_GROUP_ROWS; // 4
static constexpr int VT_SLICE_BYTES = VT_NUM_GROUPS * VT_GROUP_STRIDE; // 544
static constexpr int VT_TILE_BYTES = VT_NUM_SLICES * VT_SLICE_BYTES; // 17408
static constexpr uint32_t BUFFER_RESOURCE_CONFIG = 0x00020000;

using bf16 = hip_bfloat16;
using floatx4 = float __attribute__((ext_vector_type(4)));
using floatx16 = float __attribute__((ext_vector_type(16)));
using intx4 = int __attribute__((ext_vector_type(4)));
using intx8 = int __attribute__((ext_vector_type(8)));
using as3_uint32_ptr = uint32_t __attribute__((address_space(3)))*;

extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
    intx4 rsrc,
    as3_uint32_ptr lds_ptr,
    int size,
    int voffset,
    int soffset,
    int offset,
    int aux) __asm("llvm.amdgcn.raw.buffer.load.lds");

struct alignas(16) U128 {
    uint32_t x0, x1, x2, x3;
};

union U64Bytes {
    uint8_t b[8];
    uint64_t u64;
};

struct buffer_resource {
    uint64_t ptr;
    uint32_t range;
    uint32_t config;
};

__device__ __forceinline__ float hw_fp8_to_f32(uint32_t packed) {
    return __builtin_amdgcn_cvt_f32_fp8(packed, 0);
}

__device__ __forceinline__ intx4 make_srsrc(const void* ptr, uint32_t range_bytes) {
    buffer_resource rsrc = {
        reinterpret_cast<uint64_t>(ptr),
        range_bytes, BUFFER_RESOURCE_CONFIG};
    return *reinterpret_cast<const intx4*>(&rsrc);
}

__device__ __forceinline__ as3_uint32_ptr make_wave_lds_ptr(uintptr_t p) {
    uint32_t lane0 = __builtin_amdgcn_readfirstlane(static_cast<uint32_t>(p));
    return reinterpret_cast<as3_uint32_ptr>(static_cast<uintptr_t>(lane0));
}

__device__ __forceinline__ U128 ld16u(const uint8_t* p) {
    return *reinterpret_cast<const U128*>(p);
}

__device__ __forceinline__ uint64_t ld_u64(const uint8_t* p) {
    return *reinterpret_cast<const uint64_t*>(p);
}

__device__ __forceinline__ uint64_t ds_read_tr8_u64(const uint8_t* p)
{
#define __LDS_ADDR __attribute__((address_space(3)))
    typedef __attribute__((__vector_size__(2 * sizeof(int)))) int llvm_i32x2_t;
#pragma clang diagnostic push
#pragma clang diagnostic ignored "-Wold-style-cast"
    const auto p_lds = (__LDS_ADDR uint8_t*)(const_cast<uint8_t*>(p));
#pragma clang diagnostic pop
    auto lds_ptr = reinterpret_cast<__LDS_ADDR llvm_i32x2_t*>(p_lds);
    auto bits    = __builtin_amdgcn_ds_read_tr8_b64_v2i32(lds_ptr);
    return *reinterpret_cast<uint64_t*>(&bits);
#undef __LDS_ADDR
}

__device__ __forceinline__ void st16u(uint8_t* p, U128 v) {
    *reinterpret_cast<U128*>(p) = v;
}

__device__ __forceinline__ float to_f(bf16 x) { return static_cast<float>(x); }
__device__ __forceinline__ bf16 to_b(float x) { return bf16(x); }
__device__ __forceinline__ float fast_exp(float x) {
    // Inline asm: 2 instructions instead of compiler's 6 (skip range reduction).
    // Input x is always in [-54, 0] (attention score diffs bounded by FP8 range),
    // so x*log2(e) in [-78, 0], well within v_exp_f32 safe range [-126, 128].
    float r;
    asm("v_mul_f32 %0, 0x3fb8aa3b, %1\n\t"  // r = x * log2(e)
        "v_exp_f32 %0, %0"                    // r = 2^r = exp(x)
        : "=v"(r) : "v"(x));
    return r;
}

__device__ __forceinline__ int q_off(int qi, int h, int d) {
    return (qi * NUM_HEADS + h) * D_QK + d;
}

__device__ __forceinline__ int out_off(int qi, int h, int d) {
    return (qi * NUM_HEADS + h) * D_V + d;
}

__device__ __forceinline__ int pml_off(int b, int s, int h, int ns) {
    return (b * NUM_HEADS + h) * ns + s;
}

__device__ __forceinline__ int po_off(int b, int s, int h, int d, int ns) {
    return ((b * NUM_HEADS + h) * ns + s) * D_V + d;
}

__device__ __forceinline__ int lq8(int h, int d) { return h * D_QK + d; }

__device__ __forceinline__ int kv2_lds_offset(int row, int d)
{
    const int block = d / K_COL_BLOCK;
    const int d_in_block = d % K_COL_BLOCK;
    const int half = row / 16;
    const int row16 = row % 16;
    const int row_phy = (row16 / 2) * 4 + (row16 % 2);
    return block * K_NUM_BYTES_PER_BLOCK +
           half * 128 +
           (row_phy / 4) * K_NUM_BYTES_PER_SUBBLOCK +
           (row_phy % 4) * K_NUM_BYTES_PER_ROW +
           d_in_block;
}

__device__ __forceinline__ float wreduce_max(float v) {
    v = fmaxf(v, __shfl_down(v, 32));
    v = fmaxf(v, __shfl_down(v, 16));
    v = fmaxf(v, __shfl_down(v, 8));
    v = fmaxf(v, __shfl_down(v, 4));
    v = fmaxf(v, __shfl_down(v, 2));
    v = fmaxf(v, __shfl_down(v, 1));
    return v;
}

__device__ __forceinline__ float wbcast(float v) { return __shfl(v, 0); }

// DPP row_ror butterfly reduction within 16-lane rows.
// row_ror:N rotates right by N within each 16-lane row (wraps around).
// Uses fused v_max_f32_dpp / v_add_f32_dpp to avoid pipeline hazards.
// s_nop between steps ensures the result commits before the next DPP read.

// DPP row_ror butterfly reduction within 16-lane rows.
// Separate v_mov_b32_dpp + regular ALU to avoid RAW hazard.
// s_nop 1 between steps: 1 intervening ALU + 2 nop cycles = 3 cycles,
// sufficient for single-value DPP chain on gfx950.
__device__ __forceinline__ float wave16_max(float v) {
    float t;
    asm volatile(
        "v_mov_b32_dpp %1, %0 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_max_f32 %0, %0, %1\n\t"
        "s_nop 1\n\t"
        "v_mov_b32_dpp %1, %0 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_max_f32 %0, %0, %1\n\t"
        "s_nop 1\n\t"
        "v_mov_b32_dpp %1, %0 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_max_f32 %0, %0, %1\n\t"
        "s_nop 1\n\t"
        "v_mov_b32_dpp %1, %0 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_max_f32 %0, %0, %1"
        : "+v"(v), "=&v"(t));
    return v;
}

__device__ __forceinline__ float wave16_sum(float v) {
    float t;
    asm volatile(
        "v_mov_b32_dpp %1, %0 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_add_f32 %0, %0, %1\n\t"
        "s_nop 1\n\t"
        "v_mov_b32_dpp %1, %0 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_add_f32 %0, %0, %1\n\t"
        "s_nop 1\n\t"
        "v_mov_b32_dpp %1, %0 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_add_f32 %0, %0, %1\n\t"
        "s_nop 1\n\t"
        "v_mov_b32_dpp %1, %0 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_add_f32 %0, %0, %1"
        : "+v"(v), "=&v"(t));
    return v;
}

// DPP row_ror butterfly max/sum for 4 values simultaneously.
// 4 intervening ALU ops between DPP read and next use of same register
// provides sufficient latency (4-5 cycles on gfx950) — no s_nop needed.
__device__ __forceinline__ void wave16_max4(float &a, float &b, float &c, float &d) {
    float t0, t1, t2, t3;
    asm volatile(
        "v_mov_b32_dpp %4, %0 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %5, %1 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %6, %2 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %7, %3 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_max_f32 %0, %0, %4\n\t"
        "v_max_f32 %1, %1, %5\n\t"
        "v_max_f32 %2, %2, %6\n\t"
        "v_max_f32 %3, %3, %7\n\t"
        "v_mov_b32_dpp %4, %0 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %5, %1 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %6, %2 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %7, %3 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_max_f32 %0, %0, %4\n\t"
        "v_max_f32 %1, %1, %5\n\t"
        "v_max_f32 %2, %2, %6\n\t"
        "v_max_f32 %3, %3, %7\n\t"
        "v_mov_b32_dpp %4, %0 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %5, %1 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %6, %2 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %7, %3 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_max_f32 %0, %0, %4\n\t"
        "v_max_f32 %1, %1, %5\n\t"
        "v_max_f32 %2, %2, %6\n\t"
        "v_max_f32 %3, %3, %7\n\t"
        "v_mov_b32_dpp %4, %0 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %5, %1 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %6, %2 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %7, %3 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_max_f32 %0, %0, %4\n\t"
        "v_max_f32 %1, %1, %5\n\t"
        "v_max_f32 %2, %2, %6\n\t"
        "v_max_f32 %3, %3, %7"
        : "+v"(a), "+v"(b), "+v"(c), "+v"(d),
          "=&v"(t0), "=&v"(t1), "=&v"(t2), "=&v"(t3));
}

__device__ __forceinline__ void wave16_sum4(float &a, float &b, float &c, float &d) {
    float t0, t1, t2, t3;
    asm volatile(
        "v_mov_b32_dpp %4, %0 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %5, %1 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %6, %2 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %7, %3 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
        "v_add_f32 %0, %0, %4\n\t"
        "v_add_f32 %1, %1, %5\n\t"
        "v_add_f32 %2, %2, %6\n\t"
        "v_add_f32 %3, %3, %7\n\t"
        "v_mov_b32_dpp %4, %0 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %5, %1 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %6, %2 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %7, %3 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
        "v_add_f32 %0, %0, %4\n\t"
        "v_add_f32 %1, %1, %5\n\t"
        "v_add_f32 %2, %2, %6\n\t"
        "v_add_f32 %3, %3, %7\n\t"
        "v_mov_b32_dpp %4, %0 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %5, %1 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %6, %2 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %7, %3 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
        "v_add_f32 %0, %0, %4\n\t"
        "v_add_f32 %1, %1, %5\n\t"
        "v_add_f32 %2, %2, %6\n\t"
        "v_add_f32 %3, %3, %7\n\t"
        "v_mov_b32_dpp %4, %0 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %5, %1 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %6, %2 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_mov_b32_dpp %7, %3 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
        "v_add_f32 %0, %0, %4\n\t"
        "v_add_f32 %1, %1, %5\n\t"
        "v_add_f32 %2, %2, %6\n\t"
        "v_add_f32 %3, %3, %7"
        : "+v"(a), "+v"(b), "+v"(c), "+v"(d),
          "=&v"(t0), "=&v"(t1), "=&v"(t2), "=&v"(t3));
}

__device__ __forceinline__ float wave32_max(float v) {
    v = fmaxf(v, __shfl_down(v, 16, 32));
    v = fmaxf(v, __shfl_down(v, 8, 32));
    v = fmaxf(v, __shfl_down(v, 4, 32));
    v = fmaxf(v, __shfl_down(v, 2, 32));
    v = fmaxf(v, __shfl_down(v, 1, 32));
    return __shfl(v, 0, 32);
}

__device__ __forceinline__ float wave32_sum(float v) {
    v += __shfl_down(v, 16, 32);
    v += __shfl_down(v, 8, 32);
    v += __shfl_down(v, 4, 32);
    v += __shfl_down(v, 2, 32);
    v += __shfl_down(v, 1, 32);
    return __shfl(v, 0, 32);
}

struct SM {
    float m, l;
};

__device__ __forceinline__ void sm_init(SM& s) {
    s.m = NEG_INF;
    s.l = 0.f;
}

__device__ __forceinline__ void sm_upd(SM& s, float sc, float& a, float& b) {
    float mn = fmaxf(s.m, sc);
    a = fast_exp(s.m - mn);
    b = fast_exp(sc - mn);
    s.l = a * s.l + b;
    s.m = mn;
}

__device__ inline void build_fp8_lut(float* lut, float scale, int tid) {
    if (tid < 256) {
        lut[tid] = hw_fp8_to_f32(static_cast<uint32_t>(tid)) * scale;
    }
}

__device__ __forceinline__ uint8_t cvt_fp8_scalar(float x) {
    uint32_t w = 0;
    x = fminf(fmaxf(x, -static_cast<float>(FP8_MAX)), static_cast<float>(FP8_MAX));
    w = __builtin_amdgcn_cvt_pk_fp8_f32(x, 0.0f, w, 0);
    return static_cast<uint8_t>(w & 0xffu);
}

__device__ __forceinline__ uint32_t cvt_fp8x4(float a, float b, float c, float d) {
    uint32_t out = 0;
    out = __builtin_amdgcn_cvt_pk_fp8_f32(a, b, out, 0);
    out = __builtin_amdgcn_cvt_pk_fp8_f32(c, d, out, 1);
    return out;
}

__device__ inline void stage_q_fp8(
    const bf16* __restrict__ q,
    uint8_t* q8,
    float* q_scales,
    int qi,
    int wave,
    int lane)
{
    const int hb = wave * HEADS_PER_WAVE;
    float local_max = 0.f;
    for (int i = lane; i < HEADS_PER_WAVE * D_QK; i += WAVE_SIZE) {
        int h = hb + (i / D_QK);
        int d = i % D_QK;
        local_max = fmaxf(local_max, fabsf(to_f(q[q_off(qi, h, d)])));
    }

    float max_abs = wbcast(wreduce_max(local_max));
    float q_scale = fmaxf(
        max_abs / static_cast<float>(FP8_MAX),
        1.0f / static_cast<float>(FP8_MAX));
    float inv_q_scale = 1.0f / q_scale;

    if (lane == 0) {
        q_scales[wave] = q_scale;
    }

    for (int i = lane * 4; i < HEADS_PER_WAVE * D_QK; i += WAVE_SIZE * 4) {
        int remain = HEADS_PER_WAVE * D_QK - i;
        if (remain >= 4) {
            int h0 = hb + ((i + 0) / D_QK);
            int h1 = hb + ((i + 1) / D_QK);
            int h2 = hb + ((i + 2) / D_QK);
            int h3 = hb + ((i + 3) / D_QK);
            int d0 = (i + 0) % D_QK;
            int d1 = (i + 1) % D_QK;
            int d2 = (i + 2) % D_QK;
            int d3 = (i + 3) % D_QK;

            uint32_t packed = cvt_fp8x4(
                to_f(q[q_off(qi, h0, d0)]) * inv_q_scale,
                to_f(q[q_off(qi, h1, d1)]) * inv_q_scale,
                to_f(q[q_off(qi, h2, d2)]) * inv_q_scale,
                to_f(q[q_off(qi, h3, d3)]) * inv_q_scale);
            *reinterpret_cast<uint32_t*>(&q8[hb * D_QK + i]) = packed;
        } else {
            for (int j = 0; j < remain; ++j) {
                int idx = i + j;
                int h = hb + (idx / D_QK);
                int d = idx % D_QK;
                q8[hb * D_QK + idx] = cvt_fp8_scalar(to_f(q[q_off(qi, h, d)]) * inv_q_scale);
            }
        }
    }
}

template <int COL_OFFSET>
__device__ __forceinline__ void direct_load_k_block(
    intx4 srsrc,
    uintptr_t p_lds_k_warp_base,
    int row,
    int col_base)
{
    constexpr int k_block_idx = COL_OFFSET / K_COL_BLOCK;
    constexpr uintptr_t k_lds_block_base =
        k_block_idx * K_NUM_BYTES_PER_BLOCK - COL_OFFSET;
    const int voffset = row * D_QK + col_base;
    llvm_amdgcn_raw_buffer_load_lds(
        srsrc,
        make_wave_lds_ptr(p_lds_k_warp_base + k_lds_block_base),
        4,
        voffset,
        0,
        COL_OFFSET,
        0);
}

__device__ inline void stage_k_tile_kv2(
    intx4 srsrc,
    uint8_t* klds,
    int base_token,
    int rows,
    int wave,
    int lane)
{
    const int col_base = (lane & 15) * 4;
    #pragma unroll
    for (int emu = 0; emu < 2; ++emu) {
        const int warp_idx = wave + emu * 4;
        const int row_base = (lane / 32) * 16 + ((lane / 16) & 1) + warp_idx * 2;
        const int row = (row_base < rows) ? (base_token + row_base) : -1;
        if (row < 0) continue;  // Skip DMA for invalid rows
        const uintptr_t p_lds_k_warp_base =
            reinterpret_cast<uintptr_t>(klds)
            + warp_idx * K_NUM_BYTES_PER_SUBBLOCK;
        direct_load_k_block<0>(srsrc, p_lds_k_warp_base, row, col_base);
        direct_load_k_block<64>(srsrc, p_lds_k_warp_base, row, col_base);
        direct_load_k_block<128>(srsrc, p_lds_k_warp_base, row, col_base);
        direct_load_k_block<192>(srsrc, p_lds_k_warp_base, row, col_base);
        direct_load_k_block<256>(srsrc, p_lds_k_warp_base, row, col_base);
        direct_load_k_block<320>(srsrc, p_lds_k_warp_base, row, col_base);
        direct_load_k_block<384>(srsrc, p_lds_k_warp_base, row, col_base);
        direct_load_k_block<448>(srsrc, p_lds_k_warp_base, row, col_base);
        direct_load_k_block<512>(srsrc, p_lds_k_warp_base, row, col_base);
    }
}

__device__ inline void stage_v_tile_linear(
    intx4 srsrc,
    uint8_t* vlds,
    int base_token,
    int rows,
    int wave,
    int lane)
{
    const int seg = lane & 31;
    const int which = lane >> 5;
    for (int pair = wave; pair < TILE_ROWS / 2; pair += 4) {
        const int row0 = pair * 2;
        const int local_row = row0 + which;
        const int global_row = (local_row < rows) ? (base_token + local_row) : -1;
        const int voffset = (global_row >= 0) ? (global_row * D_QK + seg * 16) : 0x80000000;
        const uintptr_t p_lds_pair = reinterpret_cast<uintptr_t>(vlds + row0 * D_V);
        llvm_amdgcn_raw_buffer_load_lds(
            srsrc,
            make_wave_lds_ptr(p_lds_pair),
            16,
            voffset,
            0,
            0,
            0);
    }
}

__device__ inline void zero_tail_tiles(uint8_t* klds, uint8_t* vlds, int rows, int tid)
{
    if (rows >= TILE_ROWS) {
        return;
    }

    for (int row = rows; row < TILE_ROWS; ++row) {
        for (int d = tid; d < D_QK; d += BLOCK_THREADS) {
            const int block = d / K_COL_BLOCK;
            const int d_in_block = d % K_COL_BLOCK;
            const int half = row / 16;
            const int row16 = row % 16;
            const int row_phy = (row16 / 2) * 4 + (row16 % 2);
            const int offs = block * K_NUM_BYTES_PER_BLOCK +
                             half * 128 +
                             (row_phy / 4) * K_NUM_BYTES_PER_SUBBLOCK +
                             (row_phy % 4) * K_NUM_BYTES_PER_ROW +
                             d_in_block;
            klds[offs] = 0;
        }
        for (int d = tid; d < D_V; d += BLOCK_THREADS) {
            vlds[row * D_V + d] = 0;
        }
    }
}

__device__ __forceinline__ uint64_t pack_q_mfma(
    const uint8_t* q8,
    int hb,
    int lane,
    int k_block)
{
    int row = lane & 15;
    int real_h = hb + (row & 3);
    int k_base = k_block + ((lane >> 4) * 8);
    return *reinterpret_cast<const uint64_t*>(&q8[lq8(real_h, k_base)]);
}

__device__ __forceinline__ uint64_t load_k_frag_kv2(
    const uint8_t* klds,
    int lane,
    int row_offset,
    int k_block)
{
    const int row = lane & 15;
    const int row_phy = (row / 2) * 4 + (row % 2);
    const int col = (lane >> 4) * 8;
    const int fixed = (row_offset / 16) * 128
        + (k_block % K_COL_BLOCK)
        + (k_block / K_COL_BLOCK) * K_NUM_BYTES_PER_BLOCK;
    const uint8_t* p = klds +
                       (row_phy / 4) * K_NUM_BYTES_PER_SUBBLOCK +
                       (row_phy % 4) * K_NUM_BYTES_PER_ROW +
                       col +
                       fixed;
    return *reinterpret_cast<const uint64_t*>(p);
}

__device__ __forceinline__ floatx4 mfma_fp8_16x16x32(long a, long b, floatx4 c) {
    return __builtin_amdgcn_mfma_f32_16x16x32_fp8_fp8(a, b, c, 0, 0, 0);
}

__device__ __forceinline__ floatx16 mfma_fp8_32x32x16(long a, long b, floatx16 c) {
    return __builtin_amdgcn_mfma_f32_32x32x16_fp8_fp8(a, b, c, 0, 0, 0);
}

// Scaled MFMA: 16x16x128 with fp8 (E4M3=type 0), 4x more K-reduction per MFMA
// a,b: v8i (32 bytes = 32 fp8 elements per lane)
// scale_a, scale_b: per-block scaling factors (int, passed as sgpr)
__device__ __forceinline__ floatx4 mfma_scale_fp8_16x16x128_noscale(
    intx8 a, intx8 b, floatx4 c) {
    // opsel=1 bypasses the scale VGPR and uses implicit scale=1 (no scaling)
    return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        a, b, c, /*Atype=*/0, /*Btype=*/0, /*opsel_a=*/1, 0, /*opsel_b=*/1, 0);
}
__device__ __forceinline__ floatx4 mfma_scale_fp8_16x16x128(
    intx8 a, intx8 b, floatx4 c, int scale_a, int scale_b) {
    return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        a, b, c, /*Atype=*/0, /*Btype=*/0, /*opsel_a=*/0, scale_a, /*opsel_b=*/0, scale_b);
}

template <int TILE>
__device__ inline void attn_direct_kv2_fp8(
    const uint8_t* q8,
    const float* q_scales,
    const float* lut,
    float sc,
    const uint8_t* __restrict__ kvg,
    uint8_t* klds,
    uint8_t* vlds,
    int hb,
    int wave,
    int lane,
    int rs,
    int re,
    float* a0,
    float* a1,
    float* a2,
    float* a3,
    SM* st)
{
    const intx4 srsrc = make_srsrc(kvg, 0xffffffffu);
    const float score_scale = q_scales[wave] * sc * SM_SCALE;

    for (int tb = rs; tb < re; tb += TILE) {
        const int rows = min(TILE, re - tb);
        stage_k_tile_kv2(srsrc, klds, tb, rows, wave, lane);
        stage_v_tile_linear(srsrc, vlds, tb, rows, wave, lane);
        __builtin_amdgcn_s_waitcnt(0);
        // Yield briefly while the previous tile DMA drains.
        asm volatile("s_sleep 1");
        __syncthreads();
        zero_tail_tiles(klds, vlds, rows, threadIdx.x);
        __syncthreads();

        floatx4 acc_lo = {0.f, 0.f, 0.f, 0.f};
        floatx4 acc_hi = {0.f, 0.f, 0.f, 0.f};
        #pragma unroll
        for (int k_block = 0; k_block < D_QK; k_block += 32) {
            long a_frag = static_cast<long>(pack_q_mfma(q8, hb, lane, k_block));
            long b_lo = static_cast<long>(load_k_frag_kv2(klds, lane, 0, k_block));
            long b_hi = static_cast<long>(load_k_frag_kv2(klds, lane, 16, k_block));
            acc_lo = mfma_fp8_16x16x32(a_frag, b_lo, acc_lo);
            acc_hi = mfma_fp8_16x16x32(a_frag, b_hi, acc_hi);
        }

        acc_lo.x *= score_scale;
        acc_lo.y *= score_scale;
        acc_lo.z *= score_scale;
        acc_lo.w *= score_scale;
        acc_hi.x *= score_scale;
        acc_hi.y *= score_scale;
        acc_hi.z *= score_scale;
        acc_hi.w *= score_scale;

        const int rows_lo = rows < 16 ? rows : 16;
        for (int r = 0; r < rows_lo; ++r) {
            float s0 = __shfl(acc_lo.x, r);
            float s1 = __shfl(acc_lo.y, r);
            float s2 = __shfl(acc_lo.z, r);
            float s3 = __shfl(acc_lo.w, r);

            float al[4], bt[4];
            sm_upd(st[0], s0, al[0], bt[0]);
            sm_upd(st[1], s1, al[1], bt[1]);
            sm_upd(st[2], s2, al[2], bt[2]);
            sm_upd(st[3], s3, al[3], bt[3]);

            #pragma unroll
            for (int j = 0; j < V_ELEMS; ++j) {
                int vd = lane + j * WAVE_SIZE;
                float vv = lut[vlds[r * D_V + vd]];
                a0[j] = al[0] * a0[j] + bt[0] * vv;
                a1[j] = al[1] * a1[j] + bt[1] * vv;
                a2[j] = al[2] * a2[j] + bt[2] * vv;
                a3[j] = al[3] * a3[j] + bt[3] * vv;
            }
        }

        for (int r = 16; r < rows; ++r) {
            int rr = r - 16;
            float s0 = __shfl(acc_hi.x, rr);
            float s1 = __shfl(acc_hi.y, rr);
            float s2 = __shfl(acc_hi.z, rr);
            float s3 = __shfl(acc_hi.w, rr);

            float al[4], bt[4];
            sm_upd(st[0], s0, al[0], bt[0]);
            sm_upd(st[1], s1, al[1], bt[1]);
            sm_upd(st[2], s2, al[2], bt[2]);
            sm_upd(st[3], s3, al[3], bt[3]);

            #pragma unroll
            for (int j = 0; j < V_ELEMS; ++j) {
                int vd = lane + j * WAVE_SIZE;
                float vv = lut[vlds[r * D_V + vd]];
                a0[j] = al[0] * a0[j] + bt[0] * vv;
                a1[j] = al[1] * a1[j] + bt[1] * vv;
                a2[j] = al[2] * a2[j] + bt[2] * vv;
                a3[j] = al[3] * a3[j] + bt[3] * vv;
            }
        }
        __syncthreads();
    }
}

static constexpr int DUET_BLOCK_THREADS = 512;
static constexpr int DUET_WAVES_PER_BLOCK = DUET_BLOCK_THREADS / WAVE_SIZE;
static constexpr int DUET_D_SLICE = D_V / DUET_WAVES_PER_BLOCK;
static constexpr int DUET_PRODUCER_WAVES = 2;
static constexpr int NONPRODUCER_WAVES = DUET_WAVES_PER_BLOCK - DUET_PRODUCER_WAVES;
static constexpr int SCORE_K_BLOCKS = D_QK / 32;

// Parallel Q staging: all 512 threads (32 per head, wave32 reduction)
__device__ inline void stage_q_fp8_per_head(
    const bf16* __restrict__ q,
    uint8_t* q8,
    float* q_scales,
    int qi,
    int wave,
    int lane)
{
    // 8 waves × 2 heads/wave = 16 heads, 32 threads per head
    const int h = wave * 2 + (lane >> 5);
    const int lt = lane & 31;

    // Parallel max reduction: 32 threads each read 18 elements
    float local_max = 0.0f;
    #pragma unroll 4
    for (int d = lt; d < D_QK; d += 32) {
        local_max = fmaxf(local_max, fabsf(to_f(q[q_off(qi, h, d)])));
    }
    float max_abs = wave32_max(local_max);

    const float q_scale = fmaxf(max_abs / static_cast<float>(FP8_MAX),
                                 1.0f / static_cast<float>(FP8_MAX));
    const float inv_q_scale = 1.0f / q_scale;

    if (lt == 0) {
        q_scales[h] = q_scale;
    }

    // Parallel quantize: 4 bytes at a time, 32 threads covering 576 elements
    #pragma unroll 5
    for (int d = lt * 4; d < D_QK; d += 128) {
        const uint32_t packed = cvt_fp8x4(
            to_f(q[q_off(qi, h, d + 0)]) * inv_q_scale,
            to_f(q[q_off(qi, h, d + 1)]) * inv_q_scale,
            to_f(q[q_off(qi, h, d + 2)]) * inv_q_scale,
            to_f(q[q_off(qi, h, d + 3)]) * inv_q_scale);
        *reinterpret_cast<uint32_t*>(&q8[h * D_QK + d]) = packed;
    }
}







// Extract V from K-LDS → VT-LDS during softmax (no LDS contention with QK MFMA)
// Only called by non-producer waves while producers compute softmax
__device__ inline void extract_v_from_klds_softmax(
    const uint8_t* __restrict__ klds,
    uint8_t* vt,
    int rows,
    int wave,
    int lane)
{
    const int local_wave = wave - DUET_PRODUCER_WAVES;
    if (local_wave < 0 || local_wave >= NONPRODUCER_WAVES) return;

    const int worker_id = local_wave * WAVE_SIZE + lane;
    const int num_workers = NONPRODUCER_WAVES * WAVE_SIZE;

    // Row-major VT with group padding: vt[slice*544 + group*136 + tok_in_group*16 + d]
    // Contiguous 16-byte writes per token (2x uint64_t), padding breaks bank conflicts
    for (int wu = worker_id; wu < TILE_ROWS * VT_NUM_SLICES; wu += num_workers) {
        const int tok = wu / VT_NUM_SLICES;
        const int slice = wu % VT_NUM_SLICES;
        const int d_base = slice * VT_DV_SLICE;
        const int tok_group = tok >> 3;
        const int tok_in_group = tok & 7;
        uint8_t* dst = vt + slice * VT_SLICE_BYTES + tok_group * VT_GROUP_STRIDE
                      + tok_in_group * VT_DV_SLICE;

        if (tok >= rows) {
            *reinterpret_cast<uint64_t*>(dst) = 0;
            *reinterpret_cast<uint64_t*>(dst + 8) = 0;
        } else {
            const int src_off = kv2_lds_offset(tok, d_base);
            *reinterpret_cast<uint64_t*>(dst) =
                *reinterpret_cast<const uint64_t*>(klds + src_off);
            *reinterpret_cast<uint64_t*>(dst + 8) =
                *reinterpret_cast<const uint64_t*>(klds + src_off + 8);
        }
    }
}

// Extract V from K-LDS → VT buffer for a batch of 16 VT slices
// slice_start: global VT slice to start from (0 for batch 0, 16 for batch 1)
// Output always written to local positions 0-15 in vt buffer
__device__ inline void extract_v_from_klds_batch(
    const uint8_t* __restrict__ klds,
    uint8_t* vt,
    int rows,
    int worker_id,
    int num_workers,
    int slice_start)
{
    constexpr int BATCH_SLICES = 16;
    for (int wu = worker_id; wu < TILE_ROWS * BATCH_SLICES; wu += num_workers) {
        const int tok = wu / BATCH_SLICES;
        const int local_slice = wu % BATCH_SLICES;
        const int global_slice = slice_start + local_slice;
        const int d_base = global_slice * VT_DV_SLICE;
        const int tok_group = tok >> 3;
        const int tok_in_group = tok & 7;
        uint8_t* dst = vt + local_slice * VT_SLICE_BYTES + tok_group * VT_GROUP_STRIDE
                      + tok_in_group * VT_DV_SLICE;

        if (tok >= rows) {
            *reinterpret_cast<uint64_t*>(dst) = 0;
            *reinterpret_cast<uint64_t*>(dst + 8) = 0;
        } else {
            const int src_off = kv2_lds_offset(tok, d_base);
            *reinterpret_cast<uint64_t*>(dst) =
                *reinterpret_cast<const uint64_t*>(klds + src_off);
            *reinterpret_cast<uint64_t*>(dst + 8) =
                *reinterpret_cast<const uint64_t*>(klds + src_off + 8);
        }
    }
}

__device__ __forceinline__ uint64_t load_v_frag_tr8(
    const uint8_t* vt,
    int vt_slice,
    int lane)
{
    // Row-major VT with group padding: 8 byte reads at stride 16
    const int col = lane & 15;
    const int group = lane >> 4;
    const uint8_t* base = vt + vt_slice * VT_SLICE_BYTES + group * VT_GROUP_STRIDE;
    uint64_t result;
    uint8_t* rb = reinterpret_cast<uint8_t*>(&result);
    rb[0] = base[0 * 16 + col];
    rb[1] = base[1 * 16 + col];
    rb[2] = base[2 * 16 + col];
    rb[3] = base[3 * 16 + col];
    rb[4] = base[4 * 16 + col];
    rb[5] = base[5 * 16 + col];
    rb[6] = base[6 * 16 + col];
    rb[7] = base[7 * 16 + col];
    return result;
}

// Read V directly from klds using dword loads + v_perm_b32 for fast byte packing
// Dword loads enable 4-lane broadcast (1 LDS cycle vs 4 for byte reads)
// Total: 8 ds_read_b32 + ~10 ALU = ~26 cycles (fits in 64-cycle MFMA window)
__device__ __forceinline__ uint64_t load_v_from_klds_bytes(
    const uint8_t* klds, int vt_slice, int lane)
{
    const int col = lane & 15;
    const int group = lane >> 4;
    const int d = vt_slice * VT_DV_SLICE + col;
    const int group_base = (group >= 2 ? 128 : 0)
                         + ((group & 1) ? 4 * K_NUM_BYTES_PER_SUBBLOCK : 0);
    const uint8_t* rb = klds + (d >> 6) * K_NUM_BYTES_PER_BLOCK + group_base + (d & 63);
    const uint8_t b0 = rb[0];
    const uint8_t b1 = rb[K_NUM_BYTES_PER_ROW];
    const uint8_t b2 = rb[K_NUM_BYTES_PER_SUBBLOCK];
    const uint8_t b3 = rb[K_NUM_BYTES_PER_SUBBLOCK + K_NUM_BYTES_PER_ROW];
    const uint8_t b4 = rb[2 * K_NUM_BYTES_PER_SUBBLOCK];
    const uint8_t b5 = rb[2 * K_NUM_BYTES_PER_SUBBLOCK + K_NUM_BYTES_PER_ROW];
    const uint8_t b6 = rb[3 * K_NUM_BYTES_PER_SUBBLOCK];
    const uint8_t b7 = rb[3 * K_NUM_BYTES_PER_SUBBLOCK + K_NUM_BYTES_PER_ROW];
    return (uint64_t)b0 | ((uint64_t)b1 << 8) | ((uint64_t)b2 << 16) | ((uint64_t)b3 << 24)
         | ((uint64_t)b4 << 32) | ((uint64_t)b5 << 40) | ((uint64_t)b6 << 48) | ((uint64_t)b7 << 56);
}

// Load 128-col K fragment from K-LDS for scaled MFMA 16x16x128
// Returns intx8 (32 bytes) for the B operand of mfma_scale_f32_16x16x128
__device__ __forceinline__ intx8 load_k128_frag(
    const uint8_t* klds, int lane, int tok_base, int mfma_blk128)
{
    const int _row = lane & 15;
    const int grp = lane >> 4;
    const int _rphy = (_row / 2) * 4 + (_row % 2);
    const int half = tok_base / 16;
    const int k_block = mfma_blk128 * 2 + grp / 2;
    const int col_in_kblock = (grp & 1) * 32;
    const uint8_t* base = klds
        + k_block * K_NUM_BYTES_PER_BLOCK
        + half * 128
        + (_rphy / 4) * K_NUM_BYTES_PER_SUBBLOCK
        + (_rphy % 4) * K_NUM_BYTES_PER_ROW
        + col_in_kblock;
    intx8 r;
    reinterpret_cast<uint64_t*>(&r)[0] = *reinterpret_cast<const uint64_t*>(base);
    reinterpret_cast<uint64_t*>(&r)[1] = *reinterpret_cast<const uint64_t*>(base + 8);
    reinterpret_cast<uint64_t*>(&r)[2] = *reinterpret_cast<const uint64_t*>(base + 16);
    reinterpret_cast<uint64_t*>(&r)[3] = *reinterpret_cast<const uint64_t*>(base + 24);
    return r;
}

// Load V fragment directly from K-LDS using transposed read (no separate vt buffer)
__device__ __forceinline__ uint64_t load_v_from_klds_tr8(
    const uint8_t* klds, int vt_slice, int lane)
{
    const int lc = lane & 15;
    const int lg = lane >> 4;
    const int row = lg * 8 + (lc >> 1);
    const int d = vt_slice * VT_DV_SLICE + (lc & 1) * 8;
    const int block = d >> 6;
    const int d_in_block = d & 63;
    const int half = row >> 4;
    const int row16 = row & 15;
    const int row_phy = (row16 >> 1) * 4 + (row16 & 1);
    const int off = block * K_NUM_BYTES_PER_BLOCK
                  + half * 128
                  + (row_phy >> 2) * K_NUM_BYTES_PER_SUBBLOCK
                  + (row_phy & 3) * K_NUM_BYTES_PER_ROW
                  + d_in_block;
    return ds_read_tr8_u64(klds + off);
}

template <int COL_OFFSET>
__device__ __forceinline__ void direct_load_k_block_duet(
    intx4 srsrc,
    uintptr_t p_lds_k_warp_base,
    int row,
    int col_base)
{
    constexpr int k_block_idx = COL_OFFSET / K_COL_BLOCK;
    constexpr uintptr_t k_lds_block_base =
        k_block_idx * K_NUM_BYTES_PER_BLOCK - COL_OFFSET;
    const int voffset = row * D_QK + col_base;
    llvm_amdgcn_raw_buffer_load_lds(
        srsrc,
        make_wave_lds_ptr(p_lds_k_warp_base + k_lds_block_base),
        4,
        voffset,
        0,
        COL_OFFSET,
        0);
}

__device__ inline void stage_k_tile_kv2_duet(
    intx4 srsrc,
    uint8_t* klds,
    int base_token,
    int rows,
    int wave,
    int lane)
{
    const int col_base = (lane & 15) * 4;
    const int warp_idx = wave;
    const int row_base = (lane / 32) * 16 + ((lane / 16) & 1) + warp_idx * 2;
    const int row = (row_base < rows) ? (base_token + row_base) : -1;
    if (row < 0) return;
    const uintptr_t p_lds_k_warp_base =
        reinterpret_cast<uintptr_t>(klds)
        + warp_idx * K_NUM_BYTES_PER_SUBBLOCK;
    direct_load_k_block_duet<0>(srsrc, p_lds_k_warp_base, row, col_base);
    direct_load_k_block_duet<64>(srsrc, p_lds_k_warp_base, row, col_base);
    direct_load_k_block_duet<128>(srsrc, p_lds_k_warp_base, row, col_base);
    direct_load_k_block_duet<192>(srsrc, p_lds_k_warp_base, row, col_base);
    direct_load_k_block_duet<256>(srsrc, p_lds_k_warp_base, row, col_base);
    direct_load_k_block_duet<320>(srsrc, p_lds_k_warp_base, row, col_base);
    direct_load_k_block_duet<384>(srsrc, p_lds_k_warp_base, row, col_base);
    direct_load_k_block_duet<448>(srsrc, p_lds_k_warp_base, row, col_base);
    direct_load_k_block_duet<512>(srsrc, p_lds_k_warp_base, row, col_base);
}

// Coalesced VMEM K-tile loader: HBM -> registers -> LDS
// Each wave reads one row at a time, all 64 lanes access same row at consecutive 16B chunks.
// D_QK=576, 576/16=36 chunks per row. Lanes 0-35 active, 36-63 idle.
// Ensures coalesced HBM access (all active lanes within one cache line region).
__device__ inline void vmem_load_k_tile(
    const uint8_t* __restrict__ kv, uint8_t* __restrict__ klds,
    int base_token, int rows, int local_wave, int num_waves, int lane)
{
    for (int row = local_wave; row < TILE_ROWS; row += num_waves) {
        const bool valid = (row < rows);
        const int hbm_base = valid ? ((base_token + row) * D_QK) : 0;
        const int d = lane << 4;  // lane * 16: lanes 0-35 cover d=0,16,...,560

        if (d < D_QK) {
            U128 data = {};
            if (valid) {
                data = *reinterpret_cast<const U128*>(&kv[hbm_base + d]);
            }
            const int lds_off = kv2_lds_offset(row, d);
            // Two 8B LDS writes for 16B total
            *reinterpret_cast<uint64_t*>(&klds[lds_off]) =
                static_cast<uint64_t>(data.x0) | (static_cast<uint64_t>(data.x1) << 32);
            *reinterpret_cast<uint64_t*>(&klds[lds_off + 8]) =
                static_cast<uint64_t>(data.x2) | (static_cast<uint64_t>(data.x3) << 32);
        }
    }
}

__device__ inline void zero_tail_k_tile(uint8_t* klds, int rows, int tid)
{
    if (rows >= TILE_ROWS) {
        return;
    }

    for (int row = rows; row < TILE_ROWS; ++row) {
        for (int d = tid; d < D_QK; d += DUET_BLOCK_THREADS) {
            const int block = d / K_COL_BLOCK;
            const int d_in_block = d % K_COL_BLOCK;
            const int half = row / 16;
            const int row16 = row % 16;
            const int row_phy = (row16 / 2) * 4 + (row16 % 2);
            const int offs = block * K_NUM_BYTES_PER_BLOCK +
                             half * 128 +
                             (row_phy / 4) * K_NUM_BYTES_PER_SUBBLOCK +
                             (row_phy % 4) * K_NUM_BYTES_PER_ROW +
                             d_in_block;
            klds[offs] = 0;
        }
    }
}

__device__ __forceinline__ uint64_t pack_q_mfma_full(
    const uint8_t* q8,
    int lane,
    int k_block)
{
    const int row = lane & 15;
    const int k_base = k_block + ((lane >> 4) * 8);
    return *reinterpret_cast<const uint64_t*>(&q8[lq8(row, k_base)]);
}






// ===================== light + inline Q quantization v2 ==================
// Takes bf16 Q, quantizes to FP8 into LDS at kernel start (all waves cooperate).
// Then loads Q FP8 from LDS into registers (same register footprint as mla_s1_light).
// Eliminates the separate Q quantization kernel launch (~3us savings on ranked).
template <bool DIRECT_OUT, bool FUSE_S2>
__global__ __launch_bounds__(DUET_BLOCK_THREADS, 2)
void mla_s1_light_iq2(
    const uint16_t* __restrict__ q_bf16,
    const uint8_t* __restrict__ kv,
    const float* __restrict__ kv_scale_ptr,
    const int32_t* __restrict__ qo,
    const int32_t* __restrict__ kvi,
    float* __restrict__ pm,
    float* __restrict__ pl,
    bf16* __restrict__ po,
    bf16* __restrict__ out,
    int bs,
    int ns)
{
    const float kv_scale = *kv_scale_ptr;
    const intx4 srsrc = make_srsrc(kv, 0xffffffffu);
    const int bid = blockIdx.x;
    const int b = DIRECT_OUT ? bid : (bid / ns);
    const int sid = DIRECT_OUT ? 0 : (bid % ns);
    if (b >= bs) return;

    const int tid = threadIdx.x;
    const int wave = tid >> 6;
    const int lane = tid & 63;
    const int lane_col = lane & 15;
    const int lane_group = lane >> 4;
    const int row_base = lane_group * 4;

    const int qs = qo[b];
    const int qe = qo[b + 1];
    if (qe - qs != 1) return;
    const int qi = qs;
    const int kvs = kvi[b];
    const int kve = kvi[b + 1];
    const int kvlen = kve - kvs;

    // Single-buffered K tiles for occ=2 (LDS < 32KB)
    __shared__ uint8_t klds_buf[1][K_TILE_BYTES];
    __shared__ uint8_t p8[NUM_HEADS * TILE_ROWS];
    __shared__ float half_max[2][NUM_HEADS];
    __shared__ float half_sum[2][NUM_HEADS];
    __shared__ uint8_t q_fp8_lds[NUM_HEADS * D_QK + 128];
    float* q_sc_lds = reinterpret_cast<float*>(&q_fp8_lds[NUM_HEADS * D_QK]);
    __shared__ int fused_last;
    int* __done_count = nullptr;
    if constexpr (FUSE_S2) {
        __done_count = reinterpret_cast<int*>(pm + bs * NUM_HEADS * ns);
    }

    float local_max[4] = {NEG_INF, NEG_INF, NEG_INF, NEG_INF};
    float local_sum[4] = {0.0f, 0.0f, 0.0f, 0.0f};
    const bool producer = wave < DUET_PRODUCER_WAVES;

    const int ss = DIRECT_OUT ? kvs : (kvs + (kvlen * sid) / ns);
    const int se = DIRECT_OUT ? kve : (kvs + (kvlen * (sid + 1)) / ns);

    if (se <= ss) {
        if constexpr (DIRECT_OUT) {
            for (int idx = tid; idx < NUM_HEADS * D_V; idx += DUET_BLOCK_THREADS) {
                out[out_off(qi, idx / D_V, idx % D_V)] = to_b(0.0f);
            }
            return;
        }
        if (tid < NUM_HEADS) {
            pm[pml_off(b, sid, tid, ns)] = NEG_INF;
            pl[pml_off(b, sid, tid, ns)] = 0.0f;
        }
        for (int idx = tid; idx < NUM_HEADS * D_V; idx += DUET_BLOCK_THREADS) {
            po[po_off(b, sid, idx / D_V, idx % D_V, ns)] = to_b(0.0f);
        }
        if constexpr (!FUSE_S2) return;
        // Release: drain all stores before signaling
        asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
        if (tid == 0) fused_last = atomicAdd(&__done_count[b], 1);
        __syncthreads();
        if (fused_last != ns - 1) return;
        // Acquire: only last CTA needs full fence to see other CTAs' stores
        __threadfence();
        goto s2_tail_iq2;
    }
    { // scope: all vars here invisible at s2_tail_iq2, so goto is valid
    // Start first tile DMA EARLY — overlaps with Q quantization below
    const int first_rows = min(TILE_ROWS, se - ss);
    stage_k_tile_kv2_duet(srsrc, klds_buf[0], ss, first_rows, wave, lane);

    // ---- Inline Q quantization: all 8 waves, 2 heads per wave ----
    // Runs concurrently with buffer_load_lds DMA above (different LDS regions)
    {
        const int hh0 = wave * 2;
        #pragma unroll
        for (int hoff = 0; hoff < 2; ++hoff) {
            const int h = hh0 + hoff;
            const int gbase = (qi * NUM_HEADS + h) * D_QK;
            float la = 0.0f;
            float vals[9];
            #pragma unroll
            for (int i = 0; i < 9; ++i) {
                float fv = __uint_as_float((uint32_t)q_bf16[gbase + lane + i * 64] << 16);
                vals[i] = fv;
                la = fmaxf(la, fabsf(fv));
            }
            float amax = wbcast(wreduce_max(la));
            amax = fmaxf(amax, 1.0f / static_cast<float>(FP8_MAX));
            float inv_s = static_cast<float>(FP8_MAX) / amax;
            if (lane == 0) q_sc_lds[h] = amax / static_cast<float>(FP8_MAX);
            #pragma unroll
            for (int i = 0; i < 4; ++i) {
                float a = fminf(fmaxf(vals[i*2] * inv_s, -static_cast<float>(FP8_MAX)), static_cast<float>(FP8_MAX));
                float bb = fminf(fmaxf(vals[i*2+1] * inv_s, -static_cast<float>(FP8_MAX)), static_cast<float>(FP8_MAX));
                uint32_t pk = 0;
                pk = __builtin_amdgcn_cvt_pk_fp8_f32(a, bb, pk, 0);
                q_fp8_lds[h * D_QK + lane + (i*2) * 64] = static_cast<uint8_t>(pk & 0xFFu);
                q_fp8_lds[h * D_QK + lane + (i*2+1) * 64] = static_cast<uint8_t>((pk >> 8) & 0xFFu);
            }
            {
                float a = fminf(fmaxf(vals[8] * inv_s, -static_cast<float>(FP8_MAX)), static_cast<float>(FP8_MAX));
                uint32_t pk = 0;
                pk = __builtin_amdgcn_cvt_pk_fp8_f32(a, 0.0f, pk, 0);
                q_fp8_lds[h * D_QK + lane + 8 * 64] = static_cast<uint8_t>(pk & 0xFFu);
            }
        }
    }
    __syncthreads();  // barrier0: Q fp8 in LDS ready

    {
    // First tile DMA already in flight from above — no need to start it here

    const int slice_base = wave * DUET_D_SLICE;
    const uint8_t* q_row = producer ? (q_fp8_lds + (lane & 15) * D_QK) : nullptr;
    const int q_grp128 = lane >> 4;
    const int q_grp32 = q_grp128 * 8;
    const int q_grp128_col = q_grp128 * 32;
    const float score_scale_mul = kv_scale * SM_SCALE;
    intx8 q0_cached = {};
    intx8 q1_cached = {};
    intx8 q2_cached = {};
    intx8 q3_cached = {};
    float scale0_cached = 0.0f;
    float scale1_cached = 0.0f;
    float scale2_cached = 0.0f;
    float scale3_cached = 0.0f;
    if (producer) {
        q0_cached = *reinterpret_cast<const intx8*>(&q_row[q_grp128_col + 0]);
        q1_cached = *reinterpret_cast<const intx8*>(&q_row[q_grp128_col + 128]);
        q2_cached = *reinterpret_cast<const intx8*>(&q_row[q_grp128_col + 256]);
        q3_cached = *reinterpret_cast<const intx8*>(&q_row[q_grp128_col + 384]);
        scale0_cached = q_sc_lds[row_base + 0] * score_scale_mul;
        scale1_cached = q_sc_lds[row_base + 1] * score_scale_mul;
        scale2_cached = q_sc_lds[row_base + 2] * score_scale_mul;
        scale3_cached = q_sc_lds[row_base + 3] * score_scale_mul;
    }

    floatx4 o_acc[4];
    #pragma unroll
    for (int i = 0; i < 4; ++i) o_acc[i] = {0.0f, 0.0f, 0.0f, 0.0f};
    uint64_t v_pre[4];

    for (int tb = ss; tb < se; tb += TILE_ROWS) {
        const int rows = min(TILE_ROWS, se - tb);
        const bool full_tile = rows == TILE_ROWS;
        const bool init_tile = tb == ss;
        uint8_t* klds = klds_buf[0];
        const int next_tb = tb + TILE_ROWS;
        const bool has_next = next_tb < se;
        const int next_rows = has_next ? min(TILE_ROWS, se - next_tb) : 0;

        __builtin_amdgcn_s_waitcnt(0);
        __syncthreads();

        floatx4 score = {0.0f, 0.0f, 0.0f, 0.0f};
        const int tok_base = wave * 16;
        const int rows_half = tok_base < rows ? min(16, rows - tok_base) : 0;

        if (producer) {
            asm volatile("s_setprio 1");
            // Compiler-managed QK scoring: no hardcoded VGPR clobbers.
            floatx4 sa = {0.0f, 0.0f, 0.0f, 0.0f};
            floatx4 sb = {0.0f, 0.0f, 0.0f, 0.0f};
            const uint64_t qr0 = *reinterpret_cast<const uint64_t*>(&q_row[512 + q_grp32]);
            const uint64_t qr1 = *reinterpret_cast<const uint64_t*>(&q_row[544 + q_grp32]);
            intx8 ka = load_k128_frag(klds, lane, tok_base, 0);
            intx8 kb_data = load_k128_frag(klds, lane, tok_base, 1);
            sa = mfma_scale_fp8_16x16x128_noscale(q0_cached, ka, sa);
            sb = mfma_scale_fp8_16x16x128_noscale(q1_cached, kb_data, sb);
            intx8 ka2 = load_k128_frag(klds, lane, tok_base, 2);
            intx8 kb_data2 = load_k128_frag(klds, lane, tok_base, 3);
            sa = mfma_scale_fp8_16x16x128_noscale(q2_cached, ka2, sa);
            sb = mfma_scale_fp8_16x16x128_noscale(q3_cached, kb_data2, sb);
            #pragma unroll
            for (int cb = 0; cb < 4; ++cb) {
                v_pre[cb] = load_v_from_klds_tr8(klds,
                    wave * (DUET_D_SLICE / VT_DV_SLICE) + cb, lane);
            }
            uint64_t kf0 = load_k_frag_kv2(klds, lane, tok_base, 512);
            uint64_t kf1 = load_k_frag_kv2(klds, lane, tok_base, 544);
            sa = mfma_fp8_16x16x32(static_cast<long>(qr0), static_cast<long>(kf0), sa);
            sb = mfma_fp8_16x16x32(static_cast<long>(qr1), static_cast<long>(kf1), sb);
            score.x = sa.x + sb.x; score.y = sa.y + sb.y;
            score.z = sa.z + sb.z; score.w = sa.w + sb.w;
            asm volatile("s_setprio 0");

            {
                const float scale0 = scale0_cached;
                const float scale1 = scale1_cached;
                const float scale2 = scale2_cached;
                const float scale3 = scale3_cached;
                float s0, s1, s2, s3;
                if (full_tile) {
                    s0 = score[0] * scale0;
                    s1 = score[1] * scale1;
                    s2 = score[2] * scale2;
                    s3 = score[3] * scale3;
                } else {
                    s0 = lane_col < rows_half ? score[0] * scale0 : NEG_INF;
                    s1 = lane_col < rows_half ? score[1] * scale1 : NEG_INF;
                    s2 = lane_col < rows_half ? score[2] * scale2 : NEG_INF;
                    s3 = lane_col < rows_half ? score[3] * scale3 : NEG_INF;
                }
                wave16_max4(s0, s1, s2, s3);
                if (lane_col == 0) {
                    half_max[wave][row_base + 0] = s0;
                    half_max[wave][row_base + 1] = s1;
                    half_max[wave][row_base + 2] = s2;
                    half_max[wave][row_base + 3] = s3;
                }
            }
        } else {
            #pragma unroll
            for (int cb = 0; cb < 4; ++cb) {
                v_pre[cb] = load_v_from_klds_tr8(klds,
                    wave * (DUET_D_SLICE / VT_DV_SLICE) + cb, lane);
            }
        }
        if (has_next) {
            stage_k_tile_kv2_duet(srsrc, klds_buf[0],
                next_tb, next_rows, wave, lane);
        }

        __syncthreads();  // barrier2: half_max ready, V loaded

        // All waves: compute alpha + rescale o_acc (overlap with producer softmax)
        // Pre-load all half_max values to allow LDS read pipelining
        // (avoids serial read-wait-compute chains)
        float hm0 = half_max[0][row_base + 0];
        float hm1 = half_max[1][row_base + 0];
        float hm2 = half_max[0][row_base + 1];
        float hm3 = half_max[1][row_base + 1];
        float hm4 = half_max[0][row_base + 2];
        float hm5 = half_max[1][row_base + 2];
        float hm6 = half_max[0][row_base + 3];
        float hm7 = half_max[1][row_base + 3];
        // Compiler fence: force all 8 LDS reads to issue before continuing.
        // Without this, compiler sinks reads near uses → 4 serial waits (~160 cy).
        // With fence: 1 batched wait (~43 cy), saving ~117 cycles per tile.
        asm volatile("" : "+v"(hm0), "+v"(hm1), "+v"(hm2), "+v"(hm3),
                          "+v"(hm4), "+v"(hm5), "+v"(hm6), "+v"(hm7));

        float alpha[4];
        if (init_tile) {
            local_max[0] = fmaxf(hm0, hm1);
            local_max[1] = fmaxf(hm2, hm3);
            local_max[2] = fmaxf(hm4, hm5);
            local_max[3] = fmaxf(hm6, hm7);
            alpha[0] = 0.0f;
            alpha[1] = 0.0f;
            alpha[2] = 0.0f;
            alpha[3] = 0.0f;
        } else if (full_tile) {
            float new_m0 = fmaxf(local_max[0], fmaxf(hm0, hm1));
            alpha[0] = fast_exp(local_max[0] - new_m0);
            local_max[0] = new_m0;
            float new_m1 = fmaxf(local_max[1], fmaxf(hm2, hm3));
            alpha[1] = fast_exp(local_max[1] - new_m1);
            local_max[1] = new_m1;
            float new_m2 = fmaxf(local_max[2], fmaxf(hm4, hm5));
            alpha[2] = fast_exp(local_max[2] - new_m2);
            local_max[2] = new_m2;
            float new_m3 = fmaxf(local_max[3], fmaxf(hm6, hm7));
            alpha[3] = fast_exp(local_max[3] - new_m3);
            local_max[3] = new_m3;
        } else {
            float new_m0 = fmaxf(local_max[0], fmaxf(hm0, hm1));
            alpha[0] = local_sum[0] > 0.0f ? fast_exp(local_max[0] - new_m0) : 0.0f;
            local_max[0] = new_m0;
            float new_m1 = fmaxf(local_max[1], fmaxf(hm2, hm3));
            alpha[1] = local_sum[1] > 0.0f ? fast_exp(local_max[1] - new_m1) : 0.0f;
            local_max[1] = new_m1;
            float new_m2 = fmaxf(local_max[2], fmaxf(hm4, hm5));
            alpha[2] = local_sum[2] > 0.0f ? fast_exp(local_max[2] - new_m2) : 0.0f;
            local_max[2] = new_m2;
            float new_m3 = fmaxf(local_max[3], fmaxf(hm6, hm7));
            alpha[3] = local_sum[3] > 0.0f ? fast_exp(local_max[3] - new_m3) : 0.0f;
            local_max[3] = new_m3;
        }

        #pragma unroll
        for (int i = 0; i < 4; ++i) {
            o_acc[i].x *= alpha[0]; o_acc[i].y *= alpha[1];
            o_acc[i].z *= alpha[2]; o_acc[i].w *= alpha[3];
        }

        if (producer) {
            const float scale0 = scale0_cached;
            const float scale1 = scale1_cached;
            const float scale2 = scale2_cached;
            const float scale3 = scale3_cached;
            if (full_tile) {
                const int p_col = tok_base + lane_col;
                const float s0 = score[0] * scale0;
                const float s1 = score[1] * scale1;
                const float s2 = score[2] * scale2;
                const float s3 = score[3] * scale3;
                float p0 = fast_exp(s0 - local_max[0]);
                float p1 = fast_exp(s1 - local_max[1]);
                float p2 = fast_exp(s2 - local_max[2]);
                float p3 = fast_exp(s3 - local_max[3]);
                const uint32_t p_pack = cvt_fp8x4(p0, p1, p2, p3);
                p8[(row_base + 0) * TILE_ROWS + p_col] = static_cast<uint8_t>(p_pack & 0xffu);
                p8[(row_base + 1) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 8) & 0xffu);
                p8[(row_base + 2) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 16) & 0xffu);
                p8[(row_base + 3) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 24) & 0xffu);
                wave16_sum4(p0, p1, p2, p3);
                if (lane_col == 0) {
                    half_sum[wave][row_base + 0] = p0;
                    half_sum[wave][row_base + 1] = p1;
                    half_sum[wave][row_base + 2] = p2;
                    half_sum[wave][row_base + 3] = p3;
                }
            } else {
                const bool valid = lane_col < rows_half;
                const int p_col = tok_base + lane_col;
                const float s0 = valid ? score[0] * scale0 : NEG_INF;
                const float s1 = valid ? score[1] * scale1 : NEG_INF;
                const float s2 = valid ? score[2] * scale2 : NEG_INF;
                const float s3 = valid ? score[3] * scale3 : NEG_INF;
                float p0 = valid ? fast_exp(s0 - local_max[0]) : 0.0f;
                float p1 = valid ? fast_exp(s1 - local_max[1]) : 0.0f;
                float p2 = valid ? fast_exp(s2 - local_max[2]) : 0.0f;
                float p3 = valid ? fast_exp(s3 - local_max[3]) : 0.0f;
                const uint32_t p_pack = cvt_fp8x4(p0, p1, p2, p3);
                p8[(row_base + 0) * TILE_ROWS + p_col] = static_cast<uint8_t>(p_pack & 0xffu);
                p8[(row_base + 1) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 8) & 0xffu);
                p8[(row_base + 2) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 16) & 0xffu);
                p8[(row_base + 3) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 24) & 0xffu);
                wave16_sum4(p0, p1, p2, p3);
                if (lane_col == 0) {
                    half_sum[wave][row_base + 0] = p0;
                    half_sum[wave][row_base + 1] = p1;
                    half_sum[wave][row_base + 2] = p2;
                    half_sum[wave][row_base + 3] = p3;
                }
            }
        }
        __syncthreads();  // barrier3: half_sum + p8 ready

        // Update local_sum (needs half_sum from barrier3)
        #pragma unroll
        for (int h = 0; h < 4; ++h) {
            const int row = row_base + h;
            local_sum[h] = alpha[h] * local_sum[h]
                + half_sum[0][row] + half_sum[1][row];
        }

        const uint64_t p_frag =
            *reinterpret_cast<const uint64_t*>(
                &p8[lane_col * TILE_ROWS + lane_group * 8]);

        #pragma unroll
        for (int col_block = 0; col_block < 4; ++col_block) {
            o_acc[col_block] = mfma_fp8_16x16x32(
                static_cast<long>(p_frag),
                static_cast<long>(v_pre[col_block]),
                o_acc[col_block]);
        }

    }

    if constexpr (!DIRECT_OUT) {
        if (wave == 0 && lane_col == 0) {
            #pragma unroll
            for (int h = 0; h < 4; ++h) {
                pm[pml_off(b, sid, row_base + h, ns)] = local_max[h];
                pl[pml_off(b, sid, row_base + h, ns)] = local_sum[h];
            }
        }
    }

    const float inv0 = local_sum[0] > 0.0f ? (kv_scale / local_sum[0]) : 0.0f;
    const float inv1 = local_sum[1] > 0.0f ? (kv_scale / local_sum[1]) : 0.0f;
    const float inv2 = local_sum[2] > 0.0f ? (kv_scale / local_sum[2]) : 0.0f;
    const float inv3 = local_sum[3] > 0.0f ? (kv_scale / local_sum[3]) : 0.0f;
    #pragma unroll
    for (int col_block = 0; col_block < 4; ++col_block) {
        const int dv = slice_base + col_block * 16 + lane_col;
        if constexpr (DIRECT_OUT) {
            out[out_off(qi, row_base + 0, dv)] = to_b(o_acc[col_block].x * inv0);
            out[out_off(qi, row_base + 1, dv)] = to_b(o_acc[col_block].y * inv1);
            out[out_off(qi, row_base + 2, dv)] = to_b(o_acc[col_block].z * inv2);
            out[out_off(qi, row_base + 3, dv)] = to_b(o_acc[col_block].w * inv3);
        } else {
            po[po_off(b, sid, row_base + 0, dv, ns)] = to_b(o_acc[col_block].x * inv0);
            po[po_off(b, sid, row_base + 1, dv, ns)] = to_b(o_acc[col_block].y * inv1);
            po[po_off(b, sid, row_base + 2, dv, ns)] = to_b(o_acc[col_block].z * inv2);
            po[po_off(b, sid, row_base + 3, dv, ns)] = to_b(o_acc[col_block].w * inv3);
        }
    }

    } // end inner scope (original block for main computation)

    } // end outer scope — first_rows, slice_base, o_acc etc dead; goto s2_tail_iq2 is valid

    if constexpr (DIRECT_OUT) return;
    if constexpr (!FUSE_S2) return;

    // ---- Fused S2: atomic completion + last-CTA reduction ----
    // Release: drain all stores before signaling (cheap — no cache invalidation)
    asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
    if (tid == 0) fused_last = atomicAdd(&__done_count[b], 1);
    __syncthreads();
    if (fused_last != ns - 1) return;
    // Acquire: only last CTA needs full fence to see other CTAs' stores
    __threadfence();

s2_tail_iq2:
    // S2 reduction: 512 threads = 16 heads × 32 threads
    // Each thread reduces ns splits for 16 contiguous V dims → 32×16 = 512 = D_V
    {
        const int s2_h  = tid >> 5;        // head index 0..15
        const int s2_lt = tid & 31;        // lane within head group 0..31
        const int vd_base = s2_lt * 16;    // 16 contiguous V dims per thread
        const int bh_ns = (b * NUM_HEADS + s2_h) * ns;

        float mm = NEG_INF;
        float ll = 0.0f;
        // 16 V-dim accumulators — reuse VGPRs freed from S1's o_acc/q_frag
        float r0=0.f, r1=0.f, r2=0.f, r3=0.f;
        float r4=0.f, r5=0.f, r6=0.f, r7=0.f;
        float r8=0.f, r9=0.f, r10=0.f, r11=0.f;
        float r12=0.f, r13=0.f, r14=0.f, r15=0.f;

        for (int s = 0; s < ns; ++s) {
            float m = pm[bh_ns + s];
            float l = pl[bh_ns + s];
            // Online softmax: merge split s into running accumulator
            float mn = fmaxf(mm, m);
            float a = fast_exp(mm - mn);
            float bw = fast_exp(m - mn) * l;
            ll = a * ll + bw;
            mm = mn;

            // Vectorized 128-bit loads: 2 × uint4 = 16 bf16 values (32 bytes)
            const bf16* pv = &po[bh_ns * D_V + s * D_V + vd_base];
            const uint4 lo4 = *reinterpret_cast<const uint4*>(pv);
            const uint4 hi4 = *reinterpret_cast<const uint4*>(pv + 8);
            // Unpack bf16 pairs from packed uint32 and FMA into accumulators
            // lo4.x = [bf16_1 | bf16_0], lo4.y = [bf16_3 | bf16_2], etc.
            r0  = __builtin_fmaf(a, r0,  bw * to_f(reinterpret_cast<const bf16*>(&lo4.x)[0]));
            r1  = __builtin_fmaf(a, r1,  bw * to_f(reinterpret_cast<const bf16*>(&lo4.x)[1]));
            r2  = __builtin_fmaf(a, r2,  bw * to_f(reinterpret_cast<const bf16*>(&lo4.y)[0]));
            r3  = __builtin_fmaf(a, r3,  bw * to_f(reinterpret_cast<const bf16*>(&lo4.y)[1]));
            r4  = __builtin_fmaf(a, r4,  bw * to_f(reinterpret_cast<const bf16*>(&lo4.z)[0]));
            r5  = __builtin_fmaf(a, r5,  bw * to_f(reinterpret_cast<const bf16*>(&lo4.z)[1]));
            r6  = __builtin_fmaf(a, r6,  bw * to_f(reinterpret_cast<const bf16*>(&lo4.w)[0]));
            r7  = __builtin_fmaf(a, r7,  bw * to_f(reinterpret_cast<const bf16*>(&lo4.w)[1]));
            r8  = __builtin_fmaf(a, r8,  bw * to_f(reinterpret_cast<const bf16*>(&hi4.x)[0]));
            r9  = __builtin_fmaf(a, r9,  bw * to_f(reinterpret_cast<const bf16*>(&hi4.x)[1]));
            r10 = __builtin_fmaf(a, r10, bw * to_f(reinterpret_cast<const bf16*>(&hi4.y)[0]));
            r11 = __builtin_fmaf(a, r11, bw * to_f(reinterpret_cast<const bf16*>(&hi4.y)[1]));
            r12 = __builtin_fmaf(a, r12, bw * to_f(reinterpret_cast<const bf16*>(&hi4.z)[0]));
            r13 = __builtin_fmaf(a, r13, bw * to_f(reinterpret_cast<const bf16*>(&hi4.z)[1]));
            r14 = __builtin_fmaf(a, r14, bw * to_f(reinterpret_cast<const bf16*>(&hi4.w)[0]));
            r15 = __builtin_fmaf(a, r15, bw * to_f(reinterpret_cast<const bf16*>(&hi4.w)[1]));
            }
        // Normalize and write final output
        float inv = ll > 0.0f ? (1.0f / ll) : 0.0f;
        bf16* odst = &out[out_off(qi, s2_h, vd_base)];
        // Vectorized 128-bit stores via uint4
        uint4 out_lo, out_hi;
        reinterpret_cast<bf16*>(&out_lo.x)[0] = to_b(r0 * inv);
        reinterpret_cast<bf16*>(&out_lo.x)[1] = to_b(r1 * inv);
        reinterpret_cast<bf16*>(&out_lo.y)[0] = to_b(r2 * inv);
        reinterpret_cast<bf16*>(&out_lo.y)[1] = to_b(r3 * inv);
        reinterpret_cast<bf16*>(&out_lo.z)[0] = to_b(r4 * inv);
        reinterpret_cast<bf16*>(&out_lo.z)[1] = to_b(r5 * inv);
        reinterpret_cast<bf16*>(&out_lo.w)[0] = to_b(r6 * inv);
        reinterpret_cast<bf16*>(&out_lo.w)[1] = to_b(r7 * inv);
        reinterpret_cast<bf16*>(&out_hi.x)[0] = to_b(r8 * inv);
        reinterpret_cast<bf16*>(&out_hi.x)[1] = to_b(r9 * inv);
        reinterpret_cast<bf16*>(&out_hi.y)[0] = to_b(r10 * inv);
        reinterpret_cast<bf16*>(&out_hi.y)[1] = to_b(r11 * inv);
        reinterpret_cast<bf16*>(&out_hi.z)[0] = to_b(r12 * inv);
        reinterpret_cast<bf16*>(&out_hi.z)[1] = to_b(r13 * inv);
        reinterpret_cast<bf16*>(&out_hi.w)[0] = to_b(r14 * inv);
        reinterpret_cast<bf16*>(&out_hi.w)[1] = to_b(r15 * inv);
        *reinterpret_cast<uint4*>(odst)     = out_lo;
        *reinterpret_cast<uint4*>(odst + 8) = out_hi;

        // Self-reset done_count for next invocation
        if (tid == 0) __done_count[b] = 0;
    }
}


// Fast S2: 1 head per block, contiguous V reads for vectorized loads.
// Grid: (bs, NUM_HEADS), Block: 64. Best for low-ns tails where launch cost matters.
__global__ __launch_bounds__(64, 16)
void mla_s2_fast(
    const float* __restrict__ pm,
    const float* __restrict__ pl,
    const bf16* __restrict__ po,
    bf16* __restrict__ out,
    const int32_t* __restrict__ qo,
    int bs,
    int ns)
{
    const int b = blockIdx.x;
    const int h = blockIdx.y;
    if (b >= bs) return;

    const int lane = threadIdx.x;
    const int qs = qo[b];
    const int qe = qo[b + 1];
    if (qe - qs != 1) return;
    const int qi = qs;

    const int vd_base = lane * V_ELEMS;
    float mm = NEG_INF;
    float ll = 0.0f;
    float r[V_ELEMS];
    #pragma unroll
    for (int j = 0; j < V_ELEMS; ++j) r[j] = 0.0f;

    const int po_bh = ((b * NUM_HEADS + h) * ns) * D_V;
    for (int s = 0; s < ns; ++s) {
        float m = pm[(b * NUM_HEADS + h) * ns + s];
        float l = pl[(b * NUM_HEADS + h) * ns + s];
        float mn = fmaxf(mm, m);
        float a = fast_exp(mm - mn);
        float bw = fast_exp(m - mn) * l;
        ll = a * ll + bw;
        mm = mn;

        const bf16* po_ptr = &po[po_bh + s * D_V + vd_base];
        #pragma unroll
        for (int j = 0; j < V_ELEMS; ++j) {
            r[j] = a * r[j] + bw * to_f(po_ptr[j]);
        }
    }

    float inv = ll > 0.0f ? (1.0f / ll) : 0.0f;
    #pragma unroll
    for (int j = 0; j < V_ELEMS; ++j) {
        out[out_off(qi, h, vd_base + j)] = to_b(r[j] * inv);
    }
}

__global__ __launch_bounds__(512, 2)
void mla_s2_vpar2(
    const float* __restrict__ pm,
    const float* __restrict__ pl,
    const bf16* __restrict__ po,
    bf16* __restrict__ out,
    const int32_t* __restrict__ qo,
    int bs,
    int ns)
{
    const int b = blockIdx.x;
    const int h = blockIdx.y;
    const int vb = blockIdx.z;
    if (b >= bs) return;

    const int tid = threadIdx.x;
    const int group = tid / 64;
    const int lane = tid % 64;
    const int num_groups = blockDim.x / 64;

    const int qs = qo[b];
    if (qo[b + 1] - qs != 1) return;
    const int qi = qs;

    const int vd_base = vb * 128 + lane * 2;
    const int bh_ns = (b * NUM_HEADS + h) * ns;

    const int splits_per_group = (ns + num_groups - 1) / num_groups;
    const int s_start = min(group * splits_per_group, ns);
    const int s_end = min(s_start + splits_per_group, ns);

    float mm = NEG_INF;
    float ll = 0.0f;
    float r0 = 0.0f, r1 = 0.0f;

    for (int s = s_start; s < s_end; ++s) {
        float m = pm[bh_ns + s];
        float l = pl[bh_ns + s];
        float mn = fmaxf(mm, m);
        float a = fast_exp(mm - mn);
        float bw = fast_exp(m - mn) * l;
        ll = a * ll + bw;
        mm = mn;
        uint32_t packed = *reinterpret_cast<const uint32_t*>(
            &po[bh_ns * D_V + s * D_V + vd_base]);
        r0 = a * r0 + bw * to_f(*reinterpret_cast<const bf16*>(&packed));
        r1 = a * r1 + bw * to_f(*(reinterpret_cast<const bf16*>(&packed) + 1));
    }

    __shared__ float smem[8][64][4];
    smem[group][lane][0] = mm;
    smem[group][lane][1] = ll;
    smem[group][lane][2] = r0;
    smem[group][lane][3] = r1;
    __syncthreads();

    if (group == 0) {
        for (int g = 1; g < num_groups; ++g) {
            if (g * splits_per_group >= ns) break;
            float m = smem[g][lane][0];
            float l = smem[g][lane][1];
            float mn = fmaxf(mm, m);
            float a = fast_exp(mm - mn);
            float bv = fast_exp(m - mn);
            ll = a * ll + bv * l;
            mm = mn;
            r0 = a * r0 + bv * smem[g][lane][2];
            r1 = a * r1 + bv * smem[g][lane][3];
        }
        float inv = ll > 0.0f ? (1.0f / ll) : 0.0f;
        out[out_off(qi, h, vd_base)] = to_b(r0 * inv);
        out[out_off(qi, h, vd_base + 1)] = to_b(r1 * inv);
    }
}

__global__ __launch_bounds__(64, 4)
void mla_s2_lean(
    const float* __restrict__ pm,
    const float* __restrict__ pl,
    const bf16* __restrict__ po,
    bf16* __restrict__ out,
    const int32_t* __restrict__ qo,
    int bs,
    int ns)
{
    const int b = blockIdx.x;
    const int h = blockIdx.y;
    const int vb = blockIdx.z;
    if (b >= bs) return;

    const int lane = threadIdx.x;
    const int qs = qo[b];
    if (qo[b + 1] - qs != 1) return;
    const int qi = qs;
    const int vd_base = vb * 128 + lane * 2;
    const int bh_ns = (b * NUM_HEADS + h) * ns;

    float mm = NEG_INF;
    float ll = 0.0f;
    float r0 = 0.0f, r1 = 0.0f;

    for (int s = 0; s < ns; ++s) {
        float m = pm[bh_ns + s];
        float l = pl[bh_ns + s];
        float mn = fmaxf(mm, m);
        float a = fast_exp(mm - mn);
        float bw = fast_exp(m - mn) * l;
        ll = a * ll + bw;
        mm = mn;
        uint32_t packed = *reinterpret_cast<const uint32_t*>(
            &po[bh_ns * D_V + s * D_V + vd_base]);
        r0 = a * r0 + bw * to_f(*reinterpret_cast<const bf16*>(&packed));
        r1 = a * r1 + bw * to_f(*(reinterpret_cast<const bf16*>(&packed) + 1));
    }

    float inv = ll > 0.0f ? (1.0f / ll) : 0.0f;
    out[out_off(qi, h, vd_base)] = to_b(r0 * inv);
    out[out_off(qi, h, vd_base + 1)] = to_b(r1 * inv);
}

}  // namespace mla

// Raw C-exported launch functions for ctypes fast path (bypasses pybind11 overhead)
extern "C" {

void launch_iq2_s2_light_raw(
    const void* q_bf16,
    const void* kv, const void* kv_scale,
    const void* qo, const void* kvi,
    void* pm, void* pl, void* po,
    void* out,
    int bs, int ns)
{
    if (ns == 1) {
        // Single-split decode writes final output directly; no reduction bookkeeping needed.
        hipLaunchKernelGGL(
            (mla::mla_s1_light_iq2<true, false>),
            dim3(bs), dim3(mla::DUET_BLOCK_THREADS), 0, 0,
            reinterpret_cast<const uint16_t*>(q_bf16),
            reinterpret_cast<const uint8_t*>(kv), reinterpret_cast<const float*>(kv_scale),
            reinterpret_cast<const int32_t*>(qo),
            reinterpret_cast<const int32_t*>(kvi),
            reinterpret_cast<float*>(pm),
            reinterpret_cast<float*>(pl),
            reinterpret_cast<mla::bf16*>(po),
            reinterpret_cast<mla::bf16*>(out),
            bs, 1);
        return;
    }

    if (ns > 1 && ns <= 4) {
        // For small split counts, a separate 1-wave S2 beats the fused last-CTA tail.
        hipLaunchKernelGGL(
            (mla::mla_s1_light_iq2<false, false>),
            dim3(bs * ns), dim3(mla::DUET_BLOCK_THREADS), 0, 0,
            reinterpret_cast<const uint16_t*>(q_bf16),
            reinterpret_cast<const uint8_t*>(kv), reinterpret_cast<const float*>(kv_scale),
            reinterpret_cast<const int32_t*>(qo),
            reinterpret_cast<const int32_t*>(kvi),
            reinterpret_cast<float*>(pm),
            reinterpret_cast<float*>(pl),
            reinterpret_cast<mla::bf16*>(po),
            reinterpret_cast<mla::bf16*>(out),
            bs, ns);
        hipLaunchKernelGGL(
            mla::mla_s2_fast,
            dim3(bs, mla::NUM_HEADS), dim3(64), 0, 0,
            reinterpret_cast<const float*>(pm),
            reinterpret_cast<const float*>(pl),
            reinterpret_cast<const mla::bf16*>(po),
            reinterpret_cast<mla::bf16*>(out),
            reinterpret_cast<const int32_t*>(qo),
            bs, ns);
        return;
    }

    if (ns > 4) {
        // Higher split counts favor a standalone reduction kernel over last-CTA fusion.
        hipLaunchKernelGGL(
            (mla::mla_s1_light_iq2<false, false>),
            dim3(bs * ns), dim3(mla::DUET_BLOCK_THREADS), 0, 0,
            reinterpret_cast<const uint16_t*>(q_bf16),
            reinterpret_cast<const uint8_t*>(kv), reinterpret_cast<const float*>(kv_scale),
            reinterpret_cast<const int32_t*>(qo),
            reinterpret_cast<const int32_t*>(kvi),
            reinterpret_cast<float*>(pm),
            reinterpret_cast<float*>(pl),
            reinterpret_cast<mla::bf16*>(po),
            reinterpret_cast<mla::bf16*>(out),
            bs, ns);
        if (bs <= 8) {
            const int vpar_block = min(ns, 8) * 64;
            hipLaunchKernelGGL(
                mla::mla_s2_vpar2,
                dim3(bs, mla::NUM_HEADS, mla::D_V / 128), dim3(vpar_block), 0, 0,
                reinterpret_cast<const float*>(pm),
                reinterpret_cast<const float*>(pl),
                reinterpret_cast<const mla::bf16*>(po),
                reinterpret_cast<mla::bf16*>(out),
                reinterpret_cast<const int32_t*>(qo),
                bs, ns);
        } else {
            hipLaunchKernelGGL(
                mla::mla_s2_lean,
                dim3(bs, mla::NUM_HEADS, mla::D_V / 128), dim3(64), 0, 0,
                reinterpret_cast<const float*>(pm),
                reinterpret_cast<const float*>(pl),
                reinterpret_cast<const mla::bf16*>(po),
                reinterpret_cast<mla::bf16*>(out),
                reinterpret_cast<const int32_t*>(qo),
                bs, ns);
        }
        return;
    }

    // Fallback fused S1+S2 path, with only the done_count tail memset.
    const size_t pm_head_floats = static_cast<size_t>(bs) * mla::NUM_HEADS * ns;
    hipMemsetAsync(
        reinterpret_cast<int*>(reinterpret_cast<float*>(pm) + pm_head_floats),
        0,
        static_cast<size_t>(bs) * sizeof(int),
        0);
    hipLaunchKernelGGL(
        (mla::mla_s1_light_iq2<false, true>),
        dim3(bs * ns), dim3(mla::DUET_BLOCK_THREADS), 0, 0,
        reinterpret_cast<const uint16_t*>(q_bf16),
        reinterpret_cast<const uint8_t*>(kv), reinterpret_cast<const float*>(kv_scale),
        reinterpret_cast<const int32_t*>(qo),
        reinterpret_cast<const int32_t*>(kvi),
        reinterpret_cast<float*>(pm),
        reinterpret_cast<float*>(pl),
        reinterpret_cast<mla::bf16*>(po),
        reinterpret_cast<mla::bf16*>(out),
        bs, ns);
}

}  // extern "C"

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    // All dispatch goes through ctypes; pybind exports are unused
}
"""


@lru_cache(maxsize=1)
def _ext():
    return load_inline(
        name="mla_fp8hw_quacksoftract_nosleep",
        cpp_sources="",
        cuda_sources=HIP_SRC,
        functions=None,
        extra_cuda_cflags=[
            "-O3", "-ffast-math", "--offload-arch=gfx950",
            "-mllvm", "-enable-post-misched=1",
            "-mllvm", "--lsr-drop-solution=1",
            "-mllvm", "-amdgpu-early-inline-all=true",
            "-mllvm", "-amdgpu-function-calls=false",
            "-mllvm", "-amdgpu-max-memory-clause=64",
        ],
        with_cuda=True,
        verbose=False,
    )


def _hip_src_dsread2ns1():
    src = HIP_SRC
    src = src.replace(
        "using bf16 = hip_bfloat16;\n",
        "using bf16 = hip_bfloat16;\n"
        "using floatx2 = float __attribute__((ext_vector_type(2)));\n",
        1,
    )
    src = src.replace(
        "using intx4 = int __attribute__((ext_vector_type(4)));\n",
        "using intx4 = int __attribute__((ext_vector_type(4)));\n"
        "using intx2 = int __attribute__((ext_vector_type(2)));\n",
        1,
    )
    src = src.replace(
        "__device__ __forceinline__ void st16u(uint8_t* p, U128 v) {\n"
        "    *reinterpret_cast<U128*>(p) = v;\n"
        "}\n",
        "__device__ __forceinline__ void st16u(uint8_t* p, U128 v) {\n"
        "    *reinterpret_cast<U128*>(p) = v;\n"
        "}\n"
        "\n"
        "__device__ __forceinline__ floatx2 ds_read2_f32_pair(const float* p)\n"
        "{\n"
        "#define __LDS_ADDR __attribute__((address_space(3)))\n"
        "#pragma clang diagnostic push\n"
        "#pragma clang diagnostic ignored \"-Wold-style-cast\"\n"
        "    const auto p_lds = (__LDS_ADDR const float*)(p);\n"
        "#pragma clang diagnostic pop\n"
        "    intx2 bits;\n"
        "    const uint32_t addr = static_cast<uint32_t>(reinterpret_cast<uintptr_t>(p_lds));\n"
        "    asm volatile(\n"
        "        \"ds_read2_b32 %0, %1 offset0:0 offset1:16\"\n"
        "        : \"=v\"(bits)\n"
        "        : \"v\"(addr));\n"
        "    return *reinterpret_cast<floatx2*>(&bits);\n"
        "#undef __LDS_ADDR\n"
        "}\n",
        1,
    )
    src = src.replace(
        "        float hm0 = half_max[0][row_base + 0];\n"
        "        float hm1 = half_max[1][row_base + 0];\n"
        "        float hm2 = half_max[0][row_base + 1];\n"
        "        float hm3 = half_max[1][row_base + 1];\n"
        "        float hm4 = half_max[0][row_base + 2];\n"
        "        float hm5 = half_max[1][row_base + 2];\n"
        "        float hm6 = half_max[0][row_base + 3];\n"
        "        float hm7 = half_max[1][row_base + 3];\n",
        "        float hm0, hm1, hm2, hm3, hm4, hm5, hm6, hm7;\n"
        "        if constexpr (DIRECT_OUT) {\n"
        "            const floatx2 hm01 = ds_read2_f32_pair(&half_max[0][row_base + 0]);\n"
        "            const floatx2 hm23 = ds_read2_f32_pair(&half_max[0][row_base + 1]);\n"
        "            const floatx2 hm45 = ds_read2_f32_pair(&half_max[0][row_base + 2]);\n"
        "            const floatx2 hm67 = ds_read2_f32_pair(&half_max[0][row_base + 3]);\n"
        "            hm0 = hm01.x; hm1 = hm01.y;\n"
        "            hm2 = hm23.x; hm3 = hm23.y;\n"
        "            hm4 = hm45.x; hm5 = hm45.y;\n"
        "            hm6 = hm67.x; hm7 = hm67.y;\n"
        "        } else {\n"
        "            hm0 = half_max[0][row_base + 0];\n"
        "            hm1 = half_max[1][row_base + 0];\n"
        "            hm2 = half_max[0][row_base + 1];\n"
        "            hm3 = half_max[1][row_base + 1];\n"
        "            hm4 = half_max[0][row_base + 2];\n"
        "            hm5 = half_max[1][row_base + 2];\n"
        "            hm6 = half_max[0][row_base + 3];\n"
        "            hm7 = half_max[1][row_base + 3];\n"
        "        }\n",
        1,
    )
    src = src.replace(
        "            local_sum[h] = alpha[h] * local_sum[h]\n"
        "                + half_sum[0][row] + half_sum[1][row];\n",
        "            if constexpr (DIRECT_OUT) {\n"
        "                const floatx2 hs = ds_read2_f32_pair(&half_sum[0][row]);\n"
        "                local_sum[h] = alpha[h] * local_sum[h] + hs.x + hs.y;\n"
        "            } else {\n"
        "                local_sum[h] = alpha[h] * local_sum[h]\n"
        "                    + half_sum[0][row] + half_sum[1][row];\n"
        "            }\n",
        1,
    )
    return src


@lru_cache(maxsize=1)
def _ext_dsread2ns1():
    return load_inline(
        name="mla_fp8hw_quacksoftract_nosleep_dualns1",
        cpp_sources="",
        cuda_sources=_hip_src_dsread2ns1(),
        functions=None,
        extra_cuda_cflags=[
            "-O3", "-ffast-math", "--offload-arch=gfx950",
            "-mllvm", "-enable-post-misched=1",
            "-mllvm", "--lsr-drop-solution=1",
            "-mllvm", "-amdgpu-early-inline-all=true",
            "-mllvm", "-amdgpu-function-calls=false",
            "-mllvm", "-amdgpu-max-memory-clause=64",
        ],
        with_cuda=True,
        verbose=False,
    )

def _get_bufs(batch_size, ns, device):
    # pm has extra space at tail for done_count (bs ints) used by fused S2
    # Layout: [bs*NUM_HEADS*ns floats] + [bs ints padded as floats].
    # The payload is fully overwritten in-kernel; only the tail counter is zeroed by HIP.
    pm_extra = batch_size  # sizeof(int)==sizeof(float)==4, need bs slots
    return (
        torch.empty((batch_size * NUM_HEADS * ns + pm_extra,), device=device, dtype=torch.float32),
        torch.empty((batch_size, NUM_HEADS, ns), device=device, dtype=torch.float32),
        torch.empty(
            (batch_size, NUM_HEADS, ns, V_HEAD_DIM),
            device=device, dtype=torch.bfloat16),
    )


def _get_out(batch_size, device):
    return torch.empty(
        (batch_size, NUM_HEADS, V_HEAD_DIM),
        device=device, dtype=torch.bfloat16)


_SPLIT_TABLE_256CU = {
    (4, 1024): 32,
    (4, 8192): 64,
    (32, 1024): 8,
    (32, 8192): 8,
    (64, 1024): 4,
    (64, 8192): 4,
    (256, 1024): 1,
    (256, 8192): 2,
}
# MI355X has ~512 CUs — need higher splits to fill them
_SPLIT_TABLE_512CU = {
    (4, 1024): 32,
    (4, 8192): 64,
    (32, 1024): 8,
    (32, 8192): 8,
    (64, 1024): 4,
    (64, 8192): 4,
    (256, 1024): 1,
    (256, 8192): 2,
}

def _detect_split_table():
    try:
        props = torch.cuda.get_device_properties(0)
        cu_count = props.multi_processor_count
        if cu_count >= 400:
            return _SPLIT_TABLE_512CU
    except Exception:
        pass
    return _SPLIT_TABLE_256CU

_SPLIT_TABLE = _detect_split_table()

def _pick_splits_duet(bs, kvl):
    key = (bs, kvl)
    if key in _SPLIT_TABLE:
        return _SPLIT_TABLE[key]
    if kvl <= 2048:
        ns_cu = max(1, 256 // bs)
        ns_kv = max(1, kvl // 512)
        ns = max(ns_cu, ns_kv)
        ns = min(ns, max(1, kvl // 128))
        p = 1
        while p * 2 <= ns:
            p *= 2
        return max(1, min(p, 64))
    ns_cu = min(max(1, 512 // bs), 32)
    ns_cu = min(ns_cu, kvl // 32) if kvl >= 32 else 1
    kv_div = 2048 if bs >= 64 else 1024
    ns_kv = max(1, kvl // kv_div)
    ns = max(ns_cu, ns_kv)
    p = 1
    while p * 2 <= ns:
        p *= 2
    return max(1, min(p, 64))



import ctypes as _ct

_cached_ext = None
_raw_iq2_s2 = None
_cached_ext_dsread2ns1 = None
_raw_iq2_s2_dsread2ns1 = None
def _init_raw_dispatch(ext):
    """Load the compiled .so with ctypes for fast kernel dispatch."""
    so = ext.__file__
    lib = _ct.CDLL(so)
    P = _ct.c_void_p
    lib.launch_iq2_s2_light_raw.restype = None
    lib.launch_iq2_s2_light_raw.argtypes = [
        P, P, P, P, P, P, P, P, P, _ct.c_int, _ct.c_int]
    return lib.launch_iq2_s2_light_raw


def custom_kernel(data):
    global _cached_ext, _raw_iq2_s2
    global _cached_ext_dsread2ns1, _raw_iq2_s2_dsread2ns1
    q, kv_data, qo_indptr, kv_indptr, config = data

    bs = int(config["batch_size"])
    kvl = int(config["kv_seq_len"])

    ns = _SPLIT_TABLE.get((bs, kvl)) or _pick_splits_duet(bs, kvl)

    if _raw_iq2_s2 is None:
        _cached_ext = _ext()
        _raw_iq2_s2 = _init_raw_dispatch(_cached_ext)

    pm, pl, po = _get_bufs(bs, ns, q.device)
    out = _get_out(bs, q.device)

    q_c = q.contiguous()
    kv_t = kv_data["fp8"][0].contiguous()
    sc_t = kv_data["fp8"][1].float()
    qo_c = qo_indptr.contiguous()
    kvi_c = kv_indptr.contiguous()

    raw_fn = _raw_iq2_s2
    if ns == 1:
        if _raw_iq2_s2_dsread2ns1 is None:
            _cached_ext_dsread2ns1 = _ext_dsread2ns1()
            _raw_iq2_s2_dsread2ns1 = _init_raw_dispatch(_cached_ext_dsread2ns1)
        raw_fn = _raw_iq2_s2_dsread2ns1

    raw_fn(
        q_c.data_ptr(),
        kv_t.data_ptr(),
        sc_t.data_ptr(),
        qo_c.data_ptr(),
        kvi_c.data_ptr(),
        pm.data_ptr(),
        pl.data_ptr(),
        po.data_ptr(),
        out.data_ptr(),
        bs,
        ns,
    )
    return out

scrolls · 2110 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON