Skip to content
KernelIndex
Search⌘K

submission 650401

div22 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution_511.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-650401?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
8.17µs
#35 of 1143
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:176cc2a98048914f21ac47cc0eb0d2736c9911d5cae7a908fdc8af3749080525
license declaredunknown
license concludedunknown
authorsdiv22
imported2026-08-15

Techniques

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

fp4int a_lds_offs[G5_A_PER_THREAD]; // FP4 LDS destination
shared-memoryconst int4_v _w0 = *reinterpret_cast<const int4_v*>((smem_ptr) + (lds_off)); \
split-ktemplate<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,
tile-m = 0constexpr auto full_m = (BM >= 16) ? C_M / BM : 0;
tile-n = 32constexpr auto BN = 32;

Kernel source

solution_511.py3117 lines
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
from typing import Tuple
import torch

HIP_SRC = r"""

#include <hip/hip_runtime.h>
#include <stdint.h>

// ============================================================================
// Pure C-style approach: NO templates. Each benchmark shape gets its own
// kernel function with all dimensions hardcoded as constants.
//
// s292: Replace global_load_dwordx4 with buffer_load_dwordx4 in PATH B
//       - buffer_load uses SGPR buffer resource (64-bit base) + SGPR soffset
//         + 32-bit VGPR voffset. Eliminates ALL 64-bit address math (v_lshl_add_u64).
//       - "s" constraints force rsrc/soffset into SGPRs
//       - Only per-lane 32-bit offset remains in VGPR (computed once, constant)
// ============================================================================

using int4_v   = int   __attribute__((ext_vector_type(4)));
using float4_v = float __attribute__((ext_vector_type(4)));
using bf16x2   = __bf16 __attribute__((ext_vector_type(2)));

#define FP4_E2M1 4

__device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
    int4_v a, int4_v b, float4_v c,
    int cbsz, int blgp, int op_sel_a, int scale_a, int op_sel_b, int scale_b
) __asm("llvm.amdgcn.mfma.scale.f32.16x16x128.f8f6f4.v4i32.v4i32");

__device__ __forceinline__ int4_v load16(const uint8_t* __restrict__ p) {
    return *reinterpret_cast<const int4_v*>(p);
}

__device__ __forceinline__ int4_v load16_nt(const uint8_t* __restrict__ p) {
    return __builtin_nontemporal_load(reinterpret_cast<const int4_v*>(p));
}

// Fast bf16 conversion — bitwise round-to-nearest-even
__device__ __forceinline__ uint16_t fast_f32_to_bf16(float f) {
    unsigned u;
    __builtin_memcpy(&u, &f, 4);
    unsigned rnd = u + 0x7FFFu + ((u >> 16) & 1u);
    return static_cast<uint16_t>(rnd >> 16);
}

// Original float_to_bf16 for correctness fallback
__device__ __forceinline__ uint16_t float_to_bf16(float f) {
    return fast_f32_to_bf16(f);
}

// ============================================================================
// Buffer resource descriptor helpers for buffer_load_dwordx4
// ============================================================================

// Build a 128-bit buffer resource descriptor (V#) for a global pointer.
// Format: { base_lo32, base_hi32, num_records=0xFFFFFFFF, flags }
// flags: DST_SEL_X/Y/Z/W=0, NUM_FORMAT=UINT, DATA_FORMAT=32, ADD_TID_ENABLE=0
// For gfx950: stride=0, cache swizzle=0, swizzle_enable=0
// Word3: 0x00020000 = NUM_FORMAT_UINT(1)<<20 | DATA_FORMAT_32(4)<<16
//        Actually for raw buffer: word3 = 0x00020000 is fine (stride=0, OOB_SELECT=2)
//        OOB_SELECT=2 means "return 0 on out-of-bounds" which is safe
// For raw buffer loads: word3 bits [30:28] = OOB_SELECT, we want 2 for "return 0"
// Actually simpler: word3 = 0x20000000 (OOB_SELECT=2, rest zero)
// Wait, let's use the proper format:
//   word3[14:12] = DST_SEL_X..W (0=zero, 1=one, 4=sel_x etc.)
//   For raw buffer: word3 = 0x00000000 works with stride=0
//   But we need num_format + data_format for typed, or use raw mode
//   For RAW buffer (no format conversion): set word3[30:28]=OOB_SELECT=2
//   word3 = (2 << 28) = 0x20000000
// Actually the simplest correct approach for gfx9:
//   word0 = base & 0xFFFFFFFF
//   word1 = base >> 32
//   word2 = 0xFFFFFFFF (num_records, max range)
//   word3 = 0x00027000 for gfx9 (NUM_FORMAT=1, DATA_FORMAT=4, stride=0)
//          ... but this is typed. For raw: just 0x00020000?
// Let me use the approach from known working code:
//   word3 = 0x00020000 (stride=0, swizzle=0, ADD_TID=0, format fields)
// Actually for buffer_load_dwordx4 with offen, the simplest is:
//   {base_lo, base_hi, 0xFFFFFFFF, 0x00020000}
// This gives stride=0, num_records=UINT_MAX, basic format.

struct __attribute__((aligned(16))) BufferRsrc {
    uint32_t w0, w1, w2, w3;
};

__device__ __forceinline__ BufferRsrc make_buffer_rsrc(const void* ptr) {
    uint64_t addr = reinterpret_cast<uint64_t>(ptr);
    return BufferRsrc{
        static_cast<uint32_t>(addr),
        static_cast<uint32_t>(addr >> 32),
        0xFFFFFFFFu,
        0x00020000u
    };
}


// ============================================================================
// COMMON MACROS for quant logic
// ============================================================================

#define QUANT_GROUP_32(smem_ptr, lds_off, out_fp4, out_scale) \
{ \
    const int4_v _w0 = *reinterpret_cast<const int4_v*>((smem_ptr) + (lds_off)); \
    const int4_v _w1 = *reinterpret_cast<const int4_v*>((smem_ptr) + (lds_off) + 16); \
    const int4_v _w2 = *reinterpret_cast<const int4_v*>((smem_ptr) + (lds_off) + 32); \
    const int4_v _w3 = *reinterpret_cast<const int4_v*>((smem_ptr) + (lds_off) + 48); \
    const uint32_t* _p0 = reinterpret_cast<const uint32_t*>(&_w0); \
    const uint32_t* _p1 = reinterpret_cast<const uint32_t*>(&_w1); \
    const uint32_t* _p2 = reinterpret_cast<const uint32_t*>(&_w2); \
    const uint32_t* _p3 = reinterpret_cast<const uint32_t*>(&_w3); \
    float _absMax = 1e-10f; \
    { \
        auto _AM = [&](uint32_t pair) { \
            __bf16 lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
            __bf16 hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair >> 16)); \
            float flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
            float fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
            _absMax = (flo > _absMax) ? flo : _absMax; \
            _absMax = (fhi > _absMax) ? fhi : _absMax; \
        }; \
        _AM(_p0[0]); _AM(_p0[1]); _AM(_p0[2]); _AM(_p0[3]); \
        _AM(_p1[0]); _AM(_p1[1]); _AM(_p1[2]); _AM(_p1[3]); \
        _AM(_p2[0]); _AM(_p2[1]); _AM(_p2[2]); _AM(_p2[3]); \
        _AM(_p3[0]); _AM(_p3[1]); _AM(_p3[2]); _AM(_p3[3]); \
    } \
    uint32_t _u32 = __builtin_bit_cast(uint32_t, _absMax); \
    uint32_t _amax_exp = ((_u32 + 0x200000u) >> 23) & 0xFFu; \
    uint32_t _inv_exp = (_amax_exp >= 2u) ? (_amax_exp - 2u) : 0u; \
    (out_scale) = static_cast<int>(_inv_exp); \
    float _hw_scale = __builtin_bit_cast(float, _inv_exp << 23); \
    unsigned _d0 = 0u, _d1 = 0u, _d2 = 0u, _d3 = 0u; \
    _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[0]), _hw_scale, 0); \
    _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[1]), _hw_scale, 1); \
    _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[2]), _hw_scale, 2); \
    _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[3]), _hw_scale, 3); \
    _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[0]), _hw_scale, 0); \
    _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[1]), _hw_scale, 1); \
    _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[2]), _hw_scale, 2); \
    _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[3]), _hw_scale, 3); \
    _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[0]), _hw_scale, 0); \
    _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[1]), _hw_scale, 1); \
    _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[2]), _hw_scale, 2); \
    _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[3]), _hw_scale, 3); \
    _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[0]), _hw_scale, 0); \
    _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[1]), _hw_scale, 1); \
    _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[2]), _hw_scale, 2); \
    _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[3]), _hw_scale, 3); \
    (out_fp4) = int4_v{static_cast<int>(_d0), static_cast<int>(_d1), \
                       static_cast<int>(_d2), static_cast<int>(_d3)}; \
}


// ============================================================================
// SHAPE 1: M=4, N=2880, K=512 — Fused quant+GEMM, BN=64
// WAVES_M=1, WAVES_N=4, CKT=4, 256 threads
// ============================================================================
__global__ void __launch_bounds__(256, 2)
gemm_m4_n2880_k512(
    const __bf16* __restrict__ A_bf16,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    uint16_t*      __restrict__ C_final)
{
    // Hardcoded constants
    #define S1_M 4
    #define S1_K 512
    #define S1_N 2880
    #define S1_CKT 4
    #define S1_BN 64
    #define S1_WAVES_M 1
    #define S1_WAVES_N 4
    #define S1_NWARPS 4
    #define S1_SCALEN 16

    const int tid = threadIdx.x;
    const int lane = tid % 64;
    const int wave_id = tid / 64;
    const int lrow = lane % 16;
    const int kgrp = lane / 16;
    const int wave_m = wave_id / S1_WAVES_N;  // always 0
    const int wave_n = wave_id % S1_WAVES_N;

    const int tile_m = 0;  // wave_m * 16 = 0
    const int tile_n = static_cast<int>(blockIdx.x) * S1_BN + wave_n * 16;
    const int my_row = lrow;  // tile_m + lrow

    // Load A bf16 into shared memory (4 * 512 * 2 = 4096 bytes)
    extern __shared__ uint8_t smem[];
    {
        // 256 threads * 16 bytes = 4096 bytes per iter; need 1 iter
        const int offset = tid * 16;
        if (offset + 16 <= S1_M * S1_K * 2) {
            *reinterpret_cast<int4_v*>(smem + offset) =
                *reinterpret_cast<const int4_v*>(
                    reinterpret_cast<const uint8_t*>(A_bf16) + offset);
        }
    }
    __syncthreads();

    // Quant A into registers
    int4_v a_fp4_regs[S1_CKT];
    int a_scale_regs[S1_CKT];

    // M=4, so rows 0..3 valid, rows 4..15 invalid
    const bool a_row_valid = (my_row < S1_M);

    #pragma unroll
    for (int kt = 0; kt < S1_CKT; ++kt) {
        const int kg = kt * 4 + kgrp;
        if (a_row_valid) {
            const int lds_off = my_row * S1_K * 2 + kg * 64;
            QUANT_GROUP_32(smem, lds_off, a_fp4_regs[kt], a_scale_regs[kt])
        } else {
            a_fp4_regs[kt] = int4_v{0, 0, 0, 0};
            a_scale_regs[kt] = 127;
        }
    }

    float4_v acc{0.f, 0.f, 0.f, 0.f};

    const int n_tile = tile_n / 16;
    const long bsh_n_stride = (long)(S1_K / 64) * 512;  // 4096
    const uint8_t* bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
    const int k_half_off = (kgrp & 1) * 256;
    const int k_blk_base = kgrp >> 1;

    const int gn = tile_n + lrow;
    const int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
                          + (gn >> 5) * (32 * S1_SCALEN);

    // s483: Prefetch all B tiles + scales into registers before MFMA loop
    int4_v b_regs_s1[S1_CKT];
    int sb_regs_s1[S1_CKT];
    #pragma unroll
    for (int kt = 0; kt < S1_CKT; ++kt) {
        b_regs_s1[kt] = load16_nt(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
        sb_regs_s1[kt] = static_cast<int>(*(Bssh + bssh_base
                        + (kt & 1) * 2 + (kt >> 1) * 256));
    }

    #pragma unroll
    for (int kt = 0; kt < S1_CKT; ++kt) {
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a_fp4_regs[kt], b_regs_s1[kt], acc, FP4_E2M1, FP4_E2M1,
            0, a_scale_regs[kt], 0, sb_regs_s1[kt]);
    }

    const int out_col = tile_n + lrow;
    const int out_row_base = kgrp * 4;
    if (out_col >= S1_N) return;

    uint16_t* c_out = C_final + (long)out_row_base * S1_N + out_col;
    #pragma unroll
    for (int i = 0; i < 4; ++i, c_out += S1_N) {
        if (out_row_base + i < S1_M)
            *c_out = fast_f32_to_bf16(acc[i]);
    }

    #undef S1_M
    #undef S1_K
    #undef S1_N
    #undef S1_CKT
    #undef S1_BN
    #undef S1_WAVES_M
    #undef S1_WAVES_N
    #undef S1_NWARPS
    #undef S1_SCALEN
}


// ============================================================================
// SHAPE 2: M=16, N=2112, K=7168 — Fused quant+GEMM splitK
// splitK=14, KPS=4, BN=64, WAVES_M=1, WAVES_N=4
// ============================================================================
__global__ void __launch_bounds__(256, 2)
gemm_m16_n2112_k7168(
    const __bf16* __restrict__ A_bf16,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    float*         __restrict__ C_partial)
{
    #define S2_M 16
    #define S2_K 7168
    #define S2_N 2112
    #define S2_CKT 4
    #define S2_BN 64
    #define S2_KPS 4
    #define S2_NUM_KSPLIT 14
    #define S2_SCALEN 224

    const int tid = threadIdx.x;
    const int lane = tid % 64;
    const int wave_id = tid / 64;
    const int lrow = lane % 16;
    const int kgrp = lane / 16;
    const int wave_n = wave_id % 4;
    const int ks_idx = static_cast<int>(blockIdx.z);

    const int tile_n = static_cast<int>(blockIdx.x) * S2_BN + wave_n * 16;
    const int my_row = lrow;

    const int ks_start = ks_idx * S2_KPS;

    // s483: Direct HBM quant for M=16 — skip LDS, quant A directly from global memory.
    // 4 N-waves read the same A rows → L2 handles the sharing (16KB per split).
    const uint8_t* _a_byte = reinterpret_cast<const uint8_t*>(A_bf16);
    int4_v a_fp4_regs[S2_CKT];
    int a_scale_regs[S2_CKT];

    #pragma unroll
    for (int kt = 0; kt < S2_CKT; ++kt) {
        const int kg = kt * 4 + kgrp;
        const long global_off = (long)my_row * S2_K * 2 + (long)(ks_start * 128 + kg * 32) * 2;
        QUANT_GROUP_32(_a_byte, global_off, a_fp4_regs[kt], a_scale_regs[kt])
    }

    float4_v acc{0.f, 0.f, 0.f, 0.f};
    const int n_tile = tile_n / 16;
    const long bsh_n_stride = (long)(S2_K / 64) * 512;
    const uint8_t* bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
    const int k_half_off = (kgrp & 1) * 256;
    const int k_blk_base = kgrp >> 1;

    const int gn = tile_n + lrow;
    const int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
                          + (gn >> 5) * (32 * S2_SCALEN);

    // s483: Prefetch all B tiles + scales into registers before MFMA loop
    int4_v b_regs_s2[S2_CKT];
    int sb_regs_s2[S2_CKT];
    #pragma unroll
    for (int kt = 0; kt < S2_CKT; ++kt) {
        const int global_kt = ks_start + kt;
        b_regs_s2[kt] = load16_nt(bsh_lane_base + (long)(global_kt * 2 + k_blk_base) * 512 + k_half_off);
        sb_regs_s2[kt] = static_cast<int>(*(Bssh + bssh_base
                        + (global_kt & 1) * 2 + (global_kt >> 1) * 256));
    }

    // MFMA loop — all operands in registers, no memory stalls
    #pragma unroll
    for (int kt = 0; kt < S2_CKT; ++kt) {
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a_fp4_regs[kt], b_regs_s2[kt], acc, FP4_E2M1, FP4_E2M1,
            0, a_scale_regs[kt], 0, sb_regs_s2[kt]);
    }

    const int out_col = tile_n + lrow;
    const int out_row_base = kgrp * 4;
    if (out_col >= S2_N) return;

    float* c_out = C_partial + (long)ks_idx * S2_M * S2_N + (long)out_row_base * S2_N + out_col;
    #pragma unroll
    for (int i = 0; i < 4; ++i, c_out += S2_N) {
        if (out_row_base + i < S2_M)
            *c_out = acc[i];
    }

    #undef S2_M
    #undef S2_K
    #undef S2_N
    #undef S2_CKT
    #undef S2_BN
    #undef S2_KPS
    #undef S2_NUM_KSPLIT
    #undef S2_SCALEN
}


// Reduce kernel for M=16, N=2112, splitK=14
__global__ void reduce_m16_n2112_k14(
    const float*  __restrict__ C_partial,
    uint16_t*     __restrict__ C_out)
{
    #define R_M 16
    #define R_N 2112
    #define R_NSPLIT 14

    const int col = blockIdx.x * 32 + threadIdx.x;
    const int row = blockIdx.y * 16 + threadIdx.y;
    if (row >= R_M || col >= R_N) return;

    const long mn = (long)row * R_N + col;
    const long mn_stride = (long)R_M * R_N;
    float sum = 0.f;
    const float* ptr = C_partial + mn;
    #pragma unroll
    for (int k = 0; k < R_NSPLIT; ++k)
        sum += ptr[k * mn_stride];

    C_out[mn] = fast_f32_to_bf16(sum);

    #undef R_M
    #undef R_N
    #undef R_NSPLIT
}


// ============================================================================
// SHAPE 3: M=32, N=4096, K=512 — Fused quant+GEMM, BN=32
// s511: 128 threads (2 warps), 2D grid for 2× CU utilization (256 vs 128 blocks)
// WAVES_M=1 (via grid.y=2), WAVES_N=2, CKT=4
// ============================================================================
__global__ void __launch_bounds__(128, 4)
gemm_m32_n4096_k512(
    const __bf16* __restrict__ A_bf16,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    uint16_t*      __restrict__ C_final)
{
    #define S3_M 32
    #define S3_K 512
    #define S3_N 4096
    #define S3_CKT 4
    #define S3_BN 32
    #define S3_WAVES_M 1
    #define S3_WAVES_N 2
    #define S3_SCALEN 16

    const int tid = threadIdx.x;
    const int lane = tid % 64;
    const int wave_id = tid / 64;
    const int lrow = lane % 16;
    const int kgrp = lane / 16;
    const int wave_n = wave_id % S3_WAVES_N;

    const int tile_m = static_cast<int>(blockIdx.y) * 16;  // s511: M-tile from grid.y
    const int tile_n = static_cast<int>(blockIdx.x) * S3_BN + wave_n * 16;
    const int my_row = tile_m + lrow;

    // s480: Direct HBM quant for M=32 — skip LDS, quant A from global memory
    const uint8_t* _a_byte = reinterpret_cast<const uint8_t*>(A_bf16);
    int4_v a_fp4_regs[S3_CKT];
    int a_scale_regs[S3_CKT];

    #pragma unroll
    for (int kt = 0; kt < S3_CKT; ++kt) {
        const int kg = kt * 4 + kgrp;
        const long global_off = (long)my_row * S3_K * 2 + (long)kg * 64;
        QUANT_GROUP_32(_a_byte, global_off, a_fp4_regs[kt], a_scale_regs[kt])
    }

    float4_v acc{0.f, 0.f, 0.f, 0.f};

    const int n_tile = tile_n / 16;
    const long bsh_n_stride = (long)(S3_K / 64) * 512;
    const uint8_t* bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
    const int k_half_off = (kgrp & 1) * 256;
    const int k_blk_base = kgrp >> 1;

    const int gn = tile_n + lrow;
    const int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
                          + (gn >> 5) * (32 * S3_SCALEN);

    // s483: Prefetch all B tiles + scales into registers before MFMA loop
    int4_v b_regs_s3[S3_CKT];
    int sb_regs_s3[S3_CKT];
    #pragma unroll
    for (int kt = 0; kt < S3_CKT; ++kt) {
        b_regs_s3[kt] = load16_nt(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
        sb_regs_s3[kt] = static_cast<int>(*(Bssh + bssh_base
                        + (kt & 1) * 2 + (kt >> 1) * 256));
    }

    #pragma unroll
    for (int kt = 0; kt < S3_CKT; ++kt) {
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a_fp4_regs[kt], b_regs_s3[kt], acc, FP4_E2M1, FP4_E2M1,
            0, a_scale_regs[kt], 0, sb_regs_s3[kt]);
    }

    const int out_col = tile_n + lrow;
    const int out_row_base = tile_m + kgrp * 4;
    if (out_col >= S3_N) return;

    uint16_t* c_out = C_final + (long)out_row_base * S3_N + out_col;
    // M=32, BM=32: all rows always valid
    #pragma unroll
    for (int i = 0; i < 4; ++i, c_out += S3_N) {
        *c_out = fast_f32_to_bf16(acc[i]);
    }

    #undef S3_M
    #undef S3_K
    #undef S3_N
    #undef S3_CKT
    #undef S3_BN
    #undef S3_WAVES_M
    #undef S3_WAVES_N
    #undef S3_SCALEN
}


// ============================================================================
// SHAPE 4: M=32, N=2880, K=512 — Fused quant+GEMM, BN=32
// s511: 128 threads, 2D grid (180 blocks vs 90 = 2× CU utilization)
// ============================================================================
__global__ void __launch_bounds__(128, 4)
gemm_m32_n2880_k512(
    const __bf16* __restrict__ A_bf16,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    uint16_t*      __restrict__ C_final)
{
    #define S4_M 32
    #define S4_K 512
    #define S4_N 2880
    #define S4_CKT 4
    #define S4_BN 32
    #define S4_WAVES_M 1
    #define S4_WAVES_N 2
    #define S4_SCALEN 16

    const int tid = threadIdx.x;
    const int lane = tid % 64;
    const int wave_id = tid / 64;
    const int lrow = lane % 16;
    const int kgrp = lane / 16;
    const int wave_n = wave_id % S4_WAVES_N;

    const int tile_m = static_cast<int>(blockIdx.y) * 16;  // s511: M-tile from grid.y
    const int tile_n = static_cast<int>(blockIdx.x) * S4_BN + wave_n * 16;
    const int my_row = tile_m + lrow;

    // s480: Direct HBM quant for M=32 — skip LDS, quant A from global memory
    const uint8_t* _a_byte = reinterpret_cast<const uint8_t*>(A_bf16);
    int4_v a_fp4_regs[S4_CKT];
    int a_scale_regs[S4_CKT];

    #pragma unroll
    for (int kt = 0; kt < S4_CKT; ++kt) {
        const int kg = kt * 4 + kgrp;
        const long global_off = (long)my_row * S4_K * 2 + (long)kg * 64;
        QUANT_GROUP_32(_a_byte, global_off, a_fp4_regs[kt], a_scale_regs[kt])
    }

    float4_v acc{0.f, 0.f, 0.f, 0.f};

    const int n_tile = tile_n / 16;
    const long bsh_n_stride = (long)(S4_K / 64) * 512;
    const uint8_t* bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
    const int k_half_off = (kgrp & 1) * 256;
    const int k_blk_base = kgrp >> 1;

    const int gn = tile_n + lrow;
    const int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
                          + (gn >> 5) * (32 * S4_SCALEN);

    // s483: Prefetch all B tiles + scales into registers before MFMA loop
    int4_v b_regs_s4[S4_CKT];
    int sb_regs_s4[S4_CKT];
    #pragma unroll
    for (int kt = 0; kt < S4_CKT; ++kt) {
        b_regs_s4[kt] = load16_nt(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
        sb_regs_s4[kt] = static_cast<int>(*(Bssh + bssh_base
                        + (kt & 1) * 2 + (kt >> 1) * 256));
    }

    #pragma unroll
    for (int kt = 0; kt < S4_CKT; ++kt) {
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a_fp4_regs[kt], b_regs_s4[kt], acc, FP4_E2M1, FP4_E2M1,
            0, a_scale_regs[kt], 0, sb_regs_s4[kt]);
    }

    const int out_col = tile_n + lrow;
    const int out_row_base = tile_m + kgrp * 4;
    if (out_col >= S4_N) return;

    uint16_t* c_out = C_final + (long)out_row_base * S4_N + out_col;
    #pragma unroll
    for (int i = 0; i < 4; ++i, c_out += S4_N) {
        *c_out = fast_f32_to_bf16(acc[i]);
    }

    #undef S4_M
    #undef S4_K
    #undef S4_N
    #undef S4_CKT
    #undef S4_BN
    #undef S4_WAVES_M
    #undef S4_WAVES_N
    #undef S4_SCALEN
}


// ============================================================================
// SHAPE 5: M=64, N=7168, K=2048 — Separate quant + GEMM PATH B
// Quant: 128 threads
// GEMM: BM=32, BN=32, NWARPS=4, CKT=16 (PATH B with LDS double buffering)
// ============================================================================

// Quant kernel for M=64, K=2048
__global__ void __launch_bounds__(128, 4)
quant_m64_k2048(
    const __bf16* __restrict__ A_bf16,
    uint8_t*      __restrict__ A_fp4,
    uint8_t*      __restrict__ A_scale)
{
    #define Q5_M 64
    #define Q5_K 2048
    #define Q5_KS 64
    #define Q5_K2 1024

    const int group = blockIdx.x * 128 + threadIdx.x;
    const int row   = group / Q5_KS;
    const int kg    = group % Q5_KS;

    if (row >= Q5_M) return;

    const long row_k = (long)row * Q5_K;
    const __bf16* src = A_bf16 + row_k + kg * 32;

    const int4_v w0 = *reinterpret_cast<const int4_v*>(src);
    const int4_v w1 = *reinterpret_cast<const int4_v*>(src + 8);
    const int4_v w2 = *reinterpret_cast<const int4_v*>(src + 16);
    const int4_v w3 = *reinterpret_cast<const int4_v*>(src + 24);

    const uint32_t* p0 = reinterpret_cast<const uint32_t*>(&w0);
    const uint32_t* p1 = reinterpret_cast<const uint32_t*>(&w1);
    const uint32_t* p2 = reinterpret_cast<const uint32_t*>(&w2);
    const uint32_t* p3 = reinterpret_cast<const uint32_t*>(&w3);

    float absMax = 1e-10f;
    #define AMAX_P(pair) { \
        __bf16 lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
        __bf16 hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
        float flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
        float fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
        absMax = (flo > absMax) ? flo : absMax; \
        absMax = (fhi > absMax) ? fhi : absMax; \
    }
    AMAX_P(p0[0]) AMAX_P(p0[1]) AMAX_P(p0[2]) AMAX_P(p0[3])
    AMAX_P(p1[0]) AMAX_P(p1[1]) AMAX_P(p1[2]) AMAX_P(p1[3])
    AMAX_P(p2[0]) AMAX_P(p2[1]) AMAX_P(p2[2]) AMAX_P(p2[3])
    AMAX_P(p3[0]) AMAX_P(p3[1]) AMAX_P(p3[2]) AMAX_P(p3[3])
    #undef AMAX_P

    uint32_t u32 = __builtin_bit_cast(uint32_t, absMax);
    uint32_t amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
    uint32_t inv_exp  = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
    A_scale[row_k / 32 + kg] = static_cast<uint8_t>(inv_exp);

    float hw_scale = __builtin_bit_cast(float, inv_exp << 23);

    #define CVT_Q(d, pair, sel) \
        d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)

    unsigned d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
    CVT_Q(d0, p0[0], 0); CVT_Q(d0, p0[1], 1); CVT_Q(d0, p0[2], 2); CVT_Q(d0, p0[3], 3);
    CVT_Q(d1, p1[0], 0); CVT_Q(d1, p1[1], 1); CVT_Q(d1, p1[2], 2); CVT_Q(d1, p1[3], 3);
    CVT_Q(d2, p2[0], 0); CVT_Q(d2, p2[1], 1); CVT_Q(d2, p2[2], 2); CVT_Q(d2, p2[3], 3);
    CVT_Q(d3, p3[0], 0); CVT_Q(d3, p3[1], 1); CVT_Q(d3, p3[2], 2); CVT_Q(d3, p3[3], 3);
    #undef CVT_Q

    int4_v* dst = reinterpret_cast<int4_v*>(A_fp4 + row_k / 2 + kg * 16);
    *dst = int4_v{static_cast<int>(d0), static_cast<int>(d1),
                  static_cast<int>(d2), static_cast<int>(d3)};

    #undef Q5_M
    #undef Q5_K
    #undef Q5_KS
    #undef Q5_K2
}

// Quant kernel for M=256, K=1536
__global__ void __launch_bounds__(128, 2)
quant_m256_k1536(
    const __bf16* __restrict__ A_bf16,
    uint8_t*      __restrict__ A_fp4,
    uint8_t*      __restrict__ A_scale)
{
    #define Q6_M 256
    #define Q6_K 1536
    #define Q6_KS 48
    #define Q6_K2 768

    const int group = blockIdx.x * 128 + threadIdx.x;
    const int row   = group / Q6_KS;
    const int kg    = group % Q6_KS;

    if (row >= Q6_M) return;

    const long row_k = (long)row * Q6_K;
    const __bf16* src = A_bf16 + row_k + kg * 32;

    const int4_v w0 = *reinterpret_cast<const int4_v*>(src);
    const int4_v w1 = *reinterpret_cast<const int4_v*>(src + 8);
    const int4_v w2 = *reinterpret_cast<const int4_v*>(src + 16);
    const int4_v w3 = *reinterpret_cast<const int4_v*>(src + 24);

    const uint32_t* p0 = reinterpret_cast<const uint32_t*>(&w0);
    const uint32_t* p1 = reinterpret_cast<const uint32_t*>(&w1);
    const uint32_t* p2 = reinterpret_cast<const uint32_t*>(&w2);
    const uint32_t* p3 = reinterpret_cast<const uint32_t*>(&w3);

    float absMax = 1e-10f;
    #define AMAX_P6(pair) { \
        __bf16 lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
        __bf16 hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
        float flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
        float fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
        absMax = (flo > absMax) ? flo : absMax; \
        absMax = (fhi > absMax) ? fhi : absMax; \
    }
    AMAX_P6(p0[0]) AMAX_P6(p0[1]) AMAX_P6(p0[2]) AMAX_P6(p0[3])
    AMAX_P6(p1[0]) AMAX_P6(p1[1]) AMAX_P6(p1[2]) AMAX_P6(p1[3])
    AMAX_P6(p2[0]) AMAX_P6(p2[1]) AMAX_P6(p2[2]) AMAX_P6(p2[3])
    AMAX_P6(p3[0]) AMAX_P6(p3[1]) AMAX_P6(p3[2]) AMAX_P6(p3[3])
    #undef AMAX_P6

    uint32_t u32 = __builtin_bit_cast(uint32_t, absMax);
    uint32_t amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
    uint32_t inv_exp  = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
    A_scale[row_k / 32 + kg] = static_cast<uint8_t>(inv_exp);

    float hw_scale = __builtin_bit_cast(float, inv_exp << 23);

    #define CVT_Q6(d, pair, sel) \
        d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)

    unsigned d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
    CVT_Q6(d0, p0[0], 0); CVT_Q6(d0, p0[1], 1); CVT_Q6(d0, p0[2], 2); CVT_Q6(d0, p0[3], 3);
    CVT_Q6(d1, p1[0], 0); CVT_Q6(d1, p1[1], 1); CVT_Q6(d1, p1[2], 2); CVT_Q6(d1, p1[3], 3);
    CVT_Q6(d2, p2[0], 0); CVT_Q6(d2, p2[1], 1); CVT_Q6(d2, p2[2], 2); CVT_Q6(d2, p2[3], 3);
    CVT_Q6(d3, p3[0], 0); CVT_Q6(d3, p3[1], 1); CVT_Q6(d3, p3[2], 2); CVT_Q6(d3, p3[3], 3);
    #undef CVT_Q6

    int4_v* dst = reinterpret_cast<int4_v*>(A_fp4 + row_k / 2 + kg * 16);
    *dst = int4_v{static_cast<int>(d0), static_cast<int>(d1),
                  static_cast<int>(d2), static_cast<int>(d3)};

    #undef Q6_M
    #undef Q6_K
    #undef Q6_KS
    #undef Q6_K2
}


// ============================================================================
// GEMM PATH B kernels for shapes 5 and 6
// s292: Replace global_load_dwordx4 with buffer_load_dwordx4 in ISSUE_LOADS
//
// buffer_load_dwordx4 v[dst:dst+3], v[voff], s[rsrc:rsrc+3], s[soff] offen
//   - rsrc: 128-bit buffer resource descriptor in SGPRs (base addr + config)
//   - soff: scalar offset in SGPR (wave-uniform K iteration offset)
//   - voff: 32-bit per-lane offset in VGPR (computed once, constant across K loop)
//
// This eliminates ALL v_lshl_add_u64 instructions that were needed for
// 64-bit pointer arithmetic with global_load_dwordx4.
//
// Shape 5: M=64, K=2048, N=7168, BM=32, BN=32, NWARPS=4
//   WAVES_M=2, WAVES_N=2, CHUNK_K=8, CKT=16, NUM_CHUNKS=2, TAIL=0
//   LDS layout constants (CHUNK_K=8):
//     A_ROW = 8*64+16 = 528
//     A_TILE = 16*528 = 8448
//     A_SIZE = 2*8448 = 16896
//     B_KGRP_STRIDE = 256
//     B_KT_STRIDE = 1280
//     B_TILE = 8*1280 = 10240
//     B_SIZE = 2*10240 = 20480
//     BUF_SIZE = 16896+20480 = 37376
//     TOTAL_LDS = 2*37376 = 74752
//   LoadCounts (NTHREADS=256):
//     A_TOTAL = 2*16*8*4 = 1024, A_PER_THREAD = 4
//     B_TOTAL = 2*8*4*16 = 1024, B_PER_THREAD = 4
//
// Shape 6: M=256, K=1536, N=3072, BM=64, BN=32, NWARPS=8
//   WAVES_M=4, WAVES_N=2, CHUNK_K=4, CKT=12, NUM_CHUNKS=3, TAIL=0
//   LDS layout:
//     A_ROW = 272
//     A_TILE = 4352
//     A_SIZE = 4*4352 = 17408
//     B_SIZE = 2*5120 = 10240
//     BUF_SIZE = 17408+10240 = 27648
//     TOTAL_LDS = 2*27648 = 55296
//   LoadCounts (NTHREADS=512):
//     A_TOTAL = 4*16*4*4 = 1024, A_PER_THREAD = 2
//     B_TOTAL = 2*4*4*16 = 512, B_PER_THREAD = 1
// ============================================================================

// Helper: LDS index calculations as inline functions
__device__ __forceinline__ int lds_a_idx(int wm, int row, int k_idx,
                                          int A_TILE, int A_ROW) {
    return wm * A_TILE + row * A_ROW + k_idx * 16;
}

__device__ __forceinline__ int lds_b_idx(int wn, int kt, int kgrp, int lrow,
                                          int B_TILE, int B_KT_STRIDE,
                                          int B_KGRP_STRIDE) {
    return wn * B_TILE + kt * B_KT_STRIDE + kgrp * B_KGRP_STRIDE + lrow * 16;
}


// ============================================================================
// BUFFER_LOAD_DWORDX4 inline ASM helper macro
// Loads 16 bytes (dwordx4) using buffer resource descriptor.
// rsrc = s[4] buffer resource, soff = s[1] scalar offset, voff = v[1] VGPR offset
// ============================================================================
// Buffer load via LLVM intrinsic (handles SGPR placement automatically)
__device__ int4_v __llvm_buffer_load_v4i32(int4_v rsrc, int voff, int soff, int aux)
    __asm("llvm.amdgcn.raw.buffer.load.v4i32");

#define BUFFER_LOAD_DWORDX4(dst, voff, rsrc, soff) \
    dst = __llvm_buffer_load_v4i32(*reinterpret_cast<const int4_v*>(&(rsrc)), static_cast<int>(voff), soff, 0)

#define BUFFER_LOAD_DWORDX4_NT(dst, voff, rsrc, soff) \
    dst = __llvm_buffer_load_v4i32(*reinterpret_cast<const int4_v*>(&(rsrc)), static_cast<int>(voff), soff, 2)


// GEMM for Shape 5: M=64, N=7168, K=2048
// s466: BM=16, BN=64 — 20% less HBM traffic per block (128KB vs 160KB)
// WAVES_M=1, WAVES_N=4. Grid: {112, 4}. OCC=3 (LDS-limited, sufficient).
__global__ void __launch_bounds__(256, 3)
__attribute__((amdgpu_flat_work_group_size(256, 256)))
gemm_m64_n7168_k2048(
    const __bf16*  __restrict__ A_bf16,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    uint16_t*      __restrict__ C_final)
{
    // All constants hardcoded
    #define G5_M 64
    #define G5_N 7168
    #define G5_K 2048
    #define G5_BM 16
    #define G5_BN 64
    #define G5_WAVES_M 1
    #define G5_WAVES_N 4
    #define G5_NWARPS 4
    // CK=4 for pipeline overlap
    #define G5_CHUNK_K 4
    #define G5_CKT 16
    #define G5_NUM_CHUNKS 4
    #define G5_SCALEN 64
    // LDS layout (CK=4: A_ROW=4*64+16=272)
    #define G5_A_ROW 272
    #define G5_A_TILE 4352
    #define G5_A_SIZE 4352
    #define G5_B_KGRP_STRIDE 256
    #define G5_B_KT_STRIDE 1280
    #define G5_B_TILE 5120
    #define G5_B_SIZE 20480
    #define G5_BUF_SIZE 24832
    // Load counts: A=1*16*16=256 (1/thread), B=4*4*4*16=1024 (4/thread)
    #define G5_NTHREADS 256
    #define G5_A_TOTAL 256
    #define G5_A_PER_THREAD 1
    #define G5_B_TOTAL 1024
    #define G5_B_PER_THREAD 4

    const long bsh_n_stride = (long)(G5_K / 64) * 512;

    const int tid    = static_cast<int>(threadIdx.x);
    const int lane   = tid % 64;
    const int wave   = tid / 64;
    const int wave_m = wave / G5_WAVES_N;
    const int wave_n = wave % G5_WAVES_N;

    const int tile_m_base = static_cast<int>(blockIdx.y) * G5_BM;
    const int tile_n_base = static_cast<int>(blockIdx.x) * G5_BN;
    const int tile_m = tile_m_base + wave_m * 16;
    const int tile_n = tile_n_base + wave_n * 16;

    const int lrow = lane % 16;
    const int kgrp = lane / 16;
    __builtin_assume(lrow >= 0 && lrow < 16);
    __builtin_assume(kgrp >= 0 && kgrp < 4);

    const int gn = tile_n + lrow;

    int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * G5_SCALEN);

    float4_v acc{0.f, 0.f, 0.f, 0.f};

    // PATH B: LDS double-buffered
    extern __shared__ uint8_t smem_raw[];
    uint8_t* buf0 = smem_raw;
    uint8_t* buf1 = smem_raw + G5_BUF_SIZE;

    // ================================================================
    // Build buffer resource descriptors (SGPR) for A_bf16 and Bsh
    // ================================================================
    const BufferRsrc rsrc_a_bf16 = make_buffer_rsrc(A_bf16);
    const BufferRsrc rsrc_bsh    = make_buffer_rsrc(Bsh);

    // ================================================================
    // Pre-compute per-thread offsets for A (bf16) and B
    // Each thread handles 4 quantization groups.
    // Each group: 32 bf16 = 64 bytes = 4 buffer_load_dwordx4
    // The bf16 base offset for group i:
    //   row * K * 2 + k_idx * 32 * 2 = row * 4096 + k_idx * 64
    // ================================================================

    // For A: we need bf16 offset (row in bf16 matrix, k_idx = group within chunk)
    unsigned a_bf16_voff[G5_A_PER_THREAD]; // base bf16 offset per group (load 4 dwordx4 from here)
    int  a_lds_offs[G5_A_PER_THREAD];      // FP4 LDS destination
    int  a_k_idx[G5_A_PER_THREAD];         // k_idx for scale tracking

    #pragma unroll
    for (int i = 0; i < G5_A_PER_THREAD; ++i) {
        int linear = i * G5_NTHREADS + tid;
        int k_idx      = linear % (G5_CHUNK_K * 4);
        int m_local    = (linear / (G5_CHUNK_K * 4)) % 16;
        int wave_m_idx = linear / (16 * G5_CHUNK_K * 4);
        int row        = tile_m_base + wave_m_idx * 16 + m_local;
        // bf16 offset: row * K * sizeof(bf16) + k_idx * 32 * sizeof(bf16)
        // = row * 4096 + k_idx * 64
        // But soffset will add ks_base * 256 (ks_base * 128_elements * 2_bytes)
        a_bf16_voff[i] = static_cast<unsigned>(row * G5_K * 2 + k_idx * 64);
        a_lds_offs[i]  = lds_a_idx(wave_m_idx, m_local, k_idx, G5_A_TILE, G5_A_ROW);
        a_k_idx[i]     = k_idx;
    }

    unsigned b_voff[G5_B_PER_THREAD];
    int  b_lds_offs[G5_B_PER_THREAD];

    #pragma unroll
    for (int i = 0; i < G5_B_PER_THREAD; ++i) {
        int linear = i * G5_NTHREADS + tid;
        int b_lrow     = linear % 16;
        int b_kgrp     = (linear / 16) % 4;
        int kt         = (linear / 64) % G5_CHUNK_K;
        int wave_n_idx = linear / (G5_CHUNK_K * 64);
        int b_tile_n   = tile_n_base + wave_n_idx * 16;
        int n_tile     = b_tile_n / 16;
        int k_half     = (b_kgrp & 1) * 256;
        int kgrp_half  = b_kgrp / 2;

        b_voff[i]      = static_cast<unsigned>((long)n_tile * bsh_n_stride + k_half
                         + (long)b_lrow * 16 + kt * 1024 + kgrp_half * 512);
        b_lds_offs[i]  = G5_A_SIZE + lds_b_idx(wave_n_idx, kt, b_kgrp, b_lrow,
                                                 G5_B_TILE, G5_B_KT_STRIDE, G5_B_KGRP_STRIDE);
    }

    int4_v b_regs[G5_B_PER_THREAD];

    // ================================================================
    // Scale storage: each thread quantizes 4 groups per chunk.
    // Store scales for use in COMPUTE phase.
    // a_scales[chunk][group_within_thread] — indexed by thread's group assignment
    // But COMPUTE_4MFMA needs scales indexed by (kt, kgrp) for the lane.
    // We'll store ALL scales for the chunk in a register array indexed by k_idx.
    // Each thread computes 4 scales (for its 4 k_idx values).
    // In COMPUTE, each lane needs scale for its own (row, k_tile, kgrp).
    // Since scales are per-row-per-group, and each lane in COMPUTE reads a
    // specific group, we need the scale for THAT group — which was computed
    // by whichever thread handled that (row, k_idx).
    // Problem: the thread that loaded group k_idx for row R is NOT necessarily
    // the same lane that will use it in COMPUTE.
    // Solution: write scales to LDS alongside FP4 data, or use global memory.
    // Simpler: write scales to a small LDS region.
    // ================================================================

    // Scale LDS: after the main LDS buffers. We need CK*4 = 32 scales per wave_m per row.
    // Actually: per chunk, 2 wave_m × 16 rows × 32 groups = 1024 scale bytes.
    // Place after BUF_SIZE in each buffer.
    // Actually that changes BUF_SIZE and might break the double buffering.
    // SIMPLER: just write scales to global memory as the quant kernel does,
    // using a small workspace. But that adds global stores + loads.
    //
    // SIMPLEST: Load A bf16, quantize, write FP4 to LDS, write scale to global A_scale.
    // But we don't have A_scale pointer... We removed it from the signature!
    //
    // Better approach: Since each lane in COMPUTE needs ONE scale value per MFMA
    // (its own row's scale for that k-tile's kgrp), and the lane's (row, kgrp)
    // is determined by (wave_m, lrow, kgrp), we can have each lane compute
    // ITS OWN scales for all k-tiles. This means: each lane loads bf16 for
    // its own row and quantizes. But the LOAD mapping is different from
    // the COMPUTE mapping.
    //
    // Alternative: restructure so that each lane in LOAD handles the same
    // (row, kgrp) as it will in COMPUTE. Then scales naturally match.
    // But the LOAD mapping distributes work evenly (1024 groups / 256 threads = 4 per thread)
    // while the COMPUTE mapping has each lane process kgrp=lane/16.
    //
    // Let me use LDS for scales: add a small scale region to LDS.
    // Per chunk: 2 wave_m * 16 rows * 32 groups = 1024 bytes. Small.
    // Add to BUF_SIZE.

    // Actually, I'll allocate scale storage in LDS within the existing padding.
    // Each A_ROW has 528 bytes but only 512 bytes of FP4 data (32 groups × 16 bytes).
    // The extra 16 bytes of padding per row are unused!
    // We have 2 wave_m × 16 rows = 32 rows, each with 16 bytes padding = 512 bytes.
    // But we need 32 scale bytes per row (one per group). 32 > 16. Doesn't fit.
    //
    // Plan B: Use a separate small LDS region for scales.
    // Per chunk scales: 32 rows × 32 groups per chunk = 1024 bytes.
    // Add after the main double buffer. Total LDS = 2 * 37376 + 1024 = 75776.
    // With OCC=2: 2 × 75776 = 151552 ≤ 163840. Still fits!
    //
    // Actually, we don't need to double-buffer the scales. The scales are produced
    // and consumed within the same buffer step. So just 1 × 1024 = 1024 bytes.
    // Total LDS = 2 * 37376 + 1024 = 75776. OCC check: 2 × 75776 = 151552 ≤ 163840. ✓

    // Scale LDS region: double-buffered (2 × 512 bytes) after the two main buffers
    // With WAVES_M=1: 1 × 16 rows × 16 groups, stride=32 → 512 bytes per buffer
    #define G5_SCALE_LDS_SIZE 512
    #define G5_SCALE_LDS_BASE0 (2 * G5_BUF_SIZE)
    #define G5_SCALE_LDS_BASE1 (2 * G5_BUF_SIZE + G5_SCALE_LDS_SIZE)
    #define G5_SCALE_LDS_STRIDE 32

    // ================================================================
    // FUSED ISSUE_LOADS: Load bf16 A, quantize to FP4, write FP4+scales to LDS.
    // Then load B as before.
    // ================================================================

    #define G5_ISSUE_LOADS_FUSED(ks_base, buf, scale_base) \
    { \
        const unsigned _b_soff = __builtin_amdgcn_readfirstlane(static_cast<unsigned>((ks_base) * 1024)); \
        /* bf16 loads use pointer-based access, no buffer resource needed */ \
        /* Issue B loads first (to HBM, high latency) */ \
        _Pragma("unroll") \
        for (int i = 0; i < G5_B_PER_THREAD; ++i) { \
            BUFFER_LOAD_DWORDX4_NT(b_regs[i], b_voff[i], rsrc_bsh, _b_soff); \
        } \
        /* For each A group: load 64 bytes bf16, quantize, write FP4 to LDS */ \
        const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
        _Pragma("unroll") \
        for (int i = 0; i < G5_A_PER_THREAD; ++i) { \
            /* Load 64 bytes of bf16 (32 bf16 values) via regular pointer loads */ \
            int4_v _w0, _w1, _w2, _w3; \
            { \
                int _lin = i * G5_NTHREADS + tid; \
                int _kidx = _lin % (G5_CHUNK_K * 4); \
                int _mloc = (_lin / (G5_CHUNK_K * 4)) % 16; \
                int _wmid = _lin / (16 * G5_CHUNK_K * 4); \
                int _row = tile_m_base + _wmid * 16 + _mloc; \
                const int4_v* _src = reinterpret_cast<const int4_v*>( \
                    reinterpret_cast<const uint8_t*>(A_bf16) + \
                    (long)_row * G5_K * 2 + (long)((ks_base) * 128 + _kidx * 32) * 2); \
                _w0 = _src[0]; \
                _w1 = _src[1]; \
                _w2 = _src[2]; \
                _w3 = _src[3]; \
            } \
            /* Quantize: find max abs, compute scale, convert to FP4 */ \
            const uint32_t* _p0 = reinterpret_cast<const uint32_t*>(&_w0); \
            const uint32_t* _p1 = reinterpret_cast<const uint32_t*>(&_w1); \
            const uint32_t* _p2 = reinterpret_cast<const uint32_t*>(&_w2); \
            const uint32_t* _p3 = reinterpret_cast<const uint32_t*>(&_w3); \
            float _absMax = 1e-10f; \
            { \
                auto _AM = [&](uint32_t pair) { \
                    __bf16 lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
                    __bf16 hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair >> 16)); \
                    float flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
                    float fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
                    _absMax = (flo > _absMax) ? flo : _absMax; \
                    _absMax = (fhi > _absMax) ? fhi : _absMax; \
                }; \
                _AM(_p0[0]); _AM(_p0[1]); _AM(_p0[2]); _AM(_p0[3]); \
                _AM(_p1[0]); _AM(_p1[1]); _AM(_p1[2]); _AM(_p1[3]); \
                _AM(_p2[0]); _AM(_p2[1]); _AM(_p2[2]); _AM(_p2[3]); \
                _AM(_p3[0]); _AM(_p3[1]); _AM(_p3[2]); _AM(_p3[3]); \
            } \
            uint32_t _u32 = __builtin_bit_cast(uint32_t, _absMax); \
            uint32_t _amax_exp = ((_u32 + 0x200000u) >> 23) & 0xFFu; \
            uint32_t _inv_exp = (_amax_exp >= 2u) ? (_amax_exp - 2u) : 0u; \
            float _hw_scale = __builtin_bit_cast(float, _inv_exp << 23); \
            unsigned _d0 = 0u, _d1 = 0u, _d2 = 0u, _d3 = 0u; \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[0]), _hw_scale, 0); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[1]), _hw_scale, 1); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[2]), _hw_scale, 2); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[3]), _hw_scale, 3); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[0]), _hw_scale, 0); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[1]), _hw_scale, 1); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[2]), _hw_scale, 2); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[3]), _hw_scale, 3); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[0]), _hw_scale, 0); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[1]), _hw_scale, 1); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[2]), _hw_scale, 2); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[3]), _hw_scale, 3); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[0]), _hw_scale, 0); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[1]), _hw_scale, 1); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[2]), _hw_scale, 2); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[3]), _hw_scale, 3); \
            int4_v _fp4 = int4_v{static_cast<int>(_d0), static_cast<int>(_d1), \
                                 static_cast<int>(_d2), static_cast<int>(_d3)}; \
            /* Write FP4 to LDS (same layout as before) */ \
            { \
                const unsigned _off = _buf_off + a_lds_offs[i]; \
                asm volatile("ds_write_b128 %0, %1 offset:%2" \
                    :: "v"(_off), "v"(_fp4), "n"(0) : "memory"); \
            } \
            /* Write scale to LDS */ \
            { \
                int _linear = i * G5_NTHREADS + tid; \
                int _k_idx = _linear % (G5_CHUNK_K * 4); \
                int _m_local = (_linear / (G5_CHUNK_K * 4)) % 16; \
                int _wave_m_idx = _linear / (16 * G5_CHUNK_K * 4); \
                unsigned _scale_off = (scale_base) + _wave_m_idx * 512 + _m_local * 32 + _k_idx; \
                smem_raw[_scale_off] = static_cast<uint8_t>(_inv_exp); \
            } \
        } \
        /* Wait for B loads, write B to LDS */ \
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); \
        _Pragma("unroll") \
        for (int i = 0; i < G5_B_PER_THREAD; ++i) { \
            const unsigned _off = _buf_off + b_lds_offs[i]; \
            asm volatile("ds_write_b128 %0, %1 offset:%2" \
                :: "v"(_off), "v"(b_regs[i]), "n"(0) : "memory"); \
        } \
    }

    /* G5_COMPUTE_4MFMA: CK-style sched_group_barrier scheduling.
     * Uses C++ builtins + sched_group_barrier for precise ds_read/MFMA interleaving.
     * The compiler can schedule VALU/address math in MFMA stall slots.
     * 0x100 = DS_READ, 0x008 = MFMA */
    #define G5_COMPUTE_4MFMA(buf, chunk_ks, kt_off, scale_base) \
    { \
        const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
        const unsigned _a_base = _buf_off; \
        const unsigned _b_base = _buf_off + G5_A_SIZE; \
        const unsigned _addr_a0 = _a_base + lds_a_idx(wave_m, lrow, ((kt_off)+0) * 4 + kgrp, G5_A_TILE, G5_A_ROW); \
        const unsigned _addr_b0 = _b_base + lds_b_idx(wave_n, (kt_off)+0, kgrp, lrow, G5_B_TILE, G5_B_KT_STRIDE, G5_B_KGRP_STRIDE); \
        const unsigned _addr_a1 = _a_base + lds_a_idx(wave_m, lrow, ((kt_off)+1) * 4 + kgrp, G5_A_TILE, G5_A_ROW); \
        const unsigned _addr_b1 = _b_base + lds_b_idx(wave_n, (kt_off)+1, kgrp, lrow, G5_B_TILE, G5_B_KT_STRIDE, G5_B_KGRP_STRIDE); \
        const unsigned _addr_a2 = _a_base + lds_a_idx(wave_m, lrow, ((kt_off)+2) * 4 + kgrp, G5_A_TILE, G5_A_ROW); \
        const unsigned _addr_b2 = _b_base + lds_b_idx(wave_n, (kt_off)+2, kgrp, lrow, G5_B_TILE, G5_B_KT_STRIDE, G5_B_KGRP_STRIDE); \
        const unsigned _addr_a3 = _a_base + lds_a_idx(wave_m, lrow, ((kt_off)+3) * 4 + kgrp, G5_A_TILE, G5_A_ROW); \
        const unsigned _addr_b3 = _b_base + lds_b_idx(wave_n, (kt_off)+3, kgrp, lrow, G5_B_TILE, G5_B_KT_STRIDE, G5_B_KGRP_STRIDE); \
        int _sa0, _sb0, _sa1, _sb1, _sa2, _sb2, _sa3, _sb3; \
        { \
            unsigned _sc_base = (scale_base) + wave_m * 512 + lrow * 32; \
            _sa0 = static_cast<int>(smem_raw[_sc_base + ((kt_off)+0) * 4 + kgrp]); \
            _sa1 = static_cast<int>(smem_raw[_sc_base + ((kt_off)+1) * 4 + kgrp]); \
            _sa2 = static_cast<int>(smem_raw[_sc_base + ((kt_off)+2) * 4 + kgrp]); \
            _sa3 = static_cast<int>(smem_raw[_sc_base + ((kt_off)+3) * 4 + kgrp]); \
        } \
        { \
            int _ksv0 = (chunk_ks) + (kt_off); \
            const uint8_t* _bp0 = Bssh + bssh_base + (_ksv0 & 1) * 2 + (_ksv0 >> 1) * 256; \
            const uint8_t* _bp1 = Bssh + bssh_base + ((_ksv0+1) & 1) * 2 + ((_ksv0+1) >> 1) * 256; \
            const uint8_t* _bp2 = Bssh + bssh_base + ((_ksv0+2) & 1) * 2 + ((_ksv0+2) >> 1) * 256; \
            const uint8_t* _bp3 = Bssh + bssh_base + ((_ksv0+3) & 1) * 2 + ((_ksv0+3) >> 1) * 256; \
            _sb0 = static_cast<int>(*_bp0); \
            _sb1 = static_cast<int>(*_bp1); \
            _sb2 = static_cast<int>(*_bp2); \
            _sb3 = static_cast<int>(*_bp3); \
        } \
        /* CK-style scheduling: sched_barrier(0) creates island, */ \
        /* sched_group_barrier controls ds_read(0x100)/MFMA(0x008) interleaving */ \
        __builtin_amdgcn_sched_barrier(0); \
        asm volatile("s_setprio 3" ::: "memory"); \
        /* Phase 1: 4 ds_reads for MFMA 0 and 1 */ \
        __builtin_amdgcn_sched_group_barrier(0x100, 4, 0); \
        int4_v _a0 = *reinterpret_cast<const int4_v*>(smem_raw + _addr_a0); \
        int4_v _b0 = *reinterpret_cast<const int4_v*>(smem_raw + _addr_b0); \
        int4_v _a1 = *reinterpret_cast<const int4_v*>(smem_raw + _addr_a1); \
        int4_v _b1 = *reinterpret_cast<const int4_v*>(smem_raw + _addr_b1); \
        /* MFMA 0 */ \
        __builtin_amdgcn_sched_group_barrier(0x008, 1, 0); \
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
            _a0, _b0, acc, 4, 4, 0, _sa0, 0, _sb0); \
        /* Phase 2: 2 ds_reads for MFMA 2 */ \
        __builtin_amdgcn_sched_group_barrier(0x100, 2, 0); \
        int4_v _a2 = *reinterpret_cast<const int4_v*>(smem_raw + _addr_a2); \
        int4_v _b2 = *reinterpret_cast<const int4_v*>(smem_raw + _addr_b2); \
        /* MFMA 1 */ \
        __builtin_amdgcn_sched_group_barrier(0x008, 1, 0); \
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
            _a1, _b1, acc, 4, 4, 0, _sa1, 0, _sb1); \
        /* Phase 3: 2 ds_reads for MFMA 3 */ \
        __builtin_amdgcn_sched_group_barrier(0x100, 2, 0); \
        int4_v _a3 = *reinterpret_cast<const int4_v*>(smem_raw + _addr_a3); \
        int4_v _b3 = *reinterpret_cast<const int4_v*>(smem_raw + _addr_b3); \
        /* MFMA 2 */ \
        __builtin_amdgcn_sched_group_barrier(0x008, 1, 0); \
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
            _a2, _b2, acc, 4, 4, 0, _sa2, 0, _sb2); \
        /* MFMA 3 */ \
        __builtin_amdgcn_sched_group_barrier(0x008, 1, 0); \
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
            _a3, _b3, acc, 4, 4, 0, _sa3, 0, _sb3); \
        asm volatile("s_setprio 0" ::: "memory"); \
        __builtin_amdgcn_sched_barrier(0); \
    }

    /* G5_COMPUTE_CHUNK: Process 4 k-tiles (CK=4) — single COMPUTE_4MFMA call */
    #define G5_COMPUTE_CHUNK(buf, chunk_ks, scale_base) \
    { \
        G5_COMPUTE_4MFMA(buf, chunk_ks, 0, scale_base); \
    }

    // s481: 4-phase pipeline — issue VMEM before COMPUTE to overlap latency
    // Persistent A bf16 registers (live across COMPUTE)
    int4_v _pst_w0, _pst_w1, _pst_w2, _pst_w3;

    // Phase macros for split pipeline
    #define G5_ISSUE_VMEM(ks_base) \
    { \
        const unsigned _b_soff = __builtin_amdgcn_readfirstlane(static_cast<unsigned>((ks_base) * 1024)); \
        _Pragma("unroll") \
        for (int i = 0; i < G5_B_PER_THREAD; ++i) { \
            BUFFER_LOAD_DWORDX4_NT(b_regs[i], b_voff[i], rsrc_bsh, _b_soff); \
        } \
        { \
            int _lin = tid; \
            int _kidx = _lin % (G5_CHUNK_K * 4); \
            int _mloc = (_lin / (G5_CHUNK_K * 4)) % 16; \
            int _wmid = _lin / (16 * G5_CHUNK_K * 4); \
            int _row = tile_m_base + _wmid * 16 + _mloc; \
            const int4_v* _src = reinterpret_cast<const int4_v*>( \
                reinterpret_cast<const uint8_t*>(A_bf16) + \
                (long)_row * G5_K * 2 + (long)((ks_base) * 128 + _kidx * 32) * 2); \
            _pst_w0 = _src[0]; \
            _pst_w1 = _src[1]; \
            _pst_w2 = _src[2]; \
            _pst_w3 = _src[3]; \
        } \
    }

    #define G5_COMPLETE_QUANT_STORE(buf, scale_base) \
    { \
        const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
        { \
            const uint32_t* _p0 = reinterpret_cast<const uint32_t*>(&_pst_w0); \
            const uint32_t* _p1 = reinterpret_cast<const uint32_t*>(&_pst_w1); \
            const uint32_t* _p2 = reinterpret_cast<const uint32_t*>(&_pst_w2); \
            const uint32_t* _p3 = reinterpret_cast<const uint32_t*>(&_pst_w3); \
            float _absMax = 1e-10f; \
            { \
                auto _AM = [&](uint32_t pair) { \
                    __bf16 lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
                    __bf16 hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair >> 16)); \
                    float flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
                    float fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
                    _absMax = (flo > _absMax) ? flo : _absMax; \
                    _absMax = (fhi > _absMax) ? fhi : _absMax; \
                }; \
                _AM(_p0[0]); _AM(_p0[1]); _AM(_p0[2]); _AM(_p0[3]); \
                _AM(_p1[0]); _AM(_p1[1]); _AM(_p1[2]); _AM(_p1[3]); \
                _AM(_p2[0]); _AM(_p2[1]); _AM(_p2[2]); _AM(_p2[3]); \
                _AM(_p3[0]); _AM(_p3[1]); _AM(_p3[2]); _AM(_p3[3]); \
            } \
            uint32_t _u32 = __builtin_bit_cast(uint32_t, _absMax); \
            uint32_t _amax_exp = ((_u32 + 0x200000u) >> 23) & 0xFFu; \
            uint32_t _inv_exp = (_amax_exp >= 2u) ? (_amax_exp - 2u) : 0u; \
            float _hw_scale = __builtin_bit_cast(float, _inv_exp << 23); \
            unsigned _d0 = 0u, _d1 = 0u, _d2 = 0u, _d3 = 0u; \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[0]), _hw_scale, 0); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[1]), _hw_scale, 1); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[2]), _hw_scale, 2); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[3]), _hw_scale, 3); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[0]), _hw_scale, 0); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[1]), _hw_scale, 1); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[2]), _hw_scale, 2); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[3]), _hw_scale, 3); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[0]), _hw_scale, 0); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[1]), _hw_scale, 1); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[2]), _hw_scale, 2); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[3]), _hw_scale, 3); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[0]), _hw_scale, 0); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[1]), _hw_scale, 1); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[2]), _hw_scale, 2); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[3]), _hw_scale, 3); \
            int4_v _fp4 = int4_v{static_cast<int>(_d0), static_cast<int>(_d1), \
                                 static_cast<int>(_d2), static_cast<int>(_d3)}; \
            { \
                const unsigned _off = _buf_off + a_lds_offs[0]; \
                asm volatile("ds_write_b128 %0, %1 offset:%2" \
                    :: "v"(_off), "v"(_fp4), "n"(0) : "memory"); \
            } \
            { \
                int _linear = tid; \
                int _k_idx = _linear % (G5_CHUNK_K * 4); \
                int _m_local = (_linear / (G5_CHUNK_K * 4)) % 16; \
                int _wave_m_idx = _linear / (16 * G5_CHUNK_K * 4); \
                unsigned _scale_off = (scale_base) + _wave_m_idx * 512 + _m_local * 32 + _k_idx; \
                smem_raw[_scale_off] = static_cast<uint8_t>(_inv_exp); \
            } \
        } \
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); \
        _Pragma("unroll") \
        for (int i = 0; i < G5_B_PER_THREAD; ++i) { \
            const unsigned _off = _buf_off + b_lds_offs[i]; \
            asm volatile("ds_write_b128 %0, %1 offset:%2" \
                :: "v"(_off), "v"(b_regs[i]), "n"(0) : "memory"); \
        } \
    }

    int cur_ks = 0;
    unsigned cur_scale_base = G5_SCALE_LDS_BASE0;
    unsigned nxt_scale_base = G5_SCALE_LDS_BASE1;

    // Initial prefetch: full load for chunk 0
    G5_ISSUE_LOADS_FUSED(cur_ks, buf0, cur_scale_base);
    asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
    __syncthreads();

    uint8_t* cur_buf = buf0;
    uint8_t* nxt_buf = buf1;

    // 4-phase pipeline loop: VMEM→COMPUTE→QUANT_STORE→sync
    #pragma unroll
    for (int c = 0; c < G5_NUM_CHUNKS - 1; ++c) {
        G5_ISSUE_VMEM(cur_ks + G5_CHUNK_K);              // Phase 1: issue VMEM (fast)
        G5_COMPUTE_CHUNK(cur_buf, cur_ks, cur_scale_base); // Phase 2: compute (overlaps VMEM)
        G5_COMPLETE_QUANT_STORE(nxt_buf, nxt_scale_base);  // Phase 3: wait + quant + store
        asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
        __syncthreads();                                    // Phase 4: sync

        uint8_t* tmp = cur_buf;
        cur_buf = nxt_buf;
        nxt_buf = tmp;
        unsigned tmp_scale = cur_scale_base;
        cur_scale_base = nxt_scale_base;
        nxt_scale_base = tmp_scale;
        cur_ks += G5_CHUNK_K;
    }

    G5_COMPUTE_CHUNK(cur_buf, cur_ks, cur_scale_base);

    #undef G5_ISSUE_VMEM
    #undef G5_COMPLETE_QUANT_STORE

    #undef G5_ISSUE_LOADS_FUSED
    #undef G5_COMPUTE_4MFMA
    #undef G5_COMPUTE_CHUNK

    // Write output — all tiles full, all rows valid
    const int out_col      = tile_n + lrow;
    const int out_row_base = tile_m + kgrp * 4;

    uint16_t* c_out = C_final + (long)out_row_base * G5_N + out_col;
    #pragma unroll
    for (int i = 0; i < 4; ++i, c_out += G5_N) {
        *c_out = fast_f32_to_bf16(acc[i]);
    }

    #undef G5_M
    #undef G5_N
    #undef G5_K
    #undef G5_BM
    #undef G5_BN
    #undef G5_WAVES_M
    #undef G5_WAVES_N
    #undef G5_NWARPS
    #undef G5_CHUNK_K
    #undef G5_CKT
    #undef G5_NUM_CHUNKS
    #undef G5_SCALEN
    #undef G5_A_ROW
    #undef G5_A_TILE
    #undef G5_A_SIZE
    #undef G5_B_KGRP_STRIDE
    #undef G5_B_KT_STRIDE
    #undef G5_B_TILE
    #undef G5_B_SIZE
    #undef G5_BUF_SIZE
    #undef G5_NTHREADS
    #undef G5_A_TOTAL
    #undef G5_A_PER_THREAD
    #undef G5_B_TOTAL
    #undef G5_B_PER_THREAD
    #undef G5_SCALE_LDS_SIZE
    #undef G5_SCALE_LDS_BASE0
    #undef G5_SCALE_LDS_BASE1
    #undef G5_SCALE_LDS_STRIDE
}


// GEMM for Shape 6: M=256, N=3072, K=1536
// s468: BM=16, BN=64 with G5-style scale handling (stride 32, byte B-scale loads)
__global__ void __launch_bounds__(256, 3)
__attribute__((amdgpu_flat_work_group_size(256, 256)))
gemm_m256_n3072_k1536(
    const __bf16*  __restrict__ A_bf16,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    uint16_t*      __restrict__ C_final)
{
    #define G6_M 256
    #define G6_N 3072
    #define G6_K 1536
    #define G6_BM 16
    #define G6_BN 64
    #define G6_WAVES_M 1
    #define G6_WAVES_N 4
    #define G6_NWARPS 4
    #define G6_CHUNK_K 4
    #define G6_CKT 12
    #define G6_NUM_CHUNKS 3
    #define G6_SCALEN 48
    // LDS layout (CK=4: A_ROW=4*64+16=272)
    #define G6_A_ROW 272
    #define G6_A_TILE 4352
    #define G6_A_SIZE 4352
    #define G6_B_KGRP_STRIDE 256
    #define G6_B_KT_STRIDE 1280
    #define G6_B_TILE 5120
    #define G6_B_SIZE 20480
    #define G6_BUF_SIZE 24832
    // Load counts: A=1*16*16=256 (1/thread), B=4*4*4*16=1024 (4/thread)
    #define G6_NTHREADS 256
    #define G6_A_TOTAL 256
    #define G6_A_PER_THREAD 1
    #define G6_B_TOTAL 1024
    #define G6_B_PER_THREAD 4

    const long bsh_n_stride = (long)(G6_K / 64) * 512;

    const int tid    = static_cast<int>(threadIdx.x);
    const int lane   = tid % 64;
    const int wave   = tid / 64;
    const int wave_m = wave / G6_WAVES_N;
    const int wave_n = wave % G6_WAVES_N;

    const int tile_m_base = static_cast<int>(blockIdx.y) * G6_BM;
    const int tile_n_base = static_cast<int>(blockIdx.x) * G6_BN;
    const int tile_m = tile_m_base + wave_m * 16;
    const int tile_n = tile_n_base + wave_n * 16;

    const int lrow = lane % 16;
    const int kgrp = lane / 16;
    __builtin_assume(lrow >= 0 && lrow < 16);
    __builtin_assume(kgrp >= 0 && kgrp < 4);

    const int gn = tile_n + lrow;

    int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * G6_SCALEN);

    float4_v acc{0.f, 0.f, 0.f, 0.f};

    extern __shared__ uint8_t smem_raw[];
    uint8_t* buf0 = smem_raw;
    uint8_t* buf1 = smem_raw + G6_BUF_SIZE;

    // Build buffer resource descriptor for Bsh (SGPR)
    const BufferRsrc rsrc_bsh = make_buffer_rsrc(Bsh);

    // Pre-compute per-thread A LDS offsets (for FP4 writes)
    int  a_lds_offs[G6_A_PER_THREAD];

    #pragma unroll
    for (int i = 0; i < G6_A_PER_THREAD; ++i) {
        int linear = i * G6_NTHREADS + tid;
        int k_idx      = linear % (G6_CHUNK_K * 4);
        int m_local    = (linear / (G6_CHUNK_K * 4)) % 16;
        int wave_m_idx = linear / (16 * G6_CHUNK_K * 4);
        a_lds_offs[i] = lds_a_idx(wave_m_idx, m_local, k_idx, G6_A_TILE, G6_A_ROW);
    }

    unsigned b_voff[G6_B_PER_THREAD];
    int  b_lds_offs[G6_B_PER_THREAD];

    #pragma unroll
    for (int i = 0; i < G6_B_PER_THREAD; ++i) {
        int linear = i * G6_NTHREADS + tid;
        int b_lrow     = linear % 16;
        int b_kgrp     = (linear / 16) % 4;
        int kt         = (linear / 64) % G6_CHUNK_K;
        int wave_n_idx = linear / (G6_CHUNK_K * 64);
        int b_tile_n   = tile_n_base + wave_n_idx * 16;
        int n_tile     = b_tile_n / 16;
        int k_half     = (b_kgrp & 1) * 256;
        int kgrp_half  = b_kgrp / 2;

        b_voff[i]      = static_cast<unsigned>((long)n_tile * bsh_n_stride + k_half
                         + (long)b_lrow * 16 + kt * 1024 + kgrp_half * 512);
        b_lds_offs[i]  = G6_A_SIZE + lds_b_idx(wave_n_idx, kt, b_kgrp, b_lrow,
                                                 G6_B_TILE, G6_B_KT_STRIDE, G6_B_KGRP_STRIDE);
    }

    int4_v b_regs[G6_B_PER_THREAD];

    // Scale LDS region: G5-style stride 32 per row, 512 per wave_m
    // WAVES_M=1: max offset = 15*32+15 = 495 → 512 bytes per buffer
    #define G6_SCALE_LDS_SIZE 512
    #define G6_SCALE_LDS_BASE0 (2 * G6_BUF_SIZE)
    #define G6_SCALE_LDS_BASE1 (2 * G6_BUF_SIZE + G6_SCALE_LDS_SIZE)

    // ================================================================
    // FUSED ISSUE_LOADS: Load bf16 A, quantize to FP4, write FP4+scales to LDS.
    // Then load B as before.
    // ================================================================

    #define G6_ISSUE_LOADS_FUSED(ks_base, buf, scale_base) \
    { \
        const unsigned _b_soff = __builtin_amdgcn_readfirstlane(static_cast<unsigned>((ks_base) * 1024)); \
        /* Issue B loads first (to HBM, high latency) */ \
        _Pragma("unroll") \
        for (int i = 0; i < G6_B_PER_THREAD; ++i) { \
            BUFFER_LOAD_DWORDX4_NT(b_regs[i], b_voff[i], rsrc_bsh, _b_soff); \
        } \
        /* For each A group: load 64 bytes bf16, quantize, write FP4 to LDS */ \
        const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
        _Pragma("unroll") \
        for (int i = 0; i < G6_A_PER_THREAD; ++i) { \
            /* Load 64 bytes of bf16 (32 bf16 values) via regular pointer loads */ \
            int4_v _w0, _w1, _w2, _w3; \
            int _lin = i * G6_NTHREADS + tid; \
            int _kidx = _lin % (G6_CHUNK_K * 4); \
            int _mloc = (_lin / (G6_CHUNK_K * 4)) % 16; \
            int _wmid = _lin / (16 * G6_CHUNK_K * 4); \
            int _row = tile_m_base + _wmid * 16 + _mloc; \
            { \
                const int4_v* _src = reinterpret_cast<const int4_v*>( \
                    reinterpret_cast<const uint8_t*>(A_bf16) + \
                    (long)_row * G6_K * 2 + (long)((ks_base) * 128 + _kidx * 32) * 2); \
                _w0 = _src[0]; \
                _w1 = _src[1]; \
                _w2 = _src[2]; \
                _w3 = _src[3]; \
            } \
            /* Quantize: find max abs, compute scale, convert to FP4 */ \
            const uint32_t* _p0 = reinterpret_cast<const uint32_t*>(&_w0); \
            const uint32_t* _p1 = reinterpret_cast<const uint32_t*>(&_w1); \
            const uint32_t* _p2 = reinterpret_cast<const uint32_t*>(&_w2); \
            const uint32_t* _p3 = reinterpret_cast<const uint32_t*>(&_w3); \
            float _absMax = 1e-10f; \
            { \
                auto _AM = [&](uint32_t pair) { \
                    __bf16 lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
                    __bf16 hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair >> 16)); \
                    float flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
                    float fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
                    _absMax = (flo > _absMax) ? flo : _absMax; \
                    _absMax = (fhi > _absMax) ? fhi : _absMax; \
                }; \
                _AM(_p0[0]); _AM(_p0[1]); _AM(_p0[2]); _AM(_p0[3]); \
                _AM(_p1[0]); _AM(_p1[1]); _AM(_p1[2]); _AM(_p1[3]); \
                _AM(_p2[0]); _AM(_p2[1]); _AM(_p2[2]); _AM(_p2[3]); \
                _AM(_p3[0]); _AM(_p3[1]); _AM(_p3[2]); _AM(_p3[3]); \
            } \
            uint32_t _u32 = __builtin_bit_cast(uint32_t, _absMax); \
            uint32_t _amax_exp = ((_u32 + 0x200000u) >> 23) & 0xFFu; \
            uint32_t _inv_exp = (_amax_exp >= 2u) ? (_amax_exp - 2u) : 0u; \
            float _hw_scale = __builtin_bit_cast(float, _inv_exp << 23); \
            unsigned _d0 = 0u, _d1 = 0u, _d2 = 0u, _d3 = 0u; \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[0]), _hw_scale, 0); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[1]), _hw_scale, 1); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[2]), _hw_scale, 2); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[3]), _hw_scale, 3); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[0]), _hw_scale, 0); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[1]), _hw_scale, 1); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[2]), _hw_scale, 2); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[3]), _hw_scale, 3); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[0]), _hw_scale, 0); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[1]), _hw_scale, 1); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[2]), _hw_scale, 2); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[3]), _hw_scale, 3); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[0]), _hw_scale, 0); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[1]), _hw_scale, 1); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[2]), _hw_scale, 2); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[3]), _hw_scale, 3); \
            int4_v _fp4 = int4_v{static_cast<int>(_d0), static_cast<int>(_d1), \
                                 static_cast<int>(_d2), static_cast<int>(_d3)}; \
            /* Write FP4 to LDS (same layout as before) */ \
            { \
                const unsigned _off = _buf_off + a_lds_offs[i]; \
                asm volatile("ds_write_b128 %0, %1 offset:%2" \
                    :: "v"(_off), "v"(_fp4), "n"(0) : "memory"); \
            } \
            /* Write scale to LDS: G5-style stride=32 per row, 512 per wave_m */ \
            { \
                unsigned _scale_off = (scale_base) + _wmid * 512 + _mloc * 32 + _kidx; \
                smem_raw[_scale_off] = static_cast<uint8_t>(_inv_exp); \
            } \
        } \
        /* Wait for B loads, write B to LDS */ \
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); \
        _Pragma("unroll") \
        for (int i = 0; i < G6_B_PER_THREAD; ++i) { \
            const unsigned _off = _buf_off + b_lds_offs[i]; \
            asm volatile("ds_write_b128 %0, %1 offset:%2" \
                :: "v"(_off), "v"(b_regs[i]), "n"(0) : "memory"); \
        } \
    }

    /* G6_COMPUTE_CHUNK: Process 4 k-tiles using inline ASM (original s468 style) */
    #define G6_COMPUTE_CHUNK(buf, chunk_ks, scale_base) \
    { \
        const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
        const unsigned _a_base = _buf_off; \
        const unsigned _b_base = _buf_off + G6_A_SIZE; \
        const unsigned _addr_a0 = _a_base + lds_a_idx(wave_m, lrow, 0 * 4 + kgrp, G6_A_TILE, G6_A_ROW); \
        const unsigned _addr_b0 = _b_base + lds_b_idx(wave_n, 0, kgrp, lrow, G6_B_TILE, G6_B_KT_STRIDE, G6_B_KGRP_STRIDE); \
        const unsigned _addr_a1 = _a_base + lds_a_idx(wave_m, lrow, 1 * 4 + kgrp, G6_A_TILE, G6_A_ROW); \
        const unsigned _addr_b1 = _b_base + lds_b_idx(wave_n, 1, kgrp, lrow, G6_B_TILE, G6_B_KT_STRIDE, G6_B_KGRP_STRIDE); \
        const unsigned _addr_a2 = _a_base + lds_a_idx(wave_m, lrow, 2 * 4 + kgrp, G6_A_TILE, G6_A_ROW); \
        const unsigned _addr_b2 = _b_base + lds_b_idx(wave_n, 2, kgrp, lrow, G6_B_TILE, G6_B_KT_STRIDE, G6_B_KGRP_STRIDE); \
        const unsigned _addr_a3 = _a_base + lds_a_idx(wave_m, lrow, 3 * 4 + kgrp, G6_A_TILE, G6_A_ROW); \
        const unsigned _addr_b3 = _b_base + lds_b_idx(wave_n, 3, kgrp, lrow, G6_B_TILE, G6_B_KT_STRIDE, G6_B_KGRP_STRIDE); \
        int _sa0, _sb0, _sa1, _sb1, _sa2, _sb2, _sa3, _sb3; \
        { \
            unsigned _sc_base = (scale_base) + wave_m * 512 + lrow * 32; \
            _sa0 = static_cast<int>(smem_raw[_sc_base + 0 * 4 + kgrp]); \
            _sa1 = static_cast<int>(smem_raw[_sc_base + 1 * 4 + kgrp]); \
            _sa2 = static_cast<int>(smem_raw[_sc_base + 2 * 4 + kgrp]); \
            _sa3 = static_cast<int>(smem_raw[_sc_base + 3 * 4 + kgrp]); \
        } \
        { \
            int _ksv0 = (chunk_ks); \
            const uint8_t* _bp0 = Bssh + bssh_base + (_ksv0 & 1) * 2 + (_ksv0 >> 1) * 256; \
            const uint8_t* _bp1 = Bssh + bssh_base + ((_ksv0+1) & 1) * 2 + ((_ksv0+1) >> 1) * 256; \
            const uint8_t* _bp2 = Bssh + bssh_base + ((_ksv0+2) & 1) * 2 + ((_ksv0+2) >> 1) * 256; \
            const uint8_t* _bp3 = Bssh + bssh_base + ((_ksv0+3) & 1) * 2 + ((_ksv0+3) >> 1) * 256; \
            _sb0 = static_cast<int>(*_bp0); \
            _sb1 = static_cast<int>(*_bp1); \
            _sb2 = static_cast<int>(*_bp2); \
            _sb3 = static_cast<int>(*_bp3); \
        } \
        int4_v _ae, _be, _ao, _bo; \
        __builtin_amdgcn_sched_barrier(0x020); \
        asm volatile( \
            "s_setprio 3                                                 \n\t" \
            "ds_read_b128 %[ae], %[aa0]                                  \n\t" \
            "ds_read_b128 %[be], %[ab0]                                  \n\t" \
            "ds_read_b128 %[ao], %[aa1]                                  \n\t" \
            "ds_read_b128 %[bo], %[ab1]                                  \n\t" \
            "s_waitcnt lgkmcnt(2)                                        \n\t" \
            "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ae], %[be], %[acc], %[sa0], %[sb0] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
            "ds_read_b128 %[ae], %[aa2]                                  \n\t" \
            "ds_read_b128 %[be], %[ab2]                                  \n\t" \
            "s_waitcnt lgkmcnt(2)                                        \n\t" \
            "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ao], %[bo], %[acc], %[sa1], %[sb1] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
            "ds_read_b128 %[ao], %[aa3]                                  \n\t" \
            "ds_read_b128 %[bo], %[ab3]                                  \n\t" \
            "s_waitcnt lgkmcnt(2)                                        \n\t" \
            "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ae], %[be], %[acc], %[sa2], %[sb2] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
            "s_waitcnt lgkmcnt(0)                                        \n\t" \
            "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ao], %[bo], %[acc], %[sa3], %[sb3] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
            "s_setprio 0                                                 \n\t" \
            : [acc] "+v"(acc), \
              [ae] "=&v"(_ae), [be] "=&v"(_be), \
              [ao] "=&v"(_ao), [bo] "=&v"(_bo) \
            : [aa0] "v"(_addr_a0), [ab0] "v"(_addr_b0), \
              [aa1] "v"(_addr_a1), [ab1] "v"(_addr_b1), \
              [aa2] "v"(_addr_a2), [ab2] "v"(_addr_b2), \
              [aa3] "v"(_addr_a3), [ab3] "v"(_addr_b3), \
              [sa0] "v"(_sa0), [sb0] "v"(_sb0), \
              [sa1] "v"(_sa1), [sb1] "v"(_sb1), \
              [sa2] "v"(_sa2), [sb2] "v"(_sb2), \
              [sa3] "v"(_sa3), [sb3] "v"(_sb3) \
        ); \
        __builtin_amdgcn_sched_barrier(0x020); \
    }

    // s482: 4-phase pipeline for G6 — issue VMEM before COMPUTE
    int4_v _g6_pst_w0, _g6_pst_w1, _g6_pst_w2, _g6_pst_w3;

    #define G6_ISSUE_VMEM(ks_base) \
    { \
        const unsigned _b_soff = __builtin_amdgcn_readfirstlane(static_cast<unsigned>((ks_base) * 1024)); \
        _Pragma("unroll") \
        for (int i = 0; i < G6_B_PER_THREAD; ++i) { \
            BUFFER_LOAD_DWORDX4_NT(b_regs[i], b_voff[i], rsrc_bsh, _b_soff); \
        } \
        { \
            int _lin = tid; \
            int _kidx = _lin % (G6_CHUNK_K * 4); \
            int _mloc = (_lin / (G6_CHUNK_K * 4)) % 16; \
            int _wmid = _lin / (16 * G6_CHUNK_K * 4); \
            int _row = tile_m_base + _wmid * 16 + _mloc; \
            const int4_v* _src = reinterpret_cast<const int4_v*>( \
                reinterpret_cast<const uint8_t*>(A_bf16) + \
                (long)_row * G6_K * 2 + (long)((ks_base) * 128 + _kidx * 32) * 2); \
            _g6_pst_w0 = _src[0]; \
            _g6_pst_w1 = _src[1]; \
            _g6_pst_w2 = _src[2]; \
            _g6_pst_w3 = _src[3]; \
        } \
    }

    #define G6_COMPLETE_QUANT_STORE(buf, scale_base) \
    { \
        const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
        { \
            const uint32_t* _p0 = reinterpret_cast<const uint32_t*>(&_g6_pst_w0); \
            const uint32_t* _p1 = reinterpret_cast<const uint32_t*>(&_g6_pst_w1); \
            const uint32_t* _p2 = reinterpret_cast<const uint32_t*>(&_g6_pst_w2); \
            const uint32_t* _p3 = reinterpret_cast<const uint32_t*>(&_g6_pst_w3); \
            float _absMax = 1e-10f; \
            { \
                auto _AM = [&](uint32_t pair) { \
                    __bf16 lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
                    __bf16 hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair >> 16)); \
                    float flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
                    float fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
                    _absMax = (flo > _absMax) ? flo : _absMax; \
                    _absMax = (fhi > _absMax) ? fhi : _absMax; \
                }; \
                _AM(_p0[0]); _AM(_p0[1]); _AM(_p0[2]); _AM(_p0[3]); \
                _AM(_p1[0]); _AM(_p1[1]); _AM(_p1[2]); _AM(_p1[3]); \
                _AM(_p2[0]); _AM(_p2[1]); _AM(_p2[2]); _AM(_p2[3]); \
                _AM(_p3[0]); _AM(_p3[1]); _AM(_p3[2]); _AM(_p3[3]); \
            } \
            uint32_t _u32 = __builtin_bit_cast(uint32_t, _absMax); \
            uint32_t _amax_exp = ((_u32 + 0x200000u) >> 23) & 0xFFu; \
            uint32_t _inv_exp = (_amax_exp >= 2u) ? (_amax_exp - 2u) : 0u; \
            float _hw_scale = __builtin_bit_cast(float, _inv_exp << 23); \
            unsigned _d0 = 0u, _d1 = 0u, _d2 = 0u, _d3 = 0u; \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[0]), _hw_scale, 0); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[1]), _hw_scale, 1); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[2]), _hw_scale, 2); \
            _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[3]), _hw_scale, 3); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[0]), _hw_scale, 0); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[1]), _hw_scale, 1); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[2]), _hw_scale, 2); \
            _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[3]), _hw_scale, 3); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[0]), _hw_scale, 0); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[1]), _hw_scale, 1); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[2]), _hw_scale, 2); \
            _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[3]), _hw_scale, 3); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[0]), _hw_scale, 0); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[1]), _hw_scale, 1); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[2]), _hw_scale, 2); \
            _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[3]), _hw_scale, 3); \
            int4_v _fp4 = int4_v{static_cast<int>(_d0), static_cast<int>(_d1), \
                                 static_cast<int>(_d2), static_cast<int>(_d3)}; \
            { \
                const unsigned _off = _buf_off + a_lds_offs[0]; \
                asm volatile("ds_write_b128 %0, %1 offset:%2" \
                    :: "v"(_off), "v"(_fp4), "n"(0) : "memory"); \
            } \
            { \
                int _linear = tid; \
                int _k_idx = _linear % (G6_CHUNK_K * 4); \
                int _m_local = (_linear / (G6_CHUNK_K * 4)) % 16; \
                int _wave_m_idx = _linear / (16 * G6_CHUNK_K * 4); \
                unsigned _scale_off = (scale_base) + _wave_m_idx * 512 + _m_local * 32 + _k_idx; \
                smem_raw[_scale_off] = static_cast<uint8_t>(_inv_exp); \
            } \
        } \
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); \
        _Pragma("unroll") \
        for (int i = 0; i < G6_B_PER_THREAD; ++i) { \
            const unsigned _off = _buf_off + b_lds_offs[i]; \
            asm volatile("ds_write_b128 %0, %1 offset:%2" \
                :: "v"(_off), "v"(b_regs[i]), "n"(0) : "memory"); \
        } \
    }

    int cur_ks = 0;
    unsigned cur_scale_base = G6_SCALE_LDS_BASE0;
    unsigned nxt_scale_base = G6_SCALE_LDS_BASE1;

    // Initial prefetch: full load for chunk 0
    G6_ISSUE_LOADS_FUSED(cur_ks, buf0, cur_scale_base);
    asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
    __syncthreads();

    uint8_t* cur_buf = buf0;
    uint8_t* nxt_buf = buf1;

    // 4-phase pipeline loop: VMEM→COMPUTE→QUANT_STORE→sync
    #pragma unroll
    for (int c = 0; c < G6_NUM_CHUNKS - 1; ++c) {
        G6_ISSUE_VMEM(cur_ks + G6_CHUNK_K);
        G6_COMPUTE_CHUNK(cur_buf, cur_ks, cur_scale_base);
        G6_COMPLETE_QUANT_STORE(nxt_buf, nxt_scale_base);
        asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
        __syncthreads();

        uint8_t* tmp = cur_buf;
        cur_buf = nxt_buf;
        nxt_buf = tmp;
        unsigned tmp_scale = cur_scale_base;
        cur_scale_base = nxt_scale_base;
        nxt_scale_base = tmp_scale;
        cur_ks += G6_CHUNK_K;
    }

    G6_COMPUTE_CHUNK(cur_buf, cur_ks, cur_scale_base);

    #undef G6_ISSUE_VMEM
    #undef G6_COMPLETE_QUANT_STORE
    #undef G6_ISSUE_LOADS_FUSED
    #undef G6_COMPUTE_CHUNK

    // Write output — all tiles full, all rows valid
    const int out_col      = tile_n + lrow;
    const int out_row_base = tile_m + kgrp * 4;

    uint16_t* c_out = C_final + (long)out_row_base * G6_N + out_col;
    #pragma unroll
    for (int i = 0; i < 4; ++i, c_out += G6_N) {
        *c_out = fast_f32_to_bf16(acc[i]);
    }

    #undef G6_M
    #undef G6_N
    #undef G6_K
    #undef G6_BM
    #undef G6_BN
    #undef G6_WAVES_M
    #undef G6_WAVES_N
    #undef G6_NWARPS
    #undef G6_CHUNK_K
    #undef G6_CKT
    #undef G6_NUM_CHUNKS
    #undef G6_SCALEN
    #undef G6_A_ROW
    #undef G6_A_TILE
    #undef G6_A_SIZE
    #undef G6_B_KGRP_STRIDE
    #undef G6_B_KT_STRIDE
    #undef G6_B_TILE
    #undef G6_B_SIZE
    #undef G6_BUF_SIZE
    #undef G6_NTHREADS
    #undef G6_A_TOTAL
    #undef G6_A_PER_THREAD
    #undef G6_B_TOTAL
    #undef G6_B_PER_THREAD
    #undef G6_SCALE_LDS_SIZE
    #undef G6_SCALE_LDS_BASE0
    #undef G6_SCALE_LDS_BASE1
}


// ============================================================================
// GENERIC TEMPLATE FALLBACK (from s194) — handles ALL shapes not hardcoded above
// ============================================================================

template<int C_M, int C_K>
__global__ void __launch_bounds__(128, (C_M * (C_K / 32) <= 256) ? 2 : 4)
mxfp4_quant(
    const __bf16* __restrict__ A_bf16,
    uint8_t*      __restrict__ A_fp4,
    uint8_t*      __restrict__ A_scale)
{
    constexpr auto KS = C_K / 32;
    constexpr auto K2 = C_K / 2;
    const auto group = blockIdx.x * 128 + threadIdx.x;
    const auto row   = group / KS;
    const auto kg    = group % KS;

    if (row >= C_M) return;

    const auto row_k = (long)row * C_K;
    const auto* src = A_bf16 + row_k + kg * 32;

    const auto w0 = *reinterpret_cast<const int4_v*>(src);
    const auto w1 = *reinterpret_cast<const int4_v*>(src + 8);
    const auto w2 = *reinterpret_cast<const int4_v*>(src + 16);
    const auto w3 = *reinterpret_cast<const int4_v*>(src + 24);

    const auto* p0 = reinterpret_cast<const uint32_t*>(&w0);
    const auto* p1 = reinterpret_cast<const uint32_t*>(&w1);
    const auto* p2 = reinterpret_cast<const uint32_t*>(&w2);
    const auto* p3 = reinterpret_cast<const uint32_t*>(&w3);

    auto absMax = 1e-10f;
    #define AMAX_PAIR(pair) { \
        const auto lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
        const auto hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
        const auto flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
        const auto fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
        absMax = (flo > absMax) ? flo : absMax; \
        absMax = (fhi > absMax) ? fhi : absMax; \
    }
    AMAX_PAIR(p0[0]) AMAX_PAIR(p0[1]) AMAX_PAIR(p0[2]) AMAX_PAIR(p0[3])
    AMAX_PAIR(p1[0]) AMAX_PAIR(p1[1]) AMAX_PAIR(p1[2]) AMAX_PAIR(p1[3])
    AMAX_PAIR(p2[0]) AMAX_PAIR(p2[1]) AMAX_PAIR(p2[2]) AMAX_PAIR(p2[3])
    AMAX_PAIR(p3[0]) AMAX_PAIR(p3[1]) AMAX_PAIR(p3[2]) AMAX_PAIR(p3[3])
    #undef AMAX_PAIR

    const auto u32 = __builtin_bit_cast(uint32_t, absMax);
    const auto amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
    const auto inv_exp  = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
    A_scale[row_k / 32 + kg] = static_cast<uint8_t>(inv_exp);

    const auto hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);

    #define CVT(d, pair, sel) \
        d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)

    auto d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
    CVT(d0, p0[0], 0); CVT(d0, p0[1], 1); CVT(d0, p0[2], 2); CVT(d0, p0[3], 3);
    CVT(d1, p1[0], 0); CVT(d1, p1[1], 1); CVT(d1, p1[2], 2); CVT(d1, p1[3], 3);
    CVT(d2, p2[0], 0); CVT(d2, p2[1], 1); CVT(d2, p2[2], 2); CVT(d2, p2[3], 3);
    CVT(d3, p3[0], 0); CVT(d3, p3[1], 1); CVT(d3, p3[2], 2); CVT(d3, p3[3], 3);
    #undef CVT

    auto* dst = reinterpret_cast<int4_v*>(A_fp4 + row_k / 2 + kg * 16);
    *dst = int4_v{static_cast<int>(d0), static_cast<int>(d1),
                  static_cast<int>(d2), static_cast<int>(d3)};
}


template<bool ALWAYS_VALID>
__device__ __forceinline__ auto load_or_zero(bool rt_valid, const uint8_t* p) {
    if constexpr (ALWAYS_VALID) return load16(p);
    else                        return rt_valid ? load16(p) : int4_v{0,0,0,0};
}


template<int WAVES_M, int WAVES_N, int CHUNK_K>
struct LdsLayout {
    static constexpr auto A_ROW   = CHUNK_K * 64 + 16;   // +16B pad per row
    static constexpr auto A_TILE  = 16 * A_ROW;
    static constexpr auto A_SIZE  = WAVES_M * A_TILE;

    static constexpr auto B_KGRP_STRIDE = 16 * 16;        // 16 lrows x 16B = 256B
    static constexpr auto B_KT_STRIDE   = 5 * B_KGRP_STRIDE;  // 5 slots (4+1 pad)
    static constexpr auto B_TILE  = CHUNK_K * B_KT_STRIDE;
    static constexpr auto B_SIZE  = WAVES_N * B_TILE;

    static constexpr auto BUF_SIZE  = A_SIZE + B_SIZE;
    static constexpr auto TOTAL_LDS = 2 * BUF_SIZE;

    __device__ __forceinline__ static constexpr auto a_off(uint8_t* const buf) { return buf; }
    __device__ __forceinline__ static constexpr auto b_off(uint8_t* const buf) { return buf + A_SIZE; }

    __device__ __forceinline__ static constexpr auto a_idx(const auto wm, const auto row, const auto k_idx) {
        return wm * A_TILE + row * A_ROW + k_idx * 16;
    }

    __device__ __forceinline__ static constexpr auto b_idx(const auto wn, const auto kt, const auto kgrp, const auto lrow) {
        return wn * B_TILE + kt * B_KT_STRIDE + kgrp * B_KGRP_STRIDE + lrow * 16;
    }
};


template<int WAVES_M, int WAVES_N, int CHUNK_K, int NWARPS>
struct LoadCounts {
    static constexpr auto NTHREADS = NWARPS * 64;
    static constexpr auto A_TOTAL = WAVES_M * 16 * CHUNK_K * 4;
    static constexpr auto B_TOTAL = WAVES_N * CHUNK_K * 4 * 16;
    static constexpr auto A_PER_THREAD = (A_TOTAL + NTHREADS - 1) / NTHREADS;
    static constexpr auto B_PER_THREAD = (B_TOTAL + NTHREADS - 1) / NTHREADS;
};


template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,
         int CKT, int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS,
         int C_M, int C_TILE_OFF_X, int C_TILE_OFF_Y>
__global__ void __launch_bounds__(NWARPS * 64, (NWARPS <= 2) ? 4 : 2)
__attribute__((amdgpu_flat_work_group_size(NWARPS * 64, NWARPS * 64)))
mxfp4_gemm(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ As,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    float*         __restrict__ C_partial,
    uint16_t*      __restrict__ C_final)
{
    static_assert(BN % 16 == 0);
    constexpr auto WAVES_M = (BM + 15) / 16;
    constexpr auto WAVES_N = BN / 16;
    static_assert(WAVES_M * WAVES_N == NWARPS);

    constexpr auto M = C_M;
    constexpr auto N = C_N;
    constexpr auto K = C_K;
    constexpr auto scaleN = C_SCALEN;
    constexpr auto K2 = K / 2;
    constexpr auto KS = K / 32;
    constexpr auto bsh_n_stride = (long)(K / 64) * 512;

    const auto ks_idx = static_cast<int>(blockIdx.z);
    const auto lane   = static_cast<int>(threadIdx.x % 64);
    const auto wave   = static_cast<int>(threadIdx.x / 64);
    const auto wave_m = wave / WAVES_N;
    const auto wave_n = wave % WAVES_N;
    const auto tid    = static_cast<int>(threadIdx.x);

    const auto tile_m_base = (static_cast<int>(blockIdx.y) + C_TILE_OFF_Y) * BM;
    const auto tile_n_base = (static_cast<int>(blockIdx.x) + C_TILE_OFF_X) * BN;
    const auto tile_m = tile_m_base + wave_m * 16;
    const auto tile_n = tile_n_base + wave_n * 16;

    constexpr auto ktiles_per_split = C_KPS;
    const auto ks_start = ks_idx * ktiles_per_split;

    const auto lrow = lane % 16;
    const auto kgrp = lane / 16;
    __builtin_assume(lrow >= 0 && lrow < 16);
    __builtin_assume(kgrp >= 0 && kgrp < 4);

    const auto gm = tile_m + lrow;
    const auto gn = tile_n + lrow;

    const auto a_rt = A_VALID | (gm < M);
    const auto b_rt = B_VALID | (gn < N);

    const uint8_t* as_row = nullptr;
    if constexpr (A_VALID) {
        as_row = As + (long)gm * KS;
    } else {
        if (a_rt) as_row = As + (long)gm * KS;
    }

    auto bssh_base = 0;
    if constexpr (B_VALID) {
        bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
    } else {
        if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
    }

    float4_v acc{0.f, 0.f, 0.f, 0.f};

    // PATH A: CKT <= 4 -- Direct global loads, no LDS
    if constexpr (CKT <= 4) {
        if (tile_m >= M || tile_n >= N) return;

        const auto n_tile = tile_n / 16;
        const auto bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;

        const auto k_half_off = (kgrp & 1) * 256;
        const auto k_blk_base = kgrp >> 1;

        constexpr auto a_kt_stride   = 64L;
        constexpr auto bsh_kt_stride = 1024L;

        const uint8_t* a_ptr    = nullptr;
        const uint8_t* bsh_ptr  = nullptr;
        const uint8_t* bssh_ptr = nullptr;

        if constexpr (A_VALID) {
            a_ptr = A + (long)gm * K2 + (long)ks_start * a_kt_stride + kgrp * 16;
        } else {
            if (a_rt) a_ptr = A + (long)gm * K2 + (long)ks_start * a_kt_stride + kgrp * 16;
        }
        if constexpr (B_VALID) {
            bsh_ptr  = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
            bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
        } else {
            if (b_rt) {
                bsh_ptr  = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
                bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
            }
        }

        const auto bssh_step0 = (ks_start & 1) ? 254 : 2;
        const auto bssh_step1 = 256 - bssh_step0;

        #define DO_MFMA(a_off, b_off, bssh_off, ks_val) \
        { \
            const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \
            const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \
            const auto ks = (ks_val) * 4 + kgrp; \
            auto sa = 0, sb = 0; \
            if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \
            else                   sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \
            if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \
            else                   sb = (b_rt & (ks < KS)) ? static_cast<int>(*(bssh_ptr + (bssh_off))) : 127; \
            acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,acc,FP4_E2M1,FP4_E2M1,0,sa,0,sb); \
        }

        static_assert(CKT > 0);
        #pragma unroll
        for (auto q = 0; q < (CKT / 4); ++q) {
            DO_MFMA(0, 0, 0, ks_start + q*4)
            DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)
            DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)
            DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, ks_start + q*4 + 3)
            a_ptr    += 4 * a_kt_stride;
            bsh_ptr  += 4 * bsh_kt_stride;
            bssh_ptr += 512;
        }
        if constexpr ((CKT % 4) >= 2) {
            DO_MFMA(0, 0, 0, ks_start + (CKT/4)*4)
            DO_MFMA(1, 1, bssh_step0, ks_start + (CKT/4)*4 + 1)
            a_ptr    += 2 * a_kt_stride;
            bsh_ptr  += 2 * bsh_kt_stride;
            bssh_ptr += 256;
        }
        if constexpr ((CKT % 2) == 1) {
            DO_MFMA(0, 0, 0, ks_start + CKT - 1)
        }
        #undef DO_MFMA

    // PATH B: CKT > 4 -- LDS double-buffered with register-staged pipelining
    } else {
        constexpr auto CHUNK_K = 4;
        constexpr auto NUM_CHUNKS = CKT / CHUNK_K;
        constexpr auto TAIL_KT = CKT % CHUNK_K;

        using Lds = LdsLayout<WAVES_M, WAVES_N, CHUNK_K>;
        using LC  = LoadCounts<WAVES_M, WAVES_N, CHUNK_K, NWARPS>;

        extern __shared__ uint8_t smem_raw[];
        auto* buf0 = smem_raw;
        auto* buf1 = smem_raw + Lds::BUF_SIZE;

        const auto tile_valid = (tile_m < M) && (tile_n < N);

        int  a_row_g[LC::A_PER_THREAD];
        int  a_k_idx_g[LC::A_PER_THREAD];
        int  a_lds_offs_g[LC::A_PER_THREAD];
        bool a_valid_g[LC::A_PER_THREAD];

        #pragma unroll
        for (auto i = 0; i < LC::A_PER_THREAD; ++i) {
            auto linear = i * LC::NTHREADS + tid;
            if (linear >= LC::A_TOTAL) {
                a_valid_g[i] = false;
                a_lds_offs_g[i] = -1;
            } else {
                auto k_idx      = linear % (CHUNK_K * 4);
                auto m_local    = (linear / (CHUNK_K * 4)) % 16;
                auto wave_m_idx = linear / (16 * CHUNK_K * 4);
                a_row_g[i]      = tile_m_base + wave_m_idx * 16 + m_local;
                a_k_idx_g[i]    = k_idx;
                a_lds_offs_g[i] = Lds::a_idx(wave_m_idx, m_local, k_idx);
                if constexpr (A_VALID) a_valid_g[i] = tile_valid;
                else                   a_valid_g[i] = tile_valid && (a_row_g[i] < M);
            }
        }

        int  b_kt_local_g[LC::B_PER_THREAD];
        long b_base_off_g[LC::B_PER_THREAD];
        int  b_kgrp_half_g[LC::B_PER_THREAD];
        int  b_lds_kgrp_g[LC::B_PER_THREAD];
        int  b_lds_lrow_g[LC::B_PER_THREAD];
        int  b_lds_wn_g[LC::B_PER_THREAD];
        int  b_lds_offs_g[LC::B_PER_THREAD];
        bool b_valid_g[LC::B_PER_THREAD];

        #pragma unroll
        for (auto i = 0; i < LC::B_PER_THREAD; ++i) {
            auto linear = i * LC::NTHREADS + tid;
            if (linear >= LC::B_TOTAL) {
                b_valid_g[i] = false;
                b_lds_offs_g[i] = -1;
            } else {
                auto b_lrow     = linear % 16;
                auto b_kgrp     = (linear / 16) % 4;
                auto kt         = (linear / 64) % CHUNK_K;
                auto wave_n_idx = linear / (CHUNK_K * 64);
                auto b_tile_n   = tile_n_base + wave_n_idx * 16;
                auto n_tile     = b_tile_n / 16;
                auto k_half     = (b_kgrp & 1) * 256;

                b_kt_local_g[i]  = kt;
                b_base_off_g[i]  = (long)n_tile * bsh_n_stride + k_half + (long)b_lrow * 16;
                b_kgrp_half_g[i] = b_kgrp / 2;
                b_lds_kgrp_g[i]  = b_kgrp;
                b_lds_lrow_g[i]  = b_lrow;
                b_lds_wn_g[i]    = wave_n_idx;
                b_lds_offs_g[i]  = Lds::b_idx(wave_n_idx, kt, b_kgrp, b_lrow);

                if constexpr (B_VALID) b_valid_g[i] = tile_valid;
                else                   b_valid_g[i] = tile_valid && (b_tile_n + b_lrow < N);
            }
        }

        int4_v a_regs_g[LC::A_PER_THREAD];
        int4_v b_regs_g[LC::B_PER_THREAD];


        #define ISSUE_LOADS_G(ks_base) \
        { \
            _Pragma("unroll") \
            for (auto i = 0; i < LC::A_PER_THREAD; ++i) { \
                a_regs_g[i] = int4_v{0,0,0,0}; \
                if (a_valid_g[i]) { \
                    auto k_byte = ((ks_base) * 4 + a_k_idx_g[i]) * 16; \
                    if constexpr (A_VALID) \
                        a_regs_g[i] = load16(A + (long)a_row_g[i] * K2 + k_byte); \
                    else if (k_byte + 16 <= K2) \
                        a_regs_g[i] = load16(A + (long)a_row_g[i] * K2 + k_byte); \
                } \
            } \
            _Pragma("unroll") \
            for (auto i = 0; i < LC::B_PER_THREAD; ++i) { \
                b_regs_g[i] = int4_v{0,0,0,0}; \
                if (b_valid_g[i]) { \
                    auto ks = (ks_base) + b_kt_local_g[i]; \
                    auto k_blk = ks * 2 + b_kgrp_half_g[i]; \
                    auto global_off = b_base_off_g[i] + (long)k_blk * 512; \
                    b_regs_g[i] = load16_nt(Bsh + global_off); \
                } \
            } \
        }

        #define STORE_TO_LDS_G(buf) \
        { \
            auto* _sa = Lds::a_off(buf); \
            auto* _sb = Lds::b_off(buf); \
            _Pragma("unroll") \
            for (auto i = 0; i < LC::A_PER_THREAD; ++i) { \
                if (a_lds_offs_g[i] >= 0) \
                    *reinterpret_cast<int4_v*>(_sa + a_lds_offs_g[i]) = a_regs_g[i]; \
            } \
            _Pragma("unroll") \
            for (auto i = 0; i < LC::B_PER_THREAD; ++i) { \
                if (b_lds_offs_g[i] >= 0) \
                    *reinterpret_cast<int4_v*>(_sb + b_lds_offs_g[i]) = b_regs_g[i]; \
            } \
        }

        #define COMPUTE_CHUNK_G(buf, chunk_ks) \
        { \
            const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
            const unsigned _a_base = _buf_off; \
            const unsigned _b_base = _buf_off + static_cast<unsigned>(Lds::A_SIZE); \
            const unsigned _addr_a0 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 0 * 4 + kgrp)); \
            const unsigned _addr_b0 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 0, kgrp, lrow)); \
            const unsigned _addr_a1 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 1 * 4 + kgrp)); \
            const unsigned _addr_b1 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 1, kgrp, lrow)); \
            const unsigned _addr_a2 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 2 * 4 + kgrp)); \
            const unsigned _addr_b2 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 2, kgrp, lrow)); \
            const unsigned _addr_a3 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 3 * 4 + kgrp)); \
            const unsigned _addr_b3 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 3, kgrp, lrow)); \
            int _sa0, _sb0, _sa1, _sb1, _sa2, _sb2, _sa3, _sb3; \
            if constexpr (A_VALID) { \
                auto _sa_base = as_row + (chunk_ks) * 4; \
                auto _sa_vec = *reinterpret_cast<const int4_v*>(_sa_base); \
                auto _sa_shift = kgrp * 8; \
                _sa0 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[0] >> _sa_shift) & 0xFF; \
                _sa1 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[1] >> _sa_shift) & 0xFF; \
                _sa2 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[2] >> _sa_shift) & 0xFF; \
                _sa3 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[3] >> _sa_shift) & 0xFF; \
            } else { \
                auto _ks0 = (chunk_ks) * 4 + kgrp; \
                _sa0 = (a_rt & (_ks0 < KS)) ? static_cast<int>(as_row[_ks0]) : 127; \
                _sa1 = (a_rt & (_ks0+4 < KS)) ? static_cast<int>(as_row[_ks0+4]) : 127; \
                _sa2 = (a_rt & (_ks0+8 < KS)) ? static_cast<int>(as_row[_ks0+8]) : 127; \
                _sa3 = (a_rt & (_ks0+12 < KS)) ? static_cast<int>(as_row[_ks0+12]) : 127; \
            } \
            { \
                auto _ksv0 = (chunk_ks); \
                auto* _bp0 = Bssh + bssh_base + (_ksv0 & 1) * 2 + (_ksv0 >> 1) * 256; \
                auto* _bp1 = Bssh + bssh_base + ((_ksv0+1) & 1) * 2 + ((_ksv0+1) >> 1) * 256; \
                auto* _bp2 = Bssh + bssh_base + ((_ksv0+2) & 1) * 2 + ((_ksv0+2) >> 1) * 256; \
                auto* _bp3 = Bssh + bssh_base + ((_ksv0+3) & 1) * 2 + ((_ksv0+3) >> 1) * 256; \
                if constexpr (B_VALID) { \
                    _sb0 = static_cast<int>(*_bp0); \
                    _sb1 = static_cast<int>(*_bp1); \
                    _sb2 = static_cast<int>(*_bp2); \
                    _sb3 = static_cast<int>(*_bp3); \
                } else { \
                    auto _ks0 = (chunk_ks) * 4 + kgrp; \
                    _sb0 = (b_rt & (_ks0 < KS)) ? static_cast<int>(*_bp0) : 127; \
                    _sb1 = (b_rt & (_ks0+4 < KS)) ? static_cast<int>(*_bp1) : 127; \
                    _sb2 = (b_rt & (_ks0+8 < KS)) ? static_cast<int>(*_bp2) : 127; \
                    _sb3 = (b_rt & (_ks0+12 < KS)) ? static_cast<int>(*_bp3) : 127; \
                } \
            } \
            int4_v _ae, _be, _ao, _bo; \
            __builtin_amdgcn_sched_barrier(0x020); \
            asm volatile( \
                "s_setprio 3                                                 \n\t" \
                "ds_read_b128 %[ae], %[aa0]                                  \n\t" \
                "ds_read_b128 %[be], %[ab0]                                  \n\t" \
                "ds_read_b128 %[ao], %[aa1]                                  \n\t" \
                "ds_read_b128 %[bo], %[ab1]                                  \n\t" \
                "s_waitcnt lgkmcnt(2)                                        \n\t" \
                "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ae], %[be], %[acc], %[sa0], %[sb0] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
                "ds_read_b128 %[ae], %[aa2]                                  \n\t" \
                "ds_read_b128 %[be], %[ab2]                                  \n\t" \
                "s_waitcnt lgkmcnt(2)                                        \n\t" \
                "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ao], %[bo], %[acc], %[sa1], %[sb1] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
                "ds_read_b128 %[ao], %[aa3]                                  \n\t" \
                "ds_read_b128 %[bo], %[ab3]                                  \n\t" \
                "s_waitcnt lgkmcnt(2)                                        \n\t" \
                "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ae], %[be], %[acc], %[sa2], %[sb2] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
                "s_waitcnt lgkmcnt(0)                                        \n\t" \
                "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ao], %[bo], %[acc], %[sa3], %[sb3] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
                "s_setprio 0                                                 \n\t" \
                : [acc] "+v"(acc), \
                  [ae] "=&v"(_ae), [be] "=&v"(_be), \
                  [ao] "=&v"(_ao), [bo] "=&v"(_bo) \
                : [aa0] "v"(_addr_a0), [ab0] "v"(_addr_b0), \
                  [aa1] "v"(_addr_a1), [ab1] "v"(_addr_b1), \
                  [aa2] "v"(_addr_a2), [ab2] "v"(_addr_b2), \
                  [aa3] "v"(_addr_a3), [ab3] "v"(_addr_b3), \
                  [sa0] "v"(_sa0), [sb0] "v"(_sb0), \
                  [sa1] "v"(_sa1), [sb1] "v"(_sb1), \
                  [sa2] "v"(_sa2), [sb2] "v"(_sb2), \
                  [sa3] "v"(_sa3), [sb3] "v"(_sb3) \
            ); \
            __builtin_amdgcn_sched_barrier(0x020); \
        }

        auto cur_ks = ks_start;
        ISSUE_LOADS_G(cur_ks);
        asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
        STORE_TO_LDS_G(buf0);
        __syncthreads();

        auto* cur_buf = buf0;
        auto* nxt_buf = buf1;

        #pragma unroll
        for (auto c = 0; c < NUM_CHUNKS - 1; ++c) {
            ISSUE_LOADS_G(cur_ks + CHUNK_K);
            COMPUTE_CHUNK_G(cur_buf, cur_ks);
            asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
            STORE_TO_LDS_G(nxt_buf);
            __syncthreads();

            auto* tmp = cur_buf;
            cur_buf = nxt_buf;
            nxt_buf = tmp;
            cur_ks += CHUNK_K;
        }

        COMPUTE_CHUNK_G(cur_buf, cur_ks);

        if constexpr (TAIL_KT > 0) {
            cur_ks += CHUNK_K;
            ISSUE_LOADS_G(cur_ks);
            asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
            STORE_TO_LDS_G(buf0);
            __syncthreads();

            auto* _ca = Lds::a_off(buf0);
            auto* _cb = Lds::b_off(buf0);
            #pragma unroll
            for (auto kt = 0; kt < TAIL_KT; ++kt) {
                auto av = *reinterpret_cast<const int4_v*>(
                    _ca + Lds::a_idx(wave_m, lrow, kt * 4 + kgrp));
                auto bv = *reinterpret_cast<const int4_v*>(
                    _cb + Lds::b_idx(wave_n, kt, kgrp, lrow));
                auto ks_val = cur_ks + kt;
                auto ks = ks_val * 4 + kgrp;
                auto sa = 0, sb = 0;
                if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]);
                else                   sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127;
                auto* bssh_p = Bssh + bssh_base + (ks_val & 1) * 2 + (ks_val >> 1) * 256;
                if constexpr (B_VALID) sb = static_cast<int>(*bssh_p);
                else                   sb = (b_rt & (ks < KS)) ? static_cast<int>(*bssh_p) : 127;
                acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                    av, bv, acc, FP4_E2M1, FP4_E2M1, 0, sa, 0, sb);
            }
        }

        #undef ISSUE_LOADS_G
        #undef STORE_TO_LDS_G
        #undef COMPUTE_CHUNK_G

        if (!tile_valid) return;
    }


    const auto out_col      = tile_n + lrow;
    const auto out_row_base = tile_m + kgrp * 4;

    if constexpr (!B_VALID) { if (out_col >= N) return; }

    constexpr auto out_rows_always_valid = A_VALID && (BM >= 16);

    if constexpr (SPLITK) {
        auto* c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;
        #pragma unroll
        for (auto i = 0; i < 4; ++i, c_out += N) {
            if constexpr (out_rows_always_valid) *c_out = acc[i];
            else if (out_row_base + i < M)       *c_out = acc[i];
        }
    } else {
        auto* c_out = C_final + (long)out_row_base * N + out_col;
        #pragma unroll
        for (auto i = 0; i < 4; ++i, c_out += N) {
            if constexpr (out_rows_always_valid) *c_out = float_to_bf16(acc[i]);
            else if (out_row_base + i < M)       *c_out = float_to_bf16(acc[i]);
        }
    }
}


template<int C_N, int C_NUM_KSPLIT, int C_M>
__global__ void mxfp4_reduce(
    const float*  __restrict__ C_partial,
    uint16_t*     __restrict__ C_out)
{
    constexpr auto M = C_M;
    constexpr auto N = C_N;
    constexpr auto mn_stride = (long)M * N;
    const auto col = blockIdx.x * 32 + threadIdx.x;
    const auto row = blockIdx.y * 16 + threadIdx.y;
    if (row >= M || col >= N) return;

    const auto mn = (long)row * N + col;
    auto sum = 0.f;
    const auto* ptr = C_partial + mn;
    #pragma unroll
    for (auto k = 0; k < C_NUM_KSPLIT; ++k)
        sum += ptr[k * mn_stride];

    C_out[mn] = float_to_bf16(sum);
}

template<int C_M, int C_K, int C_N, int C_SCALEN, int C_WAVES_N = 4>
__global__ void __launch_bounds__(((C_M + 15) / 16) * C_WAVES_N * 64, (C_WAVES_N <= 2) ? 4 : 2)
mxfp4_fused_quant_gemm(
    const __bf16* __restrict__ A_bf16,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    uint16_t*      __restrict__ C_final)
{
    constexpr auto KS = C_K / 32;
    constexpr auto K2 = C_K / 2;
    constexpr auto CKT = C_K / 128;
    constexpr auto WAVES_M = (C_M + 15) / 16;
    constexpr auto WAVES_N = C_WAVES_N;
    constexpr auto BN = WAVES_N * 16;
    constexpr auto NWARPS = WAVES_M * WAVES_N;

    const auto tid = threadIdx.x;
    const auto lane = tid % 64;
    const auto wave_id = tid / 64;
    const auto lrow = lane % 16;
    const auto kgrp = lane / 16;
    const auto wave_m = wave_id / WAVES_N;
    const auto wave_n = wave_id % WAVES_N;

    const auto tile_m = wave_m * 16;
    const auto tile_n = static_cast<int>(blockIdx.x) * BN + wave_n * 16;
    const auto my_row = tile_m + lrow;

    constexpr auto A_BF16_BYTES = C_M * C_K * 2;
    extern __shared__ uint8_t smem[];
    {
        constexpr auto BYTES_PER_ITER = NWARPS * 64 * 16;
        constexpr auto NITERS = (A_BF16_BYTES + BYTES_PER_ITER - 1) / BYTES_PER_ITER;
        #pragma unroll
        for (auto iter = 0; iter < NITERS; ++iter) {
            const auto offset = iter * BYTES_PER_ITER + tid * 16;
            if (offset + 16 <= A_BF16_BYTES) {
                *reinterpret_cast<int4_v*>(smem + offset) =
                    *reinterpret_cast<const int4_v*>(
                        reinterpret_cast<const uint8_t*>(A_bf16) + offset);
            }
        }
    }
    __syncthreads();

    int4_v a_fp4_regs[CKT];
    int a_scale_regs[CKT];

    constexpr auto a_row_always_valid = (C_M >= ((C_M + 15) / 16) * 16);
    const auto a_row_valid = a_row_always_valid ? true : (my_row < C_M);

    #pragma unroll
    for (auto kt = 0; kt < CKT; ++kt) {
        const auto kg = kt * 4 + kgrp;

        if (a_row_valid) {
            const auto lds_off = my_row * C_K * 2 + kg * 64;

            const auto w0 = *reinterpret_cast<const int4_v*>(smem + lds_off);
            const auto w1 = *reinterpret_cast<const int4_v*>(smem + lds_off + 16);
            const auto w2 = *reinterpret_cast<const int4_v*>(smem + lds_off + 32);
            const auto w3 = *reinterpret_cast<const int4_v*>(smem + lds_off + 48);

            const auto* p0 = reinterpret_cast<const uint32_t*>(&w0);
            const auto* p1 = reinterpret_cast<const uint32_t*>(&w1);
            const auto* p2 = reinterpret_cast<const uint32_t*>(&w2);
            const auto* p3 = reinterpret_cast<const uint32_t*>(&w3);

            auto absMax = 1e-10f;
            #define AMAX_F(pair) { \
                const auto lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
                const auto hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
                const auto flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
                const auto fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
                absMax = (flo > absMax) ? flo : absMax; \
                absMax = (fhi > absMax) ? fhi : absMax; \
            }
            AMAX_F(p0[0]) AMAX_F(p0[1]) AMAX_F(p0[2]) AMAX_F(p0[3])
            AMAX_F(p1[0]) AMAX_F(p1[1]) AMAX_F(p1[2]) AMAX_F(p1[3])
            AMAX_F(p2[0]) AMAX_F(p2[1]) AMAX_F(p2[2]) AMAX_F(p2[3])
            AMAX_F(p3[0]) AMAX_F(p3[1]) AMAX_F(p3[2]) AMAX_F(p3[3])
            #undef AMAX_F

            const auto u32 = __builtin_bit_cast(uint32_t, absMax);
            const auto amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
            const auto inv_exp  = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
            a_scale_regs[kt] = static_cast<int>(inv_exp);
            const auto hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);

            #define CVT_F(d, pair, sel) \
                d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
                    d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)
            auto d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
            CVT_F(d0, p0[0], 0); CVT_F(d0, p0[1], 1); CVT_F(d0, p0[2], 2); CVT_F(d0, p0[3], 3);
            CVT_F(d1, p1[0], 0); CVT_F(d1, p1[1], 1); CVT_F(d1, p1[2], 2); CVT_F(d1, p1[3], 3);
            CVT_F(d2, p2[0], 0); CVT_F(d2, p2[1], 1); CVT_F(d2, p2[2], 2); CVT_F(d2, p2[3], 3);
            CVT_F(d3, p3[0], 0); CVT_F(d3, p3[1], 1); CVT_F(d3, p3[2], 2); CVT_F(d3, p3[3], 3);
            #undef CVT_F
            a_fp4_regs[kt] = int4_v{static_cast<int>(d0), static_cast<int>(d1),
                                     static_cast<int>(d2), static_cast<int>(d3)};
        } else {
            a_fp4_regs[kt] = int4_v{0, 0, 0, 0};
            a_scale_regs[kt] = 127;
        }
    }

    float4_v acc{0.f, 0.f, 0.f, 0.f};

    const auto n_tile = tile_n / 16;
    const auto bsh_n_stride = (long)(C_K / 64) * 512;
    const auto bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
    const auto k_half_off = (kgrp & 1) * 256;
    const auto k_blk_base = kgrp >> 1;

    const auto gn = tile_n + lrow;
    constexpr auto scaleN = (C_K / 32 + 7) / 8 * 8;
    const auto bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
                           + (gn >> 5) * (32 * scaleN);

    #pragma unroll
    for (auto kt = 0; kt < CKT; ++kt) {
        const auto bv = load16(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
        const auto ks_val = kt;
        const auto sb = static_cast<int>(*(Bssh + bssh_base
                         + (ks_val & 1) * 2 + (ks_val >> 1) * 256));

        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a_fp4_regs[kt], bv, acc, FP4_E2M1, FP4_E2M1,
            0, a_scale_regs[kt], 0, sb);
    }

    const auto out_col = tile_n + lrow;
    const auto out_row_base = tile_m + kgrp * 4;
    if (out_col >= C_N) return;

    auto* c_out = C_final + (long)out_row_base * C_N + out_col;
    #pragma unroll
    for (auto i = 0; i < 4; ++i, c_out += C_N) {
        if (out_row_base + i < C_M)
            *c_out = float_to_bf16(acc[i]);
    }
}

template<int C_M, int C_K, int C_N, int C_SCALEN, int C_KPS, int C_NUM_KSPLIT>
__global__ void __launch_bounds__(((C_M + 15) / 16) * 4 * 64, 2)
mxfp4_fused_quant_gemm_splitk(
    const __bf16* __restrict__ A_bf16,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    float*         __restrict__ C_partial)
{
    constexpr auto CKT = C_KPS;
    constexpr auto WAVES_M = (C_M + 15) / 16;
    constexpr auto WAVES_N = 4;
    constexpr auto BN = WAVES_N * 16;
    constexpr auto NWARPS = WAVES_M * WAVES_N;
    constexpr auto K2 = C_K / 2;

    const auto tid = threadIdx.x;
    const auto lane = tid % 64;
    const auto wave_id = tid / 64;
    const auto lrow = lane % 16;
    const auto kgrp = lane / 16;
    const auto wave_m = wave_id / WAVES_N;
    const auto wave_n = wave_id % WAVES_N;
    const auto ks_idx = static_cast<int>(blockIdx.z);

    const auto tile_m = wave_m * 16;
    const auto tile_n = static_cast<int>(blockIdx.x) * BN + wave_n * 16;
    const auto my_row = tile_m + lrow;

    const auto ks_start = ks_idx * C_KPS;
    const auto k_byte_start = ks_start * 128;

    constexpr auto A_SLICE_ELEMS = C_M * C_KPS * 128;
    constexpr auto A_SLICE_BYTES = A_SLICE_ELEMS * 2;
    extern __shared__ uint8_t smem[];
    {
        constexpr auto BYTES_PER_ITER = NWARPS * 64 * 16;
        constexpr auto NITERS = (A_SLICE_BYTES + BYTES_PER_ITER - 1) / BYTES_PER_ITER;
        #pragma unroll
        for (auto iter = 0; iter < NITERS; ++iter) {
            const auto offset = iter * BYTES_PER_ITER + tid * 16;
            if (offset + 16 <= A_SLICE_BYTES) {
                const auto slice_elem = offset / 2;
                const auto row = slice_elem / (C_KPS * 128);
                const auto col = slice_elem % (C_KPS * 128);
                const auto global_byte_off = (long)row * C_K * 2 + (long)(k_byte_start + col) * 2;
                if (row < C_M) {
                    *reinterpret_cast<int4_v*>(smem + offset) =
                        *reinterpret_cast<const int4_v*>(
                            reinterpret_cast<const uint8_t*>(A_bf16) + global_byte_off);
                }
            }
        }
    }
    __syncthreads();

    int4_v a_fp4_regs[CKT];
    int a_scale_regs[CKT];
    constexpr auto a_row_always_valid = (C_M >= ((C_M + 15) / 16) * 16);
    const auto a_row_valid = a_row_always_valid ? true : (my_row < C_M);

    #pragma unroll
    for (auto kt = 0; kt < CKT; ++kt) {
        const auto kg = kt * 4 + kgrp;
        if (a_row_valid) {
            const auto lds_off = my_row * C_KPS * 128 * 2 + kg * 64;

            const auto w0 = *reinterpret_cast<const int4_v*>(smem + lds_off);
            const auto w1 = *reinterpret_cast<const int4_v*>(smem + lds_off + 16);
            const auto w2 = *reinterpret_cast<const int4_v*>(smem + lds_off + 32);
            const auto w3 = *reinterpret_cast<const int4_v*>(smem + lds_off + 48);

            const auto* p0 = reinterpret_cast<const uint32_t*>(&w0);
            const auto* p1 = reinterpret_cast<const uint32_t*>(&w1);
            const auto* p2 = reinterpret_cast<const uint32_t*>(&w2);
            const auto* p3 = reinterpret_cast<const uint32_t*>(&w3);

            auto absMax = 1e-10f;
            #define AMAX_S(pair) { \
                const auto lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
                const auto hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
                const auto flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
                const auto fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
                absMax = (flo > absMax) ? flo : absMax; \
                absMax = (fhi > absMax) ? fhi : absMax; \
            }
            AMAX_S(p0[0]) AMAX_S(p0[1]) AMAX_S(p0[2]) AMAX_S(p0[3])
            AMAX_S(p1[0]) AMAX_S(p1[1]) AMAX_S(p1[2]) AMAX_S(p1[3])
            AMAX_S(p2[0]) AMAX_S(p2[1]) AMAX_S(p2[2]) AMAX_S(p2[3])
            AMAX_S(p3[0]) AMAX_S(p3[1]) AMAX_S(p3[2]) AMAX_S(p3[3])
            #undef AMAX_S

            const auto u32 = __builtin_bit_cast(uint32_t, absMax);
            const auto amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
            const auto inv_exp = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
            a_scale_regs[kt] = static_cast<int>(inv_exp);
            const auto hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);

            #define CVT_S(d, pair, sel) \
                d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
                    d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)
            auto d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
            CVT_S(d0, p0[0], 0); CVT_S(d0, p0[1], 1); CVT_S(d0, p0[2], 2); CVT_S(d0, p0[3], 3);
            CVT_S(d1, p1[0], 0); CVT_S(d1, p1[1], 1); CVT_S(d1, p1[2], 2); CVT_S(d1, p1[3], 3);
            CVT_S(d2, p2[0], 0); CVT_S(d2, p2[1], 1); CVT_S(d2, p2[2], 2); CVT_S(d2, p2[3], 3);
            CVT_S(d3, p3[0], 0); CVT_S(d3, p3[1], 1); CVT_S(d3, p3[2], 2); CVT_S(d3, p3[3], 3);
            #undef CVT_S
            a_fp4_regs[kt] = int4_v{static_cast<int>(d0), static_cast<int>(d1),
                                     static_cast<int>(d2), static_cast<int>(d3)};
        } else {
            a_fp4_regs[kt] = int4_v{0, 0, 0, 0};
            a_scale_regs[kt] = 127;
        }
    }

    float4_v acc{0.f, 0.f, 0.f, 0.f};
    const auto n_tile = tile_n / 16;
    const auto bsh_n_stride = (long)(C_K / 64) * 512;
    const auto bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
    const auto k_half_off = (kgrp & 1) * 256;
    const auto k_blk_base = kgrp >> 1;

    const auto gn = tile_n + lrow;
    constexpr auto scaleN = (C_K / 32 + 7) / 8 * 8;
    const auto bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
                           + (gn >> 5) * (32 * scaleN);

    #pragma unroll
    for (auto kt = 0; kt < CKT; ++kt) {
        const auto global_kt = ks_start + kt;
        const auto bv = load16(bsh_lane_base + (long)(global_kt * 2 + k_blk_base) * 512 + k_half_off);
        const auto ks_val = global_kt;
        const auto sb = static_cast<int>(*(Bssh + bssh_base
                         + (ks_val & 1) * 2 + (ks_val >> 1) * 256));

        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a_fp4_regs[kt], bv, acc, FP4_E2M1, FP4_E2M1,
            0, a_scale_regs[kt], 0, sb);
    }

    const auto out_col = tile_n + lrow;
    const auto out_row_base = tile_m + kgrp * 4;
    if (out_col >= C_N) return;

    auto* c_out = C_partial + (long)ks_idx * C_M * C_N + (long)out_row_base * C_N + out_col;
    #pragma unroll
    for (auto i = 0; i < 4; ++i, c_out += C_N) {
        if (out_row_base + i < C_M)
            *c_out = acc[i];
    }
}


template<int C_M, int C_K>
void launch_quant_generic(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale)
{
    constexpr auto KS       = C_K / 32;
    constexpr auto n_groups = C_M * KS;
    constexpr dim3 block{128};
    constexpr dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
    mxfp4_quant<C_M, C_K><<<grid, block>>>(A_bf16, A_fp4, A_scale);
}

template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT, int C_M>
void launch_gemm_nk_generic(
    const uint8_t* A, const uint8_t* As,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final)
{
    constexpr auto do_splitk = C_NUM_KSPLIT > 1;

    constexpr auto BM     = (C_M <=  8) ?  8 : (C_M <= 32) ? 16 : (C_M <= 128) ? 32 : 64;
    constexpr auto BN     = 32;
    constexpr auto NWARPS = ((BM + 15) / 16) * (BN / 16);
    static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);

    constexpr auto WAVES_M = (BM + 15) / 16;
    constexpr auto WAVES_N = BN / 16;
    constexpr auto CHUNK_K = (C_KPS >= 4) ? 4 : C_KPS;
    constexpr auto smem_size = (C_KPS > 4)
        ? LdsLayout<WAVES_M, WAVES_N, CHUNK_K>::TOTAL_LDS
        : 0;

    constexpr auto full_m  = (BM >= 16) ? C_M / BM : 0;
    constexpr auto full_n  = C_N / BN;
    constexpr auto total_m = (C_M + BM - 1) / BM;
    constexpr auto total_n = (C_N + BN - 1) / BN;
    constexpr auto edge_m  = total_m - full_m;
    constexpr auto edge_n  = total_n - full_n;

    constexpr dim3 block{static_cast<uint32_t>(NWARPS * 64)};

    auto sub = [&]<bool AV, bool BV>() {
        constexpr auto gx = BV ? full_n : edge_n;
        constexpr auto gy = AV ? full_m : edge_m;
        constexpr auto ox = BV ? 0 : full_n;
        constexpr auto oy = AV ? 0 : full_m;
        if constexpr (gx > 0 && gy > 0) {
            constexpr auto SK = do_splitk;
            if constexpr (smem_size > 0)
                (void)hipFuncSetAttribute(
                    (const void*)mxfp4_gemm<BM,BN,NWARPS,SK,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>,
                    hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
            constexpr dim3 grid{
                static_cast<uint32_t>(gx),
                static_cast<uint32_t>(gy),
                static_cast<uint32_t>(C_NUM_KSPLIT)
            };
            if constexpr (SK)
                mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>
                    <<<grid,block,smem_size>>>(A,As,Bsh,Bssh,C_partial,nullptr);
            else
                mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>
                    <<<grid,block,smem_size>>>(A,As,Bsh,Bssh,nullptr,C_final);
        }
    };

    sub.template operator()<true,  true >();
    sub.template operator()<true,  false>();
    sub.template operator()<false, true >();
    sub.template operator()<false, false>();

    if constexpr (do_splitk) {
        constexpr dim3 rblock{32, 16};
        constexpr dim3 rgrid{
            static_cast<uint32_t>((C_N + 31) / 32),
            static_cast<uint32_t>((C_M + 15) / 16)
        };
        mxfp4_reduce<C_N, C_NUM_KSPLIT, C_M><<<rgrid, rblock>>>(C_partial, C_final);
    }
}

template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_M>
void launch_gemm_nk_7168_generic(
    const uint8_t* A, const uint8_t* As,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final)
{
    if constexpr (C_M <= 8)
        launch_gemm_nk_generic<C_N, C_K, C_SCALEN, C_TOTAL_KT, 8, 7, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
    else if constexpr (C_M <= 16)
        launch_gemm_nk_generic<C_N, C_K, C_SCALEN, C_TOTAL_KT, 4, 14, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
    else
        launch_gemm_nk_generic<C_N, C_K, C_SCALEN, C_TOTAL_KT, 56, 1, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
}

template<int C_M, int C_K, int C_N, int C_SCALEN>
void launch_fused_generic(
    const __bf16* A_bf16,
    const uint8_t* Bsh, const uint8_t* Bssh,
    uint16_t* C_final)
{
    constexpr auto WAVES_M = (C_M + 15) / 16;
    constexpr auto WAVES_N = (C_M >= 32) ? 2 : 4;
    constexpr auto BN = WAVES_N * 16;
    constexpr auto NTHREADS = WAVES_M * WAVES_N * 64;
    constexpr auto smem_size = C_M * C_K * 2;
    if constexpr (smem_size > 48 * 1024)
        (void)hipFuncSetAttribute(
            (const void*)mxfp4_fused_quant_gemm<C_M, C_K, C_N, C_SCALEN, WAVES_N>,
            hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
    constexpr dim3 grid{static_cast<uint32_t>(C_N / BN)};
    constexpr dim3 block{static_cast<uint32_t>(NTHREADS)};
    mxfp4_fused_quant_gemm<C_M, C_K, C_N, C_SCALEN, WAVES_N>
        <<<grid, block, smem_size>>>(A_bf16, Bsh, Bssh, C_final);
}

template<int C_M, int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
void launch_shape_generic(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final)
{
    constexpr auto a_bf16_bytes = C_M * C_K * 2;
    if constexpr (C_KPS <= 4 && C_NUM_KSPLIT == 1 && a_bf16_bytes <= 32 * 1024) {
        launch_fused_generic<C_M, C_K, C_N, C_SCALEN>(A_bf16, Bsh, Bssh, C_final);
        return;
    }
    launch_quant_generic<C_M, C_K>(A_bf16, A_fp4, A_scale);
    launch_gemm_nk_generic<C_N, C_K, C_SCALEN, C_TOTAL_KT, C_KPS, C_NUM_KSPLIT, C_M>(
        A_fp4, A_scale, Bsh, Bssh, C_partial, C_final);
}

template<int C_M, int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT>
void launch_shape_7168_generic(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final)
{
    constexpr auto C_KPS = (C_M <= 8) ? 8 : (C_M <= 16) ? 4 : 56;
    constexpr auto C_NUM_KSPLIT = (C_M <= 8) ? 7 : (C_M <= 16) ? 14 : 1;
    constexpr auto a_slice_bytes = C_M * C_KPS * 128 * 2;

    if constexpr (C_NUM_KSPLIT > 1 && a_slice_bytes <= 32 * 1024) {
        constexpr auto WAVES_M = (C_M + 15) / 16;
        constexpr auto WAVES_N = 4;
        constexpr auto BN = WAVES_N * 16;
        constexpr auto NTHREADS = WAVES_M * WAVES_N * 64;
        constexpr auto smem_size = a_slice_bytes;

        if constexpr (smem_size > 48 * 1024)
            (void)hipFuncSetAttribute(
                (const void*)mxfp4_fused_quant_gemm_splitk<C_M, C_K, C_N, C_SCALEN, C_KPS, C_NUM_KSPLIT>,
                hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);

        constexpr dim3 grid{static_cast<uint32_t>(C_N / BN), 1, static_cast<uint32_t>(C_NUM_KSPLIT)};
        constexpr dim3 block{static_cast<uint32_t>(NTHREADS)};
        mxfp4_fused_quant_gemm_splitk<C_M, C_K, C_N, C_SCALEN, C_KPS, C_NUM_KSPLIT>
            <<<grid, block, smem_size>>>(A_bf16, Bsh, Bssh, C_partial);

        constexpr dim3 rblock{32, 16};
        constexpr dim3 rgrid{
            static_cast<uint32_t>((C_N + 31) / 32),
            static_cast<uint32_t>((C_M + 15) / 16)
        };
        mxfp4_reduce<C_N, C_NUM_KSPLIT, C_M><<<rgrid, rblock>>>(C_partial, C_final);
    } else {
        launch_quant_generic<C_M, C_K>(A_bf16, A_fp4, A_scale);
        launch_gemm_nk_7168_generic<C_N, C_K, C_SCALEN, C_TOTAL_KT, C_M>(
            A_fp4, A_scale, Bsh, Bssh, C_partial, C_final);
    }
}

template<int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
void dispatch_shape_generic(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final, int M)
{
    #define CASE_M(val) case val: launch_shape_generic<val,C_K,C_N,C_SCALEN,C_TOTAL_KT,C_KPS,C_NUM_KSPLIT>(A_bf16,A_fp4,A_scale,Bsh,Bssh,C_partial,C_final); return
    switch (M) {
        CASE_M(4); CASE_M(8); CASE_M(16); CASE_M(32); CASE_M(64); CASE_M(256);
    }
    #undef CASE_M
}

template<int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT>
void dispatch_shape_7168_generic(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final, int M)
{
    #define CASE_M(val) case val: launch_shape_7168_generic<val,C_K,C_N,C_SCALEN,C_TOTAL_KT>(A_bf16,A_fp4,A_scale,Bsh,Bssh,C_partial,C_final); return
    switch (M) {
        CASE_M(4); CASE_M(8); CASE_M(16); CASE_M(32); CASE_M(64); CASE_M(256);
    }
    #undef CASE_M
}

extern "C" void launch_all_generic(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final,
    int M, int N, int K)
{
    if (N == 2880 && K == 512)
        dispatch_shape_generic<512, 2880, 16, 4, 4, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
    else if (N == 2112 && K == 7168)
        dispatch_shape_7168_generic<7168, 2112, 224, 56>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
    else if (N == 4096 && K == 512)
        dispatch_shape_generic<512, 4096, 16, 4, 4, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
    else if (N == 7168 && K == 2048)
        dispatch_shape_generic<2048, 7168, 64, 16, 16, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
    else if (N == 3072 && K == 1536)
        dispatch_shape_generic<1536, 3072, 48, 12, 12, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
}


// ============================================================================
// DISPATCH
// ============================================================================

extern "C" void launch_all_wrapper(
    const void* A_bf16, void* A_fp4, void* A_scale,
    const void* Bsh, const void* Bssh,
    void* C_partial, void* C_final,
    int M, int N, int K)
{
    const __bf16* a = reinterpret_cast<const __bf16*>(A_bf16);
    uint8_t* fp4 = reinterpret_cast<uint8_t*>(A_fp4);
    uint8_t* asc = reinterpret_cast<uint8_t*>(A_scale);
    const uint8_t* bsh = reinterpret_cast<const uint8_t*>(Bsh);
    const uint8_t* bssh = reinterpret_cast<const uint8_t*>(Bssh);
    float* cp = reinterpret_cast<float*>(C_partial);
    uint16_t* cf = reinterpret_cast<uint16_t*>(C_final);

    if (M == 4 && N == 2880 && K == 512) {
        // Shape 1: fused quant+GEMM, BN=64
        dim3 grid{2880 / 64};  // 45
        dim3 block{256};
        gemm_m4_n2880_k512<<<grid, block, 4 * 512 * 2>>>(a, bsh, bssh, cf);
    }
    else if (M == 16 && N == 2112 && K == 7168) {
        // Shape 2: fused splitK — s483: no LDS (direct HBM quant)
        dim3 grid{2112 / 64, 1, 14};  // 33 x 1 x 14
        dim3 block{256};
        gemm_m16_n2112_k7168<<<grid, block, 0>>>(a, bsh, bssh, cp);

        dim3 rgrid{(2112 + 31) / 32, (16 + 15) / 16};
        dim3 rblock{32, 16};
        reduce_m16_n2112_k14<<<rgrid, rblock>>>(cp, cf);
    }
    else if (M == 32 && N == 4096 && K == 512) {
        // Shape 3: fused quant+GEMM, BN=32, s511: 2D grid for 2× CU utilization
        dim3 grid{4096 / 32, 2};  // 128×2 = 256 blocks (was 128)
        dim3 block{128};           // 2 warps (was 256/4 warps)
        gemm_m32_n4096_k512<<<grid, block, 0>>>(a, bsh, bssh, cf);
    }
    else if (M == 32 && N == 2880 && K == 512) {
        // Shape 4: fused quant+GEMM, BN=32, s511: 2D grid
        dim3 grid{2880 / 32, 2};  // 90×2 = 180 blocks (was 90)
        dim3 block{128};
        gemm_m32_n2880_k512<<<grid, block, 0>>>(a, bsh, bssh, cf);
    }
    else if (M == 64 && N == 7168 && K == 2048) {
        // Shape 5: s466 BM=16,BN=64 — 20% less HBM traffic per block
        // LDS: 2 * 24832 (double-buffered GEMM) + 2 * 512 (double-buffered scales) = 50688
        constexpr int smem5 = 2 * 24832 + 2 * 512;
        (void)hipFuncSetAttribute(
            (const void*)gemm_m64_n7168_k2048,
            hipFuncAttributeMaxDynamicSharedMemorySize, smem5);
        dim3 ggrid{112, 4, 1};
        dim3 gblock{256};
        gemm_m64_n7168_k2048<<<ggrid, gblock, smem5>>>(a, bsh, bssh, cf);
    }
    else if (M == 256 && N == 3072 && K == 1536) {
        // Shape 6: s468 BM=16 BN=64, G5-style scale handling
        // LDS: 2 * 24832 + 2 * 512 = 50688
        constexpr int smem6 = 2 * 24832 + 2 * 512;
        (void)hipFuncSetAttribute(
            (const void*)gemm_m256_n3072_k1536,
            hipFuncAttributeMaxDynamicSharedMemorySize, smem6);
        dim3 ggrid{48, 16, 1};
        dim3 gblock{256};
        gemm_m256_n3072_k1536<<<ggrid, gblock, smem6>>>(a, bsh, bssh, cf);
    }
    else {
        // Generic fallback for any shape not hardcoded above
        launch_all_generic(a, fp4, asc, bsh, bssh, cp, cf, M, N, K);
    }
}
"""



