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
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.
fp4
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.shared-memory
from flydsl.utils.smem_allocator import SmemAllocator, SmemPtrvector-width = int4
is_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 Cscrolls · 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