Skip to content
KernelIndex
Search⌘K

submission 645431

mmk150 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_general_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-645431?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.59µs
#74 of 1143
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4988b381774610a188bcc2d5527cf2822cc293e639cc799796629215be0e8949
license declaredunknown
license concludedunknown
authorsmmk150
imported2026-08-15

Techniques

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

shared-memoryextern __shared__ char lds_raw[];
split-kstatic constexpr int kSplitK = 7;
tile-m = 16static constexpr int kTileM=16, kNumMTiles=kM/kTileM;
tile-n = 32static constexpr int kTileN = 32;
vector-width = uint4const uint4* a_src0 = reinterpret_cast<const uint4*>(a_row + k_group0 * 16);

Kernel source

submission_general_v1.py1500 lines
"""
submission_general_v1.py — Based on submission_general_v0.py.
No preallocated buffer pools — all tensors allocated fresh per call.

Champions:
  M=4   (4,2880,512)    : m4_mk2a (2-wave, shuffled B)
  M=16  (16,2112,7168)  : m16_varI (two-pass workspace + reduce)
  M=32  (32,2880,512)   : m32_mk4l (16x16x128 MFMA, split-M, kernarg preload)
  M=32  (32,4096,512)   : m32_mk4l (same kernel, different N)
  M=64  (64,7168,2048)  : e641v1_exp S2 (hand-unrolled 8-iter, AGPR, B-scale LDS)
  M=256 (256,3072,1536) : mk3g S4 (v_max3 absmax, deferred B scale)
"""

from __future__ import annotations

import os
import tempfile

import aiter
import torch
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

# =========================================================================
# C++ bridge
# =========================================================================
_CPP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <cstdint>

// ── M=4 extern ──
extern "C" void launch_m4_mk2a(
    const uint16_t* A, const uint8_t* B_shuffle, const uint8_t* B_scale_sh,
    uint16_t* C);

