submission 754280
Hamza · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 443 lines, June 9 Researcher Reciprocity License v1.0.
hybrid-GEMM.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754280?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:ef2dd64c2dd04691674105f19eddc7e06c63593843571abfecdaf66684edcbf4
license declaredunknown
license concludedunknown
authorsHamza
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
hybrid-GEMM.py443 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Fused BF16-to-FP4 quantization + scaled GEMM with pre-shuffled weights.
Hardware-accelerated quantization via v_cvt_scalef32_pk_fp4_bf16.
Per-shape tuned tile/split-K parameters for MI355X (256 CUs).
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
import sys as _sys
import gc as _gc
torch.set_grad_enabled(False)
_gc.disable()
_sys.setswitchinterval(1.0)
@triton.jit
def _quant_block_fp4(
inp_bf16,
TILE_M: tl.constexpr,
TILE_K: tl.constexpr,
):
"""Convert BF16 tile to packed MXFP4 using hardware pk_fp4 instruction."""
GRP_SZ: tl.constexpr = 32
N_GROUPS: tl.constexpr = TILE_K // GRP_SZ
vals_f32 = inp_bf16.to(tl.float32).reshape(TILE_M, N_GROUPS, GRP_SZ)
# Block-wise absolute max with rounding to nearest power-of-2
peak = tl.max(tl.abs(vals_f32), axis=-1, keep_dims=True)
peak = peak.to(tl.int32, bitcast=True)
peak = (peak + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
log2_peak = ((peak >> 23) & 0xFF).to(tl.int32) - 127
exp_unbiased = log2_peak - 2
exp_unbiased = tl.minimum(tl.maximum(exp_unbiased, -127), 127)
scale_e8m0 = exp_unbiased.to(tl.uint8) + 127
# Build IEEE754 float divisor: 2^unbiased for hw instruction
div_bits = (exp_unbiased.to(tl.int32) + 127).to(tl.uint32) << 23
divisor = div_bits.to(tl.float32, bitcast=True) # [M, N_GROUPS, 1]
# Expand divisor to per-pair level
div_full = tl.broadcast_to(divisor, (TILE_M, N_GROUPS, GRP_SZ))
div_full = div_full.reshape(TILE_M, TILE_K)
div_pairs = div_full.reshape(TILE_M, TILE_K // 2, 2)
div_even, _ = tl.split(div_pairs)
div_per_pair = div_even.reshape(TILE_M, TILE_K // 2)
# Interleave adjacent BF16 elements into uint32 for hw conversion
inp_u16 = inp_bf16.to(tl.uint16, bitcast=True).reshape(TILE_M, TILE_K // 2, 2)
lo_half, hi_half = tl.split(inp_u16)
packed_u32 = lo_half.to(tl.uint32) | (hi_half.to(tl.uint32) << 16)
packed_u32 = packed_u32.reshape(TILE_M, TILE_K // 2)
# Hardware FP4 pack-convert
fp4_raw = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
"=v, v, v",
[packed_u32, div_per_pair],
dtype=tl.uint32,
is_pure=True,
pack=1,
)
fp4_bytes = (fp4_raw & 0xFF).to(tl.uint8)
fp4_bytes = fp4_bytes.reshape(TILE_M, TILE_K // 2)
return fp4_bytes, scale_e8m0.reshape(TILE_M, N_GROUPS)
@triton.heuristics(
{
"ALIGNED_K": lambda args: (args["K"] % (args["TILE_K"] // 2) == 0)
and (args["SK_TILE"] % args["TILE_K"] == 0)
and (args["K"] % (args["SK_TILE"] // 2) == 0),
}
)
@triton.jit
def _matmul_fused_kernel(
inp_ptr, wt_ptr, out_ptr, wsc_ptr,
M, N, K,
stride_im, stride_ik,
stride_wn, stride_wk,
stride_ok, stride_om, stride_on,
stride_sn, stride_sk,
TILE_M: tl.constexpr,
TILE_N: tl.constexpr,
TILE_K: tl.constexpr,
GROUP_M: tl.constexpr,
N_SPLITS: tl.constexpr,
SK_TILE: tl.constexpr,
ALIGNED_K: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
load_modifier: tl.constexpr,
):
tl.assume(stride_im > 0)
tl.assume(stride_ik > 0)
tl.assume(stride_wn > 0)
tl.assume(stride_wk > 0)
tl.assume(stride_om > 0)
tl.assume(stride_on > 0)
tl.assume(stride_sn > 0)
tl.assume(stride_sk > 0)
SCALE_GRP: tl.constexpr = 32
n_tiles_m = tl.cdiv(M, TILE_M)
n_tiles_n = tl.cdiv(N, TILE_N)
flat_pid = tl.program_id(axis=0)
split_id = flat_pid % N_SPLITS
tile_pid = flat_pid // N_SPLITS
# Tile assignment: grouped swizzle for single-split, linear for multi-split
if N_SPLITS == 1:
tiles_per_grp = GROUP_M * n_tiles_n
grp = tile_pid // tiles_per_grp
first_m = grp * GROUP_M
grp_sz = min(n_tiles_m - first_m, GROUP_M)
tile_m = first_m + ((tile_pid % tiles_per_grp) % grp_sz)
tile_n = (tile_pid % tiles_per_grp) // grp_sz
else:
tile_m = tile_pid // n_tiles_n
tile_n = tile_pid % n_tiles_n
tl.assume(tile_m >= 0)
tl.assume(tile_n >= 0)
tl.assume(split_id >= 0)
if (split_id * SK_TILE // 2) < K:
k_iters = tl.cdiv(SK_TILE // 2, TILE_K // 2)
# A: BF16 input [M, 2*K]
row_a = (tile_m * TILE_M + tl.arange(0, TILE_M)) % M
col_a = split_id * SK_TILE + tl.arange(0, TILE_K)
ptrs_a = inp_ptr + (row_a[:, None] * stride_im + col_a[None, :] * stride_ik)
# B: pre-shuffled FP4 weights [N//16, K_packed*16]
shuf_range = tl.arange(0, (TILE_K // 2) * 16)
shuf_base = split_id * (SK_TILE // 2) * 16 + shuf_range
row_b = (tile_n * (TILE_N // 16) + tl.arange(0, TILE_N // 16)) % (N // 16)
ptrs_b = wt_ptr + (row_b[:, None] * stride_wn + shuf_base[None, :] * stride_wk)
# B scales: shuffled E8M0 layout
row_s = (tile_n * TILE_N + tl.arange(0, TILE_N // 32) * 32)
col_s = (split_id * (SK_TILE // SCALE_GRP) * 32) + tl.arange(
0, TILE_K // SCALE_GRP * 32
)
ptrs_s = wsc_ptr + row_s[:, None] * stride_sn + col_s[None, :] * stride_sk
acc = tl.zeros((TILE_M, TILE_N), dtype=tl.float32)
for ki in range(split_id * k_iters, (split_id + 1) * k_iters):
# Issue all loads before compute for memory-level parallelism
if ALIGNED_K:
a_tile = tl.load(ptrs_a, eviction_policy="evict_last")
s_raw = tl.load(ptrs_s, cache_modifier=load_modifier)
b_raw = tl.load(ptrs_b, cache_modifier=load_modifier)
else:
k_off = (ki - split_id * k_iters) * TILE_K
a_tile = tl.load(
ptrs_a,
mask=tl.arange(0, TILE_K)[None, :] < (2 * K - split_id * SK_TILE - k_off),
other=0.0,
eviction_policy="evict_last",
)
s_raw = tl.load(ptrs_s, cache_modifier=load_modifier)
b_raw = tl.load(
ptrs_b,
mask=shuf_range[None, :] < ((K - (split_id * (SK_TILE // 2) + (ki - split_id * k_iters) * (TILE_K // 2))) * 16),
other=0,
cache_modifier=load_modifier,
)
# On-the-fly A quantization
inp_q, inp_sc = _quant_block_fp4(a_tile, TILE_M, TILE_K)
# Reconstruct B scale layout from shuffled storage
b_sc = (
s_raw
.reshape(
TILE_N // 32,
TILE_K // SCALE_GRP // 8,
4, 16, 2, 2, 1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(TILE_N, TILE_K // SCALE_GRP)
)
# Reconstruct B tile from shuffled storage
b_tile = (
b_raw.reshape(1, TILE_N // 16, TILE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(TILE_N, TILE_K // 2)
.trans(1, 0)
)
acc = tl.dot_scaled(
inp_q, inp_sc, "e2m1", b_tile, b_sc, "e2m1", acc,
fast_math=True,
)
ptrs_a += TILE_K * stride_ik
ptrs_b += (TILE_K // 2) * 16 * stride_wk
ptrs_s += TILE_K * stride_sk
out_vals = acc.to(out_ptr.type.element_ty)
row_o = tile_m * TILE_M + tl.arange(0, TILE_M).to(tl.int64)
col_o = tile_n * TILE_N + tl.arange(0, TILE_N).to(tl.int64)
ptrs_o = (
out_ptr
+ stride_om * row_o[:, None]
+ stride_on * col_o[None, :]
+ split_id * stride_ok
)
mask_o = (row_o[:, None] < M) & (col_o[None, :] < N)
tl.store(ptrs_o, out_vals, mask=mask_o, cache_modifier=".wt")
@triton.jit
def _partial_sum_kernel(
partials_ptr, final_ptr, M, N,
stride_pk, stride_pm, stride_pn,
stride_fm, stride_fn,
RED_M: tl.constexpr, RED_N: tl.constexpr,
TRUE_SPLITS: tl.constexpr, MAX_SPLITS: tl.constexpr,
):
"""Reduce split-K partial results into final output."""
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
rows = (pid_m * RED_M + tl.arange(0, RED_M)) % M
cols = (pid_n * RED_N + tl.arange(0, RED_N)) % N
base = (
partials_ptr
+ (rows[:, None] * stride_pm)
+ (cols[None, :] * stride_pn)
)
total = tl.load(base).to(tl.float32)
for s in tl.static_range(1, MAX_SPLITS):
if s < TRUE_SPLITS:
total += tl.load(base + s * stride_pk).to(tl.float32)
out_vals = total.to(final_ptr.type.element_ty)
out_ptrs = (
final_ptr
+ (rows[:, None] * stride_fm)
+ (cols[None, :] * stride_fn)
)
tl.store(out_ptrs, out_vals)
# ---------------------------------------------------------------------------
# Split-K adjustment for divisibility constraints
# ---------------------------------------------------------------------------
def _adjust_splitk(K, BK, n_splits):
sk_tile = (
triton.cdiv((2 * triton.cdiv(K, n_splits)), BK) * BK
)
while n_splits > 1 and BK > 16:
if (
K % (sk_tile // 2) == 0
and sk_tile % BK == 0
and K % (BK // 2) == 0
):
break
elif K % (sk_tile // 2) != 0 and n_splits > 1:
n_splits = n_splits // 2
elif sk_tile % BK != 0:
if n_splits > 1:
n_splits = n_splits // 2
elif BK > 16:
BK = BK // 2
elif K % (BK // 2) != 0 and BK > 16:
BK = BK // 2
else:
break
sk_tile = (
triton.cdiv((2 * triton.cdiv(K, n_splits)), BK) * BK
)
n_splits = triton.cdiv(K, (sk_tile // 2))
return sk_tile, BK, n_splits
# ---------------------------------------------------------------------------
# Per-shape tuned configurations (MI355X, 256 CUs)
# ---------------------------------------------------------------------------
_TUNED_PARAMS = {
(4, 2880, 512): {"TILE_M": 4, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},
(16, 2112, 7168): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 512, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 14},
(32, 4096, 512): {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1},
(32, 2880, 512): {"TILE_M": 8, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},
(64, 7168, 2048): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1},
(256, 3072, 1536): {"TILE_M": 16, "TILE_N": 256, "TILE_K": 512, "GROUP_M": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},
}
_FALLBACK_PARAMS = {
"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "GROUP_M": 1,
"num_warps": 2, "num_stages": 2, "waves_per_eu": 0,
"matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1,
}
# ---------------------------------------------------------------------------
# Caches for buffers, configs, and launch grids
# ---------------------------------------------------------------------------
_out_pool = {}
_param_pool = {}
_grid_pool = {}
_operand_pool = {}
def _get_output_bufs(m, n, n_splits, dev):
key = (m, n, n_splits)
if key not in _out_pool:
final = torch.empty((m, n), dtype=torch.bfloat16, device=dev)
partials = (
torch.empty((n_splits, m, n), dtype=torch.float32, device=dev)
if n_splits > 1 else None
)
_out_pool[key] = (final, partials)
return _out_pool[key]
def _resolve_params(m, n, k):
key = (m, n, k)
if key not in _param_pool:
params = _TUNED_PARAMS.get(key, _FALLBACK_PARAMS).copy()
k_half = k // 2
if params["N_SPLITS"] > 1:
sk_tile, tile_k, n_splits = _adjust_splitk(
k_half, params["TILE_K"], params["N_SPLITS"]
)
params["SK_TILE"] = sk_tile
params["TILE_K"] = tile_k
params["N_SPLITS"] = n_splits
else:
params["SK_TILE"] = 2 * k_half
params["N_SPLITS"] = 1
if params["TILE_K"] >= 2 * k_half:
params["TILE_K"] = triton.next_power_of_2(2 * k_half)
params["SK_TILE"] = 2 * k_half
params["N_SPLITS"] = 1
params["TILE_N"] = max(params["TILE_N"], 32)
_param_pool[key] = params
return _param_pool[key]
def _reshape_operands(B_data, B_scales, n, k_half):
"""Reshape pre-shuffled weight and scale tensors for kernel indexing."""
key = B_data.data_ptr()
if key not in _operand_pool:
b_flat = B_data.view(torch.uint8).reshape(n // 16, k_half * 16)
s_flat = B_scales.view(torch.uint8)
_operand_pool[key] = (b_flat, s_flat)
return _operand_pool[key]
def _build_grid(m, n, k, dev):
"""Precompute all launch parameters for a given problem shape."""
key = (m, n, k)
if key not in _grid_pool:
params = _resolve_params(m, n, k)
k_half = k // 2
ns = params["N_SPLITS"]
final, partials = _get_output_bufs(m, n, ns, dev)
grid = (
ns
* triton.cdiv(m, params["TILE_M"])
* triton.cdiv(n, params["TILE_N"]),
)
if ns == 1:
sk_o, sm_o, sn_o = 0, final.stride(0), final.stride(1)
else:
sk_o = partials.stride(0)
sm_o = partials.stride(1)
sn_o = partials.stride(2)
info = {
'params': params,
'k_half': k_half,
'grid': grid,
'ns': ns,
'sk_o': sk_o,
'sm_o': sm_o,
'sn_o': sn_o,
}
if ns > 1:
info['red_grid'] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
info['true_splits'] = triton.cdiv(k_half, (params["SK_TILE"] // 2))
info['max_splits'] = triton.next_power_of_2(ns)
_grid_pool[key] = info
return _grid_pool[key]
def _execute_matmul(inp_mat, wt_data, wt_scales, m, n, k):
"""Run fused quantize + GEMM with optional split-K reduction."""
info = _build_grid(m, n, k, inp_mat.device)
final, partials = _get_output_bufs(m, n, info['ns'], inp_mat.device)
b_flat, s_flat = _reshape_operands(wt_data, wt_scales, n, info['k_half'])
_matmul_fused_kernel[info['grid']](
inp_mat, b_flat,
final if info['ns'] == 1 else partials,
s_flat,
m, n, info['k_half'],
inp_mat.stride(0), inp_mat.stride(1),
b_flat.stride(0), b_flat.stride(1),
info['sk_o'], info['sm_o'], info['sn_o'],
s_flat.stride(0), s_flat.stride(1),
**info['params'],
)
if info['ns'] > 1:
_partial_sum_kernel[info['red_grid']](
partials, final, m, n,
partials.stride(0), partials.stride(1), partials.stride(2),
final.stride(0), final.stride(1),
16, 64,
info['true_splits'], info['max_splits'],
)
return final
def custom_kernel(data: input_t) -> output_t:
A = data[0]
return _execute_matmul(
A, data[3], data[4], A.shape[0], data[1].shape[0], A.shape[1]
)
scrolls · 443 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 734513.
#!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355X- # v64_wavesched — aiter update + eviction_policy + TRITON_HIP_ENABLE_WAVE_SCHEDULING=1- # Best of session 55: marginal but consistent improvement over v62+ """+ Fused BF16-to-FP4 quantization + scaled GEMM with pre-shuffled weights.+ Hardware-accelerated quantization via v_cvt_scalef32_pk_fp4_bf16.+ Per-shape tuned tile/split-K parameters for MI355X (256 CUs).+ """+ from task import input_t, output_t+ import torch+ import triton+ import triton.language as tl+ import sys as _sys+ import gc as _gc- import os as _os- import sys as _isys- import subprocess as _isp- import time as _itime- _os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")- _os.environ.setdefault("CXX", "clang++")+ torch.set_grad_enabled(False)+ _gc.disable()+ _sys.setswitchinterval(1.0)- _IT0 = _itime.time()- _pe = lambda msg: print(msg, file=_isys.stderr, flush=True)- # ============================================================- # PHASE 0: Update aiter to origin/main (has MI355X tuned configs)- # ============================================================- _AITER_DIR = '/home/runner/aiter'- _AITER_UPDATED = False- try:- _pe("[v62] Fetching origin/main...")- _r = _isp.run(['git', '-C', _AITER_DIR, 'fetch', 'origin', 'main'],- capture_output=True, text=True, timeout=60)- _pe(f" fetch: rc={_r.returncode}")+ @triton.jit+ def _quant_block_fp4(+ inp_bf16,+ TILE_M: tl.constexpr,+ TILE_K: tl.constexpr,+ ):+ """Convert BF16 tile to packed MXFP4 using hardware pk_fp4 instruction."""+ GRP_SZ: tl.constexpr = 32+ N_GROUPS: tl.constexpr = TILE_K // GRP_SZ- # Save current HEAD for rollback- _r0 = _isp.run(['git', '-C', _AITER_DIR, 'rev-parse', 'HEAD'],- capture_output=True, text=True, timeout=5)- _OLD_HEAD = _r0.stdout.strip()- _pe(f" old HEAD: {_OLD_HEAD[:12]}")+ vals_f32 = inp_bf16.to(tl.float32).reshape(TILE_M, N_GROUPS, GRP_SZ)- # Checkout origin/main- _r = _isp.run(['git', '-C', _AITER_DIR, 'checkout', 'origin/main'],- capture_output=True, text=True, timeout=30)- _pe(f" checkout origin/main: rc={_r.returncode}")- if _r.stderr.strip():- _pe(f" checkout err: {_r.stderr.strip()[:200]}")+ # Block-wise absolute max with rounding to nearest power-of-2+ peak = tl.max(tl.abs(vals_f32), axis=-1, keep_dims=True)+ peak = peak.to(tl.int32, bitcast=True)+ peak = (peak + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000+ log2_peak = ((peak >> 23) & 0xFF).to(tl.int32) - 127+ exp_unbiased = log2_peak - 2+ exp_unbiased = tl.minimum(tl.maximum(exp_unbiased, -127), 127)+ scale_e8m0 = exp_unbiased.to(tl.uint8) + 127- if _r.returncode == 0:- _r2 = _isp.run(['git', '-C', _AITER_DIR, 'log', '--oneline', '-5'],- capture_output=True, text=True, timeout=5)- _pe(f" new HEAD:\n{_r2.stdout.strip()}")- _AITER_UPDATED = True- else:- _pe(" checkout FAILED, staying on old HEAD")- except Exception as _e:- _pe(f" [aiter update] FAILED: {_e}")+ # Build IEEE754 float divisor: 2^unbiased for hw instruction+ div_bits = (exp_unbiased.to(tl.int32) + 127).to(tl.uint32) << 23+ divisor = div_bits.to(tl.float32, bitcast=True) # [M, N_GROUPS, 1]- # PHASE 0b removed — eviction_policy now applied via in-memory _unsafe_update_src (Patch 4)- _KERN_PATCHED = False+ # Expand divisor to per-pair level+ div_full = tl.broadcast_to(divisor, (TILE_M, N_GROUPS, GRP_SZ))+ div_full = div_full.reshape(TILE_M, TILE_K)+ div_pairs = div_full.reshape(TILE_M, TILE_K // 2, 2)+ div_even, _ = tl.split(div_pairs)+ div_per_pair = div_even.reshape(TILE_M, TILE_K // 2)- # ============================================================- # PHASE 0c: Read new tuned configs if available- # ============================================================- try:- _cfg_path = '/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv'- if _os.path.exists(_cfg_path):- with open(_cfg_path) as _f:- _cfg_lines = _f.readlines()- _pe(f"[v62] Tuned config: {len(_cfg_lines)} lines")- # Print first few + last few lines- for _l in _cfg_lines[:3]:- _pe(f" {_l.rstrip()}")- if len(_cfg_lines) > 6:- _pe(" ...")- for _l in _cfg_lines[-3:]:- _pe(f" {_l.rstrip()}")+ # Interleave adjacent BF16 elements into uint32 for hw conversion+ inp_u16 = inp_bf16.to(tl.uint16, bitcast=True).reshape(TILE_M, TILE_K // 2, 2)+ lo_half, hi_half = tl.split(inp_u16)+ packed_u32 = lo_half.to(tl.uint32) | (hi_half.to(tl.uint32) << 16)+ packed_u32 = packed_u32.reshape(TILE_M, TILE_K // 2)- # Check for MI355X-specific or new entries- _mi355_lines = [l for l in _cfg_lines if '256' in l.split(',')[0:1]]- _pe(f" entries with 256 CUs: {len(_mi355_lines)}")- except Exception as _e:- _pe(f" [tuned cfg] {_e}")+ # Hardware FP4 pack-convert+ fp4_raw = tl.inline_asm_elementwise(+ "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",+ "=v, v, v",+ [packed_u32, div_per_pair],+ dtype=tl.uint32,+ is_pure=True,+ pack=1,+ )+ fp4_bytes = (fp4_raw & 0xFF).to(tl.uint8)+ fp4_bytes = fp4_bytes.reshape(TILE_M, TILE_K // 2)- _pe(f"[v62] Init phase: {_itime.time()-_IT0:.1f}s, updated={_AITER_UPDATED}, patched={_KERN_PATCHED}")- del _isp, _itime, _pe, _IT0+ return fp4_bytes, scale_e8m0.reshape(TILE_M, N_GROUPS)- import uuid as _uuid- _os.environ["TRITON_CACHE_DIR"] = f"/tmp/_triton_v64_{_uuid.uuid4().hex[:8]}"- _os.environ["TRITON_HIP_ENABLE_WAVE_SCHEDULING"] = "1"- _KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"- _CSV_PATH = "/tmp/_mxfp4_mm_config.csv"- _CU = 256- _NK_FAMILIES = [- (2880, 512), (2112, 7168), (4096, 512), (7168, 2048), (3072, 1536),- (2880, 1536), (4096, 1536), (2112, 512), (2112, 2048),- (7168, 512), (7168, 1536), (7168, 7168), (3072, 512),- (3072, 7168), (3072, 2048), (4096, 2048), (4096, 7168),- (2880, 2048), (2880, 7168),- ]- _M_VALUES = [1, 2, 4, 8, 16, 32, 64, 128, 256]- _lines = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]- for _n, _k in _NK_FAMILIES:- for _m in _M_VALUES:- _tile_num = ((_m + 31) // 32) * ((_n + 127) // 128)- _cus_per_tile = _CU / max(_tile_num, 1)- _split = 0- while _cus_per_tile >= pow(2, _split + 1) and (pow(2, _split + 1) * 128) < 2 * _k:- _split += 1- _split = min(_split, 3)- _lines.append(f"{_CU},{_m},{_n},{_k},21,{_split},1.0,{_KERNEL_32x128},0,0,0.0")- with open(_CSV_PATH, "w") as _f:- _f.write("\n".join(_lines))- # Always use ONLY our CSV — prevents module_gemm_common/a4w4_asm builds (5s+ overhead)- # Our Triton preshuffle kernel bypasses the CSV entirely for actual computation- _os.environ["AITER_CONFIG_GEMM_A4W4"] = _CSV_PATH+ @triton.heuristics(+ {+ "ALIGNED_K": lambda args: (args["K"] % (args["TILE_K"] // 2) == 0)+ and (args["SK_TILE"] % args["TILE_K"] == 0)+ and (args["K"] % (args["SK_TILE"] // 2) == 0),+ }+ )+ @triton.jit+ def _matmul_fused_kernel(+ inp_ptr, wt_ptr, out_ptr, wsc_ptr,+ M, N, K,+ stride_im, stride_ik,+ stride_wn, stride_wk,+ stride_ok, stride_om, stride_on,+ stride_sn, stride_sk,+ TILE_M: tl.constexpr,+ TILE_N: tl.constexpr,+ TILE_K: tl.constexpr,+ GROUP_M: tl.constexpr,+ N_SPLITS: tl.constexpr,+ SK_TILE: tl.constexpr,+ ALIGNED_K: tl.constexpr,+ num_warps: tl.constexpr,+ num_stages: tl.constexpr,+ waves_per_eu: tl.constexpr,+ matrix_instr_nonkdim: tl.constexpr,+ load_modifier: tl.constexpr,+ ):+ tl.assume(stride_im > 0)+ tl.assume(stride_ik > 0)+ tl.assume(stride_wn > 0)+ tl.assume(stride_wk > 0)+ tl.assume(stride_om > 0)+ tl.assume(stride_on > 0)+ tl.assume(stride_sn > 0)+ tl.assume(stride_sk > 0)- import torch- torch.set_grad_enabled(False)- import triton- import triton.language as tl- import sys as _sys- import time as _time- import gc as _gc- _sys.setswitchinterval(1.0)+ SCALE_GRP: tl.constexpr = 32+ n_tiles_m = tl.cdiv(M, TILE_M)+ n_tiles_n = tl.cdiv(N, TILE_N)- # Import with rollback safety — if updated aiter breaks, revert to old HEAD- try:- from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (- _gemm_a16wfp4_preshuffle_kernel,- )- from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (- _gemm_afp4wfp4_reduce_kernel,- )- print("[v62] aiter import OK", file=_sys.stderr, flush=True)- except Exception as _import_err:- print(f"[v62] aiter import FAILED: {_import_err}, rolling back...", file=_sys.stderr, flush=True)- import subprocess as _rbsp- try:- _rbsp.run(['git', '-C', '/home/runner/aiter', 'checkout', _OLD_HEAD],- capture_output=True, text=True, timeout=15)- import importlib- # Re-import with old code- from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (- _gemm_a16wfp4_preshuffle_kernel,- )- from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (- _gemm_afp4wfp4_reduce_kernel,- )- print("[v62] rollback OK, using old aiter", file=_sys.stderr, flush=True)- _AITER_UPDATED = False- except Exception as _rb_err:- print(f"[v62] rollback FAILED: {_rb_err}", file=_sys.stderr, flush=True)- raise _import_err- del _rbsp+ flat_pid = tl.program_id(axis=0)+ split_id = flat_pid % N_SPLITS+ tile_pid = flat_pid // N_SPLITS- from task import input_t, output_t+ # Tile assignment: grouped swizzle for single-split, linear for multi-split+ if N_SPLITS == 1:+ tiles_per_grp = GROUP_M * n_tiles_n+ grp = tile_pid // tiles_per_grp+ first_m = grp * GROUP_M+ grp_sz = min(n_tiles_m - first_m, GROUP_M)+ tile_m = first_m + ((tile_pid % tiles_per_grp) % grp_sz)+ tile_n = (tile_pid % tiles_per_grp) // grp_sz+ else:+ tile_m = tile_pid // n_tiles_n+ tile_n = tile_pid % n_tiles_n- # --- Monkey-patch heuristics ---- try:- # v55: restore default GRID_MN (tile grouping for L2 locality)- _gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True- print("[patch] EVEN_K → True (GRID_MN = default)", file=_sys.stderr, flush=True)- except Exception as _e:- print(f"[patch] heuristics failed: {_e}", file=_sys.stderr, flush=True)+ tl.assume(tile_m >= 0)+ tl.assume(tile_n >= 0)+ tl.assume(split_id >= 0)- _os.environ["HIP_FORCE_DEV_KERNARG"] = "1"+ if (split_id * SK_TILE // 2) < K:+ k_iters = tl.cdiv(SK_TILE // 2, TILE_K // 2)- # --- Replace _mxfp4_quant_op with hardware FP4 conversion ---- print("[hwfp4] Replacing _mxfp4_quant_op with hardware FP4 conversion...", file=_sys.stderr, flush=True)- try:- _jit_fn = _gemm_a16wfp4_preshuffle_kernel.fn if hasattr(_gemm_a16wfp4_preshuffle_kernel, 'fn') else _gemm_a16wfp4_preshuffle_kernel- _quant_fn = _jit_fn.__globals__['_mxfp4_quant_op']- _old_qsrc = _quant_fn._src+ # A: BF16 input [M, 2*K]+ row_a = (tile_m * TILE_M + tl.arange(0, TILE_M)) % M+ col_a = split_id * SK_TILE + tl.arange(0, TILE_K)+ ptrs_a = inp_ptr + (row_a[:, None] * stride_im + col_a[None, :] * stride_ik)- # Complete replacement of _mxfp4_quant_op with hardware FP4 instruction- _new_qsrc = '''def _mxfp4_quant_op(- x,- BLOCK_SIZE_N,- BLOCK_SIZE_M,- MXFP4_QUANT_BLOCK_SIZE,- ):- """Hardware-accelerated BF16->MXFP4 using v_cvt_scalef32_pk_fp4_bf16."""- NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE- HALF_BLOCK: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE // 2+ # B: pre-shuffled FP4 weights [N//16, K_packed*16]+ shuf_range = tl.arange(0, (TILE_K // 2) * 16)+ shuf_base = split_id * (SK_TILE // 2) * 16 + shuf_range+ row_b = (tile_n * (TILE_N // 16) + tl.arange(0, TILE_N // 16)) % (N // 16)+ ptrs_b = wt_ptr + (row_b[:, None] * stride_wn + shuf_base[None, :] * stride_wk)- x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)+ # B scales: shuffled E8M0 layout+ row_s = (tile_n * TILE_N + tl.arange(0, TILE_N // 32) * 32)+ col_s = (split_id * (SK_TILE // SCALE_GRP) * 32) + tl.arange(+ 0, TILE_K // SCALE_GRP * 32+ )+ ptrs_s = wsc_ptr + row_s[:, None] * stride_sn + col_s[None, :] * stride_sk- # Compute amax per group of 32 (same as original)- amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)- amax = amax.to(tl.int32, bitcast=True)- amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000+ acc = tl.zeros((TILE_M, TILE_N), dtype=tl.float32)- # E8M0 scale computation (v19 integer bit ops)- amax_exp = (amax >> 23) & 0xFF- scale_e8m0_unbiased = (amax_exp.to(tl.int32) - 129).to(tl.float32)- scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)+ for ki in range(split_id * k_iters, (split_id + 1) * k_iters):+ # Issue all loads before compute for memory-level parallelism+ if ALIGNED_K:+ a_tile = tl.load(ptrs_a, eviction_policy="evict_last")+ s_raw = tl.load(ptrs_s, cache_modifier=load_modifier)+ b_raw = tl.load(ptrs_b, cache_modifier=load_modifier)+ else:+ k_off = (ki - split_id * k_iters) * TILE_K+ a_tile = tl.load(+ ptrs_a,+ mask=tl.arange(0, TILE_K)[None, :] < (2 * K - split_id * SK_TILE - k_off),+ other=0.0,+ eviction_policy="evict_last",+ )+ s_raw = tl.load(ptrs_s, cache_modifier=load_modifier)+ b_raw = tl.load(+ ptrs_b,+ mask=shuf_range[None, :] < ((K - (split_id * (SK_TILE // 2) + (ki - split_id * k_iters) * (TILE_K // 2))) * 16),+ other=0,+ cache_modifier=load_modifier,+ )- # E8M0 scale bytes for output- bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.float32).to(tl.uint8)+ # On-the-fly A quantization+ inp_q, inp_sc = _quant_block_fp4(a_tile, TILE_M, TILE_K)- # Hardware scale: DIVISOR (confirmed by probe: scale=0.5 gives fp4(x/0.5)=fp4(2x))- # Instruction computes: fp4 = round_to_fp4(bf16 / hw_scale)- # We want: fp4 = round(x / 2^scale_e8m0_unbiased)- # So hw_scale = 2^scale_e8m0_unbiased, constructed via IEEE 754 bit manipulation- # biased_exp = scale_unbiased + 127, clamped to [1, 254] (avoid 0 which gives float 0.0)- biased_exp_f = tl.maximum(scale_e8m0_unbiased + 127.0, 1.0)- hw_scale = (biased_exp_f.to(tl.int32).to(tl.uint32) << 23).to(tl.float32, bitcast=True)+ # Reconstruct B scale layout from shuffled storage+ b_sc = (+ s_raw+ .reshape(+ TILE_N // 32,+ TILE_K // SCALE_GRP // 8,+ 4, 16, 2, 2, 1,+ )+ .permute(0, 5, 3, 1, 4, 2, 6)+ .reshape(TILE_N, TILE_K // SCALE_GRP)+ )- # Convert to BF16 for hardware instruction (x may be float32 from auto-promotion)- x_bf16 = x.to(tl.bfloat16)- x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HALF_BLOCK, 2)- evens, odds = tl.split(x_pairs) # each [BM, NQ, HALF_BLOCK]- lo = evens.to(tl.uint16, bitcast=True).to(tl.uint32)- hi = odds.to(tl.uint16, bitcast=True).to(tl.uint32)- packed_bf16 = lo | (hi << 16) # [BM, NQ, HALF_BLOCK]+ # Reconstruct B tile from shuffled storage+ b_tile = (+ b_raw.reshape(1, TILE_N // 16, TILE_K // 64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5)+ .reshape(TILE_N, TILE_K // 2)+ .trans(1, 0)+ )- # Hardware FP4 conversion!- # hw_scale [BM, NQ, 1] broadcasts to [BM, NQ, HALF_BLOCK] implicitly- result = tl.inline_asm_elementwise(- "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",- "=v,v,v",- [packed_bf16, hw_scale],- dtype=tl.uint32,- is_pure=True,- pack=1,- )+ acc = tl.dot_scaled(+ inp_q, inp_sc, "e2m1", b_tile, b_sc, "e2m1", acc,+ fast_math=True,+ )- # Extract byte 0 (the 2 packed FP4 nibbles)- x_fp4 = (result & 0xFF).to(tl.uint8)- x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)+ ptrs_a += TILE_K * stride_ik+ ptrs_b += (TILE_K // 2) * 16 * stride_wk+ ptrs_s += TILE_K * stride_sk- return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)- '''+ out_vals = acc.to(out_ptr.type.element_ty)- if hasattr(_quant_fn, '_unsafe_update_src'):- _quant_fn._unsafe_update_src(_new_qsrc)- else:- _quant_fn._src = _new_qsrc- if hasattr(_quant_fn, 'src'):- _quant_fn.src = _new_qsrc- if hasattr(_quant_fn, 'hash'):- _quant_fn.hash = None-- # Also modify the KERNEL source to bust its Triton cache key- _old_ksrc = _jit_fn._src- # Patch 1: acc=accumulator (avoids extra zero-init)- _new_ksrc = _old_ksrc.replace(- 'accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")',- 'accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc=accumulator)'- )- # Patch 2: fast_math=True (relaxed FP precision for MFMA scheduling)- _new_ksrc = _new_ksrc.replace(- 'acc=accumulator)',- 'acc=accumulator, fast_math=True)'- )- # Patch 3: .wt store modifier (write-through — avoids L2 pollution from output writes)- _new_ksrc = _new_ksrc.replace(- 'tl.store(c_ptrs, c, mask=c_mask)',- 'tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")'- )- # Patch 4: eviction_policy for A loads (keep A in L2 for N-tile reuse)- _evict_count = 0- if 'a_bf16 = tl.load(a_ptrs)' in _new_ksrc:- _new_ksrc = _new_ksrc.replace(- 'a_bf16 = tl.load(a_ptrs)',- 'a_bf16 = tl.load(a_ptrs, eviction_policy="evict_last")'+ row_o = tile_m * TILE_M + tl.arange(0, TILE_M).to(tl.int64)+ col_o = tile_n * TILE_N + tl.arange(0, TILE_N).to(tl.int64)+ ptrs_o = (+ out_ptr+ + stride_om * row_o[:, None]+ + stride_on * col_o[None, :]+ + split_id * stride_ok)- _evict_count += 1- # Also patch masked A load (non-EVEN_K path)- if 'a_bf16 = tl.load(a_ptrs,' in _new_ksrc and 'eviction_policy' not in _new_ksrc.split('a_bf16 = tl.load(a_ptrs,')[1].split(')')[0]:- # More robust: find "a_bf16 = tl.load(\n a_ptrs,\n mask="- # and insert eviction_policy before mask- import re as _re- _pat = r'(a_bf16 = tl\.load\(\s*\n\s*a_ptrs,)\s*\n(\s*mask=)'- _rep = r'\1 eviction_policy="evict_last",\n\2'- _new_ksrc2 = _re.sub(_pat, _rep, _new_ksrc)- if _new_ksrc2 != _new_ksrc:- _new_ksrc = _new_ksrc2- _evict_count += 1- print(f"[hwfp4] eviction_policy patches: {_evict_count}", file=_sys.stderr, flush=True)- _n_patches = sum([- _new_ksrc != _old_ksrc,- 'fast_math=True' in _new_ksrc,- 'cache_modifier=".wt"' in _new_ksrc,- _evict_count > 0,- ])- if _new_ksrc != _old_ksrc:- _jit_fn._unsafe_update_src(_new_ksrc)- print(f"[hwfp4] Applied hardware quant + {_n_patches} kernel patches", file=_sys.stderr, flush=True)- else:- print("[hwfp4] Applied hardware quant, kernel mod FAILED", file=_sys.stderr, flush=True)+ mask_o = (row_o[:, None] < M) & (col_o[None, :] < N)+ tl.store(ptrs_o, out_vals, mask=mask_o, cache_modifier=".wt")- # Verify- _vq = _quant_fn._src if hasattr(_quant_fn, '_src') else ''- print(f"[hwfp4] quant has inline_asm: {'inline_asm_elementwise' in _vq}",- file=_sys.stderr, flush=True)- except Exception as _e:- import traceback- print(f"[hwfp4] FAILED: {_e}", file=_sys.stderr, flush=True)- traceback.print_exc(file=_sys.stderr)- # --- HIP reduce kernel (same as v19) ---- _HIP_REDUCE_SRC = r"""- #include <hip/hip_runtime.h>+ @triton.jit+ def _partial_sum_kernel(+ partials_ptr, final_ptr, M, N,+ stride_pk, stride_pm, stride_pn,+ stride_fm, stride_fn,+ RED_M: tl.constexpr, RED_N: tl.constexpr,+ TRUE_SPLITS: tl.constexpr, MAX_SPLITS: tl.constexpr,+ ):+ """Reduce split-K partial results into final output."""+ pid_m = tl.program_id(axis=0)+ pid_n = tl.program_id(axis=1)+ rows = (pid_m * RED_M + tl.arange(0, RED_M)) % M+ cols = (pid_n * RED_N + tl.arange(0, RED_N)) % N- __device__ __forceinline__ unsigned short f32_to_bf16(float f) {- unsigned int u;- __builtin_memcpy(&u, &f, sizeof(u));- unsigned int rounding_bias = ((u >> 16) & 1) + 0x7FFFu;- return (unsigned short)((u + rounding_bias) >> 16);- }+ base = (+ partials_ptr+ + (rows[:, None] * stride_pm)+ + (cols[None, :] * stride_pn)+ )+ total = tl.load(base).to(tl.float32)+ for s in tl.static_range(1, MAX_SPLITS):+ if s < TRUE_SPLITS:+ total += tl.load(base + s * stride_pk).to(tl.float32)- template <int KSPLIT>- __global__ void reduce_k_vec4(const float* __restrict__ pp,- unsigned short* __restrict__ out, int MN) {- int idx4 = (blockIdx.x * blockDim.x + threadIdx.x) * 4;- if (idx4 + 3 < MN) {- float4 s = *reinterpret_cast<const float4*>(pp + idx4);- #pragma unroll- for (int k = 1; k < KSPLIT; k++) {- float4 v = *reinterpret_cast<const float4*>(pp + k * MN + idx4);- s.x += v.x; s.y += v.y; s.z += v.z; s.w += v.w;- }- unsigned short r0 = f32_to_bf16(s.x);- unsigned short r1 = f32_to_bf16(s.y);- unsigned short r2 = f32_to_bf16(s.z);- unsigned short r3 = f32_to_bf16(s.w);- *reinterpret_cast<unsigned long long*>(out + idx4) =- (unsigned long long)r0 | ((unsigned long long)r1 << 16) |- ((unsigned long long)r2 << 32) | ((unsigned long long)r3 << 48);- } else {- for (int i = idx4; i < MN && i < idx4 + 4; i++) {- float s = pp[i];- #pragma unroll- for (int k = 1; k < KSPLIT; k++) s += pp[k * MN + i];- out[i] = f32_to_bf16(s);- }- }- }+ out_vals = total.to(final_ptr.type.element_ty)+ out_ptrs = (+ final_ptr+ + (rows[:, None] * stride_fm)+ + (cols[None, :] * stride_fn)+ )+ tl.store(out_ptrs, out_vals)- __global__ void reduce_k_gen(const float* __restrict__ pp,- unsigned short* __restrict__ out, int MN, int ksplit) {- int idx = blockIdx.x * blockDim.x + threadIdx.x;- if (idx < MN) {- float s = pp[idx];- for (int k = 1; k < ksplit; k++) s += pp[k * MN + idx];- out[idx] = f32_to_bf16(s);- }- }- void reduce_op(torch::Tensor pp, torch::Tensor out, int M, int N, int ksplit) {- int MN = M * N;- const float* pp_ptr = pp.data_ptr<float>();- unsigned short* out_ptr = reinterpret_cast<unsigned short*>(out.data_ptr());- const int threads_v = 64;- const int elems_per_block = threads_v * 4;- const int blocks_v = (MN + elems_per_block - 1) / elems_per_block;- switch (ksplit) {- case 2: reduce_k_vec4<2><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;- case 3: reduce_k_vec4<3><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;- case 4: reduce_k_vec4<4><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;- case 7: reduce_k_vec4<7><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;- case 8: reduce_k_vec4<8><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;- default: {- const int threads = 256;- const int blocks = (MN + threads - 1) / threads;- reduce_k_gen<<<blocks, threads>>>(pp_ptr, out_ptr, MN, ksplit);- break;- }- }- }- """- _HIP_REDUCE_CPP = "void reduce_op(torch::Tensor pp, torch::Tensor out, int M, int N, int ksplit);"-- _USE_HIP_REDUCE = False- try:- from torch.utils.cpp_extension import load_inline as _load_inline- _hip_reduce_t0 = _time.time()- _hip_reduce = _load_inline(- name="mxfp4_reduce_hip",- cpp_sources=[_HIP_REDUCE_CPP],- cuda_sources=[_HIP_REDUCE_SRC],- functions=["reduce_op"],- verbose=False,- extra_cuda_cflags=["--offload-arch=gfx950", "-O3"],+ # ---------------------------------------------------------------------------+ # Split-K adjustment for divisibility constraints+ # ---------------------------------------------------------------------------+ def _adjust_splitk(K, BK, n_splits):+ sk_tile = (+ triton.cdiv((2 * triton.cdiv(K, n_splits)), BK) * BK)- _USE_HIP_REDUCE = True- print(f"[hip] reduce kernel compiled in {_time.time()-_hip_reduce_t0:.1f}s",- file=_sys.stderr, flush=True)- except Exception as _e:- print(f"[hip] reduce kernel FAILED: {_e}", file=_sys.stderr, flush=True)-- # --- Helper functions (same as v19) ---- def _get_splitk(K, BLOCK_SIZE_K, NUM_KSPLIT):- SPLITK_BLOCK_SIZE = (- triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K- )- while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:- if (K % (SPLITK_BLOCK_SIZE // 2) == 0- and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0- and K % (BLOCK_SIZE_K // 2) == 0):+ while n_splits > 1 and BK > 16:+ if (+ K % (sk_tile // 2) == 0+ and sk_tile % BK == 0+ and K % (BK // 2) == 0+ ):break- elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:- NUM_KSPLIT = NUM_KSPLIT // 2- elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:- if NUM_KSPLIT > 1:- NUM_KSPLIT = NUM_KSPLIT // 2- elif BLOCK_SIZE_K > 16:- BLOCK_SIZE_K = BLOCK_SIZE_K // 2- elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:- BLOCK_SIZE_K = BLOCK_SIZE_K // 2+ elif K % (sk_tile // 2) != 0 and n_splits > 1:+ n_splits = n_splits // 2+ elif sk_tile % BK != 0:+ if n_splits > 1:+ n_splits = n_splits // 2+ elif BK > 16:+ BK = BK // 2+ elif K % (BK // 2) != 0 and BK > 16:+ BK = BK // 2else:break- SPLITK_BLOCK_SIZE = (- triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K+ sk_tile = (+ triton.cdiv((2 * triton.cdiv(K, n_splits)), BK) * BK)- return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT+ n_splits = triton.cdiv(K, (sk_tile // 2))+ return sk_tile, BK, n_splits- _CFG_CACHE = {}+ # ---------------------------------------------------------------------------+ # Per-shape tuned configurations (MI355X, 256 CUs)+ # ---------------------------------------------------------------------------+ _TUNED_PARAMS = {+ (4, 2880, 512): {"TILE_M": 4, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},+ (16, 2112, 7168): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 512, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 14},+ (32, 4096, 512): {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1},+ (32, 2880, 512): {"TILE_M": 8, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},+ (64, 7168, 2048): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1},+ (256, 3072, 1536): {"TILE_M": 16, "TILE_N": 256, "TILE_K": 512, "GROUP_M": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},+ }- def _get_cfg(M, N, K_real):- key = (M, N, K_real)- if key in _CFG_CACHE:- return _CFG_CACHE[key]- K = K_real // 2- if M <= 32:- BLOCK_M = 8- BLOCK_N = 128- tiles_128 = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)- KSPLIT = 1- if K_real >= 4096:- KSPLIT = 7- elif K_real >= 2048:- if tiles_128 * 2 >= (_CU * 3) // 4 and tiles_128 * 2 <= _CU:- KSPLIT = 2- else:- KSPLIT = 4- elif K_real >= 1536:- if tiles_128 * 2 >= (_CU * 3) // 4 and tiles_128 * 2 <= _CU:- KSPLIT = 2- else:- KSPLIT = 3- BLOCK_K = 256 if K_real <= KSPLIT * 512 or (KSPLIT == 2 and K_real <= KSPLIT * 1024) else 512- if tiles_128 * KSPLIT < (_CU * 3) // 4:- BLOCK_N = 64- wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT- cfg = {- "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,- "waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,- }- else:- BLOCK_M = 16- if M <= 128:- tiles_bm16 = ((M + 15) // 16) * ((N + 127) // 128)- if tiles_bm16 < (_CU * 3) // 4:- BLOCK_M = 8- tiles = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)- BLOCK_N = 128- KSPLIT = 1- if _CU // 2 <= tiles <= _CU and (K_real >= 7168 or (K_real >= 2048 and BLOCK_M == 8)):- KSPLIT = 2- elif tiles < _CU // 2 and K_real > 512:- if K_real >= 4096:- if tiles * 2 >= _CU:- KSPLIT = 2- else:- KSPLIT = 7- elif K_real >= 2048:- KSPLIT = 2- elif K_real >= 1536:- KSPLIT = 3- BLOCK_K = 256 if K_real <= max(KSPLIT * 4096, 2048) else 512- if tiles * KSPLIT < (_CU * 3) // 4:- BLOCK_N = 64- wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT- cfg = {- "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,- "waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,- }+ _FALLBACK_PARAMS = {+ "TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "GROUP_M": 1,+ "num_warps": 2, "num_stages": 2, "waves_per_eu": 0,+ "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1,+ }- if cfg["NUM_KSPLIT"] > 1:- SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = _get_splitk(- K, 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:- cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)- cfg["SPLITK_BLOCK_SIZE"] = 2 * K- cfg["NUM_KSPLIT"] = 1- cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)+ # ---------------------------------------------------------------------------+ # Caches for buffers, configs, and launch grids+ # ---------------------------------------------------------------------------+ _out_pool = {}+ _param_pool = {}+ _grid_pool = {}+ _operand_pool = {}- if cfg["NUM_KSPLIT"] == 1:- cfg["SPLITK_BLOCK_SIZE"] = 2 * K- actual_ksplit = None- nk_pow2 = None- if cfg["NUM_KSPLIT"] > 1:- actual_ksplit = triton.cdiv(K, cfg["SPLITK_BLOCK_SIZE"] // 2)- nk_pow2 = triton.next_power_of_2(cfg["NUM_KSPLIT"])+ def _get_output_bufs(m, n, n_splits, dev):+ key = (m, n, n_splits)+ if key not in _out_pool:+ final = torch.empty((m, n), dtype=torch.bfloat16, device=dev)+ partials = (+ torch.empty((n_splits, m, n), dtype=torch.float32, device=dev)+ if n_splits > 1 else None+ )+ _out_pool[key] = (final, partials)+ return _out_pool[key]- num_m_tiles = triton.cdiv(M, cfg["BLOCK_SIZE_M"])- num_n_tiles = triton.cdiv(N, cfg["BLOCK_SIZE_N"])- total_tiles = num_m_tiles * num_n_tiles- grid_main = (cfg["NUM_KSPLIT"] * total_tiles,)- grid_reduce = None- if cfg["NUM_KSPLIT"] > 1:- grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 16))- result = (cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce,- K, cfg["BLOCK_SIZE_M"], cfg["BLOCK_SIZE_N"], cfg["BLOCK_SIZE_K"],- cfg["NUM_KSPLIT"], cfg["SPLITK_BLOCK_SIZE"], cfg["waves_per_eu"])- _CFG_CACHE[key] = result- return result-- # --- Nuclear pre-warming ---- _WARMUP_T0 = _time.time()- _PREWARMED_CONFIGS = {}- _NO_LSR = {}- _LSR = {}- _REDUCE = set()-- for _nw, _kw in _NK_FAMILIES:- for _mw in _M_VALUES:- _cw, _aw, _nkw, _, _, _, _, _, _, _, _, _ = _get_cfg(_mw, _nw, _kw)- _ck = (_cw["BLOCK_SIZE_M"], _cw["BLOCK_SIZE_N"], _cw["BLOCK_SIZE_K"],- _cw["NUM_KSPLIT"], _cw["SPLITK_BLOCK_SIZE"], _cw["waves_per_eu"])- if _mw <= 32 and _kw >= 1536:- _NO_LSR.setdefault(_ck, True)- elif _mw > 32:- _NO_LSR.setdefault(_ck, True) # v53: M>32 also without disable-lsr+ def _resolve_params(m, n, k):+ key = (m, n, k)+ if key not in _param_pool:+ params = _TUNED_PARAMS.get(key, _FALLBACK_PARAMS).copy()+ k_half = k // 2+ if params["N_SPLITS"] > 1:+ sk_tile, tile_k, n_splits = _adjust_splitk(+ k_half, params["TILE_K"], params["N_SPLITS"]+ )+ params["SK_TILE"] = sk_tile+ params["TILE_K"] = tile_k+ params["N_SPLITS"] = n_splitselse:- _LSR.setdefault(_ck, True)- if _aw is not None:- _REDUCE.add((_aw, _nkw))+ params["SK_TILE"] = 2 * k_half+ params["N_SPLITS"] = 1+ if params["TILE_K"] >= 2 * k_half:+ params["TILE_K"] = triton.next_power_of_2(2 * k_half)+ params["SK_TILE"] = 2 * k_half+ params["N_SPLITS"] = 1+ params["TILE_N"] = max(params["TILE_N"], 32)+ _param_pool[key] = params+ return _param_pool[key]- for _k in _NO_LSR:- _LSR.pop(_k, None)- print(f"[pre-warm] {len(_NO_LSR)} no-lsr + {len(_LSR)} lsr GEMM, {len(_REDUCE)} reduce configs",- file=_sys.stderr, flush=True)+ def _reshape_operands(B_data, B_scales, n, k_half):+ """Reshape pre-shuffled weight and scale tensors for kernel indexing."""+ key = B_data.data_ptr()+ if key not in _operand_pool:+ b_flat = B_data.view(torch.uint8).reshape(n // 16, k_half * 16)+ s_flat = B_scales.view(torch.uint8)+ _operand_pool[key] = (b_flat, s_flat)+ return _operand_pool[key]- _wA = torch.zeros(32, 8192, dtype=torch.bfloat16, device="cuda")- _wBw = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")- _wBs = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")- _wypp = torch.zeros(16, 32, 256, dtype=torch.float32, device="cuda")- _wy = torch.zeros(32, 256, dtype=torch.bfloat16, device="cuda")- def _pw(bm, bn, bk, ks, spk, wpe):- c = {"BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,- "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": ks, "SPLITK_BLOCK_SIZE": spk}- o = _wypp if ks > 1 else _wy- _gemm_a16wfp4_preshuffle_kernel[(max(ks, 1),)](- _wA, _wBw, o, _wBs, bm, bn, spk // 2,- _wA.stride(0), _wA.stride(1), _wBw.stride(0), _wBw.stride(1),- 0 if ks <= 1 else _wypp.stride(0),- _wy.stride(0) if ks <= 1 else _wypp.stride(1),- _wy.stride(1) if ks <= 1 else _wypp.stride(2),- _wBs.stride(0), _wBs.stride(1), PREQUANT=True, **c)+ def _build_grid(m, n, k, dev):+ """Precompute all launch parameters for a given problem shape."""+ key = (m, n, k)+ if key not in _grid_pool:+ params = _resolve_params(m, n, k)+ k_half = k // 2+ ns = params["N_SPLITS"]- print("[pre-warm] Phase 1: M≤32 K>=1536 (no disable-lsr)...", file=_sys.stderr, flush=True)- for _ck in sorted(_NO_LSR):- try:- _pw(*_ck)- _PREWARMED_CONFIGS[_ck] = "no-lsr"- print(f" BM={_ck[0]} BN={_ck[1]} BK={_ck[2]} KS={_ck[3]} SPK={_ck[4]} wpe={_ck[5]} ({_time.time()-_WARMUP_T0:.0f}s)",- file=_sys.stderr, flush=True)- except Exception as _e:- print(f" {_ck}: FAIL {_e}", file=_sys.stderr, flush=True)+ final, partials = _get_output_bufs(m, n, ns, dev)- _os.environ["DISABLE_LLVM_OPT"] = "disable-lsr"- print(f"[pre-warm] Phase 2: DISABLE_LLVM_OPT=disable-lsr ({_time.time()-_WARMUP_T0:.0f}s)",- file=_sys.stderr, flush=True)+ grid = (+ ns+ * triton.cdiv(m, params["TILE_M"])+ * triton.cdiv(n, params["TILE_N"]),+ )- _lsr_list = sorted(_LSR)- print(f"[pre-warm] Phase 3: {len(_lsr_list)} remaining GEMM configs...", file=_sys.stderr, flush=True)- for _idx, _ck in enumerate(_lsr_list):- if _time.time() - _WARMUP_T0 > 200:- print(f" timeout — {len(_lsr_list) - _idx} skipped", file=_sys.stderr, flush=True)- break- try:- _pw(*_ck)- _PREWARMED_CONFIGS[_ck] = "lsr"- print(f" BM={_ck[0]} BN={_ck[1]} BK={_ck[2]} KS={_ck[3]} SPK={_ck[4]} wpe={_ck[5]} ({_time.time()-_WARMUP_T0:.0f}s)",- file=_sys.stderr, flush=True)- except Exception as _e:- print(f" {_ck}: FAIL {_e}", file=_sys.stderr, flush=True)+ if ns == 1:+ sk_o, sm_o, sn_o = 0, final.stride(0), final.stride(1)+ else:+ sk_o = partials.stride(0)+ sm_o = partials.stride(1)+ sn_o = partials.stride(2)- print(f"[pre-warm] Phase 4: {len(_REDUCE)} reduce configs...", file=_sys.stderr, flush=True)- for _ak, _nk in sorted(_REDUCE):- if _time.time() - _WARMUP_T0 > 230:- print(" timeout — remaining skipped", file=_sys.stderr, flush=True)- break- try:- _gemm_afp4wfp4_reduce_kernel[(1, 1)](- _wypp, _wy, 16, 16,- _wypp.stride(0), _wypp.stride(1), _wypp.stride(2),- _wy.stride(0), _wy.stride(1), 16, 16, _ak, _nk)- except Exception:- pass+ info = {+ 'params': params,+ 'k_half': k_half,+ 'grid': grid,+ 'ns': ns,+ 'sk_o': sk_o,+ 'sm_o': sm_o,+ 'sn_o': sn_o,+ }- del _wA, _wBw, _wBs, _wypp, _wy, _pw- del _NO_LSR, _LSR, _REDUCE, _lsr_list- torch.cuda.empty_cache()- print(f"[pre-warm] Done: {len(_PREWARMED_CONFIGS)} GEMM configs in {_time.time()-_WARMUP_T0:.0f}s",- file=_sys.stderr, flush=True)+ if ns > 1:+ info['red_grid'] = (triton.cdiv(m, 16), triton.cdiv(n, 64))+ info['true_splits'] = triton.cdiv(k_half, (params["SK_TILE"] // 2))+ info['max_splits'] = triton.next_power_of_2(ns)- _gc.disable()+ _grid_pool[key] = info+ return _grid_pool[key]- # --- Runtime ---- _PRESHUFFLE_CACHE = {}- _OUT_BUF = {}- _YPP_BUF = {}- _LOGGED = set()- def _get_preshuffle_b(data):- key = data[3].data_ptr()- if key not in _PRESHUFFLE_CACHE:- N = data[3].shape[0]- K_bytes = data[3].shape[1]- sm, sn = data[4].shape- N_groups = N // 32- B_w = data[3].view(torch.uint8).reshape(N // 16, K_bytes * 16)- B_s = data[4].view(torch.uint8).reshape(sm // 32, sn * 32)[:N_groups].contiguous()- _PRESHUFFLE_CACHE[key] = (B_w, B_s, B_w.stride(0), B_s.stride(0))- return _PRESHUFFLE_CACHE[key]+ def _execute_matmul(inp_mat, wt_data, wt_scales, m, n, k):+ """Run fused quantize + GEMM with optional split-K reduction."""+ info = _build_grid(m, n, k, inp_mat.device)+ final, partials = _get_output_bufs(m, n, info['ns'], inp_mat.device)- def custom_kernel(data: input_t) -> output_t:- A = data[0]- if not A.is_contiguous():- A = A.contiguous()- _ndim = A.ndim- if _ndim == 2:- A_2d = A- M = A.shape[0]- else:- A_2d = A.view(-1, A.shape[-1])- M = A_2d.shape[0]- N = data[3].shape[0]- K_bytes = data[3].shape[1]- K_real = K_bytes * 2+ b_flat, s_flat = _reshape_operands(wt_data, wt_scales, n, info['k_half'])- cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce, K, BM, BN, BK, KS, SPK, WPE = _get_cfg(M, N, K_real)+ _matmul_fused_kernel[info['grid']](+ inp_mat, b_flat,+ final if info['ns'] == 1 else partials,+ s_flat,+ m, n, info['k_half'],+ inp_mat.stride(0), inp_mat.stride(1),+ b_flat.stride(0), b_flat.stride(1),+ info['sk_o'], info['sm_o'], info['sn_o'],+ s_flat.stride(0), s_flat.stride(1),+ **info['params'],+ )- _sk = (M, N, K_real)- if _sk not in _LOGGED:- _LOGGED.add(_sk)- print(f"[kernel] M={M} N={N} K={K_real} BM={BM} BN={BN} BK={BK} KS={KS} wpe={WPE} grid={grid_main[0]}",- file=_sys.stderr, flush=True)+ if info['ns'] > 1:+ _partial_sum_kernel[info['red_grid']](+ partials, final, m, n,+ partials.stride(0), partials.stride(1), partials.stride(2),+ final.stride(0), final.stride(1),+ 16, 64,+ info['true_splits'], info['max_splits'],+ )- okey = (M, N)- if okey not in _OUT_BUF:- _OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device="cuda")- y = _OUT_BUF[okey]+ return final- B_w, B_s, stride_bw0, stride_bs0 = _get_preshuffle_b(data)- if KS > 1:- ppkey = (nk_pow2, M, N)- if ppkey not in _YPP_BUF:- _YPP_BUF[ppkey] = torch.empty((nk_pow2, M, N), dtype=torch.float32, device="cuda")- y_pp = _YPP_BUF[ppkey]- stride_ck = M * N- stride_cm = N- else:- y_pp = None- stride_ck = 0- stride_cm = N-- _gemm_a16wfp4_preshuffle_kernel[grid_main](- A_2d, B_w,- y if y_pp is None else y_pp,- B_s, M, N, K,- K_real, 1, stride_bw0, 1,- stride_ck, stride_cm, 1,- stride_bs0, 1,- BLOCK_SIZE_M=BM, BLOCK_SIZE_N=BN, BLOCK_SIZE_K=BK,- GROUP_SIZE_M=1, NUM_KSPLIT=KS, SPLITK_BLOCK_SIZE=SPK,- num_warps=4, num_stages=2, waves_per_eu=WPE,- matrix_instr_nonkdim=16, cache_modifier=".cg",- PREQUANT=True,+ def custom_kernel(data: input_t) -> output_t:+ A = data[0]+ return _execute_matmul(+ A, data[3], data[4], A.shape[0], data[1].shape[0], A.shape[1])-- if y_pp is not None:- if _USE_HIP_REDUCE:- _hip_reduce.reduce_op(y_pp, y, M, N, actual_ksplit)- else:- _gemm_afp4wfp4_reduce_kernel[grid_reduce](- y_pp, y, M, N,- M * N, N, 1, N, 1,- 16, 16, actual_ksplit, nk_pow2,- )-- if _ndim == 2:- return y- return y.view(*A.shape[:-1], N)
scrolls · 1067 diff lines total
Best evidence level for this revision: reported
JSON