# =============================================================================
# Compilation flags (same as s292)
# =============================================================================
_CUDA_FLAGS = [
    "-O3", "-DNDEBUG", "-fno-exceptions", "-std=c++20",
    "-ffast-math", "-ffinite-math-only", "-munsafe-fp-atomics",
    "-mwavefrontsize64", "-mcumode",
    "-fgpu-flush-denormals-to-zero", "-fno-offload-uniform-block",
    "-mllvm", "-amdgpu-early-inline-all=true",
    "-mllvm", "-amdgpu-function-calls=false",
    "-mllvm", "--amdgpu-kernarg-preload-count=16",
    "-mllvm", "--lsr-drop-solution=1",
    "-mllvm", "-amdgpu-coerce-illegal-types=1",
    "-mllvm", "-amdgpu-loop-prefetch=true",
    "-mllvm", "-enable-unroll-and-jam=true",
    "-mllvm", "-unroll-threshold=500",
    "-mllvm", "-amdgpu-internalize-symbols=true",
    "-mllvm", "-greedy-reverse-local-assignment=1",
]

# =============================================================================
# CPP_SRC: C++ wrapper that handles ALL workspace management in C++.
# Zero Python in the hot path — single C++ function call.
# =============================================================================
CPP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <unordered_map>
#include <vector>