void m4_launch(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, torch::Tensor C
) {
    launch_m4_mk2a(
        reinterpret_cast<const uint16_t*>(A.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
        reinterpret_cast<uint16_t*>(C.data_ptr()));
}

// ── M=16 extern ──
extern "C" void launch_m16_varI_gemm(
    const uint16_t* A, const uint8_t* B, const uint8_t* Bs, float* ws);
extern "C" void launch_m16_varI_reduce(const float* ws, uint16_t* C);

void m16_launch(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
    torch::Tensor workspace, torch::Tensor C
) {
    launch_m16_varI_gemm(
        reinterpret_cast<const uint16_t*>(A.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
        reinterpret_cast<float*>(workspace.data_ptr()));
    launch_m16_varI_reduce(
        reinterpret_cast<const float*>(workspace.data_ptr()),
        reinterpret_cast<uint16_t*>(C.data_ptr()));
}

// ── M=32 externs ──
extern "C" void launch_m32_n2880(
    const uint16_t* A, const uint8_t* B_shuffle, const uint8_t* B_scale_sh,
    uint16_t* C);
extern "C" void launch_m32_n4096(
    const uint16_t* A, const uint8_t* B_shuffle, const uint8_t* B_scale_sh,
    uint16_t* C);

void m32_n2880_launch(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
    torch::Tensor C
) {
    launch_m32_n2880(
        reinterpret_cast<const uint16_t*>(A.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
        reinterpret_cast<uint16_t*>(C.data_ptr()));
}

void m32_n4096_launch(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
    torch::Tensor C
) {
    launch_m32_n4096(
        reinterpret_cast<const uint16_t*>(A.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
        reinterpret_cast<uint16_t*>(C.data_ptr()));
}

// ── M=64 extern ──
extern "C" void launch_e641v1_exp_64_7168_2048(
    const uint16_t* A, const uint8_t* B, const uint8_t* Bs, uint16_t* C);

void m64_launch(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, torch::Tensor C
) {
    launch_e641v1_exp_64_7168_2048(
        reinterpret_cast<const uint16_t*>(A.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
        reinterpret_cast<uint16_t*>(C.data_ptr()));
}

// ── M=256 extern ──
extern "C" void launch_mk3g_256_3072_1536(
    const uint16_t* A, const uint8_t* B, const uint8_t* Bs, uint16_t* C);

void m256_launch(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, torch::Tensor C
) {
    launch_mk3g_256_3072_1536(
        reinterpret_cast<const uint16_t*>(A.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
        reinterpret_cast<uint16_t*>(C.data_ptr()));
}
"""

# =========================================================================
# HIP source — all kernel sources concatenated
# =========================================================================
_HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <cstdint>

// =====================================================================
// KERNEL 1: M=4 (m4_mk2a — 2-wave, shuffled B)
// =====================================================================

// m4_optim_mk2a: 2-wave blocks, AGPRs, packed absmax, XCD swizzle, all loads upfront
// Shape: M=4, N=2880, K=512, TILE_N=32 (2 MFMA columns), 16x16x128 MFMA
// 90 blocks × 128 threads (2 waves) = 360 waves total (same occupancy as mk1f).
// Each wave does 2 K-iters × 2 MFMA columns = 4 MFMAs. LDS reduce 2 partials.
// Barrier syncs 2 waves instead of 4. Half the LDS traffic.

namespace m4_mk2a {

using v4i32  = int   __attribute__((ext_vector_type(4)));
using v4f32  = float __attribute__((ext_vector_type(4)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));

static constexpr int kM = 4;
static constexpr int kN = 2880;
static constexpr int kK = 512;
static constexpr int kKHalf = kK / 2;

static constexpr int kTileN = 32;

static constexpr int kNumWaves = 2;
static constexpr int kBlockSize = kNumWaves * 64;  // 128
static constexpr int kNumNTiles = kN / kTileN;     // 90 (exact, no remainder)

static constexpr int kBPanelStride = kKHalf * 16;
static constexpr int kPaddedGroups = kK / 32;

// Each wave does 2 K-iters. Wave 0: ki=0,1. Wave 1: ki=2,3.
static constexpr int kKItersPerWave = 2;

// LDS: [2 waves][kM][kTileN] = 2*4*32 = 256 floats = 1024 bytes
static constexpr int kLdsFloats = kNumWaves * kM * kTileN;
static constexpr int kLdsBytes = kLdsFloats * sizeof(float);

// XCD swizzle
static constexpr int kNXCD = 8;
static constexpr int kC = 4;
static constexpr int kBlocksPerCycle = kNXCD * kC;
static constexpr int kLimit = (kNumNTiles / kBlocksPerCycle) * kBlocksPerCycle;

__device__ __forceinline__ uint32_t bit_cast_u32(float v) {
    union { float f; uint32_t u; } x; x.f = v; return x.u;
}
__device__ __forceinline__ float bit_cast_f32(uint32_t v) {
    union { uint32_t u; float f; } x; x.u = v; return x.f;
}
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) {
    uint32_t bits = bit_cast_u32(v);
    bits += ((bits >> 16) & 1u) + 0x7FFFu;
    return static_cast<uint16_t>(bits >> 16);
}
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t max_abs) {
    const uint16_t r = static_cast<uint16_t>(max_abs + 0x20u);
    const int e = static_cast<int>((r >> 7) & 0xFFu);
    return (e <= 2) ? 0u : (e >= 255) ? 254u : static_cast<uint8_t>(e - 2);
}

__device__ __forceinline__ int scale_offset(int row, int group) {
    return (row >> 5) * (32 * kPaddedGroups)
         + (group >> 3) * 256
         + (group & 3) * 64
         + (row & 15) * 4
         + ((group >> 2) & 1) * 2
         + ((row >> 4) & 1);
}

// MFMA with AGPR accumulator
__device__ __forceinline__ v4f32 mfma_fp4_16x16(
    v4i32 a, v4i32 b, v4f32 acc, int a_scale, int b_scale) {
    asm volatile(
        "v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
        : "+v"(acc) : "v"(a), "v"(b), "v"(a_scale), "v"(b_scale));
    return acc;
}

__global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))
void kernel_m4_mk2a(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    uint16_t*       __restrict__ C)
{
    // XCD swizzle
    int bid = blockIdx.x;
    if (bid < kLimit) {
        const int xcd    = bid % kNXCD;
        const int local_ = bid / kNXCD;
        const int chunk  = local_ / kC;
        const int pos    = local_ % kC;
        bid = chunk * kBlocksPerCycle + xcd * kC + pos;
    }

    const int col_base = bid * kTileN;
    const int tid = threadIdx.x;
    const int wave_id = tid >> 6;      // 0 or 1
    const int lane = tid & 63;
    const int row = lane & 15;
    const int k_quarter = lane >> 4;   // 0..3

    extern __shared__ char lds_raw[];
    float* lds = reinterpret_cast<float*>(lds_raw);

    // Wave 0 does ki=0,1. Wave 1 does ki=2,3.
    const int ki_base = wave_id * kKItersPerWave;  // 0 or 2

    // ============================================================
    // PHASE 1: Issue ALL loads upfront
    // ============================================================

    // A loads for both K-iters (row%kM wraps for threads with row>=4)
    const uint32_t* a_row = reinterpret_cast<const uint32_t*>(
        A + (row % kM) * kK);

    const int k_group0 = ki_base * 4 + k_quarter;
    const int k_group1 = (ki_base + 1) * 4 + k_quarter;

    asm volatile("" ::: "memory");
    const uint4* a_src0 = reinterpret_cast<const uint4*>(a_row + k_group0 * 16);
    uint4 a0_c0 = a_src0[0], a0_c1 = a_src0[1], a0_c2 = a_src0[2], a0_c3 = a_src0[3];

    const uint4* a_src1 = reinterpret_cast<const uint4*>(a_row + k_group1 * 16);
    uint4 a1_c0 = a_src1[0], a1_c1 = a_src1[1], a1_c2 = a_src1[2], a1_c3 = a_src1[3];

    // B col 0 loads for both K-iters
    const int b_col0 = col_base + row;
    const uint8_t* b_base0 = B_shuffle
        + (b_col0 >> 4) * kBPanelStride + (b_col0 & 15) * 16;

    uint4 b0_ki0 = reinterpret_cast<const uint4*>(b_base0 + ki_base * 1024 + k_quarter * 256)[0];
    uint4 b0_ki1 = reinterpret_cast<const uint4*>(b_base0 + (ki_base+1) * 1024 + k_quarter * 256)[0];

    // B col 1 loads for both K-iters
    const int b_col1 = col_base + 16 + row;
    const uint8_t* b_base1 = B_shuffle
        + (b_col1 >> 4) * kBPanelStride + (b_col1 & 15) * 16;

    uint4 b1_ki0 = reinterpret_cast<const uint4*>(b_base1 + ki_base * 1024 + k_quarter * 256)[0];
    uint4 b1_ki1 = reinterpret_cast<const uint4*>(b_base1 + (ki_base+1) * 1024 + k_quarter * 256)[0];

    // B scale loads (4 total: 2 columns × 2 K-iters)
    int bsc0_ki0 = static_cast<int>(B_scale_sh[scale_offset(b_col0, k_group0)]);
    int bsc0_ki1 = static_cast<int>(B_scale_sh[scale_offset(b_col0, k_group1)]);
    int bsc1_ki0 = static_cast<int>(B_scale_sh[scale_offset(b_col1, k_group0)]);
    int bsc1_ki1 = static_cast<int>(B_scale_sh[scale_offset(b_col1, k_group1)]);

    asm volatile("" : "+v"(bsc0_ki0), "+v"(bsc0_ki1), "+v"(bsc1_ki0), "+v"(bsc1_ki1) :: "memory");

    // ============================================================
    // PHASE 2: Quantize A for both K-iters using packed u16 absmax
    // ============================================================

    // --- K-iter 0 ---
    // Packed absmax using v_pk_max_u16 pattern
    uint32_t mx0;
    {
        // Mask off sign bits to get abs bf16 values, packed as u16 pairs
        const uint32_t mask = 0x7FFF7FFFu;
        uint32_t p0 = a0_c0.x & mask, p1 = a0_c0.y & mask, p2 = a0_c0.z & mask, p3 = a0_c0.w & mask;
        uint32_t p4 = a0_c1.x & mask, p5 = a0_c1.y & mask, p6 = a0_c1.z & mask, p7 = a0_c1.w & mask;
        uint32_t p8 = a0_c2.x & mask, p9 = a0_c2.y & mask, pA = a0_c2.z & mask, pB = a0_c2.w & mask;
        uint32_t pC = a0_c3.x & mask, pD = a0_c3.y & mask, pE = a0_c3.z & mask, pF = a0_c3.w & mask;
        // Reduce with v_pk_max_u16 (compiler should emit this for packed u16 max)
        uint32_t m01, m23, m45, m67, m89, mAB, mCD, mEF;
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m01) : "v"(p0), "v"(p1));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m23) : "v"(p2), "v"(p3));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m45) : "v"(p4), "v"(p5));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m67) : "v"(p6), "v"(p7));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m89) : "v"(p8), "v"(p9));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(mAB) : "v"(pA), "v"(pB));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(mCD) : "v"(pC), "v"(pD));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(mEF) : "v"(pE), "v"(pF));
        uint32_t t0, t1, t2, t3;
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(t0) : "v"(m01), "v"(m23));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(t1) : "v"(m45), "v"(m67));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(t2) : "v"(m89), "v"(mAB));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(t3) : "v"(mCD), "v"(mEF));
        uint32_t u0, u1;
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(u0) : "v"(t0), "v"(t1));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(u1) : "v"(t2), "v"(t3));
        uint32_t final_pk;
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(final_pk) : "v"(u0), "v"(u1));
        // Reduce the 2 u16 lanes to scalar
        mx0 = (final_pk >> 16) > (final_pk & 0xFFFFu) ? (final_pk >> 16) : (final_pk & 0xFFFFu);
    }

    uint8_t scale0 = compute_e8m0_scale(static_cast<uint16_t>(mx0));
    int a_sc0 = static_cast<int>(scale0);
    float fwd0 = (scale0 == 0u) ? bit_cast_f32(0x00400000u)
               : bit_cast_f32(static_cast<uint32_t>(scale0) << 23);

    const v2bf16* p0_0 = reinterpret_cast<const v2bf16*>(&a0_c0);
    const v2bf16* p0_1 = reinterpret_cast<const v2bf16*>(&a0_c1);
    const v2bf16* p0_2 = reinterpret_cast<const v2bf16*>(&a0_c2);
    const v2bf16* p0_3 = reinterpret_cast<const v2bf16*>(&a0_c3);
    uint32_t q0[4] = {0,0,0,0};
    q0[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[0],p0_0[0],fwd0,0);
    q0[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[1],p0_1[0],fwd0,0);
    q0[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[2],p0_2[0],fwd0,0);
    q0[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[3],p0_3[0],fwd0,0);
    q0[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[0],p0_0[1],fwd0,1);
    q0[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[1],p0_1[1],fwd0,1);
    q0[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[2],p0_2[1],fwd0,1);
    q0[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[3],p0_3[1],fwd0,1);
    q0[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[0],p0_0[2],fwd0,2);
    q0[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[1],p0_1[2],fwd0,2);
    q0[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[2],p0_2[2],fwd0,2);
    q0[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[3],p0_3[2],fwd0,2);
    q0[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[0],p0_0[3],fwd0,3);
    q0[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[1],p0_1[3],fwd0,3);
    q0[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[2],p0_2[3],fwd0,3);
    q0[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[3],p0_3[3],fwd0,3);

    // --- K-iter 1 ---
    uint32_t mx1;
    {
        const uint32_t mask = 0x7FFF7FFFu;
        uint32_t p0 = a1_c0.x & mask, p1 = a1_c0.y & mask, p2 = a1_c0.z & mask, p3 = a1_c0.w & mask;
        uint32_t p4 = a1_c1.x & mask, p5 = a1_c1.y & mask, p6 = a1_c1.z & mask, p7 = a1_c1.w & mask;
        uint32_t p8 = a1_c2.x & mask, p9 = a1_c2.y & mask, pA = a1_c2.z & mask, pB = a1_c2.w & mask;
        uint32_t pC = a1_c3.x & mask, pD = a1_c3.y & mask, pE = a1_c3.z & mask, pF = a1_c3.w & mask;
        uint32_t m01, m23, m45, m67, m89, mAB, mCD, mEF;
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m01) : "v"(p0), "v"(p1));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m23) : "v"(p2), "v"(p3));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m45) : "v"(p4), "v"(p5));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m67) : "v"(p6), "v"(p7));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(m89) : "v"(p8), "v"(p9));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(mAB) : "v"(pA), "v"(pB));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(mCD) : "v"(pC), "v"(pD));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(mEF) : "v"(pE), "v"(pF));
        uint32_t t0, t1, t2, t3;
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(t0) : "v"(m01), "v"(m23));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(t1) : "v"(m45), "v"(m67));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(t2) : "v"(m89), "v"(mAB));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(t3) : "v"(mCD), "v"(mEF));
        uint32_t u0, u1;
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(u0) : "v"(t0), "v"(t1));
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(u1) : "v"(t2), "v"(t3));
        uint32_t final_pk;
        asm("v_pk_max_u16 %0, %1, %2" : "=v"(final_pk) : "v"(u0), "v"(u1));
        mx1 = (final_pk >> 16) > (final_pk & 0xFFFFu) ? (final_pk >> 16) : (final_pk & 0xFFFFu);
    }

    uint8_t scale1 = compute_e8m0_scale(static_cast<uint16_t>(mx1));
    int a_sc1 = static_cast<int>(scale1);
    float fwd1 = (scale1 == 0u) ? bit_cast_f32(0x00400000u)
               : bit_cast_f32(static_cast<uint32_t>(scale1) << 23);

    const v2bf16* p1_0 = reinterpret_cast<const v2bf16*>(&a1_c0);
    const v2bf16* p1_1 = reinterpret_cast<const v2bf16*>(&a1_c1);
    const v2bf16* p1_2 = reinterpret_cast<const v2bf16*>(&a1_c2);
    const v2bf16* p1_3 = reinterpret_cast<const v2bf16*>(&a1_c3);
    uint32_t q1[4] = {0,0,0,0};
    q1[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[0],p1_0[0],fwd1,0);
    q1[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[1],p1_1[0],fwd1,0);
    q1[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[2],p1_2[0],fwd1,0);
    q1[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[3],p1_3[0],fwd1,0);
    q1[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[0],p1_0[1],fwd1,1);
    q1[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[1],p1_1[1],fwd1,1);
    q1[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[2],p1_2[1],fwd1,1);
    q1[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[3],p1_3[1],fwd1,1);
    q1[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[0],p1_0[2],fwd1,2);
    q1[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[1],p1_1[2],fwd1,2);
    q1[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[2],p1_2[2],fwd1,2);
    q1[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[3],p1_3[2],fwd1,2);
    q1[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[0],p1_0[3],fwd1,3);
    q1[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[1],p1_1[3],fwd1,3);
    q1[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[2],p1_2[3],fwd1,3);
    q1[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[3],p1_3[3],fwd1,3);

    // Zero out A for threads with row >= kM
    v4i32 a_r0, a_r1;
    if (row < kM) {
        a_r0 = {(int)q0[0], (int)q0[1], (int)q0[2], (int)q0[3]};
        a_r1 = {(int)q1[0], (int)q1[1], (int)q1[2], (int)q1[3]};
    } else {
        a_r0 = {0,0,0,0};
        a_r1 = {0,0,0,0};
        a_sc0 = 127;
        a_sc1 = 127;
    }

    // ============================================================
    // PHASE 3: 4 MFMAs — 2 K-iters × 2 N-columns, accumulate in AGPRs
    // ============================================================

    v4i32 b0_r0 = {(int)b0_ki0.x, (int)b0_ki0.y, (int)b0_ki0.z, (int)b0_ki0.w};
    v4i32 b0_r1 = {(int)b0_ki1.x, (int)b0_ki1.y, (int)b0_ki1.z, (int)b0_ki1.w};
    v4i32 b1_r0 = {(int)b1_ki0.x, (int)b1_ki0.y, (int)b1_ki0.z, (int)b1_ki0.w};
    v4i32 b1_r1 = {(int)b1_ki1.x, (int)b1_ki1.y, (int)b1_ki1.z, (int)b1_ki1.w};

    v4f32 acc_c0 = {0.0f, 0.0f, 0.0f, 0.0f};
    v4f32 acc_c1 = {0.0f, 0.0f, 0.0f, 0.0f};

    // Grouped by accumulator — same-acc MFMAs back-to-back for 0-wait fast path
    // Col 0: K-iter 0 then K-iter 1
    acc_c0 = mfma_fp4_16x16(a_r0, b0_r0, acc_c0, a_sc0, bsc0_ki0);
    acc_c0 = mfma_fp4_16x16(a_r1, b0_r1, acc_c0, a_sc1, bsc0_ki1);
    // Col 1: K-iter 0 then K-iter 1
    acc_c1 = mfma_fp4_16x16(a_r0, b1_r0, acc_c1, a_sc0, bsc1_ki0);
    acc_c1 = mfma_fp4_16x16(a_r1, b1_r1, acc_c1, a_sc1, bsc1_ki1);

    // ============================================================
    // PHASE 4: LDS write — only rows < kM
    // ============================================================

    const int out_col_16 = lane & 15;
    const int row_base4 = (lane >> 4) * 4;
    const int lds_wave_base = wave_id * kM * kTileN;

    #pragma unroll
    for (int r = 0; r < 4; ++r) {
        const int out_row = row_base4 + r;
        if (out_row < kM) {
            lds[lds_wave_base + out_row * kTileN + out_col_16] = acc_c0[r];
            lds[lds_wave_base + out_row * kTileN + 16 + out_col_16] = acc_c1[r];
        }
    }
    __syncthreads();

    // ============================================================
    // PHASE 5: Reduce 2 partials and write output
    // 128 threads handle kM × kTileN = 4 × 32 = 128 elements (perfect 1:1)
    // ============================================================

    {
        const int local_row = tid / kTileN;
        const int local_col = tid % kTileN;
        if (local_row < kM) {
            const float v0 = lds[0 * kM * kTileN + local_row * kTileN + local_col];
            const float v1 = lds[1 * kM * kTileN + local_row * kTileN + local_col];
            const float sum = v0 + v1;
            const int out_col = col_base + local_col;
            if (out_col < kN) {
                C[local_row * kN + out_col] = float_to_bf16_rn(sum);
            }
        }
    }
}

}  // namespace m4_mk2a

extern "C" void launch_m4_mk2a(
    const uint16_t* A, const uint8_t* B_shuffle, const uint8_t* B_scale_sh,
    uint16_t* C) {
    hipLaunchKernelGGL(m4_mk2a::kernel_m4_mk2a,
        dim3(m4_mk2a::kNumNTiles), dim3(m4_mk2a::kBlockSize),
        m4_mk2a::kLdsBytes, 0,
        A, B_shuffle, B_scale_sh, C);
}

// =====================================================================
// KERNEL 2: M=16 (m16_varI — two-pass workspace + reduce)
// =====================================================================

// m16_varI: Two-pass — GEMM writes to workspace, reduce kernel sums partials
// No atomics, no torch.zeros. Each block stores directly to workspace[split_id].
// M=16, N=2112, K=7168, TILE_N=64, SK=7, 4 waves
// Workspace: [7][16][2112] f32 = 944 KB (torch.empty)
// Pass 1: 231 blocks compute and store partials
// Pass 2: small reduce kernel sums 7 partials → bf16 output

namespace m16_varI {

using v4i32  = int   __attribute__((ext_vector_type(4)));
using v4f32  = float __attribute__((ext_vector_type(4)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));

static constexpr int kM = 16;
static constexpr int kN = 2112;
static constexpr int kK = 7168;
static constexpr int kKHalf = kK / 2;
static constexpr int kPaddedGroups = 224;

static constexpr int kTileN = 64;
static constexpr int kMfmaCols = 4;
static constexpr int kSplitK = 7;
static constexpr int kKPerSplit = kK / kSplitK;  // 1024
static constexpr int kNumWaves = 4;
static constexpr int kBlockSize = kNumWaves * 64;
static constexpr int kItersPerWave = kKPerSplit / (kNumWaves * 128);  // 2

static constexpr int kNumNTiles = kN / kTileN;  // 33
static constexpr int kTotalBlocks = kNumNTiles * kSplitK;  // 231
static constexpr int kTileElements = kM * kTileN;

static constexpr int kBPanelStride = kKHalf * 16;

// LDS for 4-wave intra-block reduce: [col][row][wave][16]
static constexpr int kLdsFloatsPerCol = kM * kNumWaves * 16;
static constexpr int kLdsFloats = kMfmaCols * kLdsFloatsPerCol;
static constexpr int kLdsBytes = kLdsFloats * sizeof(float);

static constexpr int kNXCD = 8;
static constexpr int kC = 4;
static constexpr int kBlocksPerCycle = kNXCD * kC;
static constexpr int kLimit = (kTotalBlocks / kBlocksPerCycle) * kBlocksPerCycle;

__device__ __forceinline__ uint32_t bit_cast_u32(float v) { union{float f;uint32_t u;}x; x.f=v; return x.u; }
__device__ __forceinline__ float bit_cast_f32(uint32_t v) { union{uint32_t u;float f;}x; x.u=v; return x.f; }
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) { uint32_t b=bit_cast_u32(v); b+=((b>>16)&1u)+0x7FFFu; return (uint16_t)(b>>16); }
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t m) { uint16_t r=(uint16_t)(m+0x20u); int e=(int)((r>>7)&0xFFu); return (e<=2)?0u:(e>=255)?254u:(uint8_t)(e-2); }

__device__ __forceinline__ int scale_offset(int row, int group) {
    return (row>>5)*(32*kPaddedGroups)+(group>>3)*256+(group&3)*64+(row&15)*4+((group>>2)&1)*2+((row>>4)&1);
}

__device__ __forceinline__ v4f32 mfma_fp4_16x16(v4i32 a, v4i32 b, v4f32 acc, int as, int bs) {
    asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
        :"+v"(acc):"v"(a),"v"(b),"v"(as),"v"(bs)); return acc;
}

__device__ __forceinline__ void quant_a(const uint32_t* dw, v4i32& out, int& sc) {
    const uint32_t mask=0x7FFF7FFFu;
    uint32_t p0=dw[0]&mask,p1=dw[1]&mask,p2=dw[2]&mask,p3=dw[3]&mask;
    uint32_t p4=dw[4]&mask,p5=dw[5]&mask,p6=dw[6]&mask,p7=dw[7]&mask;
    uint32_t p8=dw[8]&mask,p9=dw[9]&mask,pA=dw[10]&mask,pB=dw[11]&mask;
    uint32_t pC=dw[12]&mask,pD=dw[13]&mask,pE=dw[14]&mask,pF=dw[15]&mask;
    uint32_t m01,m23,m45,m67,m89,mAB,mCD,mEF;
    asm("v_pk_max_u16 %0,%1,%2":"=v"(m01):"v"(p0),"v"(p1));asm("v_pk_max_u16 %0,%1,%2":"=v"(m23):"v"(p2),"v"(p3));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(m45):"v"(p4),"v"(p5));asm("v_pk_max_u16 %0,%1,%2":"=v"(m67):"v"(p6),"v"(p7));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(m89):"v"(p8),"v"(p9));asm("v_pk_max_u16 %0,%1,%2":"=v"(mAB):"v"(pA),"v"(pB));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(mCD):"v"(pC),"v"(pD));asm("v_pk_max_u16 %0,%1,%2":"=v"(mEF):"v"(pE),"v"(pF));
    uint32_t t0,t1,t2,t3;
    asm("v_pk_max_u16 %0,%1,%2":"=v"(t0):"v"(m01),"v"(m23));asm("v_pk_max_u16 %0,%1,%2":"=v"(t1):"v"(m45),"v"(m67));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(t2):"v"(m89),"v"(mAB));asm("v_pk_max_u16 %0,%1,%2":"=v"(t3):"v"(mCD),"v"(mEF));
    uint32_t u0,u1;
    asm("v_pk_max_u16 %0,%1,%2":"=v"(u0):"v"(t0),"v"(t1));asm("v_pk_max_u16 %0,%1,%2":"=v"(u1):"v"(t2),"v"(t3));
    uint32_t fp; asm("v_pk_max_u16 %0,%1,%2":"=v"(fp):"v"(u0),"v"(u1));
    uint32_t mx=(fp>>16)>(fp&0xFFFFu)?(fp>>16):(fp&0xFFFFu);
    uint8_t s=compute_e8m0_scale((uint16_t)mx); sc=(int)s;
    float fwd=(s==0u)?bit_cast_f32(0x00400000u):bit_cast_f32((uint32_t)s<<23);
    const v2bf16*q=reinterpret_cast<const v2bf16*>(dw); uint32_t k0=0,k1=0,k2=0,k3=0;
    k0=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k0,q[0],fwd,0);k1=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k1,q[4],fwd,0);
    k2=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k2,q[8],fwd,0);k3=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k3,q[12],fwd,0);
    k0=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k0,q[1],fwd,1);k1=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k1,q[5],fwd,1);
    k2=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k2,q[9],fwd,1);k3=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k3,q[13],fwd,1);
    k0=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k0,q[2],fwd,2);k1=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k1,q[6],fwd,2);
    k2=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k2,q[10],fwd,2);k3=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k3,q[14],fwd,2);
    k0=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k0,q[3],fwd,3);k1=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k1,q[7],fwd,3);
    k2=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k2,q[11],fwd,3);k3=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k3,q[15],fwd,3);
    out={(int)k0,(int)k1,(int)k2,(int)k3};
}

// ============================================================
// PASS 1: GEMM kernel — stores partials to workspace
// ============================================================

__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void kernel_m16_gemm(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    float*          __restrict__ workspace)  // [kSplitK][kM][kN]
{
    int xy = blockIdx.x;
    if (xy < kLimit) {
        const int xcd=xy%kNXCD, local_=xy/kNXCD, chunk=local_/kC, pos=local_%kC;
        xy = chunk*kBlocksPerCycle + xcd*kC + pos;
    }

    const int k_split = xy / kNumNTiles;
    const int n_tile  = xy % kNumNTiles;
    const int col_base = n_tile * kTileN;

    const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&63;
    const int row=lane&15, k_quarter=lane>>4;

    extern __shared__ char lds_raw[];
    float* lds = reinterpret_cast<float*>(lds_raw);

    const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + row * kK);

    const uint8_t* b_bases[kMfmaCols];
    #pragma unroll
    for (int c=0;c<kMfmaCols;++c) {
        const int b_col=col_base+c*16+row;
        b_bases[c]=B_shuffle+(b_col>>4)*kBPanelStride+(b_col&15)*16;
    }

    const int split_iter_base = k_split * (kKPerSplit / 128);
    const int ki_base = split_iter_base + wave_id * kItersPerWave;
    const int kg0 = ki_base*4+k_quarter;
    const int kg1 = (ki_base+1)*4+k_quarter;

    // All loads upfront
    asm volatile("" ::: "memory");
    uint32_t a_buf0[16], a_buf1[16];
    { const uint32_t*s=a_row+kg0*16;
      #pragma unroll
      for(int i=0;i<16;++i) a_buf0[i]=s[i]; }
    { const uint32_t*s=a_row+kg1*16;
      #pragma unroll
      for(int i=0;i<16;++i) a_buf1[i]=s[i]; }

    v4i32 b0[kMfmaCols],b1[kMfmaCols]; int bs0[kMfmaCols],bs1[kMfmaCols];
    #pragma unroll
    for(int c=0;c<kMfmaCols;++c) {
        const uint32_t*s0=reinterpret_cast<const uint32_t*>(b_bases[c]+ki_base*1024+k_quarter*256);
        b0[c]={(int)s0[0],(int)s0[1],(int)s0[2],(int)s0[3]};
        const uint32_t*s1=reinterpret_cast<const uint32_t*>(b_bases[c]+(ki_base+1)*1024+k_quarter*256);
        b1[c]={(int)s1[0],(int)s1[1],(int)s1[2],(int)s1[3]};
        const int b_col=col_base+c*16+row;
        bs0[c]=(int)B_scale_sh[scale_offset(b_col, kg0)];
        bs1[c]=(int)B_scale_sh[scale_offset(b_col, kg1)];
    }

    asm volatile("":"+v"(bs0[0]),"+v"(bs0[1]),"+v"(bs0[2]),"+v"(bs0[3]),
                     "+v"(bs1[0]),"+v"(bs1[1]),"+v"(bs1[2]),"+v"(bs1[3])::"memory");

    // Quant + MFMA
    v4f32 acc[kMfmaCols];
    #pragma unroll
    for(int c=0;c<kMfmaCols;++c) acc[c]={0,0,0,0};

    { v4i32 ar; int as; quant_a(a_buf0,ar,as);
      #pragma unroll
      for(int c=0;c<kMfmaCols;++c) acc[c]=mfma_fp4_16x16(ar,b0[c],acc[c],as,bs0[c]); }
    { v4i32 ar; int as; quant_a(a_buf1,ar,as);
      #pragma unroll
      for(int c=0;c<kMfmaCols;++c) acc[c]=mfma_fp4_16x16(ar,b1[c],acc[c],as,bs1[c]); }

    // LDS reduce 4 waves
    const int oc16=lane&15, rb4=(lane>>4)*4;
    #pragma unroll
    for(int c=0;c<kMfmaCols;++c)
        #pragma unroll
        for(int r=0;r<4;++r)
            lds[c*kLdsFloatsPerCol+(rb4+r)*kNumWaves*16+wave_id*16+oc16]=acc[c][r];
    __syncthreads();

    // Reduce and STORE to workspace (no atomics!)
    // workspace layout: [split][row][col] — split stride = kM * kN
    const int ws_base = k_split * kM * kN;
    {
        const int c = wave_id;
        #pragma unroll
        for(int r=0;r<4;++r) {
            const int out_row=rb4+r;
            if(out_row<kM) {
                const int lb=c*kLdsFloatsPerCol+out_row*kNumWaves*16+oc16;
                float sum=0.0f;
                #pragma unroll
                for(int w=0;w<kNumWaves;++w) sum+=lds[lb+w*16];
                const int gc=col_base+c*16+oc16;
                workspace[ws_base + out_row*kN + gc] = sum;
            }
        }
    }
}

// ============================================================
// PASS 2: Reduce kernel — sum 7 partials, convert to bf16
// ============================================================

static constexpr int kReduceBlock = 256;
static constexpr int kReduceGrid = (kM * kN + kReduceBlock - 1) / kReduceBlock;  // 132

__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void kernel_m16_reduce(
    const float*    __restrict__ workspace,  // [7][16][2112]
    uint16_t*       __restrict__ C)
{
    const int idx = blockIdx.x * kReduceBlock + threadIdx.x;
    if (idx >= kM * kN) return;

    float sum = 0.0f;
    #pragma unroll
    for (int s = 0; s < kSplitK; ++s)
        sum += workspace[s * kM * kN + idx];

    C[idx] = float_to_bf16_rn(sum);
}

}  // namespace m16_varI

extern "C" void launch_m16_varI_gemm(
    const uint16_t* A, const uint8_t* B, const uint8_t* Bs, float* ws) {
    hipLaunchKernelGGL(m16_varI::kernel_m16_gemm,
        dim3(m16_varI::kTotalBlocks), dim3(m16_varI::kBlockSize),
        m16_varI::kLdsBytes, 0, A, B, Bs, ws);
}

extern "C" void launch_m16_varI_reduce(const float* ws, uint16_t* C) {
    hipLaunchKernelGGL(m16_varI::kernel_m16_reduce,
        dim3(m16_varI::kReduceGrid), dim3(m16_varI::kReduceBlock),
        0, 0, ws, C);
}

// =====================================================================
// KERNEL 3: M=32 (m32_mk4l — 16x16x128 MFMA, split-M, both N values)
// =====================================================================

namespace m32_mk4l {

using v4i32 = int   __attribute__((ext_vector_type(4)));
using v4f32 = float __attribute__((ext_vector_type(4)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));

__device__ __forceinline__ uint32_t bit_cast_u32(float v) { union { float f; uint32_t u; } x; x.f=v; return x.u; }
__device__ __forceinline__ float bit_cast_f32(uint32_t v) { union { uint32_t u; float f; } x; x.u=v; return x.f; }
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) { uint32_t b=bit_cast_u32(v); b+=((b>>16)&1u)+0x7FFFu; return (uint16_t)(b>>16); }
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t m) { uint16_t r=(uint16_t)(m+0x20u); int e=(int)((r>>7)&0xFFu); return (e<=2)?0u:(e>=255)?254u:(uint8_t)(e-2); }

static constexpr int kM=32, kK=512, kKHalf=kK/2, kPaddedGroups=kK/32;
static constexpr int kWaveSize=64, kNumWaves=4, kBlockSize=kWaveSize*kNumWaves;
static constexpr int kMfmaN=16, kMfmaK=128, kTileN=32, kMfmaCols=kTileN/kMfmaN;
static constexpr int kTileM=16, kNumMTiles=kM/kTileM;
static constexpr int kBPanelStride=kKHalf*16;

__device__ __forceinline__ int scale_offset(int row, int group) {
    return (row>>5)*(32*kPaddedGroups)+(group>>3)*256+(group&3)*64+(row&15)*4+((group>>2)&1)*2+((row>>4)&1);
}
__device__ __forceinline__ v4f32 mfma_fp4_16x16(v4i32 a, v4i32 b, v4f32 acc, int as, int bs) {
    asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4":"+a"(acc):"v"(a),"v"(b),"v"(as),"v"(bs)); return acc;
}
__device__ __forceinline__ void quant_from_raw(const uint32_t* dw, v4i32& a_out, int& scale_out) {
    uint32_t bmax=0u;
    #pragma unroll
    for(int i=0;i<16;++i){uint32_t w=dw[i];uint32_t hi=(w>>16)&0x7FFFu,lo=w&0x7FFFu;bmax=(hi>bmax)?hi:bmax;bmax=(lo>bmax)?lo:bmax;}
    uint8_t scale=compute_e8m0_scale((uint16_t)bmax); scale_out=(int)scale;
    float fwd=(scale==0u)?bit_cast_f32(0x00400000u):bit_cast_f32((uint32_t)scale<<23);
    const v2bf16* p=reinterpret_cast<const v2bf16*>(dw); uint32_t pk[4]={0,0,0,0};
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],p[0],fwd,0);pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],p[4],fwd,0);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],p[8],fwd,0);pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],p[12],fwd,0);
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],p[1],fwd,1);pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],p[5],fwd,1);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],p[9],fwd,1);pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],p[13],fwd,1);
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],p[2],fwd,2);pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],p[6],fwd,2);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],p[10],fwd,2);pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],p[14],fwd,2);
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],p[3],fwd,3);pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],p[7],fwd,3);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],p[11],fwd,3);pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],p[15],fwd,3);
    a_out={int(pk[0]),int(pk[1]),int(pk[2]),int(pk[3])};
}

template <int kN>
__global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void fused_gemm_m32(const uint16_t* __restrict__ A, const uint8_t* __restrict__ B_shuffle,
    const uint8_t* __restrict__ B_scale_sh, uint16_t* __restrict__ C)
{
    constexpr int kNumNTiles = kN / kTileN;
    constexpr int kLdsFloats = kTileM * kNumWaves * kTileN;

    const int n_tile=blockIdx.x/kNumMTiles, m_tile=blockIdx.x%kNumMTiles;
    const int col_base=n_tile*kTileN, row_base=m_tile*kTileM;
    const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&63;
    const int row16=lane&15, k_quarter=lane>>4;
    extern __shared__ char lds_raw[]; float* lds=reinterpret_cast<float*>(lds_raw);

    const int ki=wave_id, k_group=ki*4+k_quarter;
    const int a_row=row_base+row16;

    // B_scale loads first
    int bsc[kMfmaCols];
    #pragma unroll
    for(int c=0;c<kMfmaCols;++c){
        int b_col=col_base+c*kMfmaN+row16;
        bsc[c]=(int)B_scale_sh[scale_offset(b_col,k_group)];
    }
    asm volatile("":"+v"(bsc[0]),"+v"(bsc[1])::"memory");

    // B data loads
    v4i32 b_regs[kMfmaCols];
    #pragma unroll
    for(int c=0;c<kMfmaCols;++c){
        int b_col=col_base+c*kMfmaN+row16;
        const uint8_t* bb=B_shuffle+(b_col>>4)*kBPanelStride+(b_col&15)*16;
        const uint32_t* bs=reinterpret_cast<const uint32_t*>(bb+ki*1024+k_quarter*256);
        b_regs[c]={int(bs[0]),int(bs[1]),int(bs[2]),int(bs[3])};
    }

    // A load
    uint32_t a_buf[16];
    {const uint32_t* as=reinterpret_cast<const uint32_t*>(A+a_row*kK+k_group*32);
     #pragma unroll
     for(int i=0;i<16;++i) a_buf[i]=as[i];}

    v4f32 acc_c0,acc_c1;
    #pragma unroll
    for(int i=0;i<4;++i){acc_c0[i]=0.f;acc_c1[i]=0.f;}
    {v4i32 ar;int as;quant_from_raw(a_buf,ar,as);
     acc_c0=mfma_fp4_16x16(ar,b_regs[0],acc_c0,as,bsc[0]);
     acc_c1=mfma_fp4_16x16(ar,b_regs[1],acc_c1,as,bsc[1]);}

    // LDS write
    const int out_col_16=lane&15, row_base4=(lane>>4)*4;
    #pragma unroll
    for(int r=0;r<4;++r){
        int out_row=row_base4+r;
        lds[out_row*(kNumWaves*kTileN)+wave_id*kTileN+out_col_16]=acc_c0[r];
        lds[out_row*(kNumWaves*kTileN)+wave_id*kTileN+out_col_16+16]=acc_c1[r];
    }
    __syncthreads();

    // Clean epilogue
    const int e0=tid, e1=tid+kBlockSize;
    const int lr0=e0/kTileN, lc0=e0%kTileN;
    const int lr1=e1/kTileN, lc1=e1%kTileN;
    const int lb0=lr0*(kNumWaves*kTileN)+lc0;
    const int lb1=lr1*(kNumWaves*kTileN)+lc1;

    float v0w0=lds[lb0],v0w1=lds[lb0+kTileN],v0w2=lds[lb0+2*kTileN],v0w3=lds[lb0+3*kTileN];
    float v1w0=lds[lb1],v1w1=lds[lb1+kTileN],v1w2=lds[lb1+2*kTileN],v1w3=lds[lb1+3*kTileN];

    float s0=v0w0+v0w1+v0w2+v0w3;
    float s1=v1w0+v1w1+v1w2+v1w3;
    C[(row_base+lr0)*kN+col_base+lc0]=float_to_bf16_rn(s0);
    C[(row_base+lr1)*kN+col_base+lc1]=float_to_bf16_rn(s1);
}

static constexpr int kLdsBytes = kTileM * kNumWaves * kTileN * (int)sizeof(float);

}  // namespace m32_mk4l

extern "C" void launch_m32_n2880(const uint16_t* A,const uint8_t* B,const uint8_t* Bs,uint16_t* C){
    constexpr int N=2880, grid=(N/m32_mk4l::kTileN)*m32_mk4l::kNumMTiles;
    hipLaunchKernelGGL((m32_mk4l::fused_gemm_m32<N>),dim3(grid),dim3(m32_mk4l::kBlockSize),m32_mk4l::kLdsBytes,0,A,B,Bs,C);
}
extern "C" void launch_m32_n4096(const uint16_t* A,const uint8_t* B,const uint8_t* Bs,uint16_t* C){
    constexpr int N=4096, grid=(N/m32_mk4l::kTileN)*m32_mk4l::kNumMTiles;
    hipLaunchKernelGGL((m32_mk4l::fused_gemm_m32<N>),dim3(grid),dim3(m32_mk4l::kBlockSize),m32_mk4l::kLdsBytes,0,A,B,Bs,C);
}


// =====================================================================
// KERNEL 4: M=64 (e641v1_exp S2 — hand-unrolled 8-iter, AGPR, B-scale LDS)
// =====================================================================

namespace e641v1_exp {

using v4i32  = int   __attribute__((ext_vector_type(4)));
using v16f32 = float __attribute__((ext_vector_type(16)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));

__device__ __forceinline__ uint32_t bit_cast_u32(float v) {
    union { float f; uint32_t u; } x; x.f = v; return x.u;
}
__device__ __forceinline__ float bit_cast_f32(uint32_t v) {
    union { uint32_t u; float f; } x; x.u = v; return x.f;
}
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) {
    uint32_t bits = bit_cast_u32(v);
    bits += ((bits >> 16) & 1u) + 0x7FFFu;
    return static_cast<uint16_t>(bits >> 16);
}
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t max_abs) {
    const uint16_t r = static_cast<uint16_t>(max_abs + 0x20u);
    const int e = static_cast<int>((r >> 7) & 0xFFu);
    return (e <= 2) ? 0u : (e >= 255) ? 254u : static_cast<uint8_t>(e - 2);
}

__device__ __forceinline__ v16f32 mfma_fp4x4(
    v4i32 a, v4i32 b, v16f32 acc, int a_scale, int b_scale) {
    asm volatile(
        "v_mfma_scale_f32_32x32x64_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
        : "+a"(acc) : "v"(a), "v"(b), "v"(a_scale), "v"(b_scale));
    return acc;
}

__device__ __forceinline__ void load_a_raw(
    const uint32_t* __restrict__ a_ptr, uint32_t* dw)
{
    *reinterpret_cast<uint4*>(&dw[0])  = *(reinterpret_cast<const uint4*>(a_ptr) + 0);
    *reinterpret_cast<uint4*>(&dw[4])  = *(reinterpret_cast<const uint4*>(a_ptr) + 1);
    *reinterpret_cast<uint4*>(&dw[8])  = *(reinterpret_cast<const uint4*>(a_ptr) + 2);
    *reinterpret_cast<uint4*>(&dw[12]) = *(reinterpret_cast<const uint4*>(a_ptr) + 3);
}

__device__ __forceinline__ void quant_from_raw(
    const uint32_t* dw, v4i32& a_out, int& scale_out)
{
    uint32_t bmax = 0u;
    #pragma unroll
    for (int i = 0; i < 16; ++i) {
        const uint32_t w = dw[i];
        const uint32_t hi = (w >> 16) & 0x7FFFu;
        const uint32_t lo = w & 0x7FFFu;
        bmax = (hi > bmax) ? hi : bmax;
        bmax = (lo > bmax) ? lo : bmax;
    }
    const uint8_t scale = compute_e8m0_scale(static_cast<uint16_t>(bmax));
    scale_out = static_cast<int>(scale);
    const float fwd = (scale == 0u) ? bit_cast_f32(0x00400000u)
                    : bit_cast_f32(static_cast<uint32_t>(scale) << 23);
    const v2bf16* pairs = reinterpret_cast<const v2bf16*>(dw);
    uint32_t pk[4] = {0,0,0,0};
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[0], fwd,0);
    pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[4], fwd,0);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[8], fwd,0);
    pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[12],fwd,0);
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[1], fwd,1);
    pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[5], fwd,1);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[9], fwd,1);
    pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[13],fwd,1);
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[2], fwd,2);
    pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[6], fwd,2);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[10],fwd,2);
    pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[14],fwd,2);
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[3], fwd,3);
    pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[7], fwd,3);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[11],fwd,3);
    pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[15],fwd,3);
    a_out = {(int)pk[0], (int)pk[1], (int)pk[2], (int)pk[3]};
}

// S2 B scale offset: PaddedGroups=64
__device__ __forceinline__ int scale_offset_s2(int row, int group) {
    return (row >> 5) * (32 * 64)
         + (group >> 3) * 256
         + (group & 3) * 64
         + (row & 15) * 4
         + ((group >> 2) & 1) * 2
         + ((row >> 4) & 1);
}

// S2 constants
static constexpr int S2_M = 64, S2_N = 7168, S2_K = 2048;
static constexpr int S2_KHalf = 1024;
static constexpr int S2_TileM = 32, S2_TileN = 64;
static constexpr int S2_NumWaves = 4, S2_BlockSize = 256;
static constexpr int S2_Iters = 8;
static constexpr int S2_NumTilesN = 112;
static constexpr int S2_NumTilesM = 2;
static constexpr int S2_TotalBlocks = 224;
static constexpr int S2_LdsBytes = 32768;
static constexpr int S2_NXCD = 8, S2_W = 2, S2_C = 14;
static constexpr int S2_BlocksPerCycle = 112;
static constexpr int S2_Limit = 224;
static constexpr int S2_TidPerGroup = 224;

#define S2_ITERATION(BUF_CUR, BUF_FAR, KI, LOAD_A_FAR, LOAD_B_NXT)          \
{                                                                              \
    const int k_group_ = (KI) * 2 + half;                                     \
                                                                               \
    int bsc0_ = static_cast<int>(lds_bscale[                                   \
        scale_offset_s2(b_col0, k_group_) - bscale_lds_base]);                 \
    int bsc1_ = static_cast<int>(lds_bscale[                                   \
        scale_offset_s2(b_col1, k_group_) - bscale_lds_base]);                 \
                                                                               \
    v4i32 areg_;                                                               \
    int   ascale_;                                                             \
    quant_from_raw(BUF_CUR, areg_, ascale_);                                   \
                                                                               \
    if constexpr (LOAD_B_NXT) {                                                \
        const int ki_nxt_ = (KI) + 1;                                         \
        const uint32_t* bnxt0_ = reinterpret_cast<const uint32_t*>(            \
            b_base0 + ki_nxt_ * 512 + half * 256);                             \
        b_nxt0 = {(int)bnxt0_[0],(int)bnxt0_[1],(int)bnxt0_[2],(int)bnxt0_[3]}; \
        const uint32_t* bnxt1_ = reinterpret_cast<const uint32_t*>(            \
            b_base1 + ki_nxt_ * 512 + half * 256);                             \
        b_nxt1 = {(int)bnxt1_[0],(int)bnxt1_[1],(int)bnxt1_[2],(int)bnxt1_[3]}; \
    }                                                                          \
                                                                               \
    acc0 = mfma_fp4x4(areg_, b_cur0, acc0, ascale_, bsc0_);                   \
    acc1 = mfma_fp4x4(areg_, b_cur1, acc1, ascale_, bsc1_);                   \
                                                                               \
    if constexpr (LOAD_A_FAR) {                                                \
        const int far_grp_ = ((KI) + 2) * 2 + half;                           \
        load_a_raw(a_row + far_grp_ * 16, BUF_FAR);                           \
    }                                                                          \
                                                                               \
    if constexpr (LOAD_B_NXT) {                                                \
        b_cur0 = b_nxt0;                                                       \
        b_cur1 = b_nxt1;                                                       \
    }                                                                          \
}

__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void fused_gemm_s2_handunrolled(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    uint16_t*       __restrict__ C)
{
    int xy = blockIdx.x;
    if (xy < S2_Limit) {
        const int xcd    = xy % S2_NXCD;
        const int local_ = xy / S2_NXCD;
        const int chunk  = local_ / S2_C;
        const int pos    = local_ % S2_C;
        xy = chunk * S2_BlocksPerCycle + xcd * S2_C + pos;
    }
    const int l      = xy % S2_TidPerGroup;
    const int m_tile = (xy / S2_TidPerGroup) * S2_W + (l % S2_W);
    const int n_tile = l / S2_W;

    const int col_base = n_tile * S2_TileN;
    const int row_base = m_tile * S2_TileM;

    const int tid     = threadIdx.x;
    const int wave_id = tid >> 6;
    const int lane    = tid & 31;
    const int half    = (tid >> 5) & 1;

    extern __shared__ char lds_bytes[];
    constexpr int lds_half = S2_NumWaves * 32 * S2_TileM;

    const int b_col0 = col_base + lane;
    const int b_col1 = col_base + 32 + lane;
    const uint8_t* b_base0 = B_shuffle
        + (b_col0 >> 4) * (S2_KHalf * 16) + (b_col0 & 15) * 16;
    const uint8_t* b_base1 = B_shuffle
        + (b_col1 >> 4) * (S2_KHalf * 16) + (b_col1 & 15) * 16;

    // Cooperative B scale preload into LDS
    {
        const int bscale_global_base_offset = (col_base >> 5) * (32 * 64);
        const uint8_t* bscale_src = B_scale_sh + bscale_global_base_offset;
        const int my_offset = tid * 16;
        #pragma unroll
        for (int bi = 0; bi < 16; ++bi) {
            reinterpret_cast<uint8_t*>(lds_bytes)[my_offset + bi] =
                bscale_src[my_offset + bi];
        }
    }
    __syncthreads();

    const uint8_t* lds_bscale = reinterpret_cast<const uint8_t*>(lds_bytes);
    const int bscale_lds_base = (col_base >> 5) * (32 * 64);

    const int my_row = row_base + lane;
    const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + my_row * S2_K);

    const int ki_base = wave_id * S2_Iters;

    v16f32 acc0, acc1;
    #pragma unroll
    for (int i = 0; i < 16; ++i) { acc0[i] = 0.0f; acc1[i] = 0.0f; }

    // Anti-thundering-herd stagger
    if (blockIdx.x & 4) asm volatile("s_sleep 1" :::);
    if (blockIdx.x & 8) asm volatile("s_sleep 1" :::);

    // Prologue: B data for iter 0
    v4i32 b_cur0, b_cur1;
    {
        const uint32_t* bs0 = reinterpret_cast<const uint32_t*>(
            b_base0 + ki_base * 512 + half * 256);
        b_cur0 = {(int)bs0[0], (int)bs0[1], (int)bs0[2], (int)bs0[3]};
        const uint32_t* bs1 = reinterpret_cast<const uint32_t*>(
            b_base1 + ki_base * 512 + half * 256);
        b_cur1 = {(int)bs1[0], (int)bs1[1], (int)bs1[2], (int)bs1[3]};
    }

    // Three STATIC A buffers
    uint32_t a_buf0[16], a_buf1[16], a_buf2[16];
    load_a_raw(a_row + (ki_base * 2 + half) * 16, a_buf0);
    load_a_raw(a_row + ((ki_base + 1) * 2 + half) * 16, a_buf1);

    v4i32 b_nxt0, b_nxt1;

    // 8 HAND-UNROLLED ITERATIONS
    S2_ITERATION(a_buf0, a_buf2, ki_base + 0, true,  true)
    S2_ITERATION(a_buf1, a_buf0, ki_base + 1, true,  true)
    S2_ITERATION(a_buf2, a_buf1, ki_base + 2, true,  true)
    S2_ITERATION(a_buf0, a_buf2, ki_base + 3, true,  true)
    S2_ITERATION(a_buf1, a_buf0, ki_base + 4, true,  true)
    S2_ITERATION(a_buf2, a_buf1, ki_base + 5, true,  true)
    S2_ITERATION(a_buf0, a_buf2, ki_base + 6, false, true)
    S2_ITERATION(a_buf1, a_buf0, ki_base + 7, false, false)

    __syncthreads();

    // Double-LDS reduction
    float* lds_reduce = reinterpret_cast<float*>(lds_bytes);
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        #pragma unroll
        for (int j = 0; j < 4; ++j) {
            const int row  = half * 4 + j + i * 8;
            const int base = row * S2_NumWaves * 32 + wave_id * 32 + lane;
            lds_reduce[base]            = acc0[i * 4 + j];
            lds_reduce[base + lds_half] = acc1[i * 4 + j];
        }
    }
    __syncthreads();

    {
        const int i = wave_id;
        #pragma unroll
        for (int j = 0; j < 4; ++j) {
            const int row        = half * 4 + j + i * 8;
            const int global_row = row_base + row;
            float sum0 = 0.0f;
            #pragma unroll
            for (int w = 0; w < S2_NumWaves; ++w)
                sum0 += lds_reduce[row * S2_NumWaves * 32 + w * 32 + lane];
            C[global_row * S2_N + col_base + lane] = float_to_bf16_rn(sum0);
            float sum1 = 0.0f;
            #pragma unroll
            for (int w = 0; w < S2_NumWaves; ++w)
                sum1 += lds_reduce[lds_half + row * S2_NumWaves * 32 + w * 32 + lane];
            C[global_row * S2_N + col_base + 32 + lane] = float_to_bf16_rn(sum1);
        }
    }
}

#undef S2_ITERATION

}  // namespace e641v1_exp

extern "C" void launch_e641v1_exp_64_7168_2048(
    const uint16_t* A, const uint8_t* B, const uint8_t* Bs, uint16_t* C) {
    hipLaunchKernelGGL(e641v1_exp::fused_gemm_s2_handunrolled,
        dim3(e641v1_exp::S2_TotalBlocks), dim3(e641v1_exp::S2_BlockSize),
        e641v1_exp::S2_LdsBytes, 0, A, B, Bs, C);
}


// =====================================================================
// KERNEL 5: M=256 (mk3g S4 — v_max3 absmax, deferred B scale)
// =====================================================================

namespace mk3g {

using v4i32  = int   __attribute__((ext_vector_type(4)));
using v16f32 = float __attribute__((ext_vector_type(16)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));

__device__ __forceinline__ uint32_t bit_cast_u32(float v) {
    union { float f; uint32_t u; } x; x.f = v; return x.u;
}
__device__ __forceinline__ float bit_cast_f32(uint32_t v) {
    union { uint32_t u; float f; } x; x.u = v; return x.f;
}
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) {
    uint32_t bits = bit_cast_u32(v);
    bits += ((bits >> 16) & 1u) + 0x7FFFu;
    return static_cast<uint16_t>(bits >> 16);
}
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t max_abs) {
    const uint16_t r = static_cast<uint16_t>(max_abs + 0x20u);
    const int e = static_cast<int>((r >> 7) & 0xFFu);
    return (e <= 2) ? 0u : (e >= 255) ? 254u : static_cast<uint8_t>(e - 2);
}

__device__ __forceinline__ v16f32 mfma_fp4x4(
    v4i32 a, v4i32 b, v16f32 acc, int a_scale, int b_scale) {
    asm volatile(
        "v_mfma_scale_f32_32x32x64_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
        : "+v"(acc) : "v"(a), "v"(b), "v"(a_scale), "v"(b_scale));
    return acc;
}

__device__ __forceinline__ void quant_a_group(
    const uint32_t* __restrict__ a_ptr, v4i32& a_out, int& scale_out)
{
    uint32_t dw[16];
    *reinterpret_cast<uint4*>(&dw[0])  = *(reinterpret_cast<const uint4*>(a_ptr) + 0);
    *reinterpret_cast<uint4*>(&dw[4])  = *(reinterpret_cast<const uint4*>(a_ptr) + 1);
    *reinterpret_cast<uint4*>(&dw[8])  = *(reinterpret_cast<const uint4*>(a_ptr) + 2);
    *reinterpret_cast<uint4*>(&dw[12]) = *(reinterpret_cast<const uint4*>(a_ptr) + 3);

    uint32_t bmax = 0u;
    #pragma unroll
    for (int i = 0; i < 16; ++i) {
        const uint32_t w = dw[i];
        const uint32_t hi = (w >> 16) & 0x7FFFu;
        const uint32_t lo = w & 0x7FFFu;
        bmax = (hi > bmax) ? hi : bmax;
        bmax = (lo > bmax) ? lo : bmax;
    }

    const uint8_t scale = compute_e8m0_scale(static_cast<uint16_t>(bmax));
    scale_out = static_cast<int>(scale);
    const float fwd = (scale == 0u) ? bit_cast_f32(0x00400000u)
                    : bit_cast_f32(static_cast<uint32_t>(scale) << 23);

    const v2bf16* pairs = reinterpret_cast<const v2bf16*>(dw);
    uint32_t pk[4] = {0,0,0,0};
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[0], fwd,0);
    pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[4], fwd,0);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[8], fwd,0);
    pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[12],fwd,0);
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[1], fwd,1);
    pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[5], fwd,1);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[9], fwd,1);
    pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[13],fwd,1);
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[2], fwd,2);
    pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[6], fwd,2);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[10],fwd,2);
    pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[14],fwd,2);
    pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[3], fwd,3);
    pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[7], fwd,3);
    pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[11],fwd,3);
    pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[15],fwd,3);
    a_out = {(int)pk[0], (int)pk[1], (int)pk[2], (int)pk[3]};
}

// Only S4 config needed
struct S4 {
    static constexpr int M = 256, N = 3072, K = 1536;
    static constexpr int KHalf = K / 2;
    static constexpr int GroupsPerRow = K / 32;
    static constexpr int PaddedGroups = (GroupsPerRow + 7) & ~7;
    static constexpr int TileM = 32, TileN = 64;
    static constexpr int NumWaves = 4;
    static constexpr int BlockSize = NumWaves * 64;
    static constexpr int MfmaItersTotal = K / 64;
    static constexpr int MfmaItersPerWave = MfmaItersTotal / NumWaves;
    static constexpr int NumTilesN = N / TileN;
    static constexpr int NumTilesM = M / TileM;
    static constexpr int TotalBlocks = NumTilesN * NumTilesM;
    static constexpr int LdsBytes = 2 * NumWaves * 32 * TileM * static_cast<int>(sizeof(float));
    static constexpr int NXCD = 8;
    static constexpr int W = 8, C = 6;
    static constexpr int BlocksPerCycle = NXCD * C;
    static constexpr int Limit = (TotalBlocks / BlocksPerCycle) * BlocksPerCycle;
    static constexpr int TidPerGroup = W * NumTilesN;
};

template<typename C>
__device__ __forceinline__ int scale_offset(int row, int group) {
    return (row >> 5) * (32 * C::PaddedGroups)
         + (group >> 3) * 256
         + (group & 3) * 64
         + (row & 15) * 4
         + ((group >> 2) & 1) * 2
         + ((row >> 4) & 1);
}

template<typename CFG>
__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void fused_gemm_mk3(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    uint16_t*       __restrict__ C)
{
    int xy = blockIdx.x;
    if (xy < CFG::Limit) {
        const int xcd    = xy % CFG::NXCD;
        const int local_ = xy / CFG::NXCD;
        const int chunk  = local_ / CFG::C;
        const int pos    = local_ % CFG::C;
        xy = chunk * CFG::BlocksPerCycle + xcd * CFG::C + pos;
    }
    const int l      = xy % CFG::TidPerGroup;
    const int m_tile = (xy / CFG::TidPerGroup) * CFG::W + (l % CFG::W);
    const int n_tile = l / CFG::W;

    const int col_base = n_tile * CFG::TileN;
    const int row_base = m_tile * CFG::TileM;

    const int tid     = threadIdx.x;
    const int wave_id = tid >> 6;
    const int lane    = tid & 31;
    const int half    = (tid >> 5) & 1;

    extern __shared__ char lds_bytes[];
    float* lds_reduce = reinterpret_cast<float*>(lds_bytes);
    constexpr int lds_half = CFG::NumWaves * 32 * CFG::TileM;

    const int b_col0 = col_base + lane;
    const int b_col1 = col_base + 32 + lane;
    const uint8_t* b_base0 = B_shuffle
        + (b_col0 >> 4) * (CFG::KHalf * 16) + (b_col0 & 15) * 16;
    const uint8_t* b_base1 = B_shuffle
        + (b_col1 >> 4) * (CFG::KHalf * 16) + (b_col1 & 15) * 16;

    const int my_row = row_base + lane;
    const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + my_row * CFG::K);

    constexpr int iters   = CFG::MfmaItersPerWave;
    const int     ki_base = wave_id * iters;

    v16f32 acc0, acc1;
    #pragma unroll
    for (int i = 0; i < 16; ++i) { acc0[i] = 0.0f; acc1[i] = 0.0f; }

    // B DATA prefetch
    v4i32 b_cur0, b_cur1;
    {
        const uint32_t* bs0 = reinterpret_cast<const uint32_t*>(
            b_base0 + ki_base * 512 + half * 256);
        b_cur0 = {(int)bs0[0], (int)bs0[1], (int)bs0[2], (int)bs0[3]};

        const uint32_t* bs1 = reinterpret_cast<const uint32_t*>(
            b_base1 + ki_base * 512 + half * 256);
        b_cur1 = {(int)bs1[0], (int)bs1[1], (int)bs1[2], (int)bs1[3]};
    }

    // Main K loop
    #pragma unroll
    for (int k = 0; k < iters; ++k) {
        const int ki      = ki_base + k;
        const int k_group = ki * 2 + half;

        // Load + quant A
        v4i32 a_reg;
        int   a_scale;
        quant_a_group(a_row + k_group * 16, a_reg, a_scale);

        // Prefetch next B DATA
        v4i32 b_nxt0 = {0,0,0,0}, b_nxt1 = {0,0,0,0};
        if (k + 1 < iters) {
            const int ki_next = ki_base + k + 1;
            const uint32_t* bs0 = reinterpret_cast<const uint32_t*>(
                b_base0 + ki_next * 512 + half * 256);
            b_nxt0 = {(int)bs0[0], (int)bs0[1], (int)bs0[2], (int)bs0[3]};

            const uint32_t* bs1 = reinterpret_cast<const uint32_t*>(
                b_base1 + ki_next * 512 + half * 256);
            b_nxt1 = {(int)bs1[0], (int)bs1[1], (int)bs1[2], (int)bs1[3]};
        }

        // B SCALE just-in-time
        int b_sc0 = static_cast<int>(B_scale_sh[
            scale_offset<CFG>(b_col0, k_group)]);
        int b_sc1 = static_cast<int>(B_scale_sh[
            scale_offset<CFG>(b_col1, k_group)]);

        // MFMA
        acc0 = mfma_fp4x4(a_reg, b_cur0, acc0, a_scale, b_sc0);
        acc1 = mfma_fp4x4(a_reg, b_cur1, acc1, a_scale, b_sc1);

        b_cur0 = b_nxt0;
        b_cur1 = b_nxt1;
    }

    // Double-LDS reduction
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        #pragma unroll
        for (int j = 0; j < 4; ++j) {
            const int row  = half * 4 + j + i * 8;
            const int base = row * CFG::NumWaves * 32 + wave_id * 32 + lane;
            lds_reduce[base]            = acc0[i * 4 + j];
            lds_reduce[base + lds_half] = acc1[i * 4 + j];
        }
    }
    __syncthreads();

    {
        const int i = wave_id;
        #pragma unroll
        for (int j = 0; j < 4; ++j) {
            const int row        = half * 4 + j + i * 8;
            const int global_row = row_base + row;
            float sum0 = 0.0f;
            #pragma unroll
            for (int w = 0; w < CFG::NumWaves; ++w)
                sum0 += lds_reduce[row * CFG::NumWaves * 32 + w * 32 + lane];
            C[global_row * CFG::N + col_base + lane] = float_to_bf16_rn(sum0);
            float sum1 = 0.0f;
            #pragma unroll
            for (int w = 0; w < CFG::NumWaves; ++w)
                sum1 += lds_reduce[lds_half + row * CFG::NumWaves * 32 + w * 32 + lane];
            C[global_row * CFG::N + col_base + 32 + lane] = float_to_bf16_rn(sum1);
        }
    }
}

}  // namespace mk3g

extern "C" void launch_mk3g_256_3072_1536(
    const uint16_t* A, const uint8_t* B, const uint8_t* Bs, uint16_t* C) {
    using C_ = mk3g::S4;
    hipLaunchKernelGGL(mk3g::fused_gemm_mk3<C_>,
        dim3(C_::TotalBlocks), dim3(C_::BlockSize), C_::LdsBytes, 0, A, B, Bs, C);
}
"""

# =========================================================================
# Python wrapper
# =========================================================================

def _extra_cuda_cflags() -> list[str]:
    flags = ["-O3", "-std=c++17"]
    arch = os.environ.get("PYTORCH_ROCM_ARCH", "gfx950")
    if all(sep not in arch for sep in (",", ";", " ")):
        flags.append(f"--offload-arch={arch}")
        # flags.append("-mllvm")
        # flags.append("-amdgpu-kernarg-preload-count=8")
    return flags


def _build_module():
    os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
    name = "submission_general_v1"
    build_dir = os.path.join(tempfile.gettempdir(), name)
    os.makedirs(build_dir, exist_ok=True)
    return load_inline(
        name=name,
        cpp_sources=_CPP_SRC,
        cuda_sources=_HIP_SRC,
        functions=[
            "m4_launch",
            "m16_launch",
            "m32_n2880_launch",
            "m32_n4096_launch",
            "m64_launch",
            "m256_launch",
        ],
        with_cuda=True,
        extra_cflags=["-O3"],
        extra_cuda_cflags=_extra_cuda_cflags(),
        build_directory=build_dir,
        verbose=False,
    )

_mod = _build_module()


def _storage_as_uint8(x):
    return x if x.dtype == torch.uint8 else x.view(torch.uint8)


def custom_kernel(data: input_t) -> output_t:
    A, _B, B_q, B_shuffle, B_scale_sh = data
    m, k = int(A.shape[0]), int(A.shape[1])
    n = int(B_shuffle.shape[0])

    if m == 4 and n == 2880 and k == 512:
        B_sh_u8 = _storage_as_uint8(B_shuffle)
        B_sc_u8 = _storage_as_uint8(B_scale_sh)
        C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _mod.m4_launch(A, B_sh_u8, B_sc_u8, C)
        return C

    if m == 16 and n == 2112 and k == 7168:
        B_sh_u8 = _storage_as_uint8(B_shuffle)
        B_sc_u8 = _storage_as_uint8(B_scale_sh)
        workspace = torch.empty((7, m, n), dtype=torch.float32, device=A.device)
        C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _mod.m16_launch(A, B_sh_u8, B_sc_u8, workspace, C)
        return C

    if m == 32 and n == 2880 and k == 512:
        B_sh_u8 = _storage_as_uint8(B_shuffle)
        B_sc_u8 = _storage_as_uint8(B_scale_sh)
        C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _mod.m32_n2880_launch(A, B_sh_u8, B_sc_u8, C)
        return C

    if m == 32 and n == 4096 and k == 512:
        B_sh_u8 = _storage_as_uint8(B_shuffle)
        B_sc_u8 = _storage_as_uint8(B_scale_sh)
        C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _mod.m32_n4096_launch(A, B_sh_u8, B_sc_u8, C)
        return C

    if m == 64 and n == 7168 and k == 2048:
        B_sh_u8 = _storage_as_uint8(B_shuffle)
        B_sc_u8 = _storage_as_uint8(B_scale_sh)
        C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _mod.m64_launch(A, B_sh_u8, B_sc_u8, C)
        return C

    if m == 256 and n == 3072 and k == 1536:
        B_sh_u8 = _storage_as_uint8(B_shuffle)
        B_sc_u8 = _storage_as_uint8(B_scale_sh)
        C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _mod.m256_launch(A, B_sh_u8, B_sc_u8, C)
        return C

    # Fallback to aiter reference
    x_fp4, bs = dynamic_mxfp4_quant(A)
    bs = e8m0_shuffle(bs)
    return aiter.gemm_a4w4(
        x_fp4.view(dtypes.fp4x2), B_shuffle,
        bs.view(dtypes.fp8_e8m0), B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 1500 lines total

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

Best evidence level for this revision: reported

JSON