submission 700494
Yufeng98 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1752 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-700494?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:c7e32c13bab7fad03900785dd16b84084230a0e85388b5bd0f6d97b19b2c5063
license declaredunknown
license concludedunknown
authorsYufeng98
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
fp4, scales = _mxfp4_quant_vcvt(a, BLOCK_K, BLOCK_M, 32)split-k
v571 (S6 fallback): BM=32,BN=32, 4 waves, splitK=1, K-step=512, 3 sync events. Grid=(96,8,1).tile-m = 32
v571 (S6 fallback): BM=32,BN=32, 4 waves, splitK=1, K-step=512, 3 sync events. Grid=(96,8,1).tile-n = 32
v571 (S6 fallback): BM=32,BN=32, 4 waves, splitK=1, K-step=512, 3 sync events. Grid=(96,8,1).Kernel source
submission.py1752 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Fused quant+GEMM for S5/S6 — mainline dispatch variants only.
v571 (S6 fallback): BM=32,BN=32, 4 waves, splitK=1, K-step=512, 3 sync events. Grid=(96,8,1).
v574 (S5 fallback): BM=32,BN=32, 4 waves, splitK=2, K-step=512. Grid=flat with (N,z,m) decode.
v575 (S5 primary): BM=32,BN=32, 4 waves, splitK=2, K-step=1024, single-pass, 3 barriers.
v576 (S6 primary): BM=32,BN=32, 4 waves, splitK=1, K-step=768, 2 sync events. Grid=(96,8,1).
S6 dispatch: v576 -> v571 -> 2-dispatch CK.
S5 dispatch: v575 -> v574 -> 2-dispatch CK.
"""
from task import input_t, output_t
import os
# HIP launch overhead reduction env vars - must be set before any HIP/torch init
os.environ['HIP_FORCE_DEV_KERNARG'] = '1'
os.environ['AMD_DIRECT_DISPATCH'] = '1'
os.environ['TRITON_LLVM_OPT_LEVEL'] = '3'
import gc as _gc
_gc.disable() # reduce GC pauses between evaluator iterations
import torch
torch.set_num_threads(1) # reduce CPU thread contention in hot path
import triton
import triton.language as tl
import subprocess
import tempfile
import sys
import json
import ctypes
# K=7168 dispatch toggle: True -> Triton KSPLIT=7 split-K with y_pp workspace (~10.7us benchmark),
# False -> CK ASM 2-dispatch (fallback, ~21.6us benchmark). Both paths produce max_err=0.0.
_USE_SPLIT_K_S2 = True
# C++ dual-launch for S5(K=2048) and S6(K=1536): routes quant+GEMM through C++ helper .so.
# Measured: hipModuleLaunchKernel ~6us from C++ = same as Python dispatch floor; zero speedup
# vs 2-dispatch CK ASM (v469 S5=14.0us with True vs v471 S5=14.2us with False -- within variance).
# Kept True: marginally faster on warm path due to reduced Python overhead (~0.2us/call).
_USE_FAST_PATH = True # S5/S6 dispatch via C++ launch_fused (single Python->C++ transition)
# ============================================================
# Triton fallback -- unchanged from v211/v435
# ============================================================
try:
from triton._utils import type_canonicalisation_dict
type_canonicalisation_dict['float4_e2m1fn_x2'] = 'u8'
type_canonicalisation_dict['float8_e8m0fnu'] = 'u8'
except Exception:
pass
# AITER version probe
try:
import subprocess as _sp
_aiter_hash = _sp.check_output(
['git', '-C', '/home/runner/aiter', 'rev-parse', '--short', 'HEAD'],
stderr=_sp.DEVNULL, timeout=5).decode().strip()
print(f'[probe] AITER hash: {_aiter_hash}', file=sys.stderr)
# Check for new gemm APIs
try:
import importlib
_aiter_mod = importlib.import_module('aiter')
_gemm_attrs = [a for a in dir(_aiter_mod) if 'gemm' in a.lower() or 'flatmm' in a.lower() or 'fused' in a.lower()]
print(f'[probe] AITER gemm-related attrs: {_gemm_attrs}', file=sys.stderr)
except Exception as _e2:
print(f'[probe] AITER attr scan failed: {_e2}', file=sys.stderr)
except Exception as _e:
print(f'[probe] AITER version check failed: {_e}', file=sys.stderr)
# Report HIP env vars
print(f'[probe] HIP_FORCE_DEV_KERNARG={os.environ.get("HIP_FORCE_DEV_KERNARG", "unset")}', file=sys.stderr)
print(f'[probe] AMD_DIRECT_DISPATCH={os.environ.get("AMD_DIRECT_DISPATCH", "unset")}', file=sys.stderr)
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
K32x128 = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E'
K64x128 = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E'
K192x128 = '_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E'
K256x128 = '_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_256x128E'
# Direct pybind11 bypass: load CK ASM module without @compile_ops wrapper overhead.
# After first call, @compile_ops overhead is minimal (type check cached), but this
# eliminates all Python wrapper logic entirely for the hot path.
_USE_DIRECT_CK = True
_direct_ck_mod = None
def _get_direct_ck():
global _direct_ck_mod
if _direct_ck_mod is None:
try:
from aiter.jit.core import get_module
_direct_ck_mod = get_module("module_gemm_a4w4_asm")
print('[direct_ck] loaded pybind11 module', file=sys.stderr)
except Exception as e:
print(f'[direct_ck] failed: {e}, using wrapped path', file=sys.stderr)
return _direct_ck_mod
@triton.jit
def _mxfp4_quant_vcvt(x, BLOCK_SIZE_K: tl.constexpr, BLOCK_SIZE_M: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr):
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_K // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax_bits = amax.to(tl.uint32, bitcast=True)
amax_rounded = (amax_bits + 0x200000) & 0xFF800000 # round up mantissa, clear it
# Direct uint32 scale computation: fwd_scale = 2^(biased_exp - 2) = amax_rounded >> 2 in float
# Since mantissa=0: (biased_exp - 2) << 23 = amax_rounded - 0x01000000
fwd_scale_bits = tl.where(amax_rounded >= 0x01000000,
amax_rounded - 0x01000000, tl.zeros_like(amax_rounded))
# Handle inf/nan (amax_rounded=0xFF800000): fwd_scale should have exp=254 -> bits=0x7F000000
fwd_scale_bits = tl.where(amax_rounded < 0xFF800000, fwd_scale_bits,
tl.full(amax_rounded.shape, 0x7F000000, dtype=tl.uint32))
fwd_scale = fwd_scale_bits.to(tl.float32, bitcast=True)
bs_e8m0 = (fwd_scale_bits >> 23).to(tl.uint8)
HALF_QBS: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE // 2
x_pairs = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HALF_QBS, 2)
x_u16 = x_pairs.to(tl.uint16, bitcast=True)
x_even, x_odd = tl.split(x_u16)
bf16x2 = x_even.to(tl.uint32) | (x_odd.to(tl.uint32) << 16)
scale_for_pairs = tl.broadcast_to(fwd_scale, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HALF_QBS])
bf16x2_flat = bf16x2.reshape(BLOCK_SIZE_M * NUM_QUANT_BLOCKS * HALF_QBS)
scale_flat = scale_for_pairs.reshape(BLOCK_SIZE_M * NUM_QUANT_BLOCKS * HALF_QBS)
fp4_raw = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2", "=v,v,v",
args=[bf16x2_flat, scale_flat], dtype=tl.uint32, is_pure=True, pack=1)
fp4_bytes = (fp4_raw & 0xFF).to(tl.uint8)
return fp4_bytes.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2), bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _quant_shuffle_kernel(a_ptr, fp4_ptr, scale_sh_ptr, M, K,
stride_am, stride_ak, stride_fm, stride_fk, stride_sm, stride_sk,
scaleN: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr,
EVEN_MK: tl.constexpr = False):
pid_m = tl.program_id(0); pid_k = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_k = pid_k * BLOCK_K + tl.arange(0, BLOCK_K)
if EVEN_MK:
a = tl.load(a_ptr + offs_m[:, None]*stride_am + offs_k[None, :]*stride_ak)
else:
mask = (offs_m[:, None] < M) & (offs_k[None, :] < K)
a = tl.load(a_ptr + offs_m[:, None]*stride_am + offs_k[None, :]*stride_ak, mask=mask, other=0.0)
fp4, scales = _mxfp4_quant_vcvt(a, BLOCK_K, BLOCK_M, 32)
offs_fk = pid_k*(BLOCK_K//2) + tl.arange(0, BLOCK_K//2)
if EVEN_MK:
tl.store(fp4_ptr + offs_m[:, None]*stride_fm + offs_fk[None, :]*stride_fk, fp4)
else:
tl.store(fp4_ptr + offs_m[:, None]*stride_fm + offs_fk[None, :]*stride_fk, fp4,
mask=(offs_m[:, None] < M) & (offs_fk[None, :] < K//2))
NUM_BLOCKS: tl.constexpr = BLOCK_K // 32
kb_idx = (pid_k*NUM_BLOCKS + tl.arange(0, NUM_BLOCKS))[None, :]
m_idx = offs_m[:, None]
d0=m_idx//32; d1=(m_idx%32)//16; d2=m_idx%16
d3=kb_idx//8; d4=(kb_idx%8)//4; d5=kb_idx%4
flat_pos = d0*(scaleN*32) + d3*256 + d5*64 + d2*4 + d4*2 + d1
if EVEN_MK:
tl.store(scale_sh_ptr + (flat_pos//scaleN)*stride_sm + (flat_pos%scaleN)*stride_sk, scales)
else:
tl.store(scale_sh_ptr + (flat_pos//scaleN)*stride_sm + (flat_pos%scaleN)*stride_sk,
scales, mask=(m_idx < M) & (kb_idx < scaleN))
# ============================================================
# hipcc HSACO compilation infrastructure (N3)
# ============================================================
_hsaco_cache = {}
_hip_for_memset = None
def _hip_zero_tensor(t):
"""Zero a GPU tensor using hipMemsetAsync on the null (0) hip handle."""
global _hip_for_memset
if _hip_for_memset is None:
_hip_for_memset = ctypes.CDLL('libamdhip64.so')
nbytes = t.numel() * t.element_size()
_hip_for_memset.hipMemsetAsync(
ctypes.c_void_p(t.data_ptr()), 0, ctypes.c_size_t(nbytes), ctypes.c_void_p(0))
def _resource_gate(hsaco_path, func_name, max_vgpr=128, max_spill=0, max_scratch=0, threads=256):
"""Run clang-offload-bundler + llvm-readelf to check resource usage of compiled HSACO.
Fail-closed on evaluator: rejects if any metadata field is missing or any threshold is exceeded.
Off evaluator (tools absent): skips gate with warning (gate will run on evaluator at submission).
threads: workgroup size for CTA/CU occupancy calculation (default 256).
Returns dict with resource info on PASS, or None on REJECT."""
import os, re
bundler = '/opt/rocm/llvm/bin/clang-offload-bundler'
readelf = '/opt/rocm/llvm/bin/llvm-readelf'
if not os.path.exists(bundler) or not os.path.exists(readelf):
print(f'[gate] llvm tools not found — skipping resource gate (not on evaluator)', file=sys.stderr)
return {'skipped': True}
try:
unbundled = hsaco_path + '.unbundled'
result = subprocess.run(
[bundler, '--unbundle', f'--inputs={hsaco_path}',
'--type=o', f'--outputs={unbundled}',
'--targets=hipv4-amdgcn-amd-amdhsa--gfx950'],
capture_output=True, text=True, timeout=30)
if result.returncode != 0:
print(f'[gate] REJECT {func_name}: unbundle failed: {result.stderr[:200]}', file=sys.stderr)
return None
result = subprocess.run(
[readelf, '--notes', unbundled],
capture_output=True, text=True, timeout=30)
if result.returncode != 0:
print(f'[gate] REJECT {func_name}: readelf failed: {result.stderr[:200]}', file=sys.stderr)
return None
notes = result.stdout
# Dump literal readelf output to stderr for AC-2 evidence archiving.
# Each non-empty line is tagged [gate-readelf-raw] so it survives log capture.
print(f'[gate-readelf-raw] {func_name}:', file=sys.stderr)
for _rl in notes.splitlines():
if _rl.strip():
print(f'[gate-readelf-raw] {_rl}', file=sys.stderr)
def _extract(pattern, text):
m = re.search(pattern, text)
return int(m.group(1)) if m else None
vgpr = _extract(r'\.vgpr_count:\s+(\d+)', notes)
sgpr_spill = _extract(r'\.sgpr_spill_count:\s+(\d+)', notes)
vgpr_spill = _extract(r'\.vgpr_spill_count:\s+(\d+)', notes)
scratch = _extract(r'\.private_segment_fixed_size:\s+(\d+)', notes)
lds = _extract(r'\.group_segment_fixed_size:\s+(\d+)', notes)
info = {'vgpr': vgpr, 'sgpr_spill': sgpr_spill, 'vgpr_spill': vgpr_spill,
'scratch': scratch, 'lds': lds}
print(f'[gate] {func_name}: vgpr={vgpr} sgpr_spill={sgpr_spill} '
f'vgpr_spill={vgpr_spill} scratch={scratch} lds={lds}', file=sys.stderr)
# Fail-closed: reject if any required field is missing from readelf output
for field_name, field_val in [('vgpr_count', vgpr), ('vgpr_spill_count', vgpr_spill),
('sgpr_spill_count', sgpr_spill),
('private_segment_fixed_size', scratch),
('group_segment_fixed_size', lds)]:
if field_val is None:
print(f'[gate] REJECT {func_name}: missing {field_name} in readelf output', file=sys.stderr)
return None
# Threshold checks
if vgpr > max_vgpr:
print(f'[gate] REJECT {func_name}: vgpr={vgpr} > {max_vgpr}', file=sys.stderr)
return None
if vgpr_spill > max_spill:
print(f'[gate] REJECT {func_name}: vgpr_spill={vgpr_spill} > {max_spill}', file=sys.stderr)
return None
if sgpr_spill > max_spill:
print(f'[gate] REJECT {func_name}: sgpr_spill={sgpr_spill} > {max_spill}', file=sys.stderr)
return None
if scratch > max_scratch:
print(f'[gate] REJECT {func_name}: scratch={scratch} > {max_scratch}', file=sys.stderr)
return None
# CTA/CU residency check: MI355X has 4 SIMDs/CU, 1024 VGPRs/SIMD.
# A CTA with `threads` threads has `waves = threads/64` wavefronts.
# Each SIMD runs `ceil(waves/4)` waves from one CTA, using vgpr_count VGPRs each.
# Per-SIMD VGPR cost per CTA = vgpr_count * ceil(waves/4).
# Max CTAs per CU = floor(1024 / per_simd_cost).
waves_per_cta = (threads + 63) // 64
vgprs_per_simd_per_cta = vgpr * ((waves_per_cta + 3) // 4)
max_ctas_by_vgpr = 1024 // vgprs_per_simd_per_cta if vgprs_per_simd_per_cta > 0 else 256
if max_ctas_by_vgpr < 1:
print(f'[gate] REJECT {func_name}: CTA/CU residency by VGPR = {max_ctas_by_vgpr} < 1', file=sys.stderr)
return None
# LDS budget check: lds_per_cta * CTAs_per_CU <= 160KB (163840 bytes)
max_ctas_by_lds = 163840 // lds if lds > 0 else 256
actual_ctas = min(max_ctas_by_vgpr, max_ctas_by_lds)
if actual_ctas < 1:
print(f'[gate] REJECT {func_name}: CTA/CU residency = {actual_ctas} < 1 '
f'(vgpr_limit={max_ctas_by_vgpr}, lds_limit={max_ctas_by_lds})', file=sys.stderr)
return None
info['max_ctas_by_vgpr'] = max_ctas_by_vgpr
info['max_ctas_by_lds'] = max_ctas_by_lds
info['actual_ctas_per_cu'] = actual_ctas
print(f'[gate] PASS {func_name}: CTA/CU={actual_ctas} '
f'(vgpr_limit={max_ctas_by_vgpr}, lds_limit={max_ctas_by_lds}, '
f'threads={threads}, waves/CTA={waves_per_cta})', file=sys.stderr)
return info
except Exception as e:
# Fail-closed: exceptions on evaluator reject the kernel
print(f'[gate] REJECT {func_name}: exception during resource check: {e}', file=sys.stderr)
return None
def _compile_hip_kernel(src, func_name, cache_key=None, resource_gate=True):
"""Compile HIP C++ source to HSACO via hipcc --genco, load via hipModule API.
Returns (hipModule, hipFunction, hipLib) or None on failure.
By default, runs llvm-readelf resource gate after compilation and rejects
kernels that exceed AC-4 thresholds (vgpr>128, spill>0, scratch>0,
CTA/CU<1, LDS budget exceeded). Pass resource_gate=False only for
non-GEMM utility kernels (e.g., native quant).
Caches compiled binary keyed by cache_key."""
if cache_key and cache_key in _hsaco_cache:
return _hsaco_cache[cache_key]
try:
import hashlib, os
# Write source to temp file
src_hash = hashlib.md5(src.encode()).hexdigest()[:8]
cache_dir = '/tmp/mxfp4_hsaco_cache'
os.makedirs(cache_dir, exist_ok=True)
hsaco_path = os.path.join(cache_dir, f'{func_name}_{src_hash}.hsaco')
src_path = os.path.join(cache_dir, f'{func_name}_{src_hash}.hip')
if not os.path.exists(hsaco_path):
with open(src_path, 'w') as f:
f.write(src)
import time as _t
t0 = _t.perf_counter()
# Include CK headers if available (for persistent/preshuffle compilation)
_ck_inc = '/home/runner/aiter/3rdparty/composable_kernel/include'
_hipcc_cmd = ['hipcc', '--genco', '--offload-arch=gfx950', '-O3']
import os as _os_h
if _os_h.path.isdir(_ck_inc):
_hipcc_cmd.extend(['-I', _ck_inc, '-std=c++17'])
_hipcc_cmd.extend(['-o', hsaco_path, src_path])
result = subprocess.run(
_hipcc_cmd,
capture_output=True, text=True, timeout=120)
dt = _t.perf_counter() - t0
if result.returncode != 0:
print(f'[hsaco] compile FAILED ({dt:.1f}s): {result.stderr[:200]}', file=sys.stderr)
return None
# hipcc may exit 0 on compile errors — check HSACO file exists and is non-empty
if not os.path.exists(hsaco_path) or os.path.getsize(hsaco_path) == 0:
print(f'[hsaco] compile produced empty/missing HSACO for {func_name}', file=sys.stderr)
return None
print(f'[hsaco] compiled {func_name} in {dt:.1f}s', file=sys.stderr)
# Resource gate: run llvm-readelf check if requested
if resource_gate:
gate_result = _resource_gate(hsaco_path, func_name)
if gate_result is None:
print(f'[hsaco] {func_name} REJECTED by resource gate', file=sys.stderr)
return None
# Load via ctypes hipModule API
hip = ctypes.CDLL('libamdhip64.so')
module = ctypes.c_void_p()
rc = hip.hipModuleLoad(ctypes.byref(module), hsaco_path.encode())
if rc != 0:
print(f'[hsaco] hipModuleLoad failed rc={rc}', file=sys.stderr)
return None
func = ctypes.c_void_p()
rc = hip.hipModuleGetFunction(ctypes.byref(func), module, func_name.encode())
if rc != 0:
print(f'[hsaco] hipModuleGetFunction failed rc={rc}', file=sys.stderr)
return None
print(f'[hsaco] loaded {func_name}', file=sys.stderr)
result = (module, func, hip)
if cache_key:
_hsaco_cache[cache_key] = result
return result
except Exception as e:
print(f'[hsaco] error: {e}', file=sys.stderr)
return None
# Native MFMA-scale test kernel: minimal kernel to verify hipcc + inline asm on evaluator
# ============================================================
NATIVE_QUANT_SRC = r"""
#include <hip/hip_runtime.h>
static __device__ __forceinline__ unsigned int float_as_uint(float f) {
unsigned int u; __builtin_memcpy(&u, &f, 4); return u;
}
static __device__ __forceinline__ float uint_as_float(unsigned int u) {
float f; __builtin_memcpy(&f, &u, 4); return f;
}
// v515: High-CTA quant kernel — 8 rows × 8 K-blocks per CTA (64 threads = 1 wavefront)
// vs v510: 32 rows × 8 K-blocks (256 threads). 4× CTA count for better CU utilization.
// S5: 16→64 CTAs (6%→25% CU util), S6: 48→192 CTAs (19%→75% CU util)
extern "C" __global__
__attribute__((amdgpu_flat_work_group_size(64,64)))
void native_quant_shuffle(
const unsigned short* __restrict__ A_bf16,
unsigned char* __restrict__ fp4_out,
unsigned char* __restrict__ scale_sh,
int M, int K, int scaleN
) {
int tid = threadIdx.x, t_row = tid >> 3, t_blk = tid & 7;
int m = (int)blockIdx.x * 8 + t_row;
int kb = (int)blockIdx.y * 8 + t_blk;
int k_base = kb * 32;
if (m >= M || k_base >= K) return;
// Load 32 bf16 values (16 pairs) — vectorized as dwordx4 where possible
const unsigned int* src = (const unsigned int*)(A_bf16 + (long)m * K + k_base);
unsigned int ap[16];
#pragma unroll 4
for (int i = 0; i < 4; i++) {
// Load 4 dwords at a time (128 bits = 8 bf16 values)
ap[i*4+0] = src[i*4+0];
ap[i*4+1] = src[i*4+1];
ap[i*4+2] = src[i*4+2];
ap[i*4+3] = src[i*4+3];
}
// Compute amax (absolute max of bf16 magnitudes)
unsigned int ab = 0;
#pragma unroll 16
for (int i = 0; i < 16; i++) {
unsigned int p = ap[i];
unsigned int lo = p & 0x7FFFu, hi = (p >> 16) & 0x7FFFu;
if (lo > ab) ab = lo; if (hi > ab) ab = hi;
}
// Integer exponent extraction (optimized -- no log2f transcendental)
unsigned int amax_bits = ((unsigned int)ab) << 16;
unsigned int rounded = (amax_bits + 0x200000u) & 0xFF800000u;
unsigned int biased_exp = (rounded >> 23) & 0xFFu;
int scale_a;
if (biased_exp <= 1) scale_a = 0;
else if (biased_exp >= 255) scale_a = 254;
else scale_a = (int)biased_exp - 2;
float fs = uint_as_float((unsigned int)scale_a << 23);
// Pack bf16 pairs to fp4 via v_cvt_scalef32_pk_fp4_bf16, store as dwords
unsigned char* dst = fp4_out + (long)m * (K/2) + k_base/2;
#pragma unroll 4
for (int d = 0; d < 4; d++) {
unsigned int dw = 0;
#pragma unroll 4
for (int p2 = 0; p2 < 4; p2++) {
unsigned int fp4b;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(fp4b) : "v"(ap[d*4+p2]), "v"(fs));
dw |= ((fp4b & 0xFFu) << (p2 * 8));
}
// 4-byte store instead of 4× 1-byte stores
*((unsigned int*)(dst + d*4)) = dw;
}
// Store shuffled scale (CK-compatible layout)
int d0=m>>5, d1=(m&31)>>4, d2=m&15;
int d3=kb>>3, d4=(kb&7)>>2, d5=kb&3;
int fp = d0*(scaleN*32) + d3*256 + d5*64 + d2*4 + d4*2 + d1;
scale_sh[(fp/scaleN)*scaleN + (fp%scaleN)] = (unsigned char)scale_a;
}
"""
_native_quant_loaded = None
# Default kernel flags (overridden by included kernel files)
_USE_FUSED_V571 = False
_USE_FUSED_V571B = False
_USE_FUSED_V571C = False
_USE_FUSED_V571DG = False
_USE_FUSED_V571GLDS = False
_USE_FUSED_V571P = False
_USE_FUSED_V571X = False
_USE_FUSED_V574 = False
_USE_FUSED_V574B = False
_USE_FUSED_V575 = False
_USE_FUSED_V575C = False
_USE_FUSED_V575DG = False
_USE_FUSED_V576 = False
def _get_native_quant():
"""Compile and cache the native quant+shuffle HSACO kernel."""
global _native_quant_loaded
if _native_quant_loaded is None:
result = _compile_hip_kernel(NATIVE_QUANT_SRC, 'native_quant_shuffle',
cache_key='nq_v515', resource_gate=False)
if result:
_native_quant_loaded = result
print('[hsaco] native_quant_shuffle ready', file=sys.stderr)
else:
print('[hsaco] native_quant_shuffle compile FAILED', file=sys.stderr)
return _native_quant_loaded
_nqck_profile_count = {}
def _run_native_quant_ck(A, B_shuffle, B_scale_sh, M, N, K, ck_kernel=None):
"""S5/S6 path: native HSACO quant + CK ASM GEMM."""
if ck_kernel is None:
ck_kernel = K32x128
key = ('nq', M, N, K, ck_kernel)
if key not in _qs_cache:
scaleN = triton.cdiv(K, 32); scaleN = triton.cdiv(scaleN, 8)*8
scale_rows = triton.cdiv(M, 32)*32
_qs_cache[key] = {
'fp4': torch.empty((M, K//2), dtype=torch.uint8, device=A.device),
'scale_sh': torch.empty((scale_rows, scaleN), dtype=torch.uint8, device=A.device),
'out': torch.empty((M, N), dtype=torch.bfloat16, device=A.device),
'scaleN': scaleN,
'grid_x': triton.cdiv(M, 8), # v515: 8 rows per CTA (was 32)
'grid_y': triton.cdiv(K, 256), # 8 blocks of 32 per thread = 256
}
sc = _qs_cache[key]
sc['fp4_view'] = sc['fp4'].view(dtypes.fp4x2)
sc['scale_sh_view'] = sc['scale_sh'].view(dtypes.fp8_e8m0)
sc = _qs_cache[key]
nq = _get_native_quant()
if nq:
_mod, _fn, _hip = nq
# Pre-cached arg pointers for quant kernel (avoid ctypes overhead on hot path)
if 'nq_args' not in sc:
_a0 = ctypes.c_void_p(sc['fp4'].data_ptr())
_a1 = ctypes.c_void_p(sc['scale_sh'].data_ptr())
_a2 = ctypes.c_int(M); _a3 = ctypes.c_int(K); _a4 = ctypes.c_int(sc['scaleN'])
sc['nq_a_slot'] = ctypes.c_void_p() # A data_ptr placeholder
sc['nq_args'] = (_a0, _a1, _a2, _a3, _a4)
_nq_all = (sc['nq_a_slot'], _a0, _a1, _a2, _a3, _a4)
sc['nq_arg_ptrs'] = (ctypes.c_void_p * 6)(
*[ctypes.cast(ctypes.pointer(a), ctypes.c_void_p) for a in _nq_all])
sc['nq_a_slot'].value = A.data_ptr()
rc = _hip.hipModuleLaunchKernel(
_fn, sc['grid_x'], sc['grid_y'], 1, 64, 1, 1, # v515: 64 threads (1 wavefront)
0, ctypes.c_void_p(0), sc['nq_arg_ptrs'], ctypes.c_void_p(0))
if rc != 0:
return _run_quant_shuffle_ck_fallback(A, B_shuffle, B_scale_sh, M, N, K, sc)
else:
return _run_quant_shuffle_ck_fallback(A, B_shuffle, B_scale_sh, M, N, K, sc)
# CK ASM GEMM -- direct ctypes launch (bypass pybind11 ~5us overhead)
ck_d = _qs_cache.get('_ck_direct')
if ck_d and 'ck_ka' not in sc:
# Build kernel arg struct (packed, matches CK's _KA layout)
# Fields: D(8) pad(8) C(8) pad(8) A(8) pad(8) B(8) pad(8)
# alpha(4) pad(12) beta(4) pad(12)
# sD0(4) pad(12) sD1(4) pad(12) sC0(4) pad(12) sC1(4) pad(12)
# sA0(4) pad(12) sA1(4) pad(12) sB0(4) pad(12) sB1(4) pad(12)
# M(4) pad(12) N(4) pad(12) K(4) pad(12)
# SA(8) pad(8) SB(8) pad(8)
# sSA0(4) pad(12) sSA1(4) pad(12) sSB0(4) pad(12) sSB1(4) pad(12)
# lks(4)
# Each field with _p2 = 8 bytes pad, _p3 = 12 bytes pad
import struct
kbs = K // 2 # K in fp4x2 bytes
scaleN_ck = sc['scaleN']
# Pre-build the arg struct template
def _mk_ck_args(D_ptr, A_ptr, B_ptr, SA_ptr, SB_ptr, M, N, K):
parts = []
def p8(v): parts.append(struct.pack('<Q', v)); parts.append(b'\x00' * 8)
def u4(v): parts.append(struct.pack('<I', v)); parts.append(b'\x00' * 12)
def f4(v): parts.append(struct.pack('<f', v)); parts.append(b'\x00' * 12)
p8(D_ptr); p8(0); p8(A_ptr); p8(B_ptr) # D, C(unused), A, B
f4(1.0); f4(0.0) # alpha, beta
u4(N); u4(1) # sD0, sD1
u4(N); u4(1) # sC0, sC1
u4(K); u4(1) # sA0, sA1 (K in nibble units = K for fp4x2)
u4(K); u4(1) # sB0, sB1
u4(M); u4(N); u4(K) # M, N, K
p8(SA_ptr); p8(SB_ptr) # SA, SB
u4(scaleN_ck); u4(1) # sSA0, sSA1
u4(scaleN_ck); u4(1) # sSB0, sSB1
parts.append(struct.pack('<i', 0)) # lks
return b''.join(parts)
sc['_mk_ck_args'] = _mk_ck_args
sc['ck_grid_x'] = triton.cdiv(N, 128)
sc['ck_grid_y'] = triton.cdiv(M, 32)
if ck_d and '_mk_ck_args' in sc:
ka_bytes = sc['_mk_ck_args'](
sc['out'].data_ptr(), sc['fp4'].data_ptr(),
B_shuffle.data_ptr(), sc['scale_sh'].data_ptr(),
B_scale_sh.data_ptr(), M, N, K)
ka_buf = (ctypes.c_char * len(ka_bytes))(*ka_bytes)
ka_ptr = ctypes.c_void_p(ctypes.addressof(ka_buf))
ka_size = ctypes.c_size_t(len(ka_bytes))
g_extra = (ctypes.c_void_p * 5)(
ctypes.c_void_p(1), ka_ptr,
ctypes.c_void_p(2), ctypes.cast(ctypes.pointer(ka_size), ctypes.c_void_p),
ctypes.c_void_p(3))
rc = ck_d['hip'].hipModuleLaunchKernel(
ck_d['fn'], sc['ck_grid_x'], sc['ck_grid_y'], 1, 256, 1, 1,
0, ctypes.c_void_p(0), ctypes.c_void_p(0),
ctypes.cast(g_extra, ctypes.c_void_p))
if rc == 0:
if _nqck_profile_count.get(('ck_direct', M, N, K), 0) < 1:
_nqck_profile_count[('ck_direct', M, N, K)] = 1
print(f'[ck_direct] launch OK M={M} N={N} K={K}', file=sys.stderr)
return sc['out']
print(f'[ck_direct] launch failed rc={rc}, falling back', file=sys.stderr)
# Fallback: pybind11 path
mod = _get_direct_ck() if _USE_DIRECT_CK else None
if mod is not None:
mod.gemm_a4w4_asm(sc['fp4_view'], B_shuffle, sc['scale_sh_view'], B_scale_sh,
sc['out'], ck_kernel, None, 1.0, 0.0, True, 0)
else:
gemm_a4w4_asm(sc['fp4_view'], B_shuffle, sc['scale_sh_view'], B_scale_sh,
sc['out'], ck_kernel, bpreshuffle=True, log2_k_split=0)
return sc['out']
def _run_quant_shuffle_ck_fallback(A, B_shuffle, B_scale_sh, M, N, K, sc):
"""Fallback: Triton quant + CK GEMM."""
BM = 32 if M >= 32 else triton.next_power_of_2(M)
BK = max(32, min(256, triton.next_power_of_2(K)))
grid = (triton.cdiv(M, BM), triton.cdiv(K, BK))
EVEN_MK = (M % BM == 0) and (K % BK == 0)
scaleN = sc['scaleN']
_quant_shuffle_kernel[grid](
A, sc['fp4'], sc['scale_sh'], M, K,
A.stride(0), A.stride(1), sc['fp4'].stride(0), sc['fp4'].stride(1),
sc['scale_sh'].stride(0), sc['scale_sh'].stride(1),
scaleN=scaleN, BLOCK_M=BM, BLOCK_K=BK, EVEN_MK=EVEN_MK)
mod = _get_direct_ck() if _USE_DIRECT_CK else None
if mod is not None:
mod.gemm_a4w4_asm(sc['fp4_view'], B_shuffle, sc['scale_sh_view'], B_scale_sh,
sc['out'], K32x128, None, 1.0, 0.0, True, 0)
else:
gemm_a4w4_asm(sc['fp4_view'], B_shuffle, sc['scale_sh_view'], B_scale_sh,
sc['out'], K32x128, bpreshuffle=True, log2_k_split=0)
return sc['out']
_fused_cache = {}
_qs_cache = {}
# ============================================================
# C++ dual-launch fast path: quant + CK GEMM via single C++ call
# ============================================================
# BL-20260319-hip-ck-quant-mismatch: B_q (raw) and B_shuffle (tile-shuffled) are separate;
# this path uses B_shuffle (from evaluator, already CK-compatible).
# BL-20260319-bscale-format-mismatch: evaluator's B_scale_sh is in CK-compatible shuffled
# layout; pass it unchanged (do NOT re-shuffle).
# getpid() guard: re-initialize per-process after fork (evaluator is multi-process).
QUANT_KERNEL_SRC = r"""
#include <hip/hip_runtime.h>
static __device__ __forceinline__ unsigned int float_as_uint(float f) {
unsigned int u; __builtin_memcpy(&u, &f, 4); return u;
}
static __device__ __forceinline__ float uint_as_float(unsigned int u) {
float f; __builtin_memcpy(&f, &u, 4); return f;
}
extern "C" __global__
__attribute__((amdgpu_flat_work_group_size(64,64)))
void quant_shuffle(
const unsigned short* __restrict__ A_bf16,
unsigned char* __restrict__ fp4_out,
unsigned char* __restrict__ scale_sh,
int M, int K, int scaleN
) {
int tid = threadIdx.x, t_row = tid >> 3, t_blk = tid & 7;
int m = (int)blockIdx.x * 8 + t_row;
int kb = (int)blockIdx.y * 8 + t_blk;
int k_base = kb * 32;
if (m >= M || k_base >= K) return;
const unsigned int* src = (const unsigned int*)(A_bf16 + (long)m * K + k_base);
unsigned int ap[16];
#pragma unroll 16
for (int i = 0; i < 16; i++) ap[i] = src[i];
unsigned int ab = 0;
#pragma unroll 16
for (int i = 0; i < 16; i++) {
unsigned int p = ap[i];
unsigned int lo = p & 0x7FFFu, hi = (p >> 16) & 0x7FFFu;
if (lo > ab) ab = lo; if (hi > ab) ab = hi;
}
int scale_a; float fs;
float amax_f = uint_as_float(((unsigned int)ab) << 16);
unsigned int abits = float_as_uint(amax_f);
if (abits == 0) { scale_a = 0; } else {
unsigned int rb = (abits + 0x200000u) & 0xFF800000u;
scale_a = (int)__builtin_floorf(__builtin_log2f(uint_as_float(rb))) - 2 + 127;
if (scale_a < 0) scale_a = 0; if (scale_a > 254) scale_a = 254;
}
fs = uint_as_float((unsigned int)scale_a << 23);
unsigned char* dst = fp4_out + (long)m * (K/2) + k_base/2;
#pragma unroll 4
for (int d = 0; d < 4; d++) {
unsigned int dw = 0;
#pragma unroll 4
for (int p2 = 0; p2 < 4; p2++) {
unsigned int fp4b;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(fp4b) : "v"(ap[d*4+p2]), "v"(fs));
dw |= ((fp4b & 0xFFu) << (p2 * 8));
}
*((unsigned int*)(dst + d*4)) = dw;
}
int d0=m>>5, d1=(m&31)>>4, d2=m&15;
int d3=kb>>3, d4=(kb&7)>>2, d5=kb&3;
int fp = d0*(scaleN*32) + d3*256 + d5*64 + d2*4 + d4*2 + d1;
scale_sh[(fp/scaleN)*scaleN + (fp%scaleN)] = (unsigned char)scale_a;
}
"""
LAUNCH_HELPER_SRC = r"""
// Forward declarations only -- no hip/hip_runtime.h to avoid hipcc device compilation.
// hipcc ALWAYS adds -x hip --offload-arch=gfx950; g++ + forward decls avoids this.
#include <cstring>
#include <cstdio>
#include <time.h>
typedef void* _hm; typedef void* _hf; typedef int _he; typedef void* _hq;
typedef void* _hg; typedef void* _hge; typedef void* _hgn;
struct _d3 { unsigned x,y,z; }; // matches dim3
struct _hknp {
_d3 blockDim;
void** extra;
void* func;
_d3 gridDim;
void** kernelParams;
unsigned sharedMemBytes;
};
extern "C" {
_he hipModuleLoad(_hm* m, const char* f);
_he hipModuleGetFunction(_hf* fn, _hm m, const char* n);
_he hipModuleLaunchKernel(_hf fn,
unsigned gx,unsigned gy,unsigned gz,unsigned bx,unsigned by,unsigned bz,
unsigned shm,_hq q,void** kp,void** ex);
_he hipMalloc(void** p, unsigned long long n);
_he hipGraphCreate(_hg* g, unsigned flags);
_he hipGraphAddKernelNode(_hgn* n, _hg g, const _hgn* deps, unsigned long long nd, const _hknp* p);
_he hipGraphInstantiate(_hge* ge, _hg g, void* errn, char* log, unsigned long long logsz);
_he hipGraphLaunch(_hge ge, _hq s);
_he hipGraphExecKernelNodeSetParams(_hge ge, _hgn n, const _hknp* p);
_he hipGraphDestroy(_hg g);
}
struct _p2{unsigned _0,_1;}; struct _p3{unsigned _0,_1,_2;};
struct __attribute__((packed)) _KA {
void*_D;_p2 _0;void*_C;_p2 _1;void*_A;_p2 _2;void*_B;_p2 _3;
float al;_p3 _4;float be;_p3 _5;
unsigned sD0;_p3 _6;unsigned sD1;_p3 _7;
unsigned sC0;_p3 _8;unsigned sC1;_p3 _9;
unsigned sA0;_p3 _a;unsigned sA1;_p3 _b;
unsigned sB0;_p3 _c;unsigned sB1;_p3 _d;
unsigned M;_p3 _e;unsigned N;_p3 _f;unsigned K;_p3 _g;
void*_SA;_p2 _h;void*_SB;_p2 _i;
unsigned sSA0;_p3 _j;unsigned sSA1;_p3 _k;
unsigned sSB0;_p3 _l;unsigned sSB1;_p3 _m;
int lks;
};
struct __attribute__((packed)) _QA {
const unsigned short* A; unsigned char* q; unsigned char* s;
int M,K,sN;
};
static _hm _qm=0,_gm=0; static _hf _qf=0,_gf=0;
static bool _rdy=false;
static int _cdiv(int a,int b){return(a+b-1)/b;}
static int _n_traced_s5=0,_n_traced_s6=0;
// Per-(M,K) shape cache with ownership metadata.
// Invariants: device=0 (evaluator single-GPU), queue=default (evaluator blocks others),
// dtype=uint8 (FP4 packed + E8M0 scale). Enforced at init_fused time.
struct _ShapeEntry {
unsigned long long ab; // device FP4 scratch: M*K/2 bytes
unsigned long long sb; // device scale scratch: sN * cdiv(M,32)*32 bytes
int M, K, sN; // shape params + precomputed sN = cdiv(cdiv(K,32),8)*8
int ab_bytes, sb_bytes; // exact allocation sizes
unsigned qgx, qgy; // precomputed quant grid
unsigned ggx, ggy; // precomputed GEMM grid (for given N)
int N_for_ggx; // N used to precompute ggx (0 = not yet set)
// Precomputed immutable _KA fields (A-side only; D/B/SB set per call)
_KA ka_template;
bool valid;
_hge gExec; _hgn qN,gN; int gN_val;
_QA qa_g; _KA ka_g; unsigned long long qa_gs,ka_gs;
void* qe[5]; void* ge[5];
};
#define MAX_ENTRIES 4
static _ShapeEntry _entries[MAX_ENTRIES];
static int _num_entries=0;
static _ShapeEntry* _find_entry(int M, int K) {
for(int i=0;i<_num_entries;i++)
if(_entries[i].valid && _entries[i].M==M && _entries[i].K==K) return &_entries[i];
return 0;
}
static _ShapeEntry* _alloc_entry(int M, int K) {
if(_num_entries>=MAX_ENTRIES) return 0;
_ShapeEntry& e=_entries[_num_entries++];
memset(&e,0,sizeof(e));
e.M=M; e.K=K;
e.sN=_cdiv(_cdiv(K,32),8)*8;
e.ab_bytes=M*(K/2);
e.sb_bytes=e.sN*(_cdiv(M,32)*32);
e.qgx=_cdiv(M,8); e.qgy=_cdiv(K,256);
e.ggy=_cdiv(M,32); e.N_for_ggx=0;
// Precompute immutable _KA template (A-side fields)
memset(&e.ka_template,0,sizeof(_KA));
e.ka_template.al=1.f;
e.ka_template.sA0=K; // nibble units
e.ka_template.sB0=K;
e.ka_template.M=M;
e.ka_template.K=K;
e.ka_template.sSA0=e.sN;
_he r=hipMalloc((void**)&e.ab,e.ab_bytes);
if(r){fprintf(stderr,"[fp]ab:%d\n",r); return 0;}
r=hipMalloc((void**)&e.sb,e.sb_bytes);
if(r){fprintf(stderr,"[fp]sb:%d\n",r); return 0;}
e.valid=true;
fprintf(stderr,"[fp_cache] (%d,%d) ab=%dB sb=%dB sN=%d qgrid=(%d,%d)\n",
M,K,e.ab_bytes,e.sb_bytes,e.sN,e.qgx,e.qgy);
return &e;
}
extern "C" {
int init_fused(const char* qco, const char* gco) {
if (_rdy) return 0;
_he r=hipModuleLoad(&_qm,qco);
if(r){fprintf(stderr,"[fp]qm:%d\n",r);return r;}
r=hipModuleGetFunction(&_qf,_qm,"quant_shuffle");
if(r){fprintf(stderr,"[fp]qf:%d\n",r);return r;}
r=hipModuleLoad(&_gm,gco);
if(r){fprintf(stderr,"[fp]gm:%d\n",r);return r;}
r=hipModuleGetFunction(&_gf,_gm,
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E");
if(r){fprintf(stderr,"[fp]gf:%d\n",r);return r;}
// Pre-allocate entries for known shapes
if(!_alloc_entry(64,2048)) return -10; // S5
if(!_alloc_entry(256,1536)) return -11; // S6
_rdy=true; return 0;
}
// Quant-only: launch quant_shuffle, copy raw A_q and A_scale_sh back to host-visible buffers
int launch_quant_only(void*A, void*out_aq, void*out_as, int M, int K, void*hip_q) {
if(!_rdy) return -1;
_ShapeEntry* e=_find_entry(M,K);
if(!e){ e=_alloc_entry(M,K); if(!e) return -2; }
_hq q=(_hq)hip_q;
_QA qa; qa.A=(const unsigned short*)A;
qa.q=(unsigned char*)e->ab; qa.s=(unsigned char*)e->sb;
qa.M=M; qa.K=K; qa.sN=e->sN;
unsigned long long qsz=sizeof(qa);
void*qcfg[]={(void*)1,&qa,(void*)2,&qsz,(void*)3};
_he r=hipModuleLaunchKernel(_qf,e->qgx,e->qgy,1,64,1,1,0,q,0,(void**)qcfg);
if(r) return r;
// Sync + copy raw quant bytes back to caller buffers (for byte audit)
// hipDeviceSynchronize forward decl not available; use hipMemcpy D2H which syncs
// hipMemcpy is available via libamdhip64
typedef _he(*_hmc)(void*,const void*,unsigned long long,int);
static _hmc _memcpy=0;
if(!_memcpy){
void*h=0;
// dlsym from already-loaded libamdhip64
typedef void*(*_dls)(void*,const char*);
_dls ds=(_dls)__builtin_expect((long)0,0); // can't dlsym without dlfcn
// Fallback: just write device pointers and let Python do hipMemcpyDtoH
}
// Write device pointers to output for Python-side copy
*(unsigned long long*)out_aq = e->ab;
*(unsigned long long*)out_as = e->sb;
return 0;
}
int launch_fused(void*A,void*Bs,void*Bsc,void*C,int M,int N,int K,int bss,void*hip_q){
if(!_rdy)return -1;
_ShapeEntry* e=_find_entry(M,K);
if(!e){ e=_alloc_entry(M,K); if(!e) return -2; }
_hq q=(_hq)hip_q;
_QA qa; qa.A=(const unsigned short*)A;
qa.q=(unsigned char*)e->ab; qa.s=(unsigned char*)e->sb;
qa.M=M; qa.K=K; qa.sN=e->sN;
unsigned long long qsz=sizeof(qa);
void*qcfg[]={(void*)1,&qa,(void*)2,&qsz,(void*)3};
int *_tcnt=(K==2048)?&_n_traced_s5:(K==1536?&_n_traced_s6:NULL);
int do_tr=(_tcnt && *_tcnt<10); struct timespec t0,t1,t2;
if(do_tr)clock_gettime(CLOCK_MONOTONIC,&t0);
_he r=hipModuleLaunchKernel(_qf,e->qgx,e->qgy,1,64,1,1,0,q,0,(void**)qcfg);
if(r)return r;
if(do_tr)clock_gettime(CLOCK_MONOTONIC,&t1);
// Build GEMM args from precomputed template + per-call fields
_KA ka=e->ka_template; // copy precomputed A-side fields
ka._D=C; ka._A=(void*)e->ab; ka._B=Bs;
ka._SA=(void*)e->sb; ka._SB=Bsc;
ka.sD0=N; ka.sC0=N; ka.N=N; ka.sSB0=bss;
unsigned long long gsz=sizeof(ka);
void*gcfg[]={(void*)1,&ka,(void*)2,&gsz,(void*)3};
// Precompute GEMM grid if N changed
if(e->N_for_ggx!=N){ e->ggx=_cdiv(N,128); e->N_for_ggx=N; }
_he ret=(int)hipModuleLaunchKernel(_gf,e->ggx,e->ggy,1,256,1,1,0,q,0,(void**)gcfg);
if(do_tr){
clock_gettime(CLOCK_MONOTONIC,&t2);
long nq=(t1.tv_sec-t0.tv_sec)*1000000000L+(t1.tv_nsec-t0.tv_nsec);
long ng=(t2.tv_sec-t1.tv_sec)*1000000000L+(t2.tv_nsec-t1.tv_nsec);
fprintf(stderr,"[fp_trace] (%d,%d,%d) call=%d quant_ns=%ld gemm_ns=%ld\n",M,N,K,*_tcnt,nq,ng);
(*_tcnt)++;
}
return ret;
}
extern "C" { _he hipGetLastError(void); _he hipDeviceSynchronize(void); }
// Graph uses kernelParams (array of arg pointers) for both AddKernelNode and SetParams.
// Quant kernel takes _QA struct as single arg; GEMM takes _KA struct as single arg.
// kernelParams = &[&struct], so one pointer per arg.
int build_gr(_ShapeEntry* e,int N,void*A,void*Bs,void*Bsc,void*C,int bss){
if(!_rdy)return-1;
if(e->gExec) e->gExec=0; // leak old (max 4 shapes)
_hg g=0;_he r=hipGraphCreate(&g,0);
if(r){fprintf(stderr,"[gr]create:%d\n",r);return r;}
// Initialize persistent quant args
e->qa_g.A=(const unsigned short*)A;
e->qa_g.q=(unsigned char*)e->ab;e->qa_g.s=(unsigned char*)e->sb;
e->qa_g.M=e->M;e->qa_g.K=e->K;e->qa_g.sN=e->sN;
// kernelParams: array of pointers to each arg field
e->qe[0]=&e->qa_g.A; e->qe[1]=&e->qa_g.q; e->qe[2]=&e->qa_g.s;
e->qe[3]=&e->qa_g.M; e->qe[4]=&e->qa_g.K;
// Note: _QA has 6 fields but we only have 5 slots in qe[5]
// Need to expand. Actually _QA = {A, q, s, M, K, sN} = 6 fields
// Use qa_gs slot to hold sN pointer
_hknp qp;memset(&qp,0,sizeof(qp));
qp.blockDim.x=64;qp.blockDim.y=1;qp.blockDim.z=1;
// Use extra format for AddKernelNode (works), then try SetParams with extra too
e->qa_gs=sizeof(_QA);
void* q_extra[]={(void*)1,&e->qa_g,(void*)2,&e->qa_gs,(void*)3};
qp.extra=q_extra; qp.func=_qf;
qp.gridDim.x=e->qgx;qp.gridDim.y=e->qgy;qp.gridDim.z=1;
r=hipGraphAddKernelNode(&e->qN,g,0,0,&qp);
if(r){fprintf(stderr,"[gr]addQ:%d\n",r);hipGraphDestroy(g);return r;}
// Initialize persistent GEMM args
e->ka_g=e->ka_template;
e->ka_g._D=C;e->ka_g._A=(void*)e->ab;e->ka_g._B=Bs;
e->ka_g._SA=(void*)e->sb;e->ka_g._SB=Bsc;
e->ka_g.sD0=N;e->ka_g.sC0=N;e->ka_g.N=N;e->ka_g.sSB0=bss;
e->ka_gs=sizeof(_KA);
void* g_extra[]={(void*)1,&e->ka_g,(void*)2,&e->ka_gs,(void*)3};
_hknp gp;memset(&gp,0,sizeof(gp));
gp.blockDim.x=256;gp.blockDim.y=1;gp.blockDim.z=1;
gp.extra=g_extra; gp.func=_gf;
gp.gridDim.x=_cdiv(N,128);gp.gridDim.y=_cdiv(e->M,32);gp.gridDim.z=1;
_hgn dep[1]={e->qN};
r=hipGraphAddKernelNode(&e->gN,g,dep,1,&gp);
if(r){fprintf(stderr,"[gr]addG:%d\n",r);hipGraphDestroy(g);return r;}
r=hipGraphInstantiate(&e->gExec,g,0,0,0);
hipGraphDestroy(g);
if(r){fprintf(stderr,"[gr]inst:%d\n",r);return r;}
e->gN_val=N;
fprintf(stderr,"[gr]built(%d,%d,%d)\n",e->M,N,e->K);
return 0;
}
// Launch: update persistent struct fields, try SetParams, then launch.
// If SetParams fails (rc=1 with extra format), rebuild the graph.
int launch_gr(void*A,void*Bs,void*Bsc,void*C,int M,int N,int K,int bss,void*hq){
if(!_rdy)return-1;
_ShapeEntry*e=_find_entry(M,K);if(!e)return-2;
if(!e->gExec||e->gN_val!=N){
int r=build_gr(e,N,A,Bs,Bsc,C,bss);
if(r){hipGetLastError();hipDeviceSynchronize();return r;}
}
// Update mutable args in persistent structs
e->qa_g.A=(const unsigned short*)A;
e->ka_g._D=C;
// Try SetParams with kernelParams format (single struct arg)
void* q_kp[1]={&e->qa_g}; // kernelParams: pointer to the struct
_hknp qp;memset(&qp,0,sizeof(qp));
qp.blockDim.x=64;qp.blockDim.y=1;qp.blockDim.z=1;
qp.func=_qf;qp.kernelParams=q_kp;
qp.gridDim.x=e->qgx;qp.gridDim.y=e->qgy;qp.gridDim.z=1;
_he r=hipGraphExecKernelNodeSetParams(e->gExec,e->qN,&qp);
if(r==0){
void* g_kp[1]={&e->ka_g};
_hknp gp;memset(&gp,0,sizeof(gp));
gp.blockDim.x=256;gp.blockDim.y=1;gp.blockDim.z=1;
gp.func=_gf;gp.kernelParams=g_kp;
gp.gridDim.x=_cdiv(N,128);gp.gridDim.y=_cdiv(e->M,32);gp.gridDim.z=1;
r=hipGraphExecKernelNodeSetParams(e->gExec,e->gN,&gp);
if(r){fprintf(stderr,"[gr]setG:%d\n",r);hipGetLastError();hipDeviceSynchronize();}
else{fprintf(stderr,"[gr]setOK\n");}
} else {
// SetParams failed -- rebuild graph with new pointers
fprintf(stderr,"[gr]setQ_kp:%d,rebuild\n",r);
hipGetLastError();
r=build_gr(e,N,A,Bs,Bsc,C,bss);
if(r){hipGetLastError();hipDeviceSynchronize();return r;}
}
r=(int)hipGraphLaunch(e->gExec,(_hq)hq);
if(r){fprintf(stderr,"[gr]launch:%d\n",r);hipGetLastError();hipDeviceSynchronize();return r;}
return 0;
}
// Export kernel function handles for HSA AQL dispatch
void get_kernel_handles(unsigned long long* qf_out, unsigned long long* gf_out,
int* qa_sz_out, int* ga_sz_out) {
*qf_out = (unsigned long long)_qf;
*gf_out = (unsigned long long)_gf;
*qa_sz_out = (int)sizeof(_QA);
*ga_sz_out = (int)sizeof(_KA);
}
// HSA AQL launch: build kernarg from existing _QA/_KA logic, dispatch via external HSA helper
// Returns: 0=success, <0=error
typedef int(*_hsa_launch_fn)(void*,int,void*,int,unsigned,unsigned,unsigned,unsigned);
static _hsa_launch_fn _hsa_launch=0;
void set_hsa_launch_fn(void* fn) { _hsa_launch = (_hsa_launch_fn)fn; }
int hsa_launch_fused(void*A,void*Bs,void*Bsc,void*C,int M,int N,int K,int bss) {
if(!_rdy || !_hsa_launch) return -1;
_ShapeEntry*e=_find_entry(M,K);if(!e){e=_alloc_entry(M,K);if(!e)return-2;}
// Build quant kernarg
_QA qa; qa.A=(const unsigned short*)A;
qa.q=(unsigned char*)e->ab; qa.s=(unsigned char*)e->sb;
qa.M=M; qa.K=K; qa.sN=e->sN;
// Build GEMM kernarg
_KA ka=e->ka_template;
ka._D=C; ka._A=(void*)e->ab; ka._B=Bs;
ka._SA=(void*)e->sb; ka._SB=Bsc;
ka.sD0=N; ka.sC0=N; ka.N=N; ka.sSB0=bss;
if(e->N_for_ggx!=N){e->ggx=_cdiv(N,128);e->N_for_ggx=N;}
// Dispatch via HSA
return _hsa_launch(&qa,(int)sizeof(qa),&ka,(int)sizeof(ka),
e->qgx,e->qgy,e->ggx,e->ggy);
}
// Return device pointers for raw quant bytes (for byte audit from Python)
void get_quant_ptrs(int M, int K, unsigned long long* out_ab, unsigned long long* out_sb,
int* out_ab_bytes, int* out_sb_bytes) {
_ShapeEntry* e=_find_entry(M,K);
if(e && e->valid){
*out_ab=e->ab; *out_sb=e->sb;
*out_ab_bytes=e->ab_bytes; *out_sb_bytes=e->sb_bytes;
} else {
*out_ab=0; *out_sb=0; *out_ab_bytes=0; *out_sb_bytes=0;
}
}
}
"""
_fast_path_state = {'pid': -1, 'ready': False, 'failed': False, 'lib': None, 'lib_k64': None}
# Output cache: keyed by (M,N,K,device_index,dtype_num). Pre-populated during warm-up
# for S5 (64,7168,2048) and S6 (256,3072,1536) to avoid torch.empty in the hot path.
# Single-device constraint: evaluator has 1 GPU (device_index=0).
# Single-queue constraint: evaluator uses default queue only; no concurrent access.
# dtype constraint: output is always bfloat16 (dtype_num=15).
_fast_path_out = {}
_fp_trace_count = {} # per (M,N,K) -- Python-side ctypes call timing
_fast_path_warmed = False
_shape_call_counts = {} # per-shape call counter for cold-start instrumentation
def _ensure_fast_path():
"""Lazy per-process init of C++ dual-launch .so. getpid() guard prevents
using stale GPU handles inherited across fork.
Two-step compile: QUANT_KERNEL_SRC --genco -> quant.co,
LAUNCH_HELPER_SRC g++ -shared -> helper.so (forward decls, no HIP headers)."""
import os as _os
pid = _os.getpid()
if _fast_path_state['pid'] != pid:
_fast_path_state.update({'pid': pid, 'ready': False, 'failed': False, 'lib': None, 'lib_k64': None})
if _fast_path_state['ready']:
return True
if _fast_path_state['failed']:
return False
try:
import time as _time
d = tempfile.mkdtemp(prefix='v542fp_')
_ck64_path = '/home/runner/aiter/hsa/gfx950/f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128.co'
_ck32_path = '/home/runner/aiter/hsa/gfx950/f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co'
_use_k64 = os.path.exists(_ck64_path)
print(f'[K64_TILE] {"64x128 FOUND -> M>=64 uses K64x128, M<64 uses K32x128" if _use_k64 else "64x128 absent -> K32x128 only"}', file=sys.stderr)
# Compile device kernel -> quant.co (shared by both helpers)
qhip = os.path.join(d, 'quant.hip'); qco = os.path.join(d, 'quant.co')
with open(qhip, 'w') as f: f.write(QUANT_KERNEL_SRC)
t0 = _time.time()
p1 = subprocess.run(
['hipcc', '--genco', '--offload-arch=gfx950', '-O2', '-o', qco, qhip],
capture_output=True, text=True, timeout=120)
t1 = round(_time.time() - t0, 1)
if p1.returncode != 0:
print(f'[fp] quant --genco err ({t1}s): {p1.stderr[-300:]}', file=sys.stderr)
_fast_path_state['failed'] = True; return False
# Helper compile helper: g++ -shared, forward decls, no HIP headers
def _compile_so(src_code, cpp_path, so_path):
with open(cpp_path, 'w') as _f: _f.write(src_code)
return subprocess.run(
['g++', '-shared', '-fPIC', '-O2', '-std=c++17',
'-L/opt/rocm/lib', '-lamdhip64', '-Wl,-rpath,/opt/rocm/lib',
'-o', so_path, cpp_path],
capture_output=True, text=True, timeout=60)
t2 = _time.time()
# Always compile K32x128 helper (safe fallback for all M)
r32 = _compile_so(LAUNCH_HELPER_SRC,
os.path.join(d, 'helper_k32.cpp'), os.path.join(d, 'helper_k32.so'))
ct = round(_time.time() - t2, 1)
if r32.returncode != 0:
print(f'[fp] k32 helper compile err ({ct}s): {r32.stderr[-300:]}', file=sys.stderr)
_fast_path_state['failed'] = True; return False
# Pre-load HIP runtime with RTLD_GLOBAL so forward-decl symbols available
for _hp in ['libamdhip64.so.2', 'libamdhip64.so',
'/opt/rocm/lib/libamdhip64.so', '/opt/rocm-7.1.0/lib/libamdhip64.so']:
try: ctypes.CDLL(_hp, ctypes.RTLD_GLOBAL); break
except Exception: pass
def _setup_lib(so_path, ck_co_path):
_lib = ctypes.CDLL(so_path)
_lib.init_fused.argtypes = [ctypes.c_char_p, ctypes.c_char_p]
_lib.init_fused.restype = ctypes.c_int
_lib.launch_fused.argtypes = ([ctypes.c_void_p]*4+[ctypes.c_int]*4+[ctypes.c_void_p])
_lib.launch_fused.restype = ctypes.c_int
_lib.get_quant_ptrs.argtypes = [ctypes.c_int, ctypes.c_int,
ctypes.POINTER(ctypes.c_ulonglong), ctypes.POINTER(ctypes.c_ulonglong),
ctypes.POINTER(ctypes.c_int), ctypes.POINTER(ctypes.c_int)]
_lib.get_quant_ptrs.restype = None
_lib.launch_gr.argtypes = ([ctypes.c_void_p]*4+[ctypes.c_int]*4+[ctypes.c_void_p])
_lib.launch_gr.restype = ctypes.c_int
_rc = _lib.init_fused(qco.encode(), ck_co_path.encode())
return _lib if _rc == 0 else None
lib_k32 = _setup_lib(os.path.join(d, 'helper_k32.so'), _ck32_path)
if lib_k32 is None:
print(f'[fp] k32 init_fused failed', file=sys.stderr)
_fast_path_state['failed'] = True; return False
# If K64x128 available, compile K64x128 helper for M>=64 (K32 still used for M<64)
lib_k64 = None
if _use_k64:
# K64x128 helper: replace kernel symbol AND GEMM grid Y (M_per_block 32->64)
_k64_src = LAUNCH_HELPER_SRC.replace(K32x128, K64x128).replace(
'e.ggy=_cdiv(M,32)', 'e.ggy=_cdiv(M,64)')
r64 = _compile_so(_k64_src,
os.path.join(d, 'helper_k64.cpp'), os.path.join(d, 'helper_k64.so'))
if r64.returncode == 0:
lib_k64 = _setup_lib(os.path.join(d, 'helper_k64.so'), _ck64_path)
if lib_k64 is None:
print(f'[fp] k64 init_fused failed, k32 only', file=sys.stderr)
else:
print(f'[fp] k64 compile err, k32 only: {r64.stderr[-200:]}', file=sys.stderr)
_fast_path_state['lib'] = lib_k32
_fast_path_state['lib_k64'] = lib_k64
_fast_path_state['ready'] = True
k64_tag = '(+k64)' if lib_k64 is not None else ''
print(f'[fp] ready (genco={t1}s helper={ct}s{k64_tag})', file=sys.stderr)
return True
except Exception as exc:
print(f'[fp] init error: {exc}', file=sys.stderr)
_fast_path_state['failed'] = True; return False
def _run_fast_path(A, B_shuffle, B_scale_sh, M, N, K):
"""Fused quant+GEMM via C++ dual-launch .so. Returns BF16 output or None on failure.
Scratch buffers: _ab (256KB) and _sb (16KB) are allocated once in init_fused() via
hipMalloc. They serve as BOTH the quant output and CK GEMM input -- no intermediate
buffer exists. Sized for max shape: S6 (256,3072,1536) needs 192KB FP4 + 12KB scale.
Ownership: single-device (evaluator device 0), default queue only,
fixed dtype (uint8 FP4 packed + uint8 E8M0 scale). No concurrent access.
"""
if not _ensure_fast_path():
return None
# K64x128 confirmed DEAD END (iter9): fewer CTAs -> lower BW utilization -> regression.
# Always use K32x128 (lib). lib_k64 compiled for diagnostic only, never routed to.
lib = _fast_path_state['lib']
import time as _time
# Output cache: keyed by (M,N,K,device_index,dtype). Pre-populated during warm-up.
dev_idx = A.device.index if A.device.index is not None else 0
key = (M, N, K, dev_idx, torch.bfloat16)
if key not in _fast_path_out:
_fast_path_out[key] = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
out = _fast_path_out[key]
try:
cnt = _fp_trace_count.get((M, N, K), 0)
do_py_trace = (cnt < 10)
if do_py_trace:
t0 = _time.perf_counter_ns()
rc = lib.launch_fused(
ctypes.c_void_p(A.data_ptr()),
ctypes.c_void_p(B_shuffle.data_ptr()),
ctypes.c_void_p(B_scale_sh.data_ptr()),
ctypes.c_void_p(out.data_ptr()),
ctypes.c_int(M), ctypes.c_int(N), ctypes.c_int(K),
ctypes.c_int(int(B_scale_sh.stride(0))),
ctypes.c_void_p(0)) # default execution context
if do_py_trace:
ctypes_ns = _time.perf_counter_ns() - t0
print(f'[fp_py] ({M},{N},{K}) call={cnt} ctypes_ns={ctypes_ns}', file=sys.stderr)
_fp_trace_count[(M, N, K)] = cnt + 1
if rc != 0:
print(f'[fp] launch_fused rc={rc}', file=sys.stderr)
return None
return out
except Exception as exc:
print(f'[fp] launch error: {exc}', file=sys.stderr)
return None
@triton.jit
def _custom_reduce_kernel(y_pp_ptr, y_ptr, M, N,
stride_pp_k, stride_pp_m, stride_pp_n, stride_ym, stride_yn,
TILE_M: tl.constexpr, TILE_N: tl.constexpr, NUM_KSPLIT: tl.constexpr):
pid_m = tl.program_id(0); pid_n = tl.program_id(1)
offs_m = pid_m*TILE_M + tl.arange(0, TILE_M)
offs_n = pid_n*TILE_N + tl.arange(0, TILE_N)
mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
acc = tl.zeros((TILE_M, TILE_N), dtype=tl.float32)
for k in tl.static_range(NUM_KSPLIT):
acc += tl.load(y_pp_ptr + k*stride_pp_k + offs_m[:, None]*stride_pp_m + offs_n[None, :]*stride_pp_n,
mask=mask, other=0.0)
tl.store(y_ptr + offs_m[:, None]*stride_ym + offs_n[None, :]*stride_yn,
acc.to(tl.bfloat16), mask=mask)
@triton.jit
def pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr):
nig = GROUP_SIZE_M*num_pid_n; gid = pid//nig; fpm = gid*GROUP_SIZE_M
gsm = tl.minimum(num_pid_m-fpm, GROUP_SIZE_M)
return fpm+((pid%nig)%gsm), (pid%nig)//gsm
@triton.heuristics({
"EVEN_K": lambda a: (a["K"]%(a["BLOCK_SIZE_K"]//2)==0) and
(a["SPLITK_BLOCK_SIZE"]%a["BLOCK_SIZE_K"]==0) and
(a["K"]%(a["SPLITK_BLOCK_SIZE"]//2)==0),
"GRID_MN": lambda a: triton.cdiv(a["M"],a["BLOCK_SIZE_M"])*triton.cdiv(a["N"],a["BLOCK_SIZE_N"])
})
@triton.jit
def _fused_preshuffle_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr, M, N, K,
stride_bn, stride_bk, stride_ck, stride_cm, stride_cn,
stride_bsn, stride_bsk,
STRIDE_AM: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr, num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr, GRID_MN: tl.constexpr,
PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
INTERNAL_SPLITS: tl.constexpr = 1,
):
"""Fused preshuffle GEMM with optional hybrid split-K.
When INTERNAL_SPLITS > 1, each CTA processes INTERNAL_SPLITS consecutive
K-partitions (each of size SPLITK_BLOCK_SIZE) and accumulates in registers
before writing a single partial. NUM_KSPLIT = external splits only."""
tl.assume(stride_bk>0); tl.assume(stride_bn>0)
tl.assume(stride_cm>0); tl.assume(stride_cn>0); tl.assume(stride_bsk>0); tl.assume(stride_bsn>0)
pid_unified = tl.program_id(axis=0); pid_k = pid_unified%NUM_KSPLIT; pid = pid_unified//NUM_KSPLIT
num_pid_m = tl.cdiv(M,BLOCK_SIZE_M); num_pid_n = tl.cdiv(N,BLOCK_SIZE_N)
if NUM_KSPLIT==1: pid_m,pid_n = pid_grid(pid,num_pid_m,num_pid_n,GROUP_SIZE_M=GROUP_SIZE_M)
else: pid_m=pid//num_pid_n; pid_n=pid%num_pid_n
tl.assume(pid_m>=0); tl.assume(pid_n>=0); tl.assume(pid_k>=0)
SCALE_GROUP_SIZE: tl.constexpr = 32
# K-iterations per original split partition
k_iter_per_split = tl.cdiv(SPLITK_BLOCK_SIZE//2, BLOCK_SIZE_K//2)
# Total K-iterations for this CTA: INTERNAL_SPLITS consecutive partitions
total_k_iters: tl.constexpr = INTERNAL_SPLITS * k_iter_per_split if INTERNAL_SPLITS > 1 else k_iter_per_split
# Base K offset: pid_k covers INTERNAL_SPLITS consecutive original partitions
k_base_idx = pid_k * INTERNAL_SPLITS
if (k_base_idx * SPLITK_BLOCK_SIZE // 2) < K:
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = k_base_idx * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m*BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + offs_am[:,None]*STRIDE_AM + offs_k_split_bf16[None,:]
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K//2)*16)
offs_k_shuffle = k_base_idx*(SPLITK_BLOCK_SIZE//2)*16 + offs_k_shuffle_arr
offs_bn = (pid_n*(BLOCK_SIZE_N//16) + tl.arange(0, BLOCK_SIZE_N//16)) % N
b_ptrs = b_ptr + offs_bn[:,None]*stride_bn + offs_k_shuffle[None,:]*stride_bk
offs_bsn = (pid_n*(BLOCK_SIZE_N//32) + tl.arange(0, (BLOCK_SIZE_N//32))) % N
offs_ks = (k_base_idx*(SPLITK_BLOCK_SIZE//SCALE_GROUP_SIZE)*32) + \
tl.arange(0, BLOCK_SIZE_K//SCALE_GROUP_SIZE*32)
b_scale_ptrs = b_scales_ptr + offs_bsn[:,None]*stride_bsn + offs_ks[None,:]*stride_bsk
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
# Loop over INTERNAL_SPLITS * k_iter_per_split K-chunks
k_start = k_base_idx * k_iter_per_split
for k in range(k_start, k_start + total_k_iters):
b_scales = (tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N//32, BLOCK_SIZE_K//SCALE_GROUP_SIZE//8, 4, 16, 2, 2, 1)
.permute(0,5,3,1,4,2,6).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K//SCALE_GROUP_SIZE))
if EVEN_K:
a_bf16 = tl.load(a_ptrs); b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (b.reshape(1,BLOCK_SIZE_N//16,BLOCK_SIZE_K//64,2,16,16).permute(0,1,4,2,3,5)
.reshape(BLOCK_SIZE_N,BLOCK_SIZE_K//2).trans(1,0))
if PREQUANT:
a, a_scales = _mxfp4_quant_vcvt(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", fast_math=True)
a_ptrs += BLOCK_SIZE_K
b_ptrs += (BLOCK_SIZE_K//2)*16*stride_bk
b_scale_ptrs += BLOCK_SIZE_K*stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = (pid_m*BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)).to(tl.int64)
offs_cn = (pid_n*BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)).to(tl.int64)
c_ptrs = c_ptr + stride_cm*offs_cm[:,None] + stride_cn*offs_cn[None,:] + pid_k*stride_ck
tl.store(c_ptrs, c, mask=(offs_cm[:,None] < M) & (offs_cn[None,:] < N))
# Explicit S2 (M=16,N=2112,K=7168) config override.
# KSPLIT=7 and BLOCK_N=128 are sweep-confirmed (workflow 23316677253): swept KSPLITx{7,8,14}
# x BLOCK_Nx{64,128}; k7n128 optimal (k8 collapses to k7 via BLOCK_SIZE_K=256 alignment).
# Full config from 244-config bm=8 direct-override sweep (v486-v490): best is
# ks=7/ns=2/wp=3/mind=16/nw=4/cm=.cg/bk=512 -> 43.99us internal timer -> S2=10.7us benchmark.
_S2_OVERRIDES = {'BLOCK_SIZE_M': 8, 'BLOCK_SIZE_N': 128, 'NUM_KSPLIT': 7,
'GROUP_SIZE_M': 8, 'num_stages': 2, 'waves_per_eu': 3,
'matrix_instr_nonkdim': 16, 'num_warps': 4,
'cache_modifier': '.cg', 'BLOCK_SIZE_K': 512}
_USE_S2_SPECIALIZED = True
_s2_spec_cache = {}
# Hybrid split-K config: INTERNAL_SPLITS consecutive K-chunks per CTA, EXTERNAL_SPLITS partials.
# Hybrid split-K tuning results (from v477-v479 ranked artifacts):
# (2,7)=hybrid: 11.4-11.8us ranked S2 — MISS (worse than 9.78us baseline)
# (4,4)=deep-hybrid: collapses to (1,14) via get_splitk alignment — no improvement
# (1,14)=baseline: 9.78us ranked S2 — BEST KNOWN CONFIG
# Conclusion: all hybrid variants are worse. Baseline (1,14) is optimal for S2.
# Multi-N fallback: would require CTA to process multiple N fragments,
# but S2's bottleneck is K-dimension parallelism not N-tile reuse.
# Keeping baseline split-K (if S2 improvement still needed).
_S2_INTERNAL_SPLITS = 1 # baseline — hybrid variants proven worse
_S2_EXTERNAL_SPLITS = 14 # baseline KSPLIT=14 is optimal
def _run_s2_specialized(A, B_shuffle, B_scale_sh, M, N, K_actual):
"""S2 (M=16, N=2112, K=7168): hybrid split-K 2-dispatch (GEMM+reduce).
With _S2_INTERNAL_SPLITS=2, _S2_EXTERNAL_SPLITS=7: each CTA iterates over
2 K-chunks (accumulates in registers), writes 1 partial. Reduce sees 7 partials."""
key = (M, N, K_actual, _S2_INTERNAL_SPLITS, _S2_EXTERNAL_SPLITS)
if key not in _s2_spec_cache:
from aiter.ops.triton.utils.gemm_config_utils import get_gemm_config
config, _ = get_gemm_config("GEMM-A16WFP4_PRESHUFFLED", M, N, K_actual)
cfg = dict(config)
# Total effective splits = INTERNAL * EXTERNAL. Use get_splitk with the total
# to determine per-partition SPLITK_BLOCK_SIZE, then the kernel loops INTERNAL times.
_total_eff_splits = _S2_INTERNAL_SPLITS * _S2_EXTERNAL_SPLITS
cfg.update({
'BLOCK_SIZE_M': 16, 'BLOCK_SIZE_N': 128,
'NUM_KSPLIT': _total_eff_splits,
'GROUP_SIZE_M': 8, 'num_stages': 2, 'waves_per_eu': 3,
'matrix_instr_nonkdim': 16, 'num_warps': 4,
'cache_modifier': '.cg', 'BLOCK_SIZE_K': 512,
})
w_ref = B_shuffle.view(torch.uint8).reshape(N//16, (K_actual//2)*16)
N_w, K_w = w_ref.shape; N_eff = N_w*16; K_eff = K_w//16
# get_splitk with the TOTAL splits to get per-partition SPLITK_BLOCK_SIZE
SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, TOTAL_KSPLIT = get_splitk(
K_eff, cfg["BLOCK_SIZE_K"], _total_eff_splits)
# External splits: grid dimension. Internal splits: loop within each CTA.
# Recompute actual external count from total after get_splitk alignment
_actual_internal = _S2_INTERNAL_SPLITS
_actual_external = max(1, TOTAL_KSPLIT // _actual_internal)
if _actual_external * _actual_internal != TOTAL_KSPLIT:
# get_splitk rounded — fall back to ext=total, int=1
_actual_external = TOTAL_KSPLIT
_actual_internal = 1
cfg["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE # per-original-partition size
cfg["BLOCK_SIZE_K"] = BLOCK_SIZE_K
cfg["NUM_KSPLIT"] = _actual_external # grid only has external CTAs
cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
y = torch.empty(M, N, dtype=torch.bfloat16, device=A.device)
y_pp = torch.empty((_actual_external, M, N), dtype=torch.float32, device=A.device)
ACTUAL_KSPLIT = _actual_external
RTILE_M = triton.next_power_of_2(M); RTILE_N = min(64, N)
reduce_grid = (triton.cdiv(M, RTILE_M), triton.cdiv(N, RTILE_N))
num_pid_m = triton.cdiv(M, cfg['BLOCK_SIZE_M'])
num_pid_n = triton.cdiv(N_eff, cfg['BLOCK_SIZE_N'])
total_wgs = _actual_external * num_pid_m * num_pid_n
grid = lambda META: (META["NUM_KSPLIT"]*triton.cdiv(M, META["BLOCK_SIZE_M"])*triton.cdiv(N_eff, META["BLOCK_SIZE_N"]),)
k_iter_per_split = triton.cdiv(SPLITK_BLOCK_SIZE//2, BLOCK_SIZE_K//2)
_s2_spec_cache[key] = {
'N_eff': N_eff, 'K_eff': K_eff, 'cfg': cfg, 'y': y, 'y_pp': y_pp,
'grid': grid, 'ACTUAL_KSPLIT': ACTUAL_KSPLIT,
'RTILE_M': RTILE_M, 'RTILE_N': RTILE_N, 'reduce_grid': reduce_grid,
'internal_splits': _actual_internal,
}
print(f'[s2_hybrid] init: BM={cfg["BLOCK_SIZE_M"]} BN={cfg["BLOCK_SIZE_N"]} '
f'internal={_actual_internal} external={_actual_external} total={TOTAL_KSPLIT} '
f'k_iter/split={k_iter_per_split} k_iter/CTA={k_iter_per_split*_actual_internal} '
f'BSK={BLOCK_SIZE_K} SPBSZ={SPLITK_BLOCK_SIZE} '
f'nw={cfg["num_warps"]} WGs={total_wgs} reduce={ACTUAL_KSPLIT}',
file=sys.stderr)
fc = _s2_spec_cache[key]
w = B_shuffle.view(torch.uint8).reshape(N//16, (K_actual//2)*16)
bs_u8 = B_scale_sh.view(torch.uint8); sm, sn = bs_u8.shape; ws = bs_u8.reshape(sm//32, sn*32)
_fused_preshuffle_kernel[fc['grid']](
A, w, fc['y_pp'], ws, M, fc['N_eff'], fc['K_eff'],
w.stride(0), w.stride(1),
fc['y_pp'].stride(0), fc['y_pp'].stride(1), fc['y_pp'].stride(2),
ws.stride(0), ws.stride(1), STRIDE_AM=K_actual,
PREQUANT=True, INTERNAL_SPLITS=fc['internal_splits'], **fc['cfg'])
_custom_reduce_kernel[fc['reduce_grid']](
fc['y_pp'], fc['y'], M, N,
fc['y_pp'].stride(0), fc['y_pp'].stride(1), fc['y_pp'].stride(2),
fc['y'].stride(0), fc['y'].stride(1),
TILE_M=fc['RTILE_M'], TILE_N=fc['RTILE_N'], NUM_KSPLIT=fc['ACTUAL_KSPLIT'])
return fc['y']
def _run_v364(A, B_shuffle, B_scale_sh, M, N, K_actual, config_name="GEMM-A16WFP4_PRESHUFFLED"):
from aiter.ops.triton.utils.gemm_config_utils import get_gemm_config
key = (M, N, K_actual, config_name)
if key not in _fused_cache:
config, _ = get_gemm_config(config_name, M, N, K_actual); cfg = dict(config)
if K_actual == 7168 and M <= 16:
cfg.update(_S2_OVERRIDES) # KSPLIT=7/BLOCK_N=128 sweep-confirmed; other fields from prior tuning
elif K_actual >= 7168:
if cfg.get('NUM_KSPLIT',1) > 7: cfg['NUM_KSPLIT']=7
cfg['num_stages'] = max(cfg.get('num_stages', 1), 2)
cfg['waves_per_eu'] = max(cfg.get('waves_per_eu', 1), 2)
if M <= 16: cfg['BLOCK_SIZE_M'] = 8
elif K_actual <= 512 and M <= 8: cfg['num_warps'] = 2
if K_actual <= 512 and M == 32: cfg['BLOCK_SIZE_M'] = 16
# Tuned: waves_per_eu=2, num_stages=2, matrix_instr_nonkdim=16 for K<=512.
# Benchmark S3: 7.01→6.86 (-0.15us). Ranked S3: 7.28→7.49 (+0.21us WORSE).
# Reverted: hints improve benchmark but hurt ranked stability.
cfg['GROUP_SIZE_M'] = 8
w_ref = B_shuffle.view(torch.uint8).reshape(N//16, (K_actual//2)*16)
N_w,K_w = w_ref.shape; N_eff=N_w*16; K_eff=K_w//16
if cfg["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE,BLOCK_SIZE_K,NUM_KSPLIT = get_splitk(K_eff,cfg["BLOCK_SIZE_K"],cfg["NUM_KSPLIT"])
cfg["SPLITK_BLOCK_SIZE"]=SPLITK_BLOCK_SIZE; cfg["BLOCK_SIZE_K"]=BLOCK_SIZE_K
cfg["NUM_KSPLIT"]=NUM_KSPLIT
if cfg["BLOCK_SIZE_K"] >= 2*K_eff:
cfg["BLOCK_SIZE_K"]=triton.next_power_of_2(2*K_eff)
cfg["SPLITK_BLOCK_SIZE"]=2*K_eff; cfg["NUM_KSPLIT"]=1
cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
y = torch.empty(M, N, dtype=torch.bfloat16, device=A.device); y_pp = None
if cfg["NUM_KSPLIT"] > 1:
y_pp = torch.empty((cfg["NUM_KSPLIT"],M,N), dtype=torch.float32, device=A.device)
ACTUAL_KSPLIT = triton.cdiv(K_eff, (cfg["SPLITK_BLOCK_SIZE"]//2))
RTILE_M = min(32,M) if M>=32 else triton.next_power_of_2(M); RTILE_N = min(64,N)
reduce_grid = (triton.cdiv(M,RTILE_M), triton.cdiv(N,RTILE_N))
else:
cfg["SPLITK_BLOCK_SIZE"]=2*K_eff; ACTUAL_KSPLIT=0; RTILE_M=0; RTILE_N=0
reduce_grid = None
grid = lambda META: (META["NUM_KSPLIT"]*triton.cdiv(M,META["BLOCK_SIZE_M"])*triton.cdiv(N_eff,META["BLOCK_SIZE_N"]),)
c_ptr = y if y_pp is None else y_pp
_fused_cache[key] = {
'N_eff':N_eff,'K_eff':K_eff,'cfg':cfg,'y':y,'y_pp':y_pp,'c_ptr':c_ptr,'grid':grid,
'stride_ck': 0 if y_pp is None else y_pp.stride(0),
'stride_cm': y.stride(0) if y_pp is None else y_pp.stride(1),
'stride_cn': y.stride(1) if y_pp is None else y_pp.stride(2),
'reduce_grid':reduce_grid,'ACTUAL_KSPLIT':ACTUAL_KSPLIT,'RTILE_M':RTILE_M,'RTILE_N':RTILE_N,
}
fc = _fused_cache[key]
# Recompute B views per call -- do NOT cache w/ws; ranked test may pass different tensors
w = B_shuffle.view(torch.uint8).reshape(N//16, (K_actual//2)*16)
bs_u8 = B_scale_sh.view(torch.uint8); sm,sn = bs_u8.shape; ws = bs_u8.reshape(sm//32, sn*32)
_fused_preshuffle_kernel[fc['grid']](
A, w, fc['c_ptr'], ws, M, fc['N_eff'], fc['K_eff'],
w.stride(0), w.stride(1),
fc['stride_ck'], fc['stride_cm'], fc['stride_cn'], ws.stride(0), ws.stride(1),
STRIDE_AM=K_actual, PREQUANT=True, **fc['cfg'])
if fc['y_pp'] is not None:
_custom_reduce_kernel[fc['reduce_grid']](
fc['y_pp'], fc['y'], M, N,
fc['y_pp'].stride(0), fc['y_pp'].stride(1), fc['y_pp'].stride(2),
fc['y'].stride(0), fc['y'].stride(1),
TILE_M=fc['RTILE_M'], TILE_N=fc['RTILE_N'], NUM_KSPLIT=fc['ACTUAL_KSPLIT'])
return fc['y']
# ============================================================
# AC-4 dispatch-time checks (workspace cap + per-shape CTA/CU average)
# ============================================================
# Per-variant kernel metadata: (BM, BN, splitK, workspace_bytes_fn, threads)
# workspace_bytes_fn: callable(M, N, K) -> int, or 0 for no workspace
# threads: workgroup size (for compile-time CTA/CU verification)
# Verified against actual kernel definitions in submission.py
def _ws_fn(BM, BN):
"""Workspace bytes for split-K=2: n_tiles * 2 * (BM*BN) * sizeof(float)"""
return lambda M, N, K: triton.cdiv(N, BN) * triton.cdiv(M, BM) * 2 * (BM * BN) * 4
_NO_WS = lambda M, N, K: 0
_FUSED_KERNEL_META = {
# S6 variants (splitK=1, direct BF16 write, no workspace)
'v576': (32, 32, 1, _NO_WS, 256), # K-step=768, ldsB[12]
'v571glds':(32, 32, 1, _NO_WS, 256), # K-step=512, global_load_lds+B_shuffle tile-coalesced
'v571': (32, 32, 1, _NO_WS, 256), # K-step=512, ldsB[8]
'v571x': (32, 32, 1, _NO_WS, 256), # K-step=512, B_shuffle correct formula VGPR+LDS
'v571c': (32, 32, 1, _NO_WS, 256), # K-step=512, B_shuffle coalesced VGPR+LDS
'v571b': (32, 32, 1, _NO_WS, 256), # K-step=512, VGPR-load LDS fix
'v571dg': (32, 32, 1, _NO_WS, 256), # K-step=512, direct-global B
# S5 variants
'v575c': (32, 32, 2, _ws_fn(32, 32), 256), # K-step=1024, B_shuffle coalesced VGPR+LDS
'v575': (32, 32, 2, _ws_fn(32, 32), 256), # K-step=1024, single-pass
'v575dg': (32, 32, 2, _ws_fn(32, 32), 256), # K-step=1024, direct-global B
'v574': (32, 32, 2, _ws_fn(32, 32), 256), # K-step=512
'v574b': (32, 32, 2, _ws_fn(32, 32), 256), # K-step=512, VGPR+LDS (no swizzle bug)
}
_ac4_dispatch_logged = set()
_WORKSPACE_CAP_BYTES = 16 * 1024 * 1024 # 16MB per DEC-4
_NUM_CUS = 256 # MI355X: 256 CUs
def _ac4_dispatch_check(variant, M, N, K):
"""Dispatch-time AC-4 check: workspace cap and per-shape CTA/CU average.
Returns True if variant is allowed for this shape, False if rejected."""
meta = _FUSED_KERNEL_META.get(variant)
if meta is None:
return True # Unknown variant, let compile-time gate handle it
BM, BN, splitK, ws_fn, _threads = meta
# Workspace cap check
ws_bytes = ws_fn(M, N, K)
if ws_bytes > _WORKSPACE_CAP_BYTES:
if (variant, M, N, K) not in _ac4_dispatch_logged:
_ac4_dispatch_logged.add((variant, M, N, K))
print(f'[ac4] REJECT {variant} M={M} N={N} K={K}: workspace={ws_bytes/1e6:.2f}MB > 16MB',
file=sys.stderr)
return False
# Per-shape CTA/CU average check: reject if < 1
num_tiles = triton.cdiv(M, BM) * triton.cdiv(N, BN) * splitK
avg_cta_per_cu = num_tiles / _NUM_CUS
if avg_cta_per_cu < 1.0:
if (variant, M, N, K) not in _ac4_dispatch_logged:
_ac4_dispatch_logged.add((variant, M, N, K))
print(f'[ac4] REJECT {variant} M={M} N={N} K={K}: avg CTA/CU={avg_cta_per_cu:.2f} < 1 '
f'(tiles={num_tiles})', file=sys.stderr)
return False
if (variant, M, N, K) not in _ac4_dispatch_logged:
_ac4_dispatch_logged.add((variant, M, N, K))
print(f'[ac4] PASS {variant} M={M} N={N} K={K}: ws={ws_bytes/1e6:.2f}MB tiles={num_tiles} '
f'avg_CTA/CU={avg_cta_per_cu:.2f}', file=sys.stderr)
return True
# ============================================================
# Main dispatch
# ============================================================
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape; N = B_q.shape[0]
if K <= 512:
return _run_v364(A, B_shuffle, B_scale_sh, M, N, K,
config_name="GEMM-AFP4WFP4_PRESHUFFLED" if M <= 32 else "GEMM-A16WFP4_PRESHUFFLED")
elif K == 7168:
return _run_s2_specialized(A, B_shuffle, B_scale_sh, M, N, K)
else:
# S6 (K=1536): v571dg → v571 → 2-dispatch CK
# v571p byte-contract probe: runs for any K=1536 shape — B is [N,K/2] independent of M.
# Crafted-input probe first, then live-data verdict; probe is diagnostic only.
if K == 1536 and _USE_FUSED_V571P and ('v571p_done', N, K) not in _ac4_dispatch_logged:
_ac4_dispatch_logged.add(('v571p_done', N, K))
try:
# Crafted-input probe: unique per-kb pattern → unambiguous verdict
verdict, _details = run_v571p_crafted_probe(N, K, B_q.device, tile_n=0)
print(f'[v571p] Crafted-input verdict: {verdict}', file=sys.stderr)
# Also run on live evaluator data for reproducibility check
run_v571p_selfcheck(B_q, B_scale_sh, N, K, tile_n=0)
live_verdict, _live = run_v571p_verdict(B_q, B_scale_sh, N, K, tile_n=0)
print(f'[v571p] Live-data verdict: {live_verdict}', file=sys.stderr)
except Exception as ex:
print(f'[v571p] probe error: {ex}', file=sys.stderr)
# Fall through to CK — probe is diagnostic only
if K == 1536 and M == 256:
# Probe mode: run v571 LDS path AND v571dg, compare outputs, emit verdict once
if (_USE_FUSED_V571 and _USE_FUSED_V571DG and
_ac4_dispatch_check('v571', M, N, K) and
_ac4_dispatch_check('v571dg', M, N, K) and
('probe_done', M, N, K) not in _ac4_dispatch_logged):
_ac4_dispatch_logged.add(('probe_done', M, N, K))
torch.cuda.synchronize()
out_lds = _run_fused_s6_v571(A, B_q, B_scale_sh, M, N, K)
torch.cuda.synchronize()
out_dg = _run_fused_s6_v571dg(A, B_q, B_scale_sh, M, N, K)
torch.cuda.synchronize()
ref = _run_native_quant_ck(A, B_shuffle, B_scale_sh, M, N, K)
torch.cuda.synchronize()
if out_lds is not None and out_dg is not None and ref is not None:
tol = 0.5
diff_lds_dg = (out_lds.float() - out_dg.float()).abs()
diff_dg_ck = (out_dg.float() - ref.float()).abs()
n_lds_dg = (diff_lds_dg > tol).sum().item()
n_dg_ck = (diff_dg_ck > tol).sum().item()
print(f'[probe] LDS vs DG: max={diff_lds_dg.max():.3f} n_diff={n_lds_dg}/{M*N}',
file=sys.stderr)
print(f'[probe] DG vs CK: max={diff_dg_ck.max():.3f} n_diff={n_dg_ck}/{M*N}',
file=sys.stderr)
threshold = int(M * N * 0.01)
if n_lds_dg > threshold:
verdict = 'bytes-mismatch'
elif n_dg_ck > threshold:
verdict = 'scale-mismatch'
else:
verdict = 'neither'
print(f'[probe] VERDICT: {verdict}', file=sys.stderr)
else:
print('[probe] VERDICT: crash/non-deterministic', file=sys.stderr)
# Production path: v571glds → v571x → v571c → v571b → v571dg (slow fallback)
if _USE_FUSED_V571GLDS and _ac4_dispatch_check('v571glds', M, N, K):
fv571glds = _run_fused_s6_v571glds(A, B_shuffle, B_scale_sh, M, N, K)
if fv571glds is not None:
return fv571glds
if _USE_FUSED_V571X and _ac4_dispatch_check('v571x', M, N, K):
fv571x = _run_fused_s6_v571x(A, B_shuffle, B_scale_sh, M, N, K)
if fv571x is not None:
return fv571x
if _USE_FUSED_V571C and _ac4_dispatch_check('v571c', M, N, K):
fv571c = _run_fused_s6_v571c(A, B_shuffle, B_scale_sh, M, N, K)
if fv571c is not None:
return fv571c
if _USE_FUSED_V571B and _ac4_dispatch_check('v571b', M, N, K):
fv571b = _run_fused_s6_v571b(A, B_q, B_scale_sh, M, N, K)
if fv571b is not None:
return fv571b
if _USE_FUSED_V571DG and _ac4_dispatch_check('v571dg', M, N, K):
fv571dg = _run_fused_s6_v571dg(A, B_q, B_scale_sh, M, N, K)
if fv571dg is not None:
return fv571dg
# S5 (K=2048, M=64): v575c (B_shuffle coalesced) → v575 → v574b → v575dg → v574 → 2-dispatch CK
if K == 2048 and M == 64:
if _USE_FUSED_V575C and _ac4_dispatch_check('v575c', M, N, K):
fv575c = _run_fused_s5_v575c(A, B_shuffle, B_scale_sh, M, N, K)
if fv575c is not None:
return fv575c
if _USE_FUSED_V575 and _ac4_dispatch_check('v575', M, N, K):
fv575 = _run_fused_s5_v575(A, B_shuffle, B_scale_sh, M, N, K)
if fv575 is not None:
return fv575
if _USE_FUSED_V574B and _ac4_dispatch_check('v574b', M, N, K):
fv574b = _run_fused_s5_v574b(A, B_q, B_scale_sh, M, N, K)
if fv574b is not None:
return fv574b
if _USE_FUSED_V575DG and _ac4_dispatch_check('v575dg', M, N, K):
fv575dg = _run_fused_s5_v575dg(A, B_q, B_scale_sh, M, N, K)
if fv575dg is not None:
return fv575dg
if _USE_FUSED_V574 and _ac4_dispatch_check('v574', M, N, K):
fv574 = _run_fused_s5_v574(A, B_shuffle, B_scale_sh, M, N, K)
if fv574 is not None:
return fv574
# Fallback: C++ fast path (quant + CK ASM 32x128, 2-dispatch)
if _USE_FAST_PATH:
fp_result = _run_fast_path(A, B_shuffle, B_scale_sh, M, N, K)
if fp_result is not None:
return fp_result
return _run_native_quant_ck(A, B_shuffle, B_scale_sh, M, N, K)
# ============================================================
# Module-level HSACO loading (import time, fast, no JIT)
# ============================================================
def _module_init():
try:
_get_native_quant()
# Pre-stage mainline S6/S5 kernels at import time (not lazy on first call)
if _USE_FUSED_V576:
_get_fused_quant_gemm_v576()
if _USE_FUSED_V571GLDS:
_get_fused_quant_gemm_v571glds()
if _USE_FUSED_V571X:
_get_fused_quant_gemm_v571x()
if _USE_FUSED_V571C:
_get_fused_quant_gemm_v571c()
if _USE_FUSED_V571B:
_get_fused_quant_gemm_v571b()
if _USE_FUSED_V571DG:
_get_fused_quant_gemm_v571dg()
if _USE_FUSED_V571:
_get_fused_quant_gemm_v571()
if _USE_FUSED_V575C:
_get_fused_quant_gemm_v575c()
if _USE_FUSED_V575DG:
_get_fused_quant_gemm_v575dg()
if _USE_FUSED_V575:
_get_fused_quant_gemm_v575()
if _USE_FUSED_V574:
_get_fused_quant_gemm_v574()
if _USE_FUSED_V574B:
_get_fused_quant_gemm_v574b()
_ck_hsaco_path = '/home/runner/aiter/hsa//gfx950/f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co'
import os
if os.path.exists(_ck_hsaco_path):
_hip_lib = ctypes.CDLL('libamdhip64.so')
_ck_mod = ctypes.c_void_p()
rc = _hip_lib.hipModuleLoad(ctypes.byref(_ck_mod), _ck_hsaco_path.encode())
if rc == 0:
_ck_fn = ctypes.c_void_p()
rc2 = _hip_lib.hipModuleGetFunction(ctypes.byref(_ck_fn), _ck_mod, K32x128.encode())
if rc2 == 0:
_qs_cache['_ck_direct'] = {'fn': _ck_fn, 'hip': _hip_lib, 'mod': _ck_mod}
# ASM pipeline test removed from bootstrap — adds ~0.3s import-time overhead
# that regresses ranked timing. Test on demand, not at import.
global _fast_path_warmed
_fast_path_warmed = True
# v516: Deep CK Tile exploration + compilation probe
_ck_base = '/home/runner/aiter/3rdparty/composable_kernel/include'
# List all files under gemm_quant directory
import glob as _glob
for d in ['ck_tile/ops/gemm_quant', 'ck_tile/ops/gemm/pipeline', 'ck_tile/ops/flatmm']:
full_d = os.path.join(_ck_base, d)
if os.path.isdir(full_d):
files = sorted(os.listdir(full_d))
print(f'[ck_explore] {d}/ ({len(files)} files): {", ".join(files[:20])}', file=sys.stderr)
else:
print(f'[ck_explore] {d}/ MISSING', file=sys.stderr)
# Check for flatmm (F16xMXF4) paths
for p in ['ck_tile/ops/flatmm/pipeline', 'ck_tile/ops/flatmm/kernel']:
full_p = os.path.join(_ck_base, p)
if os.path.isdir(full_p):
files = sorted(os.listdir(full_p))
print(f'[ck_explore] {p}/ ({len(files)} files): {", ".join(files[:20])}', file=sys.stderr)
# Check specific types mentioned by Codex
for p in ['ck_tile/ops/gemm_quant/pipeline/gemm_mxfp4_pipeline_ag_bg_cr_v3.hpp',
'ck_tile/ops/gemm/pipeline/gemm_pipeline_ag_bg_cr_comp_v4.hpp',
'ck_tile/host/kernel_launch.hpp',
'ck_tile/core/numeric/pk_fp4.hpp']:
full = os.path.join(_ck_base, p)
if os.path.exists(full):
# Read first 30 lines to find key types
with open(full, 'r') as f:
lines = f.readlines()[:30]
print(f'[ck_explore] {p} ({len(lines)} head lines):', file=sys.stderr)
for l in lines:
if 'template' in l.lower() or 'struct' in l.lower() or 'class' in l.lower() or 'using' in l.lower():
print(f' {l.rstrip()}', file=sys.stderr)
# Compilation probe: try to compile minimal CK Tile include
_ck_probe_src = r'''
#include <hip/hip_runtime.h>
#include "ck_tile/core.hpp"
extern "C" __global__ void ck_tile_probe() {
// If this compiles, CK Tile core headers work with hipcc --genco
}
'''
_probe_dir = '/tmp/ck_tile_probe'
os.makedirs(_probe_dir, exist_ok=True)
_probe_src = os.path.join(_probe_dir, 'probe.hip')
_probe_co = os.path.join(_probe_dir, 'probe.co')
with open(_probe_src, 'w') as f:
f.write(_ck_probe_src)
import time as _time
t0 = _time.time()
_probe_result = subprocess.run(
['hipcc', '--genco', '--offload-arch=gfx950', '-O2', '-std=c++20',
f'-I{_ck_base}', '-o', _probe_co, _probe_src],
capture_output=True, text=True, timeout=120)
dt = round(_time.time() - t0, 1)
if _probe_result.returncode == 0:
co_size = os.path.getsize(_probe_co)
print(f'[ck_compile] core.hpp probe OK ({dt}s, {co_size} bytes)', file=sys.stderr)
# Try more complex: include gemm headers
_ck_gemm_src = r'''
#include <hip/hip_runtime.h>
#include "ck_tile/core.hpp"
#include "ck_tile/ops/gemm.hpp"
extern "C" __global__ void ck_tile_gemm_probe() {}
'''
_gemm_src = os.path.join(_probe_dir, 'gemm_probe.hip')
_gemm_co = os.path.join(_probe_dir, 'gemm_probe.co')
with open(_gemm_src, 'w') as f:
f.write(_ck_gemm_src)
t1 = _time.time()
_gemm_result = subprocess.run(
['hipcc', '--genco', '--offload-arch=gfx950', '-O2', '-std=c++20',
f'-I{_ck_base}', '-o', _gemm_co, _gemm_src],
capture_output=True, text=True, timeout=120)
dt2 = round(_time.time() - t1, 1)
if _gemm_result.returncode == 0:
co2_size = os.path.getsize(_gemm_co)
print(f'[ck_compile] gemm.hpp probe OK ({dt2}s, {co2_size} bytes)', file=sys.stderr)
else:
print(f'[ck_compile] gemm.hpp probe FAILED ({dt2}s): {_gemm_result.stderr[-300:]}', file=sys.stderr)
else:
print(f'[ck_compile] core.hpp probe FAILED ({dt}s): {_probe_result.stderr[-300:]}', file=sys.stderr)
except Exception:
pass
pass # Correctness probe requires evaluator data; diagnostic comparison done in custom_kernel
_module_init()
scrolls · 1752 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