// Forward declare the HIP launch wrapper (defined in HIP_SRC / .so)
extern "C" void launch_all_wrapper(
    const void* A_bf16, void* A_fp4, void* A_scale,
    const void* Bsh, const void* Bssh,
    void* C_partial, void* C_final,
    int M, int N, int K);

// ============================================================================
// Static workspace cache — persists across calls, zero allocation in hot path
// ============================================================================

struct WorkspaceEntry {
    at::Tensor fp4;     // [M, K/2] uint8
    at::Tensor scale;   // [M, K/32] uint8
    at::Tensor out;     // [M, N] bf16
    at::Tensor partial; // [ns, M, N] f32 or empty
    void* fp4_ptr;
    void* scale_ptr;
    void* out_ptr;
    void* partial_ptr;
};

static std::unordered_map<int64_t, WorkspaceEntry> ws_cache;
static at::Tensor cached_bsh, cached_bssh;
static int64_t cached_bsh_key = 0, cached_bssh_key = 0;
static void* cached_bsh_ptr = nullptr;
static void* cached_bssh_ptr = nullptr;

static int compute_num_ksplits(int M, int K) {
    int t = K / 128;
    if (t < 14) return 1;
    return (M <= 8) ? 7 : ((M <= 16) ? 14 : 1);
}

