Skip to content
KernelIndex
Search⌘K

submission 688083

BillHuang2001 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-688083?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
21.8µs
#740 of 1143
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6085114c388259e9e5f95b05436da6312db6859385f02df54af844f6ed72cbb0
license declaredunknown
license concludedunknown
authorsBillHuang2001
imported2026-08-26

Techniques

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

fp4FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
shared-memoryfrom flydsl.utils.smem_allocator import SmemAllocator, SmemPtr
vector-width = int4is_int4 = in_dtype == "int4"

Kernel source

submission_v2.py1111 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
Optimized for zero-copy memory dispatch using native CDNA hardware bounds clamping.
"""
import sys
import os
import fcntl
import functools
import importlib
import importlib.metadata
import importlib.util
import pkgutil
import shutil
import subprocess
import tempfile
import traceback
import torch

from task import input_t, output_t
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle


_REQUIRED_FLYDSL_VERSION = "0.1.1.dev409"
_RUNTIME_DEPS_ROOT = os.path.join(tempfile.gettempdir(), "submission_v2_runtime_deps")
_RUNTIME_SITE_PACKAGES = os.path.join(
    _RUNTIME_DEPS_ROOT,
    f"py{sys.version_info.major}{sys.version_info.minor}",
    f"flydsl-{_REQUIRED_FLYDSL_VERSION}",
)
_RUNTIME_INSTALL_LOCK = os.path.join(_RUNTIME_DEPS_ROOT, "flydsl-install.lock")
_RUNTIME_INSTALL_SENTINEL = os.path.join(_RUNTIME_SITE_PACKAGES, ".install-complete")


def _stderr(msg):
    print(f"[submission_v2 debug] {msg}", file=sys.stderr, flush=True)


def _prepend_runtime_site_packages():
    if _RUNTIME_SITE_PACKAGES in sys.path:
        sys.path.remove(_RUNTIME_SITE_PACKAGES)
    sys.path.insert(0, _RUNTIME_SITE_PACKAGES)


def _clear_flydsl_modules():
    for module_name in list(sys.modules):
        if module_name == "flydsl" or module_name.startswith("flydsl."):
            sys.modules.pop(module_name, None)


def _flydsl_runtime_ready():
    _prepend_runtime_site_packages()
    importlib.invalidate_caches()
    spec = importlib.util.find_spec("flydsl.expr")
    if spec is None:
        return False
    try:
        _clear_flydsl_modules()
        flydsl_pkg = importlib.import_module("flydsl")
    except Exception:
        return False
    version = getattr(flydsl_pkg, "__version__", None)
    return version == _REQUIRED_FLYDSL_VERSION


def _ensure_flydsl_runtime():
    os.makedirs(_RUNTIME_DEPS_ROOT, exist_ok=True)
    os.makedirs(_RUNTIME_SITE_PACKAGES, exist_ok=True)
    _prepend_runtime_site_packages()

    if os.path.exists(_RUNTIME_INSTALL_SENTINEL) and _flydsl_runtime_ready():
        _stderr(
            f"using cached runtime FlyDSL from {_RUNTIME_SITE_PACKAGES} version={_REQUIRED_FLYDSL_VERSION}"
        )
        return

    with open(_RUNTIME_INSTALL_LOCK, "w") as lock_file:
        fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX)

        if os.path.exists(_RUNTIME_INSTALL_SENTINEL) and _flydsl_runtime_ready():
            _stderr(
                f"runtime FlyDSL became available while waiting for lock version={_REQUIRED_FLYDSL_VERSION}"
            )
            return

        _stderr(
            f"installing flydsl=={_REQUIRED_FLYDSL_VERSION} into isolated target {_RUNTIME_SITE_PACKAGES}"
        )
        shutil.rmtree(_RUNTIME_SITE_PACKAGES, ignore_errors=True)
        os.makedirs(_RUNTIME_SITE_PACKAGES, exist_ok=True)

        cmd = [
            sys.executable,
            "-m",
            "pip",
            "install",
            "--disable-pip-version-check",
            "--no-input",
            "--upgrade",
            "--ignore-installed",
            "--target",
            _RUNTIME_SITE_PACKAGES,
            f"flydsl=={_REQUIRED_FLYDSL_VERSION}",
        ]
        _stderr(f"pip command={' '.join(cmd)}")
        proc = subprocess.run(cmd, capture_output=True, text=True)
        _stderr(f"pip returncode={proc.returncode}")
        if proc.stdout:
            _stderr(f"pip stdout:\n{proc.stdout}")
        if proc.stderr:
            _stderr(f"pip stderr:\n{proc.stderr}")
        if proc.returncode != 0:
            raise RuntimeError(
                f"failed to install flydsl=={_REQUIRED_FLYDSL_VERSION} into {_RUNTIME_SITE_PACKAGES}"
            )

        _prepend_runtime_site_packages()
        importlib.invalidate_caches()
        _clear_flydsl_modules()
        if not _flydsl_runtime_ready():
            raise RuntimeError(
                f"installed flydsl does not expose flydsl.expr or version {_REQUIRED_FLYDSL_VERSION}"
            )

        with open(_RUNTIME_INSTALL_SENTINEL, "w") as sentinel_file:
            sentinel_file.write(_REQUIRED_FLYDSL_VERSION)
        _stderr(f"runtime FlyDSL install verified version={_REQUIRED_FLYDSL_VERSION}")

_ensure_flydsl_runtime()

import flydsl.compiler as flyc
import flydsl.expr as fx
from flydsl.compiler.kernel_function import CompilationContext
from flydsl.expr import range_constexpr
from flydsl.runtime.device import get_rocm_arch as get_hip_arch
from flydsl.utils.smem_allocator import SmemAllocator, SmemPtr
from flydsl._mlir import ir
from flydsl.expr import arith, vector, gpu, buffer_ops, rocdl
from flydsl.expr.typing import T

# ---------------------------------------------------------------------------
# Inlined MFMA Preshuffle Pipeline Helpers
# ---------------------------------------------------------------------------

def crd2idx(crd, layout):
    result = fx.crd2idx(crd, layout)
    scalar = fx.get_scalar(result)
    if isinstance(scalar, ir.Value) and not isinstance(scalar.type, ir.IndexType):
        scalar = arith.IndexCastOp(T.index, scalar).result
    return scalar

def swizzle_xor16(row, col, k_blocks16):
    rem = row % k_blocks16
    return col ^ (rem * 16)

def _buffer_load_vec(buffer_ops, vector, rsrc, idx, *, elem_type, vec_elems, elem_bytes, offset_in_bytes):
    elem_size = int(elem_bytes)
    load_bytes = int(vec_elems) * elem_size
    vec_width = load_bytes // 4
    if offset_in_bytes:
        idx_i32 = idx // 4
    elif elem_bytes == 2:
        idx_i32 = (idx * 2) // 4
    else:
        idx_i32 = idx
    i32_val = buffer_ops.buffer_load(rsrc, idx_i32, vec_width=vec_width, dtype=T.i32)
    if vec_width == 1:
        i32_vec = vector.from_elements(T.vec(1, T.i32), [i32_val])
    else:
        i32_vec = i32_val
    return vector.bitcast(T.vec(int(vec_elems), elem_type), i32_vec)

def _i8x4_in_i32_to_bf16x4_i64(val_i32, arith, vector, scale_val=None):
    vec1_i32_t = T.vec(1, T.i32)
    vec2_i32 = T.i32x2
    vec4_i8 = T.i8x4
    vec1_i64 = T.vec(1, T.i64)
    v1 = vector.from_elements(vec1_i32_t, [val_i32])
    i8x4 = vector.bitcast(vec4_i8, v1)
    f32_vals = []
    for i in range(4):
        val_i8 = vector.extract(i8x4, static_position=[i], dynamic_position=[])
        v = arith.sitofp(T.f32, val_i8)
        if scale_val is not None:
            v = v * scale_val
        f32_vals.append(v)
    c16 = fx.Int32(16)
    c_ffff0000 = fx.Int32(0xFFFF0000)
    bits0 = arith.bitcast(T.i32, f32_vals[0])
    bits1 = arith.bitcast(T.i32, f32_vals[1])
    bits2 = arith.bitcast(T.i32, f32_vals[2])
    bits3 = arith.bitcast(T.i32, f32_vals[3])
    i32_lo = (bits0 >> c16) | (bits1 & c_ffff0000)
    i32_hi = (bits2 >> c16) | (bits3 & c_ffff0000)
    v2 = vector.from_elements(vec2_i32, [i32_lo, i32_hi])
    v64 = vector.bitcast(vec1_i64, v2)
    return vector.extract(v64, static_position=[0], dynamic_position=[])

def load_b_pack_k32(buffer_ops, arith, vector, *, arg_b, b_rsrc, layout_b, base_k, ki_step, n_blk, n_intra, lane_div_16, elem_type, kpack_bytes=16, elem_bytes=1, unpack_int4=False):
    c64 = fx.Index(64)
    base_k_bytes = base_k * arith.constant(int(elem_bytes), index=True)
    k0_base = base_k_bytes // c64
    k0 = k0_base + arith.constant(ki_step // 2, index=True)
    k1 = lane_div_16
    half_bytes = kpack_bytes // 2
    k2_base = arith.constant((ki_step % 2) * half_bytes, index=True)
    coord_pack = (n_blk, k0, k1, n_intra, fx.Index(0))
    idx_pack = crd2idx(coord_pack, layout_b)

    if unpack_int4:
        idx_bytes = idx_pack + k2_base
        b4 = _buffer_load_vec(buffer_ops, vector, b_rsrc, idx_bytes, elem_type=elem_type, vec_elems=4, elem_bytes=1, offset_in_bytes=True)
        packed32 = vector.extract(vector.bitcast(T.vec(1, T.i32), b4), static_position=[0], dynamic_position=[])
        c_08080808 = fx.Int32(0x08080808)
        c_0f0f0f0f = fx.Int32(0x0F0F0F0F)
        c_1e = fx.Int32(0x1E)
        c_4_i32 = fx.Int32(4)
        s0 = (packed32 & c_08080808) * c_1e
        even = (packed32 & c_0f0f0f0f) | s0
        t = packed32 >> c_4_i32
        s1 = (t & c_08080808) * c_1e
        odd = (t & c_0f0f0f0f) | s1
        v2 = vector.from_elements(T.vec(2, T.i32), [even, odd])
        v64 = vector.bitcast(T.vec(1, T.i64), v2)
        return vector.extract(v64, static_position=[0], dynamic_position=[])

    vec_elems = kpack_bytes // int(elem_bytes)
    b16 = _buffer_load_vec(buffer_ops, vector, b_rsrc, idx_pack, elem_type=elem_type, vec_elems=vec_elems, elem_bytes=elem_bytes, offset_in_bytes=(elem_bytes == 1))
    b_i32x4 = vector.bitcast(T.i32x4, b16)
    half = ki_step % 2
    if half == 0:
        d0 = vector.extract(b_i32x4, static_position=[0], dynamic_position=[])
        d1 = vector.extract(b_i32x4, static_position=[1], dynamic_position=[])
    else:
        d0 = vector.extract(b_i32x4, static_position=[2], dynamic_position=[])
        d1 = vector.extract(b_i32x4, static_position=[3], dynamic_position=[])
    v2 = vector.from_elements(T.vec(2, T.i32), [d0, d1])
    v64 = vector.bitcast(T.vec(1, T.i64), v2)
    return vector.extract(v64, static_position=[0], dynamic_position=[])

def tile_chunk_coord_i32(arith, *, tx_i32_base, i, total_threads, layout_tile_div4, chunk_i32=4):
    chunk_off_i32 = arith.constant(i * total_threads * chunk_i32, index=True)
    tile_idx_i32 = tx_i32_base + chunk_off_i32
    coord_local = fx.idx2crd(tile_idx_i32, layout_tile_div4)
    row_local = fx.get(coord_local, 0)
    col_local_i32 = fx.get(coord_local, 1)
    return row_local, col_local_i32

def buffer_copy_gmem16_dwordx4(buffer_ops, vector, *, elem_type, idx_i32, rsrc, vec_elems=16, elem_bytes=1):
    return _buffer_load_vec(buffer_ops, vector, rsrc, idx_i32, elem_type=elem_type, vec_elems=vec_elems, elem_bytes=elem_bytes, offset_in_bytes=False)

def lds_load_pack_k32(arith, vector, *, lds_memref, layout_lds, k_blocks16, curr_row_a_lds, col_base, half, lds_base, ck_lds128, vec16_ty, vec8_ty, vec2_i64_ty, vec1_i64_ty):
    col_base_swz = swizzle_xor16(curr_row_a_lds, col_base, k_blocks16)
    if ck_lds128:
        coord_a16 = (curr_row_a_lds, col_base_swz)
        idx_a16 = crd2idx(coord_a16, layout_lds) + lds_base
        loaded_a16 = vector.load_op(vec16_ty, lds_memref, [idx_a16])
        a_vec128 = vector.bitcast(vec2_i64_ty, loaded_a16)
        return vector.extract(a_vec128, static_position=[half], dynamic_position=[])
    else:
        col_swizzled = col_base_swz + (half * 8)
        coord_a = (curr_row_a_lds, col_swizzled)
        idx_a = crd2idx(coord_a, layout_lds) + lds_base
        loaded_a8 = vector.load_op(vec8_ty, lds_memref, [idx_a])
        a_vec64 = vector.bitcast(vec1_i64_ty, loaded_a8)
        return vector.extract(a_vec64, static_position=[0], dynamic_position=[])

# ---------------------------------------------------------------------------
# Inlined Preshuffle GEMM Compiler Logic
# ---------------------------------------------------------------------------

_TILE_PRELOAD_TABLE = {
    (16, 64, 256):  (2, 2), (16, 64, 512):  (8, 8), (16, 128, 256): (2, 2), (16, 128, 512): (2, 2),
    (16, 192, 256): (2, 2), (16, 256, 256): (2, 2), (16, 256, 512): (2, 2), (16, 512, 256): (2, 2),
    (32, 64, 128):  (6, 6), (32, 64, 256):  (6, 6), (32, 64, 512):  (2, 2), (32, 128, 128): (6, 6),
    (32, 128, 256): (6, 6), (32, 192, 128): (6, 6), (32, 192, 256): (6, 6), (32, 256, 128): (6, 6),
    (32, 256, 256): (6, 6), (48, 64, 128):  (8, 8), (48, 64, 256):  (2, 2), (48, 128, 256): (6, 6),
    (48, 192, 256): (6, 6), (48, 256, 256): (6, 6), (64, 64, 128):  (4, 4), (64, 64, 256):  (4, 4),
    (64, 128, 128): (8, 8), (64, 128, 256): (8, 8), (64, 192, 128): (8, 8), (64, 192, 256): (8, 8),
    (64, 256, 64):  (8, 8), (64, 256, 128): (8, 8), (64, 256, 256): (8, 8), (80, 64, 256):  (4, 4),
    (80, 128, 256): (8, 8), (80, 192, 256): (8, 8), (80, 256, 256): (8, 8), (96, 64, 128):  (6, 6),
    (96, 64, 256):  (6, 6), (96, 128, 128): (8, 8), (96, 128, 256): (6, 6), (96, 192, 128): (8, 8),
    (96, 192, 256): (8, 8), (96, 256, 128): (8, 8), (96, 256, 256): (8, 8), (112, 64, 256):  (8, 8),
    (112, 128, 256): (4, 4), (112, 192, 256): (8, 8), (112, 256, 256): (8, 8), (128, 64, 128):  (6, 6),
    (128, 64, 256):  (8, 8), (128, 128, 64):  (4, 4), (128, 128, 128): (8, 8), (128, 128, 256): (4, 4),
    (128, 192, 128): (8, 8), (128, 192, 256): (8, 8), (128, 256, 128): (6, 6), (128, 256, 256): (4, 4),
    (160, 192, 128): (8, 8), (192, 64, 128):  (6, 6), (192, 128, 128): (6, 6), (224, 64, 128):  (4, 4),
    (224, 128, 128): (6, 6), (224, 192, 128): (6, 6), (256, 64, 128):  (4, 4), (256, 128, 128): (6, 6),
    (256, 192, 128): (6, 6), (256, 256, 128): (4, 4),
}

_TILE_PRELOAD_DEFAULT = (0, 0)

def _get_preload(tile_m, tile_n, tile_k):
    return _TILE_PRELOAD_TABLE.get((int(tile_m), int(tile_n), int(tile_k)), _TILE_PRELOAD_DEFAULT)

def compile_preshuffle_gemm_a8(*, M=0, N=0, K, tile_m, tile_n, tile_k, in_dtype="fp8", out_dtype="fp16", lds_stage=2, use_cshuffle_epilog=False, waves_per_eu=None, use_async_copy=False, dsrd_preload=-1, dvmem_preload=-1):
    if dsrd_preload < 0 or dvmem_preload < 0:
        if in_dtype in ("fp8", "int8") and str(get_hip_arch()) == "gfx950":
            computed_dsrd, computed_dvmem = _get_preload(tile_m, tile_n, tile_k)
        else:
            computed_dsrd, computed_dvmem = _TILE_PRELOAD_DEFAULT
        if dsrd_preload < 0: dsrd_preload = computed_dsrd
        if dvmem_preload < 0: dvmem_preload = computed_dvmem

    _out_is_bf16 = out_dtype == "bf16"
    is_fp4 = in_dtype == "fp4"
    is_int4 = in_dtype == "int4"
    is_int8 = (in_dtype == "int8") or is_int4
    is_f16 = in_dtype == "fp16"
    is_bf16 = in_dtype == "bf16"
    is_f16_or_bf16 = is_f16 or is_bf16
    elem_bytes = 1 if (in_dtype in ("fp8", "int8", "int4", "fp4")) else 2
    a_elem_vec_pack = 2 if is_fp4 else 1
    b_elem_vec_pack = 2 if is_fp4 else 1
    tile_k_bytes = int(tile_k) * int(elem_bytes)

    gpu_arch = get_hip_arch()
    allocator_pong = SmemAllocator(None, arch=gpu_arch, global_sym_name="smem0")
    allocator_ping = SmemAllocator(None, arch=gpu_arch, global_sym_name="smem1")

    total_threads = 256
    bytes_a_per_tile = int(tile_m) * int(tile_k) * int(elem_bytes) // a_elem_vec_pack
    bytes_per_thread_a = bytes_a_per_tile // total_threads
    a_load_bytes = 16
    a_async_load_bytes = 4 if gpu_arch == "gfx942" else 16
    a_async_load_dword = a_async_load_bytes // 4
    bytes_b_per_tile = int(tile_n) * int(tile_k) * int(elem_bytes) // b_elem_vec_pack
    bytes_per_thread_b = bytes_b_per_tile // total_threads
    b_load_bytes = 16
    num_b_loads = bytes_per_thread_b // b_load_bytes
    wave_size = 64
    num_a_lds_load = bytes_a_per_tile // wave_size // a_load_bytes
    num_a_async_loads = bytes_per_thread_a // a_async_load_bytes
    _is_gfx950 = str(gpu_arch).startswith("gfx950")
    _is_gfx942 = str(gpu_arch).startswith("gfx942")
    lds_stride_bytes = tile_k_bytes

    def _elem_type():
        if is_f16: return T.f16
        if is_bf16: return T.bf16
        if is_fp4: return T.i8
        return T.i8 if is_int8 else T.f8

    def _vec16_type():
        if is_f16: return T.f16x8
        if is_bf16: return T.bf16x8
        if is_fp4: return T.i8x16
        return T.i8x16 if is_int8 else T.f8x16

    def _out_elem():
        return T.bf16 if _out_is_bf16 else T.f16

    lds_tile_bytes = int(tile_m) * int(lds_stride_bytes) // a_elem_vec_pack
    lds_out_bytes = 0 
    buffer_size_bytes = max(lds_tile_bytes, lds_out_bytes // lds_stage)
    buffer_size_elems = buffer_size_bytes if elem_bytes == 1 else (buffer_size_bytes // 2)

    lds_pong_offset = allocator_pong._align(allocator_pong.ptr, 16)
    allocator_pong.ptr = lds_pong_offset + buffer_size_elems * elem_bytes

    lds_ping_offset = allocator_ping._align(allocator_ping.ptr, 16)
    allocator_ping.ptr = lds_ping_offset + buffer_size_elems * elem_bytes

    @flyc.kernel
    def kernel_gemm(arg_c: fx.Tensor, arg_a: fx.Tensor, arg_b: fx.Tensor, arg_scale_a: fx.Tensor, arg_scale_b: fx.Tensor, i32_m: fx.Int32, i32_n: fx.Int32, i32_b_nbytes: fx.Int32, i32_sb_nbytes: fx.Int32):
        c_m = arith.index_cast(T.index, i32_m)
        c_n = arith.index_cast(T.index, i32_n)
        acc_init = arith.constant_vector(0, T.i32x4) if is_int8 else arith.constant_vector(0.0, T.f32x4)
        _k_div4_factor = (K * elem_bytes) // 4 // a_elem_vec_pack
        kpack_bytes = 8 if is_int4 else 16
        kpack_elems = kpack_bytes if elem_bytes == 1 else kpack_bytes // elem_bytes
        k_bytes_b = K * elem_bytes // b_elem_vec_pack
        n0_val = N // 16
        k0_val = k_bytes_b // 64
        _stride_nlane = kpack_elems
        _stride_klane = 16 * _stride_nlane
        _stride_k0 = 4 * _stride_klane
        _stride_n0 = k0_val * _stride_k0
        layout_b = fx.make_layout((n0_val, k0_val, 4, 16, kpack_elems), (_stride_n0, _stride_k0, _stride_klane, _stride_nlane, 1))

        lds_k_dim = tile_k // a_elem_vec_pack
        shape_lds = fx.make_shape(tile_m, lds_k_dim)
        stride_lds = fx.make_stride(lds_k_dim, 1)
        layout_lds = fx.make_layout(shape_lds, stride_lds)
        k_blocks16 = arith.index(tile_k_bytes // a_elem_vec_pack // 16)

        tx = gpu.thread_id("x")
        bx = gpu.block_id("x")
        by = gpu.block_id("y")

        base_ptr_pong = allocator_pong.get_base()
        base_ptr_ping = allocator_ping.get_base()

        lds_a_pong = SmemPtr(base_ptr_pong, lds_pong_offset, _elem_type(), shape=(tile_m * tile_k,)).get()
        lds_a_ping = SmemPtr(base_ptr_ping, lds_ping_offset, _elem_type(), shape=(tile_m * tile_k,)).get()

        # Runtime byte sizes for OOB protection (matches reference preshuffle_gemm.py)
        _a_nrec = arith.index_cast(T.i64, c_m * arith.index(K * elem_bytes // a_elem_vec_pack))
        _c_nrec = arith.index_cast(T.i64, c_m * c_n * arith.index(2))
        a_rsrc = buffer_ops.create_buffer_resource(arg_a, max_size=False, num_records_bytes=_a_nrec)
        c_rsrc = buffer_ops.create_buffer_resource(arg_c, max_size=False, num_records_bytes=_c_nrec)
        _needs_per_token_scale = not is_f16_or_bf16 and not is_fp4
        scale_a_rsrc = None if (is_f16_or_bf16) else buffer_ops.create_buffer_resource(arg_scale_a, max_size=False)
        b_rsrc = buffer_ops.create_buffer_resource(arg_b, max_size=False, num_records_bytes=i32_b_nbytes)
        scale_b_rsrc = None if (is_f16_or_bf16) else buffer_ops.create_buffer_resource(arg_scale_b, max_size=False, num_records_bytes=i32_sb_nbytes)

        bx_m = bx * tile_m
        by_n = by * tile_n

        layout_wave_lane = fx.make_layout((4, wave_size), (64, 1))
        coord_wave_lane = fx.idx2crd(tx, layout_wave_lane)
        wave_id = fx.get(coord_wave_lane, 0)
        lane_id = fx.get(coord_wave_lane, 1)

        layout_lane16 = fx.make_layout((4, 16), (16, 1))
        coord_lane16 = fx.idx2crd(lane_id, layout_lane16)
        lane_div_16 = fx.get(coord_lane16, 0)
        lane_mod_16 = fx.get(coord_lane16, 1)

        row_a_lds = lane_mod_16
        kpack_elems = 16 if elem_bytes == 1 else 8
        col_offset_base = lane_div_16 * kpack_elems
        col_offset_base_bytes = (col_offset_base if elem_bytes == 1 else col_offset_base * elem_bytes)

        m_repeat = tile_m // 16
        k_unroll = tile_k_bytes // a_elem_vec_pack // 64

        num_waves = 4
        n_per_wave = tile_n // num_waves
        num_acc_n = n_per_wave // 16

        n_tile_base = wave_id * n_per_wave
        n_intra_list = []
        n_blk_list = []
        for i in range_constexpr(num_acc_n):
            global_n = by_n + n_tile_base + (i * 16) + lane_mod_16
            n_blk_list.append(global_n // 16)
            n_intra_list.append(global_n % 16)

        c64_b = 64
        _b_stride_n0_c = fx.Index(_stride_n0)
        _b_stride_k0_c = fx.Index(_stride_k0)
        _b_stride_klane_c = fx.Index(_stride_klane)
        _b_stride_nlane_c = fx.Index(_stride_nlane)

        def _extract_b_packs(b16):
            b_i64x2 = vector.bitcast(T.i64x2, b16)
            b0_i64 = vector.extract(b_i64x2, static_position=[0], dynamic_position=[])
            b1_i64 = vector.extract(b_i64x2, static_position=[1], dynamic_position=[])
            if not is_f16_or_bf16:
                return b0_i64, b1_i64
            b0_v1 = vector.from_elements(T.vec(1, T.i64), [b0_i64])
            b1_v1 = vector.from_elements(T.vec(1, T.i64), [b1_i64])
            if is_f16:
                return vector.bitcast(T.f16x4, b0_v1), vector.bitcast(T.f16x4, b1_v1)
            return vector.bitcast(T.i16x4, b0_v1), vector.bitcast(T.i16x4, b1_v1)

        def load_b_packs_k64(base_k, ku: int, ni: int):
            base_k_bytes = base_k * elem_bytes
            k0 = base_k_bytes // c64_b + ku
            idx_pack = n_blk_list[ni] * _b_stride_n0_c + k0 * _b_stride_k0_c + lane_div_16 * _b_stride_klane_c + n_intra_list[ni] * _b_stride_nlane_c
            vec_elems = 16 if elem_bytes == 1 else 8
            b16 = _buffer_load_vec(buffer_ops, vector, b_rsrc, idx_pack, elem_type=_elem_type(), vec_elems=vec_elems, elem_bytes=elem_bytes, offset_in_bytes=(elem_bytes == 1))
            return _extract_b_packs(b16)

        def load_b_tile(base_k):
            packs0_per_ku = [[] for _ in range(k_unroll)]
            packs1_per_ku = [[] for _ in range(k_unroll)]
            for ni in range_constexpr(num_acc_n):
                for ku in range_constexpr(k_unroll):
                    b0, b1 = load_b_packs_k64(base_k, ku, ni)
                    packs0_per_ku[ku].append(b0)
                    packs1_per_ku[ku].append(b1)
            b_tile = []
            for ku in range_constexpr(k_unroll):
                b_tile.append((packs0_per_ku[ku], packs1_per_ku[ku]))
            return b_tile

        lds_base_zero = fx.Index(0)
        _lds_k_dim_c = fx.Index(lds_k_dim)

        def lds_load_16b(curr_row_a_lds, col_base, lds_buffer):
            col_base_swz_bytes = swizzle_xor16(curr_row_a_lds, col_base, k_blocks16)
            col_base_swz = col_base_swz_bytes if elem_bytes == 1 else (col_base_swz_bytes // 2)
            idx_a16 = curr_row_a_lds * _lds_k_dim_c + col_base_swz
            return vector.load_op(_vec16_type(), lds_buffer, [idx_a16])

        def lds_load_packs_k64(curr_row_a_lds, col_base, lds_buffer):
            loaded_a16 = lds_load_16b(curr_row_a_lds, col_base, lds_buffer)
            a_i64x2 = vector.bitcast(T.i64x2, loaded_a16)
            a0_i64 = vector.extract(a_i64x2, static_position=[0], dynamic_position=[])
            a1_i64 = vector.extract(a_i64x2, static_position=[1], dynamic_position=[])
            if not is_f16_or_bf16:
                return a0_i64, a1_i64
            a0_v1 = vector.from_elements(T.vec(1, T.i64), [a0_i64])
            a1_v1 = vector.from_elements(T.vec(1, T.i64), [a1_i64])
            if is_f16:
                return vector.bitcast(T.f16x4, a0_v1), vector.bitcast(T.f16x4, a1_v1)
            return vector.bitcast(T.i16x4, a0_v1), vector.bitcast(T.i16x4, a1_v1)

        num_a_loads = bytes_per_thread_a // a_load_bytes
        tile_k_dwords = (tile_k * 2) // 4 if elem_bytes == 2 else tile_k // 4 // a_elem_vec_pack
        layout_a_tile_div4 = fx.make_layout((tile_m, tile_k_dwords), (tile_k_dwords, 1))
        c4 = fx.Index(4)
        tx_i32_base = tx * c4

        def a_tile_chunk_coord_i32(i: int):
            return tile_chunk_coord_i32(arith, tx_i32_base=tx_i32_base, i=i, total_threads=total_threads, layout_tile_div4=layout_a_tile_div4)

        def load_a_tile(base_k_div4):
            parts = []
            for i in range_constexpr(num_a_loads):
                row_a_local, col_a_local_i32 = a_tile_chunk_coord_i32(i)
                row_a_global = bx_m + row_a_local
                idx_i32 = row_a_global * _k_div4_factor + (base_k_div4 + col_a_local_i32)
                idx_elem = idx_i32 if elem_bytes == 1 else idx_i32 * 2
                a_16B = buffer_copy_gmem16_dwordx4(buffer_ops, vector, elem_type=_elem_type(), idx_i32=idx_elem, rsrc=a_rsrc, vec_elems=(16 if elem_bytes == 1 else 8), elem_bytes=elem_bytes)
                parts.append(vector.bitcast(T.i32x4, a_16B))
            return parts

        def store_a_tile_to_lds(vec_a_parts, lds_buffer):
            for i in range_constexpr(num_a_loads):
                row_a_local, col_a_local_i32 = a_tile_chunk_coord_i32(i)
                col_local_bytes = col_a_local_i32 * c4
                col_swz_bytes = swizzle_xor16(row_a_local, col_local_bytes, k_blocks16)
                col_swz = col_swz_bytes if elem_bytes == 1 else col_swz_bytes // 2
                idx0 = row_a_local * _lds_k_dim_c + col_swz + lds_base_zero
                v16 = vector.bitcast(_vec16_type(), vec_a_parts[i])
                vector.store(v16, lds_buffer, [idx0])

        tx_i32_async_base = tx * a_async_load_dword
        k_bytes_factor = K * elem_bytes // a_elem_vec_pack

        def a_tile_chunk_coord_i32_async(i: int):
            return tile_chunk_coord_i32(arith, tx_i32_base=tx_i32_async_base, i=i, total_threads=total_threads, layout_tile_div4=layout_a_tile_div4, chunk_i32=a_async_load_dword)

        def dma_a_tile_to_lds(base_k_div4, lds_buffer):
            from flydsl._mlir.dialects import memref as memref_dialect
            dma_bytes = a_async_load_bytes
            wave_offset = rocdl.readfirstlane(T.i64, arith.index_cast(T.i64, wave_id * arith.constant(wave_size * dma_bytes, index=True)))
            for i in range_constexpr(num_a_async_loads):
                row_a_local, col_a_local_i32 = a_tile_chunk_coord_i32_async(i)
                col_a_local_sw = swizzle_xor16(row_a_local, col_a_local_i32 * c4, k_blocks16)
                row_a_global = bx_m + row_a_local
                global_byte_idx = row_a_global * k_bytes_factor + (base_k_div4 * c4 + col_a_local_sw)
                global_offset = arith.index_cast(T.i32, global_byte_idx)
                if i == 0:
                    lds_base = memref_dialect.extract_aligned_pointer_as_index(lds_buffer)
                    lds_ptr_base = buffer_ops.create_llvm_ptr(arith.index_cast(T.i64, lds_base), address_space=3)
                    lds_ptr = buffer_ops.get_element_ptr(lds_ptr_base, wave_offset)
                else:
                    lds_ptr = buffer_ops.get_element_ptr(lds_ptr, static_byte_offset=total_threads * dma_bytes)
                rocdl.raw_ptr_buffer_load_lds(a_rsrc, lds_ptr, arith.constant(dma_bytes, type=T.i32), global_offset, arith.constant(0, type=T.i32), arith.constant(0, type=T.i32), arith.constant(1, type=T.i32))

        def prefetch_a_to_lds(base_k, lds_buffer):
            base_k_div4 = base_k // 4 // a_elem_vec_pack
            dma_a_tile_to_lds(base_k_div4, lds_buffer)

        def prefetch_a_tile(base_k):
            base_k_bytes = base_k * elem_bytes // a_elem_vec_pack
            base_k_div4 = base_k_bytes // 4
            return load_a_tile(base_k_div4)

        def prefetch_b_tile(base_k):
            base_k_packed = base_k // b_elem_vec_pack if b_elem_vec_pack > 1 else base_k
            return load_b_tile(base_k_packed)

        _fp4_tilek128 = False
        if is_fp4:
            _fp4_pack_M_outer = 2
            _fp4_pack_N_outer = 2
            _fp4_pack_K_outer = 2
            _fp4_tilek128 = int(tile_k) == 128
            _fp4_scale_chunk_k = 32 * 4 * _fp4_pack_K_outer
            _K1_outer = K // (32 * 4 * _fp4_pack_K_outer)
            _k_unroll_packed_outer = 1 if _fp4_tilek128 else (k_unroll // _fp4_pack_K_outer)
            _m_repeat_packed_outer = m_repeat // _fp4_pack_M_outer
            _num_acc_n_packed_outer = num_acc_n // _fp4_pack_N_outer
            _fp4_scale_k_stride = tile_k // (32 * 4 * _fp4_pack_K_outer)
            _fp4_use_scheduler = (tile_m >= 64)

            _scale_lane_elem_off = lane_div_16 * fx.Index(16) + lane_mod_16
            _scale_row_stride_elems = _K1_outer * 64

            _scale_a_base_elems = []
            for mi in range_constexpr(_m_repeat_packed_outer):
                mni_a = fx.Index(mi) + bx_m // arith.index(_fp4_pack_M_outer * 16)
                _scale_a_base_elems.append(mni_a * arith.index(_scale_row_stride_elems) + _scale_lane_elem_off)

            _scale_b_base_elems = []
            for ni in range_constexpr(_num_acc_n_packed_outer):
                mni_b = fx.Index(ni) + (by_n + n_tile_base) // arith.index(_fp4_pack_N_outer * 16)
                _scale_b_base_elems.append(mni_b * arith.index(_scale_row_stride_elems) + _scale_lane_elem_off)

            _stride_k0_elems = 64

            def load_fp4_scales(base_k_scale_idx):
                a_scales, b_scales = [], []
                base_k_elem_off = base_k_scale_idx * fx.Index(_stride_k0_elems)
                for ku in range_constexpr(_k_unroll_packed_outer):
                    ku_elem_off = base_k_elem_off + fx.Index(ku * _stride_k0_elems)
                    for ni in range_constexpr(_num_acc_n_packed_outer):
                        b_scales.append(buffer_ops.buffer_load(scale_b_rsrc, _scale_b_base_elems[ni] + ku_elem_off, vec_width=1, dtype=T.i32))
                    for mi in range_constexpr(_m_repeat_packed_outer):
                        a_scales.append(buffer_ops.buffer_load(scale_a_rsrc, _scale_a_base_elems[mi] + ku_elem_off, vec_width=1, dtype=T.i32))
                return a_scales, b_scales

            def load_fp4_scale_chunk(base_k):
                return load_fp4_scales(base_k // fx.Index(_fp4_scale_chunk_k))

        def compute_tile(accs_in, b_tile_in, lds_buffer, *, is_last_tile=False, a0_prefetch=None, fp4_scales=None, fp4_scale_half=0):
            scales_pf = {}
            if is_last_tile and (not is_f16_or_bf16):
                s_b_vals = []
                for ni in range_constexpr(num_acc_n):
                    col_g = by_n + n_tile_base + (ni * 16) + lane_mod_16
                    s_b_vals.append(buffer_ops.buffer_load(scale_b_rsrc, col_g, vec_width=1, dtype=T.f32))
                scales_pf["s_b_vals"] = s_b_vals
                scales_pf["s_a_vecs"] = []
                row_off_base = lane_div_16 * 4
                for mi in range_constexpr(m_repeat):
                    row_base_m = bx_m + (mi * 16)
                    row_g_base = row_base_m + row_off_base
                    s_a_vec = buffer_ops.buffer_load(scale_a_rsrc, row_g_base, vec_width=4, dtype=T.f32)
                    scales_pf["s_a_vecs"].append(vector.bitcast(T.f32x4, s_a_vec))

            current_accs_list = list(accs_in)
            mfma_res_ty = T.f32x4
            c0_i64 = arith.constant(0, type=T.i64)

            _fp4_cbsz = 4 if is_fp4 else 0
            _fp4_blgp = 4 if is_fp4 else 0
            _fp4_pack_M = 2 if is_fp4 else 1
            _fp4_pack_N = 2 if is_fp4 else 1
            _fp4_pack_K = 2 if is_fp4 else 1

            def pack_i64x4_to_i32x8(x0, x1, x2, x3):
                v4 = vector.from_elements(T.vec(4, T.i64), [x0, x1, x2, x3])
                return vector.bitcast(T.vec(8, T.i32), v4)

            if is_fp4:
                _fp4_a_sc, _fp4_b_sc = fp4_scales if fp4_scales else ([], [])
                ku128_iters = 1 if _fp4_tilek128 else _k_unroll_packed_outer
                ikxdl_iters = 1 if _fp4_tilek128 else _fp4_pack_K
                for ku128 in range_constexpr(ku128_iters):
                    a_scale_base = 0 if _fp4_tilek128 else ku128 * _m_repeat_packed_outer
                    b_scale_base = 0 if _fp4_tilek128 else ku128 * _num_acc_n_packed_outer
                    for mi_p in range_constexpr(_m_repeat_packed_outer):
                        a_scale_val = _fp4_a_sc[a_scale_base + mi_p]
                        for ni_p in range_constexpr(_num_acc_n_packed_outer):
                            b_scale_val = _fp4_b_sc[b_scale_base + ni_p]
                            for ikxdl in range_constexpr(ikxdl_iters):
                                k_idx = 0 if _fp4_tilek128 else ku128 * _fp4_pack_K + ikxdl
                                b_packs0, b_packs1 = b_tile_in[k_idx]
                                col_base = col_offset_base_bytes if _fp4_tilek128 else (col_offset_base_bytes + arith.index((k_idx * 128) // a_elem_vec_pack))
                                scale_k_sel = fp4_scale_half if _fp4_tilek128 else ikxdl
                                for imxdl in range_constexpr(_fp4_pack_M):
                                    mi_idx = mi_p * _fp4_pack_M + imxdl
                                    curr_row_a_lds = row_a_lds + (mi_idx * 16)
                                    if (a0_prefetch is not None) and (k_idx == 0) and (mi_idx == 0):
                                        a0, a1 = a0_prefetch
                                    else:
                                        a0, a1 = lds_load_packs_k64(curr_row_a_lds, col_base, lds_buffer)
                                    a128 = pack_i64x4_to_i32x8(a0, a1, c0_i64, c0_i64)
                                    for inxdl in range_constexpr(_fp4_pack_N):
                                        ni_idx = ni_p * _fp4_pack_N + inxdl
                                        b0 = b_packs0[ni_idx]
                                        b1 = b_packs1[ni_idx]
                                        b128 = pack_i64x4_to_i32x8(b0, b1, c0_i64, c0_i64)
                                        acc_idx = mi_idx * num_acc_n + ni_idx
                                        if not _fp4_use_scheduler:
                                            rocdl.sched_barrier(0)
                                        current_accs_list[acc_idx] = rocdl.mfma_scale_f32_16x16x128_f8f6f4(
                                            mfma_res_ty,
                                            [a128, b128, current_accs_list[acc_idx], _fp4_cbsz, _fp4_blgp, scale_k_sel * _fp4_pack_M + imxdl, a_scale_val, scale_k_sel * _fp4_pack_N + inxdl, b_scale_val],
                                        )
            return current_accs_list, scales_pf

        def store_output(final_accs, scales):
            if is_f16_or_bf16 or is_fp4:
                s_b_vals = None
                s_a_vecs = None
            else:
                s_b_vals = scales["s_b_vals"]
                s_a_vecs = scales["s_a_vecs"]

            def body_row(*, mi, ii, row):
                if _needs_per_token_scale:
                    s_a_vec4 = s_a_vecs[mi]
                    s_a = vector.extract(s_a_vec4, static_position=[ii], dynamic_position=[])
                
                # PERFECT CDNA 16x16 wave64 output mapping:
                # The 4 output values returned to each thread map to the SAME col, spaced 4 rows apart.
                col_base = by_n + n_tile_base + lane_mod_16
                idx_base = row * c_n + col_base
                
                for ni in range_constexpr(num_acc_n):
                    acc_idx = mi * num_acc_n + ni
                    acc = final_accs[acc_idx]
                    val = vector.extract(acc, static_position=[ii], dynamic_position=[])
                    if is_int8:
                        val = arith.sitofp(T.f32, val)
                    if is_f16_or_bf16 or is_fp4:
                        val_s = val
                    elif _needs_per_token_scale:
                        val_s = (val * s_a) * s_b_vals[ni]
                    else:
                        val_s = val
                    
                    val_f16 = arith.trunc_f(_out_elem(), val_s)
                    idx_out = idx_base + (ni * 16)
                    # offset_is_bytes=False natively manages multiplying byte offsets correctly by precision size
                    buffer_ops.buffer_store(val_f16, c_rsrc, idx_out)

            for _mi in range_constexpr(m_repeat):
                for _ii in range_constexpr(4):
                    # CORRECT ROW SCATTER ORIENTATION
                    _row = bx_m + (_mi * 16) + (lane_div_16 * 4) + _ii
                    body_row(mi=_mi, ii=_ii, row=_row)

        rocdl.sched_barrier(0)

        def hot_loop_scheduler():
            def _build_scheduler(numer: int, denom: int):
                if denom <= 0: return []
                if numer <= 0: return [0] * denom
                out = []
                prev = 0
                for i in range_constexpr(denom):
                    cur = ((i + 1) * numer + (denom - 1)) // denom
                    out.append(cur - prev)
                    prev = cur
                return out

            mfma_group = num_acc_n
            element_k_per_mfma = 128 if _is_gfx950 else 32
            num_mfma_per_tile_k = tile_k // element_k_per_mfma
            mfma_total = num_mfma_per_tile_k * m_repeat * mfma_group
            num_ds_load = num_a_lds_load
            dswr_tail = num_a_loads
            dstr_advance = 2
            if dswr_tail > mfma_total: dswr_tail = mfma_total
            num_gmem_loads = num_b_loads + num_a_async_loads
            if is_fp4 and tile_k != 128:
                num_fp4_scale_k_groups = 1 if int(tile_k) == 128 else (k_unroll // 2)
                num_a_scale_loads = num_fp4_scale_k_groups * (m_repeat // 2)
                num_b_scale_loads = num_fp4_scale_k_groups * (num_acc_n // 2)
                num_gmem_loads += num_a_scale_loads + num_b_scale_loads
            
            dsrd_preload_eff = min(int(dsrd_preload), num_ds_load)
            dvmem_preload_eff = min(int(dvmem_preload), num_gmem_loads)
            vmem_remaining = num_gmem_loads - dvmem_preload_eff
            dsrd_remaining = num_ds_load - dsrd_preload_eff
            if vmem_remaining > 0 and vmem_remaining < mfma_total:
                vmem_schedule = (_build_scheduler(vmem_remaining, vmem_remaining) + [0] * (mfma_total - vmem_remaining))
            else:
                vmem_schedule = _build_scheduler(vmem_remaining, mfma_total)
            dsrd_schedule = _build_scheduler(dsrd_remaining, mfma_total)
            dswr_start = max(mfma_total - dswr_tail - dstr_advance, 0)
            last_dsrd_mfma_idx = -1
            for sched_idx in range_constexpr(mfma_total):
                if dsrd_schedule[sched_idx]: last_dsrd_mfma_idx = sched_idx
            dswr_start = max(dswr_start, last_dsrd_mfma_idx + 1)
            idx_ds_read = dsrd_preload_eff
            idx_gmem_load = dvmem_preload_eff
            idx_ds_write = 0
            if dvmem_preload_eff: rocdl.sched_vmem(dvmem_preload_eff)
            if dsrd_preload_eff: rocdl.sched_dsrd(dsrd_preload_eff)
            for mfma_idx in range_constexpr(mfma_total):
                rocdl.sched_mfma(1)
                n_dsrd = dsrd_schedule[mfma_idx]
                if n_dsrd and (idx_ds_read < num_ds_load):
                    if idx_ds_read + n_dsrd > num_ds_load: n_dsrd = num_ds_load - idx_ds_read
                    if n_dsrd:
                        rocdl.sched_dsrd(n_dsrd)
                        idx_ds_read += n_dsrd
                n_vmem = vmem_schedule[mfma_idx]
                if n_vmem and (idx_gmem_load < num_gmem_loads):
                    if idx_gmem_load + n_vmem > num_gmem_loads: n_vmem = num_gmem_loads - idx_gmem_load
                    if n_vmem:
                        rocdl.sched_vmem(n_vmem)
                        idx_gmem_load += n_vmem
                if (not use_async_copy) and (idx_ds_write < dswr_tail) and (mfma_idx >= dswr_start):
                    rocdl.sched_dswr(1)
                    idx_ds_write += 1
            if (not use_async_copy) and (idx_ds_write < num_a_loads):
                rocdl.sched_dswr(num_a_loads - idx_ds_write)

        rocdl.sched_barrier(0)

        def _flatten_b_tile(bt):
            flat = []
            for packs0, packs1 in bt:
                flat.extend(packs0)
                flat.extend(packs1)
            return flat

        def _unflatten_b_tile(flat):
            bt = []
            idx = 0
            for _ in range_constexpr(k_unroll):
                p0 = [flat[idx + ni] for ni in range_constexpr(num_acc_n)]
                idx += num_acc_n
                p1 = [flat[idx + ni] for ni in range_constexpr(num_acc_n)]
                idx += num_acc_n
                bt.append((p0, p1))
            return bt

        n_accs = num_acc_n * m_repeat
        n_btile = k_unroll * 2 * num_acc_n
        n_a0pf = 2
        if is_fp4:
            n_fp4_asc = _k_unroll_packed_outer * _m_repeat_packed_outer
            n_fp4_bsc = _k_unroll_packed_outer * _num_acc_n_packed_outer

        def _pack_state(accs_l, bt_flat, a0pf, fp4_scales=None):
            state = list(accs_l) + list(bt_flat) + [a0pf[0], a0pf[1]]
            if is_fp4:
                a_scales, b_scales = fp4_scales
                state.extend(a_scales)
                state.extend(b_scales)
            return state

        def _unpack_state(vals):
            accs_l = list(vals[:n_accs])
            bt_flat = list(vals[n_accs:n_accs + n_btile])
            a0pf = (vals[n_accs + n_btile], vals[n_accs + n_btile + 1])
            if not is_fp4: return accs_l, bt_flat, a0pf, None
            sc_base = n_accs + n_btile + n_a0pf
            a_scales = list(vals[sc_base:sc_base + n_fp4_asc])
            b_scales = list(vals[sc_base + n_fp4_asc:sc_base + n_fp4_asc + n_fp4_bsc])
            return accs_l, bt_flat, a0pf, (a_scales, b_scales)

        def _build_pingpong_body(k_iv, inner_state):
            accs_in, bt_flat_in, a0pf_in, fp4_scales_pong_in = _unpack_state(inner_state)
            b_tile_pong_in = _unflatten_b_tile(bt_flat_in)

            if _fp4_tilek128:
                next_k1 = k_iv + tile_k
                if use_async_copy:
                    prefetch_a_to_lds(next_k1, lds_a_ping)
                else:
                    a_tile_ping = prefetch_a_tile(next_k1)
                b_tile_ping = prefetch_b_tile(next_k1)
                accs_in, _ = compute_tile(accs_in, b_tile_pong_in, lds_a_pong, a0_prefetch=a0pf_in, fp4_scales=fp4_scales_pong_in, fp4_scale_half=0)
                if not use_async_copy:
                    store_a_tile_to_lds(a_tile_ping, lds_a_ping)
                hot_loop_scheduler()
                rocdl.s_waitcnt(num_b_loads)
                gpu.barrier()
                a0_prefetch_ping = prefetch_a0_pack(lds_a_ping)
                next_k2 = k_iv + (tile_k * 2)
                _sc_ping = load_fp4_scale_chunk(next_k2) if is_fp4 else None
                rocdl.sched_barrier(0)
                if use_async_copy:
                    prefetch_a_to_lds(next_k2, lds_a_pong)
                else:
                    a_tile_pong = prefetch_a_tile(next_k2)
                b_tile_pong_new = prefetch_b_tile(next_k2)
                accs_in, _ = compute_tile(accs_in, b_tile_ping, lds_a_ping, a0_prefetch=a0_prefetch_ping, fp4_scales=fp4_scales_pong_in, fp4_scale_half=1)
                if not use_async_copy:
                    store_a_tile_to_lds(a_tile_pong, lds_a_pong)
                hot_loop_scheduler()
                rocdl.s_waitcnt(num_b_loads)
                gpu.barrier()
                a0_prefetch_pong_new = prefetch_a0_pack(lds_a_pong)
                return _pack_state(accs_in, _flatten_b_tile(b_tile_pong_new), a0_prefetch_pong_new, _sc_ping)

            next_k1 = k_iv + tile_k
            if use_async_copy:
                prefetch_a_to_lds(next_k1, lds_a_ping)
            else:
                a_tile = prefetch_a_tile(next_k1)
            _sc_ping = load_fp4_scale_chunk(k_iv + fx.Index(tile_k)) if is_fp4 else None
            b_tile_ping = prefetch_b_tile(next_k1)
            accs_in, _ = compute_tile(accs_in, b_tile_pong_in, lds_a_pong, a0_prefetch=a0pf_in, fp4_scales=fp4_scales_pong_in)
            if not use_async_copy:
                store_a_tile_to_lds(a_tile, lds_a_ping)
            hot_loop_scheduler()
            rocdl.s_waitcnt(num_b_loads)
            gpu.barrier()
            a0_prefetch_ping = prefetch_a0_pack(lds_a_ping)

            next_k2 = k_iv + (tile_k * 2)
            if use_async_copy:
                prefetch_a_to_lds(next_k2, lds_a_pong)
            else:
                a_tile = prefetch_a_tile(next_k2)
            _sc_pong = load_fp4_scale_chunk(k_iv + (tile_k * 2)) if is_fp4 else None
            b_tile_pong_new = prefetch_b_tile(next_k2)
            accs_in, _ = compute_tile(accs_in, b_tile_ping, lds_a_ping, a0_prefetch=a0_prefetch_ping, fp4_scales=_sc_ping)
            if not use_async_copy:
                store_a_tile_to_lds(a_tile, lds_a_pong)
            hot_loop_scheduler()
            rocdl.s_waitcnt(num_b_loads)
            gpu.barrier()
            a0_prefetch_pong_new = prefetch_a0_pack(lds_a_pong)
            return _pack_state(accs_in, _flatten_b_tile(b_tile_pong_new), a0_prefetch_pong_new, _sc_pong)

        def prefetch_a0_pack(lds_buffer):
            return lds_load_packs_k64(row_a_lds, col_offset_base_bytes, lds_buffer)

        k0 = fx.Index(0)
        b_tile0 = prefetch_b_tile(k0)
        if use_async_copy:
            prefetch_a_to_lds(k0, lds_a_pong)
        else:
            store_a_tile_to_lds(prefetch_a_tile(k0), lds_a_pong)
        gpu.barrier()
        accs = [acc_init] * n_accs
        a0_prefetch_pong = prefetch_a0_pack(lds_a_pong)
        fp4_scales0 = load_fp4_scale_chunk(fx.Index(0)) if is_fp4 else None

        num_tiles = K // tile_k
        if _fp4_tilek128:
            if (num_tiles % 2) == 1:
                c_k_main = K - tile_k
                init_state = _pack_state(accs, _flatten_b_tile(b_tile0), a0_prefetch_pong, fp4_scales0)
                results = init_state
                for iv, inner in range(0, c_k_main, tile_k * 2, init=init_state):
                    results = yield _build_pingpong_body(iv, inner)
                accs, bt_flat, a0pf, fp4_scales_final = _unpack_state(results)
                b_tile_pong_final = _unflatten_b_tile(bt_flat)
                final_accs, scales = compute_tile(accs, b_tile_pong_final, lds_a_pong, is_last_tile=not is_fp4, a0_prefetch=a0pf, fp4_scales=fp4_scales_final, fp4_scale_half=0)
            else:
                c_k_stop = K - (tile_k * 3)
                init_state = _pack_state(accs, _flatten_b_tile(b_tile0), a0_prefetch_pong, fp4_scales0)
                results = init_state
                for iv, inner in range(0, c_k_stop, tile_k * 2, init=init_state):
                    results = yield _build_pingpong_body(iv, inner)
                accs, bt_flat, a0pf, fp4_scales_ep = _unpack_state(results)
                b_tile_pong_ep = _unflatten_b_tile(bt_flat)
                last_k = arith.index(K - tile_k)
                b_tile_ping = prefetch_b_tile(last_k)
                if use_async_copy:
                    prefetch_a_to_lds(last_k, lds_a_ping)
                else:
                    a_regs_ping = prefetch_a_tile(last_k)
                accs, _ = compute_tile(accs, b_tile_pong_ep, lds_a_pong, a0_prefetch=a0pf, fp4_scales=fp4_scales_ep, fp4_scale_half=0)
                if not use_async_copy:
                    store_a_tile_to_lds(a_regs_ping, lds_a_ping)
                rocdl.s_waitcnt(num_b_loads)
                gpu.barrier()
                a0_prefetch_ping = prefetch_a0_pack(lds_a_ping)
                final_accs, scales = compute_tile(accs, b_tile_ping, lds_a_ping, is_last_tile=not is_fp4, a0_prefetch=a0_prefetch_ping, fp4_scales=fp4_scales_ep, fp4_scale_half=1)
        elif (num_tiles % 2) == 1:
            c_k_main = K - tile_k
            init_state = _pack_state(accs, _flatten_b_tile(b_tile0), a0_prefetch_pong, fp4_scales0)
            results = init_state
            for iv, inner in range(0, c_k_main, tile_k * 2, init=init_state):
                results = yield _build_pingpong_body(iv, inner)
            accs, bt_flat, a0pf, fp4_scales_final = _unpack_state(results)
            b_tile_pong_final = _unflatten_b_tile(bt_flat)
            final_accs, scales = compute_tile(accs, b_tile_pong_final, lds_a_pong, is_last_tile=not is_fp4, a0_prefetch=a0pf, fp4_scales=fp4_scales_final)
        else:
            c_k_stop = K - (tile_k * 3)
            init_state = _pack_state(accs, _flatten_b_tile(b_tile0), a0_prefetch_pong, fp4_scales0)
            results = init_state
            for iv, inner in range(0, c_k_stop, tile_k * 2, init=init_state):
                results = yield _build_pingpong_body(iv, inner)
            accs, bt_flat, a0pf, fp4_scales_ep = _unpack_state(results)
            b_tile_pong_ep = _unflatten_b_tile(bt_flat)

            last_k = arith.index(K - tile_k)
            b_tile_ping = prefetch_b_tile(last_k)
            if use_async_copy:
                prefetch_a_to_lds(last_k, lds_a_ping)
            else:
                a_regs_ping = prefetch_a_tile(last_k)
            _sc_last = load_fp4_scale_chunk(last_k) if is_fp4 else None
            accs, _ = compute_tile(accs, b_tile_pong_ep, lds_a_pong, a0_prefetch=a0pf, fp4_scales=fp4_scales_ep)
            if not use_async_copy:
                store_a_tile_to_lds(a_regs_ping, lds_a_ping)
            hot_loop_scheduler()
            rocdl.s_waitcnt(num_b_loads)
            gpu.barrier()
            a0_prefetch_ping = prefetch_a0_pack(lds_a_ping)
            final_accs, scales = compute_tile(accs, b_tile_ping, lds_a_ping, is_last_tile=not is_fp4, a0_prefetch=a0_prefetch_ping, fp4_scales=_sc_last)
        store_output(final_accs, scales)

    @flyc.jit
    def launch_gemm(arg_c: fx.Tensor, arg_a: fx.Tensor, arg_b: fx.Tensor, arg_scale_a: fx.Tensor, arg_scale_b: fx.Tensor, i32_m: fx.Int32, i32_n: fx.Int32, i32_b_nbytes: fx.Int32, i32_sb_nbytes: fx.Int32):
        allocator_pong.finalized = False
        allocator_ping.finalized = False
        ctx = CompilationContext.get_current()
        with ir.InsertionPoint(ctx.gpu_module_body):
            allocator_pong.finalize()
            allocator_ping.finalize()
        gx = (i32_m + (tile_m - 1)) // tile_m
        gy = i32_n // tile_n
        launcher = kernel_gemm(arg_c, arg_a, arg_b, arg_scale_a, arg_scale_b, i32_m, i32_n, i32_b_nbytes, i32_sb_nbytes)
        if waves_per_eu is not None:
            _wpe = int(waves_per_eu)
            if _wpe >= 1:
                for op in ctx.gpu_module_body.operations:
                    if hasattr(op, 'attributes') and op.OPERATION_NAME == "gpu.func":
                        op.attributes["rocdl.waves_per_eu"] = ir.IntegerAttr.get(T.i32, _wpe)
        
        launcher.launch(grid=(gx, gy, 1), block=(256, 1, 1))

    return launch_gemm

# ---------------------------------------------------------------------------
# Integration Code
# ---------------------------------------------------------------------------

@functools.lru_cache(maxsize=None)
def _compile_cached(pad_m, pad_n, k, tile_m, tile_n, tile_k):
    return compile_preshuffle_gemm_a8(
        M=pad_m,
        N=pad_n,
        K=k,
        tile_m=tile_m,
        tile_n=tile_n,
        tile_k=tile_k,
        in_dtype="fp4",
        out_dtype="bf16",
        lds_stage=2, 
        use_cshuffle_epilog=False,
        waves_per_eu=2,
    )

def get_gemm_kernel(m, n, k):
    """
    Compiles and caches the FlyDSL gemm kernel per shape via _compile_cached.
    Uses tile_n=128 for FP4 (must be multiple of 128) to maximize CU occupancy.
    """
    tile_n = 128  # FP4 constraint: must be multiple of 128; 128 maximizes CU utilization
        
    if m <= 32:
        tile_m = 32
    elif m <= 64:
        tile_m = 64
    else:
        tile_m = 64
        
    tile_k = 256 if k >= 256 and k % 256 == 0 else 128
    
    pad_m = (m + tile_m - 1) // tile_m * tile_m
    pad_n = (n + tile_n - 1) // tile_n * tile_n

    launch_fn = _compile_cached(pad_m, pad_n, k, tile_m, tile_n, tile_k)
    return launch_fn, pad_m, pad_n

# Amortize compile lag on import
try:
    KNOWN_SHAPES = [
        (4, 2880, 512),
        (16, 2112, 7168),
        (32, 4096, 512),
        (32, 2880, 512),
        (64, 7168, 2048),
        (256, 3072, 1536)
    ]
    for _m, _n, _k in KNOWN_SHAPES:
        get_gemm_kernel(_m, _n, _k)
except Exception:
    pass


# ---------------------------------------------------------------------------
# Quant + shuffle preparation (fused by torch.compile)
# ---------------------------------------------------------------------------

def _prepare_a_fn(A):
    """Quantize A to MXFP4 and shuffle scales."""
    A = A.contiguous()
    A_q, A_scale_raw = dynamic_mxfp4_quant(A)
    A_scale_sh = e8m0_shuffle(A_scale_raw)
    return A_q.reshape(-1), A_scale_sh.reshape(-1)

try:
    _prepare_a = torch.compile(_prepare_a_fn, fullgraph=False)
except Exception:
    _prepare_a = _prepare_a_fn

_b_cache = {}


def custom_kernel(data: input_t) -> output_t:
    """
    MXFP4 quant A + FP4 GEMM via FlyDSL mfma_scale_f32_16x16x128_f8f6f4.
    torch.compile fuses quant+shuffle; SRD clamping handles OOB for A and B.
    """
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B.shape[0]

    launch_fn, pad_m, pad_n = get_gemm_kernel(m, n, k)

    # Quant + shuffle fused by torch.compile (no F.pad — SRD clamps A OOB reads to 0)
    A_q_flat, A_scale_flat = _prepare_a(A)

    # Cache B flat views (keyed by data pointer for safety)
    b_key = B_shuffle.data_ptr()
    if b_key not in _b_cache:
        _bf = B_shuffle.view(torch.uint8).reshape(-1)
        _bsf = B_scale_sh.view(torch.uint8).reshape(-1)
        _b_cache[b_key] = (_bf, _bsf, _bf.numel(), _bsf.numel())
    B_sf, B_scf, b_nb, sb_nb = _b_cache[b_key]

    # Output: only m rows needed (SRD drops writes beyond row m)
    C = torch.empty((m, pad_n), dtype=torch.bfloat16, device=A.device)

    # i32_m = actual m (SRD clamping for A reads & C writes); i32_n = pad_n (C stride)
    launch_fn(C.view(-1), A_q_flat, B_sf, A_scale_flat, B_scf, m, pad_n, b_nb, sb_nb)

    return C[:, :n] if pad_n > n else C
scrolls · 1111 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