at::Tensor custom_kernel(std::vector<at::Tensor> data) {
    // data = [A, _, Q, S, SS]
    auto A  = data[0];
    auto Q  = data[2];
    auto S  = data[3];
    auto SS = data[4];

    int M  = static_cast<int>(A.size(0));
    int Kv = static_cast<int>(A.size(1));
    int N  = static_cast<int>(Q.size(0));

    // Ensure bf16 + contiguous
    if (A.scalar_type() != at::kBFloat16)
        A = A.to(at::kBFloat16);
    if (!A.is_contiguous())
        A = A.contiguous();

    // Cache B scale tensors (only recompute on pointer change)
    auto sp  = reinterpret_cast<int64_t>(S.data_ptr());
    auto ssp = reinterpret_cast<int64_t>(SS.data_ptr());
    if (sp != cached_bsh_key || ssp != cached_bssh_key) {
        cached_bsh = S.view(torch::kUInt8);
        if (!cached_bsh.is_contiguous()) cached_bsh = cached_bsh.contiguous();
        cached_bssh = SS.view(torch::kUInt8);
        if (!cached_bssh.is_contiguous()) cached_bssh = cached_bssh.contiguous();
        cached_bsh_key  = sp;
        cached_bssh_key = ssp;
        cached_bsh_ptr  = cached_bsh.data_ptr();
        cached_bssh_ptr = cached_bssh.data_ptr();
    }

    // Workspace cache lookup — key = (M << 32) | (N << 16) | K
    int64_t key = ((int64_t)M << 32) | ((int64_t)N << 16) | (int64_t)Kv;
    auto it = ws_cache.find(key);
    if (it == ws_cache.end()) {
        int ns = compute_num_ksplits(M, Kv);
        auto opts = A.options();
        WorkspaceEntry ws;
        ws.fp4     = torch::empty({M, Kv / 2}, opts.dtype(torch::kUInt8));
        ws.scale   = torch::empty({M, Kv / 32}, opts.dtype(torch::kUInt8));
        ws.out     = torch::empty({M, N}, opts.dtype(torch::kBFloat16));
        ws.partial = (ns > 1)
            ? torch::empty({ns, M, N}, opts.dtype(torch::kFloat32))
            : torch::Tensor();
        ws.fp4_ptr     = ws.fp4.data_ptr();
        ws.scale_ptr   = ws.scale.data_ptr();
        ws.out_ptr     = ws.out.data_ptr();
        ws.partial_ptr = ws.partial.defined() ? ws.partial.data_ptr() : nullptr;
        it = ws_cache.emplace(key, std::move(ws)).first;
    }

    auto& ws = it->second;

    launch_all_wrapper(
        A.data_ptr(),
        ws.fp4_ptr, ws.scale_ptr,
        cached_bsh_ptr, cached_bssh_ptr,
        ws.partial_ptr, ws.out_ptr,
        M, N, Kv);

    return ws.out;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("custom_kernel", &custom_kernel, "MXFP4 GEMM with zero-Python hot path");
}
"""

# =============================================================================
from torch.utils.cpp_extension import load_inline
_module = load_inline(
    name="mxfp4_gemm",
    cpp_sources=[CPP_SRC],
    cuda_sources=[HIP_SRC],
    extra_cflags=["-O3", "-std=c++20", "-DNDEBUG"],
    extra_cuda_cflags=_CUDA_FLAGS,
    verbose=False,
)

def custom_kernel(data):
    return _module.custom_kernel(list(data))
scrolls · 3117 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 638649.

⋯ 230 unchanged lines
const int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
+ (gn >> 5) * (32 * S1_SCALEN);
+ // s483: Prefetch all B tiles + scales into registers before MFMA loop
+ int4_v b_regs_s1[S1_CKT];
+ int sb_regs_s1[S1_CKT];
#pragma unroll
for (int kt = 0; kt < S1_CKT; ++kt) {
- const int4_v bv = load16(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
- const int sb = static_cast<int>(*(Bssh + bssh_base
+ b_regs_s1[kt] = load16_nt(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
+ sb_regs_s1[kt] = static_cast<int>(*(Bssh + bssh_base
+ (kt & 1) * 2 + (kt >> 1) * 256));
+ }
+
+ #pragma unroll
+ for (int kt = 0; kt < S1_CKT; ++kt) {
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
- a_fp4_regs[kt], bv, acc, FP4_E2M1, FP4_E2M1,
- 0, a_scale_regs[kt], 0, sb);
+ a_fp4_regs[kt], b_regs_s1[kt], acc, FP4_E2M1, FP4_E2M1,
+ 0, a_scale_regs[kt], 0, sb_regs_s1[kt]);
}
const int out_col = tile_n + lrow;
⋯ 51 unchanged lines
const int my_row = lrow;
const int ks_start = ks_idx * S2_KPS;
- const int k_byte_start = ks_start * 128;
- // Load A K-slice into LDS: 16 rows * 4*128 bf16 = 16*512*2 = 16384 bytes
- extern __shared__ uint8_t smem[];
- {
- // 256 threads * 16 = 4096 bytes/iter, need ceil(16384/4096) = 4 iters
- #pragma unroll
- for (int iter = 0; iter < 4; ++iter) {
- const int offset = iter * 4096 + tid * 16;
- if (offset + 16 <= S2_M * S2_KPS * 128 * 2) {
- const int slice_elem = offset / 2;
- const int row = slice_elem / (S2_KPS * 128);
- const int col = slice_elem % (S2_KPS * 128);
- const long global_byte_off = (long)row * S2_K * 2 + (long)(k_byte_start + col) * 2;
- if (row < S2_M) {
- *reinterpret_cast<int4_v*>(smem + offset) =
- *reinterpret_cast<const int4_v*>(
- reinterpret_cast<const uint8_t*>(A_bf16) + global_byte_off);
- }
- }
- }
- }
- __syncthreads();
-
- // M=16: all rows valid (16 rows, 16 per wave)
+ // s483: Direct HBM quant for M=16 — skip LDS, quant A directly from global memory.
+ // 4 N-waves read the same A rows → L2 handles the sharing (16KB per split).
+ const uint8_t* _a_byte = reinterpret_cast<const uint8_t*>(A_bf16);
int4_v a_fp4_regs[S2_CKT];
int a_scale_regs[S2_CKT];
#pragma unroll
for (int kt = 0; kt < S2_CKT; ++kt) {
const int kg = kt * 4 + kgrp;
- const int lds_off = my_row * S2_KPS * 128 * 2 + kg * 64;
- QUANT_GROUP_32(smem, lds_off, a_fp4_regs[kt], a_scale_regs[kt])
+ const long global_off = (long)my_row * S2_K * 2 + (long)(ks_start * 128 + kg * 32) * 2;
+ QUANT_GROUP_32(_a_byte, global_off, a_fp4_regs[kt], a_scale_regs[kt])
}
float4_v acc{0.f, 0.f, 0.f, 0.f};
⋯ 7 unchanged lines
const int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
+ (gn >> 5) * (32 * S2_SCALEN);
+ // s483: Prefetch all B tiles + scales into registers before MFMA loop
+ int4_v b_regs_s2[S2_CKT];
+ int sb_regs_s2[S2_CKT];
#pragma unroll
for (int kt = 0; kt < S2_CKT; ++kt) {
const int global_kt = ks_start + kt;
- const int4_v bv = load16(bsh_lane_base + (long)(global_kt * 2 + k_blk_base) * 512 + k_half_off);
- const int sb = static_cast<int>(*(Bssh + bssh_base
+ b_regs_s2[kt] = load16_nt(bsh_lane_base + (long)(global_kt * 2 + k_blk_base) * 512 + k_half_off);
+ sb_regs_s2[kt] = static_cast<int>(*(Bssh + bssh_base
+ (global_kt & 1) * 2 + (global_kt >> 1) * 256));
+ }
+
+ // MFMA loop — all operands in registers, no memory stalls
+ #pragma unroll
+ for (int kt = 0; kt < S2_CKT; ++kt) {
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
- a_fp4_regs[kt], bv, acc, FP4_E2M1, FP4_E2M1,
- 0, a_scale_regs[kt], 0, sb);
+ a_fp4_regs[kt], b_regs_s2[kt], acc, FP4_E2M1, FP4_E2M1,
+ 0, a_scale_regs[kt], 0, sb_regs_s2[kt]);
}
const int out_col = tile_n + lrow;
⋯ 49 unchanged lines
// ============================================================================
// SHAPE 3: M=32, N=4096, K=512 — Fused quant+GEMM, BN=32
- // WAVES_M=2, WAVES_N=2, CKT=4, 256 threads
+ // s511: 128 threads (2 warps), 2D grid for 2× CU utilization (256 vs 128 blocks)
+ // WAVES_M=1 (via grid.y=2), WAVES_N=2, CKT=4
// ============================================================================
- __global__ void __launch_bounds__(256, 2)
+ __global__ void __launch_bounds__(128, 4)
gemm_m32_n4096_k512(
const __bf16* __restrict__ A_bf16,
const uint8_t* __restrict__ Bsh,
⋯ 5 unchanged lines
#define S3_N 4096
#define S3_CKT 4
#define S3_BN 32
- #define S3_WAVES_M 2
+ #define S3_WAVES_M 1
#define S3_WAVES_N 2
#define S3_SCALEN 16
⋯ 2 unchanged lines
const int wave_id = tid / 64;
const int lrow = lane % 16;
const int kgrp = lane / 16;
- const int wave_m = wave_id / S3_WAVES_N;
const int wave_n = wave_id % S3_WAVES_N;
- const int tile_m = wave_m * 16;
+ const int tile_m = static_cast<int>(blockIdx.y) * 16; // s511: M-tile from grid.y
const int tile_n = static_cast<int>(blockIdx.x) * S3_BN + wave_n * 16;
const int my_row = tile_m + lrow;
⋯ 21 unchanged lines
const int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
+ (gn >> 5) * (32 * S3_SCALEN);
+ // s483: Prefetch all B tiles + scales into registers before MFMA loop
+ int4_v b_regs_s3[S3_CKT];
+ int sb_regs_s3[S3_CKT];
#pragma unroll
for (int kt = 0; kt < S3_CKT; ++kt) {
- const int4_v bv = load16(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
- const int sb = static_cast<int>(*(Bssh + bssh_base
+ b_regs_s3[kt] = load16_nt(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
+ sb_regs_s3[kt] = static_cast<int>(*(Bssh + bssh_base
+ (kt & 1) * 2 + (kt >> 1) * 256));
+ }
+
+ #pragma unroll
+ for (int kt = 0; kt < S3_CKT; ++kt) {
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
- a_fp4_regs[kt], bv, acc, FP4_E2M1, FP4_E2M1,
- 0, a_scale_regs[kt], 0, sb);
+ a_fp4_regs[kt], b_regs_s3[kt], acc, FP4_E2M1, FP4_E2M1,
+ 0, a_scale_regs[kt], 0, sb_regs_s3[kt]);
}
const int out_col = tile_n + lrow;
⋯ 20 unchanged lines
// ============================================================================
// SHAPE 4: M=32, N=2880, K=512 — Fused quant+GEMM, BN=32
- // Same structure as shape 3 but different N
+ // s511: 128 threads, 2D grid (180 blocks vs 90 = 2× CU utilization)
// ============================================================================
- __global__ void __launch_bounds__(256, 2)
+ __global__ void __launch_bounds__(128, 4)
gemm_m32_n2880_k512(
const __bf16* __restrict__ A_bf16,
const uint8_t* __restrict__ Bsh,
⋯ 5 unchanged lines
#define S4_N 2880
#define S4_CKT 4
#define S4_BN 32
- #define S4_WAVES_M 2
+ #define S4_WAVES_M 1
#define S4_WAVES_N 2
#define S4_SCALEN 16
⋯ 2 unchanged lines
const int wave_id = tid / 64;
const int lrow = lane % 16;
const int kgrp = lane / 16;
- const int wave_m = wave_id / S4_WAVES_N;
const int wave_n = wave_id % S4_WAVES_N;
- const int tile_m = wave_m * 16;
+ const int tile_m = static_cast<int>(blockIdx.y) * 16; // s511: M-tile from grid.y
const int tile_n = static_cast<int>(blockIdx.x) * S4_BN + wave_n * 16;
const int my_row = tile_m + lrow;
⋯ 21 unchanged lines
const int bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
+ (gn >> 5) * (32 * S4_SCALEN);
+ // s483: Prefetch all B tiles + scales into registers before MFMA loop
+ int4_v b_regs_s4[S4_CKT];
+ int sb_regs_s4[S4_CKT];
#pragma unroll
for (int kt = 0; kt < S4_CKT; ++kt) {
- const int4_v bv = load16(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
- const int sb = static_cast<int>(*(Bssh + bssh_base
+ b_regs_s4[kt] = load16_nt(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
+ sb_regs_s4[kt] = static_cast<int>(*(Bssh + bssh_base
+ (kt & 1) * 2 + (kt >> 1) * 256));
+ }
+
+ #pragma unroll
+ for (int kt = 0; kt < S4_CKT; ++kt) {
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
- a_fp4_regs[kt], bv, acc, FP4_E2M1, FP4_E2M1,
- 0, a_scale_regs[kt], 0, sb);
+ a_fp4_regs[kt], b_regs_s4[kt], acc, FP4_E2M1, FP4_E2M1,
+ 0, a_scale_regs[kt], 0, sb_regs_s4[kt]);
}
const int out_col = tile_n + lrow;
⋯ 618 unchanged lines
G5_COMPUTE_4MFMA(buf, chunk_ks, 0, scale_base); \
}
+ // s481: 4-phase pipeline — issue VMEM before COMPUTE to overlap latency
+ // Persistent A bf16 registers (live across COMPUTE)
+ int4_v _pst_w0, _pst_w1, _pst_w2, _pst_w3;
+
+ // Phase macros for split pipeline
+ #define G5_ISSUE_VMEM(ks_base) \
+ { \
+ const unsigned _b_soff = __builtin_amdgcn_readfirstlane(static_cast<unsigned>((ks_base) * 1024)); \
+ _Pragma("unroll") \
+ for (int i = 0; i < G5_B_PER_THREAD; ++i) { \
+ BUFFER_LOAD_DWORDX4_NT(b_regs[i], b_voff[i], rsrc_bsh, _b_soff); \
+ } \
+ { \
+ int _lin = tid; \
+ int _kidx = _lin % (G5_CHUNK_K * 4); \
+ int _mloc = (_lin / (G5_CHUNK_K * 4)) % 16; \
+ int _wmid = _lin / (16 * G5_CHUNK_K * 4); \
+ int _row = tile_m_base + _wmid * 16 + _mloc; \
+ const int4_v* _src = reinterpret_cast<const int4_v*>( \
+ reinterpret_cast<const uint8_t*>(A_bf16) + \
+ (long)_row * G5_K * 2 + (long)((ks_base) * 128 + _kidx * 32) * 2); \
+ _pst_w0 = _src[0]; \
+ _pst_w1 = _src[1]; \
+ _pst_w2 = _src[2]; \
+ _pst_w3 = _src[3]; \
+ } \
+ }
+
+ #define G5_COMPLETE_QUANT_STORE(buf, scale_base) \
+ { \
+ const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
+ { \
+ const uint32_t* _p0 = reinterpret_cast<const uint32_t*>(&_pst_w0); \
+ const uint32_t* _p1 = reinterpret_cast<const uint32_t*>(&_pst_w1); \
+ const uint32_t* _p2 = reinterpret_cast<const uint32_t*>(&_pst_w2); \
+ const uint32_t* _p3 = reinterpret_cast<const uint32_t*>(&_pst_w3); \
+ float _absMax = 1e-10f; \
+ { \
+ auto _AM = [&](uint32_t pair) { \
+ __bf16 lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
+ __bf16 hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair >> 16)); \
+ float flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
+ float fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
+ _absMax = (flo > _absMax) ? flo : _absMax; \
+ _absMax = (fhi > _absMax) ? fhi : _absMax; \
+ }; \
+ _AM(_p0[0]); _AM(_p0[1]); _AM(_p0[2]); _AM(_p0[3]); \
+ _AM(_p1[0]); _AM(_p1[1]); _AM(_p1[2]); _AM(_p1[3]); \
+ _AM(_p2[0]); _AM(_p2[1]); _AM(_p2[2]); _AM(_p2[3]); \
+ _AM(_p3[0]); _AM(_p3[1]); _AM(_p3[2]); _AM(_p3[3]); \
+ } \
+ uint32_t _u32 = __builtin_bit_cast(uint32_t, _absMax); \
+ uint32_t _amax_exp = ((_u32 + 0x200000u) >> 23) & 0xFFu; \
+ uint32_t _inv_exp = (_amax_exp >= 2u) ? (_amax_exp - 2u) : 0u; \
+ float _hw_scale = __builtin_bit_cast(float, _inv_exp << 23); \
+ unsigned _d0 = 0u, _d1 = 0u, _d2 = 0u, _d3 = 0u; \
+ _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[0]), _hw_scale, 0); \
+ _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[1]), _hw_scale, 1); \
+ _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[2]), _hw_scale, 2); \
+ _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[3]), _hw_scale, 3); \
+ _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[0]), _hw_scale, 0); \
+ _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[1]), _hw_scale, 1); \
+ _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[2]), _hw_scale, 2); \
+ _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[3]), _hw_scale, 3); \
+ _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[0]), _hw_scale, 0); \
+ _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[1]), _hw_scale, 1); \
+ _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[2]), _hw_scale, 2); \
+ _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[3]), _hw_scale, 3); \
+ _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[0]), _hw_scale, 0); \
+ _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[1]), _hw_scale, 1); \
+ _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[2]), _hw_scale, 2); \
+ _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[3]), _hw_scale, 3); \
+ int4_v _fp4 = int4_v{static_cast<int>(_d0), static_cast<int>(_d1), \
+ static_cast<int>(_d2), static_cast<int>(_d3)}; \
+ { \
+ const unsigned _off = _buf_off + a_lds_offs[0]; \
+ asm volatile("ds_write_b128 %0, %1 offset:%2" \
+ :: "v"(_off), "v"(_fp4), "n"(0) : "memory"); \
+ } \
+ { \
+ int _linear = tid; \
+ int _k_idx = _linear % (G5_CHUNK_K * 4); \
+ int _m_local = (_linear / (G5_CHUNK_K * 4)) % 16; \
+ int _wave_m_idx = _linear / (16 * G5_CHUNK_K * 4); \
+ unsigned _scale_off = (scale_base) + _wave_m_idx * 512 + _m_local * 32 + _k_idx; \
+ smem_raw[_scale_off] = static_cast<uint8_t>(_inv_exp); \
+ } \
+ } \
+ asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); \
+ _Pragma("unroll") \
+ for (int i = 0; i < G5_B_PER_THREAD; ++i) { \
+ const unsigned _off = _buf_off + b_lds_offs[i]; \
+ asm volatile("ds_write_b128 %0, %1 offset:%2" \
+ :: "v"(_off), "v"(b_regs[i]), "n"(0) : "memory"); \
+ } \
+ }
+
int cur_ks = 0;
unsigned cur_scale_base = G5_SCALE_LDS_BASE0;
unsigned nxt_scale_base = G5_SCALE_LDS_BASE1;
+
+ // Initial prefetch: full load for chunk 0
G5_ISSUE_LOADS_FUSED(cur_ks, buf0, cur_scale_base);
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
__syncthreads();
⋯ 1 unchanged lines
uint8_t* cur_buf = buf0;
uint8_t* nxt_buf = buf1;
- // NUM_CHUNKS=2, so 1 iteration + final
+ // 4-phase pipeline loop: VMEM→COMPUTE→QUANT_STORE→sync
#pragma unroll
for (int c = 0; c < G5_NUM_CHUNKS - 1; ++c) {
- G5_ISSUE_LOADS_FUSED(cur_ks + G5_CHUNK_K, nxt_buf, nxt_scale_base);
- G5_COMPUTE_CHUNK(cur_buf, cur_ks, cur_scale_base);
+ G5_ISSUE_VMEM(cur_ks + G5_CHUNK_K); // Phase 1: issue VMEM (fast)
+ G5_COMPUTE_CHUNK(cur_buf, cur_ks, cur_scale_base); // Phase 2: compute (overlaps VMEM)
+ G5_COMPLETE_QUANT_STORE(nxt_buf, nxt_scale_base); // Phase 3: wait + quant + store
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
- __syncthreads();
+ __syncthreads(); // Phase 4: sync
uint8_t* tmp = cur_buf;
cur_buf = nxt_buf;
⋯ 6 unchanged lines
G5_COMPUTE_CHUNK(cur_buf, cur_ks, cur_scale_base);
+ #undef G5_ISSUE_VMEM
+ #undef G5_COMPLETE_QUANT_STORE
+
#undef G5_ISSUE_LOADS_FUSED
#undef G5_COMPUTE_4MFMA
#undef G5_COMPUTE_CHUNK
⋯ 317 unchanged lines
__builtin_amdgcn_sched_barrier(0x020); \
}
+ // s482: 4-phase pipeline for G6 — issue VMEM before COMPUTE
+ int4_v _g6_pst_w0, _g6_pst_w1, _g6_pst_w2, _g6_pst_w3;
+
+ #define G6_ISSUE_VMEM(ks_base) \
+ { \
+ const unsigned _b_soff = __builtin_amdgcn_readfirstlane(static_cast<unsigned>((ks_base) * 1024)); \
+ _Pragma("unroll") \
+ for (int i = 0; i < G6_B_PER_THREAD; ++i) { \
+ BUFFER_LOAD_DWORDX4_NT(b_regs[i], b_voff[i], rsrc_bsh, _b_soff); \
+ } \
+ { \
+ int _lin = tid; \
+ int _kidx = _lin % (G6_CHUNK_K * 4); \
+ int _mloc = (_lin / (G6_CHUNK_K * 4)) % 16; \
+ int _wmid = _lin / (16 * G6_CHUNK_K * 4); \
+ int _row = tile_m_base + _wmid * 16 + _mloc; \
+ const int4_v* _src = reinterpret_cast<const int4_v*>( \
+ reinterpret_cast<const uint8_t*>(A_bf16) + \
+ (long)_row * G6_K * 2 + (long)((ks_base) * 128 + _kidx * 32) * 2); \
+ _g6_pst_w0 = _src[0]; \
+ _g6_pst_w1 = _src[1]; \
+ _g6_pst_w2 = _src[2]; \
+ _g6_pst_w3 = _src[3]; \
+ } \
+ }
+
+ #define G6_COMPLETE_QUANT_STORE(buf, scale_base) \
+ { \
+ const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
+ { \
+ const uint32_t* _p0 = reinterpret_cast<const uint32_t*>(&_g6_pst_w0); \
+ const uint32_t* _p1 = reinterpret_cast<const uint32_t*>(&_g6_pst_w1); \
+ const uint32_t* _p2 = reinterpret_cast<const uint32_t*>(&_g6_pst_w2); \
+ const uint32_t* _p3 = reinterpret_cast<const uint32_t*>(&_g6_pst_w3); \
+ float _absMax = 1e-10f; \
+ { \
+ auto _AM = [&](uint32_t pair) { \
+ __bf16 lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
+ __bf16 hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair >> 16)); \
+ float flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
+ float fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
+ _absMax = (flo > _absMax) ? flo : _absMax; \
+ _absMax = (fhi > _absMax) ? fhi : _absMax; \
+ }; \
+ _AM(_p0[0]); _AM(_p0[1]); _AM(_p0[2]); _AM(_p0[3]); \
+ _AM(_p1[0]); _AM(_p1[1]); _AM(_p1[2]); _AM(_p1[3]); \
+ _AM(_p2[0]); _AM(_p2[1]); _AM(_p2[2]); _AM(_p2[3]); \
+ _AM(_p3[0]); _AM(_p3[1]); _AM(_p3[2]); _AM(_p3[3]); \
+ } \
+ uint32_t _u32 = __builtin_bit_cast(uint32_t, _absMax); \
+ uint32_t _amax_exp = ((_u32 + 0x200000u) >> 23) & 0xFFu; \
+ uint32_t _inv_exp = (_amax_exp >= 2u) ? (_amax_exp - 2u) : 0u; \
+ float _hw_scale = __builtin_bit_cast(float, _inv_exp << 23); \
+ unsigned _d0 = 0u, _d1 = 0u, _d2 = 0u, _d3 = 0u; \
+ _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[0]), _hw_scale, 0); \
+ _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[1]), _hw_scale, 1); \
+ _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[2]), _hw_scale, 2); \
+ _d0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d0, __builtin_bit_cast(bf16x2, _p0[3]), _hw_scale, 3); \
+ _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[0]), _hw_scale, 0); \
+ _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[1]), _hw_scale, 1); \
+ _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[2]), _hw_scale, 2); \
+ _d1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d1, __builtin_bit_cast(bf16x2, _p1[3]), _hw_scale, 3); \
+ _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[0]), _hw_scale, 0); \
+ _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[1]), _hw_scale, 1); \
+ _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[2]), _hw_scale, 2); \
+ _d2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d2, __builtin_bit_cast(bf16x2, _p2[3]), _hw_scale, 3); \
+ _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[0]), _hw_scale, 0); \
+ _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[1]), _hw_scale, 1); \
+ _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[2]), _hw_scale, 2); \
+ _d3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(_d3, __builtin_bit_cast(bf16x2, _p3[3]), _hw_scale, 3); \
+ int4_v _fp4 = int4_v{static_cast<int>(_d0), static_cast<int>(_d1), \
+ static_cast<int>(_d2), static_cast<int>(_d3)}; \
+ { \
+ const unsigned _off = _buf_off + a_lds_offs[0]; \
+ asm volatile("ds_write_b128 %0, %1 offset:%2" \
+ :: "v"(_off), "v"(_fp4), "n"(0) : "memory"); \
+ } \
+ { \
+ int _linear = tid; \
+ int _k_idx = _linear % (G6_CHUNK_K * 4); \
+ int _m_local = (_linear / (G6_CHUNK_K * 4)) % 16; \
+ int _wave_m_idx = _linear / (16 * G6_CHUNK_K * 4); \
+ unsigned _scale_off = (scale_base) + _wave_m_idx * 512 + _m_local * 32 + _k_idx; \
+ smem_raw[_scale_off] = static_cast<uint8_t>(_inv_exp); \
+ } \
+ } \
+ asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); \
+ _Pragma("unroll") \
+ for (int i = 0; i < G6_B_PER_THREAD; ++i) { \
+ const unsigned _off = _buf_off + b_lds_offs[i]; \
+ asm volatile("ds_write_b128 %0, %1 offset:%2" \
+ :: "v"(_off), "v"(b_regs[i]), "n"(0) : "memory"); \
+ } \
+ }
+
int cur_ks = 0;
unsigned cur_scale_base = G6_SCALE_LDS_BASE0;
unsigned nxt_scale_base = G6_SCALE_LDS_BASE1;
+
+ // Initial prefetch: full load for chunk 0
G6_ISSUE_LOADS_FUSED(cur_ks, buf0, cur_scale_base);
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
__syncthreads();
⋯ 1 unchanged lines
uint8_t* cur_buf = buf0;
uint8_t* nxt_buf = buf1;
- // NUM_CHUNKS=3, so 2 iterations + final
+ // 4-phase pipeline loop: VMEM→COMPUTE→QUANT_STORE→sync
#pragma unroll
for (int c = 0; c < G6_NUM_CHUNKS - 1; ++c) {
- G6_ISSUE_LOADS_FUSED(cur_ks + G6_CHUNK_K, nxt_buf, nxt_scale_base);
+ G6_ISSUE_VMEM(cur_ks + G6_CHUNK_K);
G6_COMPUTE_CHUNK(cur_buf, cur_ks, cur_scale_base);
+ G6_COMPLETE_QUANT_STORE(nxt_buf, nxt_scale_base);
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
__syncthreads();
⋯ 7 unchanged lines
}
G6_COMPUTE_CHUNK(cur_buf, cur_ks, cur_scale_base);
- // No tail (12 % 4 == 0)
+ #undef G6_ISSUE_VMEM
+ #undef G6_COMPLETE_QUANT_STORE
#undef G6_ISSUE_LOADS_FUSED
#undef G6_COMPUTE_CHUNK
⋯ 1140 unchanged lines
gemm_m4_n2880_k512<<<grid, block, 4 * 512 * 2>>>(a, bsh, bssh, cf);
}
else if (M == 16 && N == 2112 && K == 7168) {
- // Shape 2: fused splitK
+ // Shape 2: fused splitK — s483: no LDS (direct HBM quant)
dim3 grid{2112 / 64, 1, 14}; // 33 x 1 x 14
dim3 block{256};
- gemm_m16_n2112_k7168<<<grid, block, 16 * 4 * 128 * 2>>>(a, bsh, bssh, cp);
+ gemm_m16_n2112_k7168<<<grid, block, 0>>>(a, bsh, bssh, cp);
dim3 rgrid{(2112 + 31) / 32, (16 + 15) / 16};
dim3 rblock{32, 16};
reduce_m16_n2112_k14<<<rgrid, rblock>>>(cp, cf);
}
else if (M == 32 && N == 4096 && K == 512) {
- // Shape 3: fused quant+GEMM, BN=32
- dim3 grid{4096 / 32}; // 128
- dim3 block{256};
- gemm_m32_n4096_k512<<<grid, block, 0>>>(a, bsh, bssh, cf); // s480: no LDS
+ // Shape 3: fused quant+GEMM, BN=32, s511: 2D grid for 2× CU utilization
+ dim3 grid{4096 / 32, 2}; // 128×2 = 256 blocks (was 128)
+ dim3 block{128}; // 2 warps (was 256/4 warps)
+ gemm_m32_n4096_k512<<<grid, block, 0>>>(a, bsh, bssh, cf);
}
else if (M == 32 && N == 2880 && K == 512) {
- // Shape 4: fused quant+GEMM, BN=32
- dim3 grid{2880 / 32}; // 90
- dim3 block{256};
- gemm_m32_n2880_k512<<<grid, block, 0>>>(a, bsh, bssh, cf); // s480: no LDS
+ // Shape 4: fused quant+GEMM, BN=32, s511: 2D grid
+ dim3 grid{2880 / 32, 2}; // 90×2 = 180 blocks (was 90)
+ dim3 block{128};
+ gemm_m32_n2880_k512<<<grid, block, 0>>>(a, bsh, bssh, cf);
}
else if (M == 64 && N == 7168 && K == 2048) {
// Shape 5: s466 BM=16,BN=64 — 20% less HBM traffic per block
scrolls · 526 diff lines total

Best evidence level for this revision: reported

JSON