submission 550508
div22 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 641 lines, June 9 Researcher Reciprocity License v1.0.
solution_new_25t_1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-550508?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:a9af7d6037fb569e7ed901880f443b1832e1295be4d51f8e010270d3815ee5cf
license declaredunknown
license concludedunknown
authorsdiv22
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
solution_new_25t_1.py641 lines
"""
MXFP4 GEMM v25t_1 — v25t + two-phase exhaustive config search on gfx950 (MI355X).
Phase 1: sweep tile params (BM, BN, BK, NUM_KSPLIT) with default (warps, stages, etc.)
Phase 2: sweep (warps, stages, wpe, gsm, cm) around best tile config.
Edit SEARCH_SHAPES below to control which benchmark shapes to search per submission.
"""
# ============================================================================
# EDIT THIS: which benchmark shapes to search this run.
# Comment/uncomment to split across multiple submissions (~5-12 min each).
# ============================================================================
SEARCH_SHAPES = {
#(4, 2880, 512), # ~3-4 min
(16, 2112, 7168), # ~8-12 min
# (32, 4096, 512), # ~5-7 min
# (32, 2880, 512), # ~5-7 min
# (64, 7168, 2048), # ~8-12 min
# (256, 3072, 1536), # ~9-13 min
}
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
import itertools
from typing import Tuple
import torch
from torch.utils.cpp_extension import load_inline
import triton
import triton.language as tl
import uuid
# ---------------------------------------------------------------------------
# C++ quant kernel (from v25d)
# ---------------------------------------------------------------------------
HIP_KERNEL = r"""
#include <hip/hip_runtime.h>
#include <stdint.h>
using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
__device__ __forceinline__ uint8_t hw_bf16x2_to_fp4x2(uint32_t bf16_pair, float scale) {
uint32_t result;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(result) : "v"(bf16_pair), "v"(scale));
return static_cast<uint8_t>(result & 0xFFu);
}
__global__ void __launch_bounds__(128, 4)
mxfp4_quant(
const __bf16* __restrict__ A_bf16,
uint8_t* __restrict__ A_fp4,
uint8_t* __restrict__ A_scale,
int M, int K)
{
const int KS = K / 32;
const int K2 = K / 2;
const int group = blockIdx.x * 128 + threadIdx.x;
const int row = group / KS;
const int kg = group % KS;
if (row >= M) return;
const auto* src = A_bf16 + (long)row * K + kg * 32;
float absMax = 1e-10f;
#pragma unroll
for (int i = 0; i < 32; ++i) {
float v = __builtin_elementwise_abs(static_cast<float>(src[i]));
absMax = (v > absMax) ? v : absMax;
}
uint32_t u32 = __builtin_bit_cast(uint32_t, absMax);
const uint32_t amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
const uint32_t inv_exp = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
A_scale[(long)row * KS + kg] = static_cast<uint8_t>(inv_exp);
const float hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);
const uint32_t* src_u32 = reinterpret_cast<const uint32_t*>(src);
auto* dst = reinterpret_cast<uint8_t*>(A_fp4 + (long)row * K2 + kg * 16);
#pragma unroll
for (int i = 0; i < 16; ++i) {
dst[i] = hw_bf16x2_to_fp4x2(src_u32[i], hw_scale);
}
}
extern "C" void launch_quant(
const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)
{
const int KS = K / 32;
const int n_groups = M * KS;
const dim3 block{128};
const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);
}
"""
CPP = r"""
#include <torch/extension.h>
#include <c10/core/DeviceGuard.h>
extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);
struct QuantWorkspace {
at::Tensor A_fp4;
at::Tensor A_scale;
int64_t last_M = -1, last_K = -1;
void ensure(int M, int K, const at::TensorOptions& opts) {
if (M == last_M && K == last_K) return;
A_fp4 = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));
A_scale = at::empty({(int64_t)M, (int64_t)(K / 32)}, opts.dtype(at::kByte));
last_M = M; last_K = K;
}
};
static QuantWorkspace g_qws;
std::vector<at::Tensor> do_quant(const at::Tensor& A) {
auto guard = at::DeviceGuard(A.device());
at::Tensor A_bf16 = (A.scalar_type() == at::kBFloat16 && A.is_contiguous())
? A : A.to(A.device(), at::kBFloat16, false, false,
at::MemoryFormat::Contiguous);
const int M = A_bf16.size(0);
const int K = A_bf16.size(1);
g_qws.ensure(M, K, A_bf16.options());
launch_quant(
reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>()),
g_qws.A_fp4.data_ptr<uint8_t>(),
g_qws.A_scale.data_ptr<uint8_t>(),
M, K);
return {g_qws.A_fp4, g_qws.A_scale};
}
"""
_ext = load_inline(
name=f"g_{uuid.uuid4().hex[:8]}",
cpp_sources=[CPP],
cuda_sources=[HIP_KERNEL],
functions=["do_quant"],
with_cuda=True,
extra_cflags=["-O3", "-std=c++20"],
extra_cuda_cflags=[
"-O3", "--offload-arch=gfx950", "-ffast-math", "-munsafe-fp-atomics",
"-std=c++20", "-mllvm", "-amdgpu-early-inline-all=true",
"-mllvm", "-amdgpu-function-calls=false", "-mwavefrontsize64",
"-mcumode", "-fgpu-flush-denormals-to-zero",
],
extra_ldflags=["-lamdhip64"],
)
# ---------------------------------------------------------------------------
# Triton helpers
# ---------------------------------------------------------------------------
@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
tall_xcds = GRID_MN % NUM_XCDS
tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
xcd = pid % NUM_XCDS
local_pid = pid // NUM_XCDS
if xcd < tall_xcds:
pid = xcd * pids_per_xcd + local_pid
else:
pid = tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid
return pid
@triton.jit
def pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr = 1):
if GROUP_SIZE_M == 1:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
else:
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + (pid % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
return pid_m, pid_n
# ---------------------------------------------------------------------------
# Triton GEMM kernel (identical to v25t)
# ---------------------------------------------------------------------------
@triton.heuristics({
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
})
@triton.jit
def _mxfp4_gemm_kernel(
a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bn, stride_bk,
stride_ck, stride_cm, stride_cn,
stride_asm, stride_ask, stride_bsn, stride_bsk,
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, cache_modifier: tl.constexpr,
):
tl.assume(stride_am > 0); tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0); tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0); tl.assume(stride_cn > 0)
tl.assume(stride_asm > 0); tl.assume(stride_ask > 0)
tl.assume(stride_bsk > 0); tl.assume(stride_bsn > 0)
GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
SCALE_GROUP_SIZE: tl.constexpr = 32
pid_unified = tl.program_id(axis=0)
pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
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)
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
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[None, :] * stride_ak
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (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_ks_a = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
a_scale_ptrs = a_scales_ptr + offs_am[:, None] * stride_asm + offs_ks_a[None, :] * stride_ask
offs_asn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
offs_ks_b = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
b_scale_ptrs = b_scales_ptr + offs_asn[:, None] * stride_bsn + offs_ks_b[None, :] * stride_bsk
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
a_scales = tl.load(a_scale_ptrs)
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 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0)
b = tl.load(b_ptrs, mask=offs_k_shuffle_arr[None, :] < (K - k * (BLOCK_SIZE_K // 2)) * 16, other=0, 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)
accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
a_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * stride_ask
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
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
@triton.jit
def _mxfp4_reduce_kernel(
c_in_ptr, c_out_ptr, M, N,
stride_c_in_k, stride_c_in_m, stride_c_in_n,
stride_c_out_m, stride_c_out_n,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
offs_k = tl.arange(0, MAX_KSPLIT)
c_in_ptrs = c_in_ptr + offs_k[:, None, None] * stride_c_in_k + offs_m[None, :, None] * stride_c_in_m + offs_n[None, None, :] * stride_c_in_n
if ACTUAL_KSPLIT == MAX_KSPLIT:
c = tl.load(c_in_ptrs)
else:
c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
c = tl.sum(c, axis=0).to(c_out_ptr.type.element_ty)
c_out_ptrs = c_out_ptr + offs_m[:, None] * stride_c_out_m + offs_n[None, :] * stride_c_out_n
tl.store(c_out_ptrs, c)
# ---------------------------------------------------------------------------
# AITER default configs (fallback / seeds)
# ---------------------------------------------------------------------------
def _cfg(bm, bn, bk, gsm, nw, ns, wpe, cm, nks):
return {
"BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
"GROUP_SIZE_M": gsm, "num_warps": nw, "num_stages": ns,
"waves_per_eu": wpe, "matrix_instr_nonkdim": 16,
"cache_modifier": cm, "NUM_KSPLIT": nks,
}
_DEFAULTS = {
(4, 2880, 512): _cfg(8, 64, 512, 1, 2, 1, 1, None, 1),
(16, 2112, 7168): _cfg(16, 32, 512, 1, 4, 1, 4, ".cg", 14),
(32, 4096, 512): _cfg(32, 64, 512, 1, 4, 1, 1, None, 1),
(32, 2880, 512): _cfg(32, 64, 512, 1, 4, 1, 1, None, 1),
(64, 7168, 2048): _cfg(32, 64, 512, 4, 2, 2, 1, None, 1),
(256,3072, 1536): _cfg(128,32, 512, 1, 4, 2, 1, None, 1),
# test shapes
(8, 2112, 7168): _cfg(8, 32, 512, 1, 2, 2, 1, ".cg", 7),
(16, 3072, 1536): _cfg(8, 32, 512, 1, 4, 2, 1, None, 1),
(64, 3072, 1536): _cfg(64, 32, 512, 1, 2, 2, 1, ".cg", 1),
(256,2880, 512): _cfg(32, 64, 512, 1, 2, 2, 1, None, 1),
}
# ---------------------------------------------------------------------------
# splitK helper (from AITER)
# ---------------------------------------------------------------------------
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):
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
else:
break
SPLITK_BLOCK_SIZE = triton.cdiv(2 * triton.cdiv(K, NUM_KSPLIT), BLOCK_SIZE_K) * BLOCK_SIZE_K
NUM_KSPLIT = triton.cdiv(K, SPLITK_BLOCK_SIZE // 2)
return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT
# ---------------------------------------------------------------------------
# Kernel runner: launches GEMM (+ reduce) with a given config
# ---------------------------------------------------------------------------
def _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp):
cfg = dict(config)
BK = cfg["BLOCK_SIZE_K"]
NUM_KSPLIT = cfg["NUM_KSPLIT"]
if BK >= 2 * K_packed:
BK = triton.next_power_of_2(2 * K_packed)
cfg["BLOCK_SIZE_K"] = BK
cfg["NUM_KSPLIT"] = 1
NUM_KSPLIT = 1
cfg["BLOCK_SIZE_K"] = max(cfg["BLOCK_SIZE_K"], 256)
BK = cfg["BLOCK_SIZE_K"]
if NUM_KSPLIT > 1:
SPLITK_BS, BK, NUM_KSPLIT = get_splitk(K_packed, BK, NUM_KSPLIT)
cfg["BLOCK_SIZE_K"] = BK
cfg["NUM_KSPLIT"] = NUM_KSPLIT
cfg["SPLITK_BLOCK_SIZE"] = SPLITK_BS
else:
cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed
c_out = y_pp if NUM_KSPLIT > 1 else y
KS = K_elem // 32
scaleN = ((KS + 7) // 8) * 8
grid = lambda META: (
META["NUM_KSPLIT"]
* triton.cdiv(M, META["BLOCK_SIZE_M"])
* triton.cdiv(N, META["BLOCK_SIZE_N"]),
)
_mxfp4_gemm_kernel[grid](
A_fp4, w, c_out, A_scale, b_scales,
M, N, K_packed,
A_fp4.stride(0), A_fp4.stride(1),
w.stride(0), w.stride(1),
0 if NUM_KSPLIT == 1 else y_pp.stride(0),
c_out.stride(-2), c_out.stride(-1),
A_scale.stride(0), A_scale.stride(1),
32 * scaleN, 1,
**cfg,
)
if NUM_KSPLIT > 1:
ACTUAL_KSPLIT = triton.cdiv(K_packed, cfg["SPLITK_BLOCK_SIZE"] // 2)
grid_r = (triton.cdiv(M, 16), triton.cdiv(N, 64))
_mxfp4_reduce_kernel[grid_r](
y_pp, y, M, N,
y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
y.stride(0), y.stride(1),
16, 64, ACTUAL_KSPLIT, triton.next_power_of_2(NUM_KSPLIT),
)
return y
# ---------------------------------------------------------------------------
# Two-phase exhaustive config search
# ---------------------------------------------------------------------------
_CONFIG_CACHE = {}
def _validate_config(config, K_packed):
BN = config["BLOCK_SIZE_N"]
BK = config["BLOCK_SIZE_K"]
if BN < 32:
return False
if BK < 64:
return False
if K_packed % (BK // 2) != 0:
return False
return True
def _time_config(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed,
n_warmup=2, n_iter=8):
NUM_KSPLIT = config.get("NUM_KSPLIT", 1)
y = torch.empty((M, N), dtype=torch.bfloat16, device=A_fp4.device)
y_pp = None
if NUM_KSPLIT > 1:
y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A_fp4.device)
for _ in range(n_warmup):
_run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp)
torch.cuda.synchronize()
start_evt = torch.cuda.Event(enable_timing=True)
end_evt = torch.cuda.Event(enable_timing=True)
start_evt.record()
for _ in range(n_iter):
_run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp)
end_evt.record()
end_evt.synchronize()
return start_evt.elapsed_time(end_evt) / n_iter * 1000 # ms -> us
def _get_param_choices(M, N, K_elem):
K_packed = K_elem // 2
bm_choices = [b for b in [8, 16, 32, 64, 128, 256] if b <= max(M * 2, 8)]
bn_choices = [b for b in [32, 64, 128, 256] if b <= N]
bk_choices = [b for b in [256, 512, 1024] if K_packed % (b // 2) == 0]
if not bk_choices:
bk_choices = [256]
max_splits = K_packed // (min(bk_choices) // 2)
sk_choices = [1]
for s in [2, 3, 4, 7, 14]:
if s <= max_splits and K_packed % s == 0:
sk_choices.append(s)
return bm_choices, bn_choices, bk_choices, sk_choices
def _search_config(M, N, K_elem, A_fp4, A_scale, w, b_scales):
K_packed = K_elem // 2
bm_choices, bn_choices, bk_choices, sk_choices = _get_param_choices(M, N, K_elem)
# Phase 1: sweep tile params (BM, BN, BK, NUM_KSPLIT) with sensible defaults
tile_combos = list(itertools.product(bm_choices, bn_choices, bk_choices, sk_choices))
p1_total = len(tile_combos)
print(f" Phase 1: {p1_total} tile combos (BM x BN x BK x SK = "
f"{len(bm_choices)}x{len(bn_choices)}x{len(bk_choices)}x{len(sk_choices)})", flush=True)
best_time = float("inf")
best_tile = None
best_config = _DEFAULTS.get((M, N, K_elem), _cfg(16, 64, 256, 4, 2, 2, 1, None, 1))
for i, (bm, bn, bk, sk) in enumerate(tile_combos):
cfg = {
"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": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": None, "NUM_KSPLIT": sk,
}
if not _validate_config(cfg, K_packed):
continue
try:
t = _time_config(cfg, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed)
except Exception:
continue
if t < best_time:
best_time = t
best_tile = (bm, bn, bk, sk)
best_config = cfg
print(f" [P1 {i+1}/{p1_total}] best={t:.1f}us BM={bm} BN={bn} BK={bk} SK={sk}", flush=True)
if best_tile is None:
print(" Phase 1: no valid tile found, using default", flush=True)
return best_config
bm, bn, bk, sk = best_tile
print(f" Phase 1 winner: BM={bm} BN={bn} BK={bk} SK={sk} = {best_time:.1f}us", flush=True)
# Phase 2: sweep (num_warps, num_stages, waves_per_eu, GROUP_SIZE_M, cache_modifier)
tune_combos = list(itertools.product(
[2, 4], # num_warps
[1, 2], # num_stages
[1, 2, 4], # waves_per_eu
[1, 4, 8], # GROUP_SIZE_M
[None, ".cg"], # cache_modifier
))
p2_total = len(tune_combos)
print(f" Phase 2: {p2_total} tune combos around best tile", flush=True)
for i, (nw, ns, wpe, gsm, cm) in enumerate(tune_combos):
cfg = {
"BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
"GROUP_SIZE_M": gsm, "num_warps": nw, "num_stages": ns,
"waves_per_eu": wpe, "matrix_instr_nonkdim": 16,
"cache_modifier": cm, "NUM_KSPLIT": sk,
}
try:
t = _time_config(cfg, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed)
except Exception:
continue
if t < best_time:
best_time = t
best_config = cfg
print(f" [P2 {i+1}/{p2_total}] best={t:.1f}us nw={nw} ns={ns} wpe={wpe} gsm={gsm} cm={cm}", flush=True)
print(f" Phase 2 winner: {best_time:.1f}us", flush=True)
return best_config
# ---------------------------------------------------------------------------
# Workspace caching
# ---------------------------------------------------------------------------
class _Workspace:
__slots__ = ["y", "y_pp", "_key"]
def __init__(self):
self.y = None; self.y_pp = None; self._key = None
def ensure(self, M, N, num_ksplit, device):
key = (M, N, num_ksplit)
if key == self._key: return
self._key = key
self.y = torch.empty((M, N), dtype=torch.bfloat16, device=device)
if num_ksplit > 1:
self.y_pp = torch.empty((num_ksplit, M, N), dtype=torch.float32, device=device)
else:
self.y_pp = None
_ws = _Workspace()
_b_cache = {"bsh_dp": 0, "bssh_dp": 0, "w": None, "b_scales": None}
# ---------------------------------------------------------------------------
# Main entry point
# ---------------------------------------------------------------------------
def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
"""MXFP4 GEMM v25t_1: two-phase exhaustive config search."""
A, _, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous().cuda()
M, K_elem = A.shape
N = B_q.shape[0]
K_packed = K_elem // 2
# Quant
quant_result = _ext.do_quant(A)
A_fp4 = quant_result[0]
A_scale = quant_result[1]
# Cache B views
bsh_dp = B_shuffle.data_ptr()
bssh_dp = B_scale_sh.data_ptr()
if bsh_dp != _b_cache["bsh_dp"] or bssh_dp != _b_cache["bssh_dp"]:
w_raw = B_shuffle.view(torch.uint8) if B_shuffle.dtype != torch.uint8 else B_shuffle
_b_cache["w"] = w_raw.reshape(N // 16, K_packed * 16).contiguous()
bs_raw = B_scale_sh.view(torch.uint8) if B_scale_sh.dtype != torch.uint8 else B_scale_sh
_b_cache["b_scales"] = bs_raw.contiguous()
_b_cache["bsh_dp"] = bsh_dp
_b_cache["bssh_dp"] = bssh_dp
w = _b_cache["w"]
b_scales = _b_cache["b_scales"]
# --- Search only for shapes listed in SEARCH_SHAPES; defaults for rest ---
shape_key = (M, N, K_elem)
if shape_key not in _CONFIG_CACHE:
if shape_key in SEARCH_SHAPES:
print(f"[v25t_1] exhaustive search for M={M} N={N} K={K_elem} ...", flush=True)
try:
best = _search_config(M, N, K_elem, A_fp4, A_scale, w, b_scales)
_CONFIG_CACHE[shape_key] = best
cm = best.get("cache_modifier", None)
print(f"[v25t_1] BEST M={M} N={N} K={K_elem}: "
f"BM={best['BLOCK_SIZE_M']} BN={best['BLOCK_SIZE_N']} "
f"BK={best['BLOCK_SIZE_K']} warps={best['num_warps']} "
f"stages={best['num_stages']} wpe={best['waves_per_eu']} "
f"GSM={best['GROUP_SIZE_M']} splitK={best['NUM_KSPLIT']} "
f"cache={cm}", flush=True)
except Exception as e:
print(f"[v25t_1] search failed: {e}, using default", flush=True)
_CONFIG_CACHE[shape_key] = _DEFAULTS.get(
shape_key, _cfg(16, 64, 256, 4, 2, 2, 1, None, 1)
)
else:
# Test shape — use defaults, no search
_CONFIG_CACHE[shape_key] = _DEFAULTS.get(
shape_key, _cfg(16, 64, 256, 4, 2, 2, 1, None, 1)
)
print(f"[v25t_1] test shape M={M} N={N} K={K_elem}, using default", flush=True)
config = _CONFIG_CACHE[shape_key]
NUM_KSPLIT = config.get("NUM_KSPLIT", 1)
_ws.ensure(M, N, NUM_KSPLIT, A.device)
return _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed,
_ws.y, _ws.y_pp)
scrolls · 641 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 546301.
"""- MXFP4 GEMM v25d — gfx950 (MI355X) optimized.+ MXFP4 GEMM v25t_1 — v25t + two-phase exhaustive config search on gfx950 (MI355X).- Changes from v25b (13.459μs):- 1. Full (N,K) template specialization + precomputed views (same as v25b)- 2. Aggressive compiler flags:- - -amdgpu-loop-prefetch: software prefetch for K-loop loads- - -enable-unroll-and-jam: fuse nested loop unrolling- - -ffinite-math-only: assume no NaN/Inf (beyond -ffast-math)- - -amdgpu-set-wave-priority: dynamic wave priority- - Increased unroll thresholds+ Phase 1: sweep tile params (BM, BN, BK, NUM_KSPLIT) with default (warps, stages, etc.)+ Phase 2: sweep (warps, stages, wpe, gsm, cm) around best tile config.++ Edit SEARCH_SHAPES below to control which benchmark shapes to search per submission."""++ # ============================================================================+ # EDIT THIS: which benchmark shapes to search this run.+ # Comment/uncomment to split across multiple submissions (~5-12 min each).+ # ============================================================================+ SEARCH_SHAPES = {+ #(4, 2880, 512), # ~3-4 min+ (16, 2112, 7168), # ~8-12 min+ # (32, 4096, 512), # ~5-7 min+ # (32, 2880, 512), # ~5-7 min+ # (64, 7168, 2048), # ~8-12 min+ # (256, 3072, 1536), # ~9-13 min+ }import osos.environ["PYTORCH_ROCM_ARCH"] = "gfx950"+ import itertoolsfrom typing import Tupleimport torchfrom torch.utils.cpp_extension import load_inline+ import triton+ import triton.language as tlimport uuid+ # ---------------------------------------------------------------------------+ # C++ quant kernel (from v25d)+ # ---------------------------------------------------------------------------+HIP_KERNEL = r"""#include <hip/hip_runtime.h>#include <stdint.h>- using int4_v = int __attribute__((ext_vector_type(4)));- using float4_v = float __attribute__((ext_vector_type(4)));- using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));+ using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));- static constexpr int FP4_E2M1 = 4;-- __device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(- int4_v a, int4_v b, float4_v c,- int cbsz, int blgp, int op_sel_a, int scale_a, int op_sel_b, int scale_b- ) __asm("llvm.amdgcn.mfma.scale.f32.16x16x128.f8f6f4.v4i32.v4i32");-- __device__ __forceinline__ int4_v load16(const uint8_t* __restrict__ p) {- return *reinterpret_cast<const int4_v*>(p);- }-- __device__ __forceinline__ uint16_t float_to_bf16(float f) {- bf16x2 v;- v[0] = static_cast<__bf16>(f);- uint16_t r;- __builtin_memcpy(&r, &v, sizeof(r));- return r;- }-__device__ __forceinline__ uint8_t hw_bf16x2_to_fp4x2(uint32_t bf16_pair, float scale) {uint32_t result;asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"⋯ 13 unchanged linesconst int group = blockIdx.x * 128 + threadIdx.x;const int row = group / KS;const int kg = group % KS;-if (row >= M) return;const auto* src = A_bf16 + (long)row * K + kg * 32;-float absMax = 1e-10f;#pragma unrollfor (int i = 0; i < 32; ++i) {⋯ 7 unchanged linesA_scale[(long)row * KS + kg] = static_cast<uint8_t>(inv_exp);const float hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);-const uint32_t* src_u32 = reinterpret_cast<const uint32_t*>(src);auto* dst = reinterpret_cast<uint8_t*>(A_fp4 + (long)row * K2 + kg * 16);#pragma unroll⋯ 2 unchanged lines}}- template<bool ALWAYS_VALID>- __device__ __forceinline__ int4_v load_or_zero(bool rt_valid, const uint8_t* p) {- if constexpr (ALWAYS_VALID) return load16(p);- else return rt_valid ? load16(p) : int4_v{0,0,0,0};+ extern "C" void launch_quant(+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)+ {+ const int KS = K / 32;+ const int n_groups = M * KS;+ const dim3 block{128};+ const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};+ mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);}+ """- // Main GEMM kernel — fully specialized on NK dimensions- template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,- int CKT, int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS>- __global__ void __launch_bounds__(NWARPS * 64, (NWARPS <= 2) ? 4 : 2)- mxfp4_gemm(- const uint8_t* __restrict__ A,- const uint8_t* __restrict__ As,- const uint8_t* __restrict__ Bsh,- const uint8_t* __restrict__ Bssh,- float* __restrict__ C_partial,- uint16_t* __restrict__ C_final,- int M,- int tile_off_x, int tile_off_y)- {- static_assert(BN % 16 == 0);- constexpr int WAVES_M = (BM + 15) / 16;- constexpr int WAVES_N = BN / 16;- static_assert(WAVES_M * WAVES_N == NWARPS);+ CPP = r"""+ #include <torch/extension.h>+ #include <c10/core/DeviceGuard.h>- constexpr int N = C_N;- constexpr int K = C_K;- constexpr int scaleN = C_SCALEN;- constexpr int K2 = K / 2;- constexpr int KS = K / 32;- constexpr long bsh_n_stride = (long)(K / 64) * 512;+ extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);- const int ks_idx = blockIdx.z;- const int lane = threadIdx.x % 64;- const int wave = threadIdx.x / 64;- const int wave_m = wave / WAVES_N;- const int wave_n = wave % WAVES_N;+ struct QuantWorkspace {+ at::Tensor A_fp4;+ at::Tensor A_scale;+ int64_t last_M = -1, last_K = -1;+ void ensure(int M, int K, const at::TensorOptions& opts) {+ if (M == last_M && K == last_K) return;+ A_fp4 = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));+ A_scale = at::empty({(int64_t)M, (int64_t)(K / 32)}, opts.dtype(at::kByte));+ last_M = M; last_K = K;+ }+ };- const int tile_m = (blockIdx.y + tile_off_y) * BM + wave_m * 16;- const int tile_n = (blockIdx.x + tile_off_x) * BN + wave_n * 16;+ static QuantWorkspace g_qws;- if (tile_m >= M || tile_n >= N) return;+ std::vector<at::Tensor> do_quant(const at::Tensor& A) {+ auto guard = at::DeviceGuard(A.device());+ at::Tensor A_bf16 = (A.scalar_type() == at::kBFloat16 && A.is_contiguous())+ ? A : A.to(A.device(), at::kBFloat16, false, false,+ at::MemoryFormat::Contiguous);+ const int M = A_bf16.size(0);+ const int K = A_bf16.size(1);+ g_qws.ensure(M, K, A_bf16.options());+ launch_quant(+ reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>()),+ g_qws.A_fp4.data_ptr<uint8_t>(),+ g_qws.A_scale.data_ptr<uint8_t>(),+ M, K);+ return {g_qws.A_fp4, g_qws.A_scale};+ }+ """- constexpr int ktiles_per_split = C_KPS;- const int ks_start = ks_idx * ktiles_per_split;+ _ext = load_inline(+ name=f"g_{uuid.uuid4().hex[:8]}",+ cpp_sources=[CPP],+ cuda_sources=[HIP_KERNEL],+ functions=["do_quant"],+ with_cuda=True,+ extra_cflags=["-O3", "-std=c++20"],+ extra_cuda_cflags=[+ "-O3", "--offload-arch=gfx950", "-ffast-math", "-munsafe-fp-atomics",+ "-std=c++20", "-mllvm", "-amdgpu-early-inline-all=true",+ "-mllvm", "-amdgpu-function-calls=false", "-mwavefrontsize64",+ "-mcumode", "-fgpu-flush-denormals-to-zero",+ ],+ extra_ldflags=["-lamdhip64"],+ )- const int lrow = lane % 16;- const int kgrp = lane / 16;- const int gm = tile_m + lrow;- const int gn = tile_n + lrow;+ # ---------------------------------------------------------------------------+ # Triton helpers+ # ---------------------------------------------------------------------------- const bool a_rt = A_VALID | (gm < M);- const bool b_rt = B_VALID | (gn < N);+ @triton.jit+ def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):+ pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS+ tall_xcds = GRID_MN % NUM_XCDS+ tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds+ xcd = pid % NUM_XCDS+ local_pid = pid // NUM_XCDS+ if xcd < tall_xcds:+ pid = xcd * pids_per_xcd + local_pid+ else:+ pid = tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid+ return pid- const int n_tile = tile_n / 16;- const auto bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;- const uint8_t* a_row = nullptr;- const uint8_t* as_row = nullptr;- if constexpr (A_VALID) {- a_row = A + (long)gm * K2;- as_row = As + (long)gm * KS;- } else {- if (a_rt) { a_row = A + (long)gm * K2;- as_row = As + (long)gm * KS; }- }+ @triton.jit+ def pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr = 1):+ if GROUP_SIZE_M == 1:+ pid_m = pid // num_pid_n+ pid_n = pid % num_pid_n+ else:+ num_pid_in_group = GROUP_SIZE_M * num_pid_n+ group_id = pid // num_pid_in_group+ first_pid_m = group_id * GROUP_SIZE_M+ group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)+ pid_m = first_pid_m + (pid % group_size_m)+ pid_n = (pid % num_pid_in_group) // group_size_m+ return pid_m, pid_n- int bssh_base = 0;- if constexpr (B_VALID) {- bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);- } else {- if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);- }- const int k_half_off = (kgrp & 1) * 256;- const int k_blk_base = kgrp >> 1;+ # ---------------------------------------------------------------------------+ # Triton GEMM kernel (identical to v25t)+ # ---------------------------------------------------------------------------- static constexpr long a_kt_stride = 64L;- static constexpr long bsh_kt_stride = 1024L;+ @triton.heuristics({+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)+ and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)+ and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),+ })+ @triton.jit+ def _mxfp4_gemm_kernel(+ a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,+ M, N, K,+ stride_am, stride_ak, stride_bn, stride_bk,+ stride_ck, stride_cm, stride_cn,+ stride_asm, stride_ask, stride_bsn, stride_bsk,+ 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, cache_modifier: tl.constexpr,+ ):+ tl.assume(stride_am > 0); tl.assume(stride_ak > 0)+ tl.assume(stride_bk > 0); tl.assume(stride_bn > 0)+ tl.assume(stride_cm > 0); tl.assume(stride_cn > 0)+ tl.assume(stride_asm > 0); tl.assume(stride_ask > 0)+ tl.assume(stride_bsk > 0); tl.assume(stride_bsn > 0)- const uint8_t* a_ptr = nullptr;- const uint8_t* bsh_ptr = nullptr;- const uint8_t* bssh_ptr = nullptr;+ GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)+ SCALE_GROUP_SIZE: tl.constexpr = 32- if constexpr (A_VALID) {- a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;- } else {- if (a_rt) a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;- }- if constexpr (B_VALID) {- bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;- bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;- } else {- if (b_rt) {- bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;- bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;- }- }+ pid_unified = tl.program_id(axis=0)+ pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)+ 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)- float4_v acc{0.f, 0.f, 0.f, 0.f};+ 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- const int bssh_step0 = (ks_start & 1) ? 254 : 2;- const int bssh_step1 = 256 - bssh_step0;+ tl.assume(pid_m >= 0); tl.assume(pid_n >= 0)- #define DO_MFMA(a_off, b_off, bssh_off, ks_val) \- { \- const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \- const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \- const int ks = (ks_val) * 4 + kgrp; \- int sa, sb; \- if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \- else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \- if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \- else sb = (b_rt & (ks < KS)) ? static_cast<int>(*(bssh_ptr + (bssh_off))) : 127; \- acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,acc,FP4_E2M1,FP4_E2M1,0,sa,0,sb); \- }+ if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:+ num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)- static_assert(CKT > 0, "Specialized kernel must have compile-time CKT");- #pragma unroll- for (int q = 0; q < (CKT / 4); ++q) {- DO_MFMA(0, 0, 0, ks_start + q*4)- DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)- DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)- DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, ks_start + q*4 + 3)- a_ptr += 4 * a_kt_stride;- bsh_ptr += 4 * bsh_kt_stride;- bssh_ptr += 512;- }- if constexpr ((CKT % 4) >= 2) {- DO_MFMA(0, 0, 0, ks_start + (CKT/4)*4)- DO_MFMA(1, 1, bssh_step0, ks_start + (CKT/4)*4 + 1)- a_ptr += 2 * a_kt_stride;- bsh_ptr += 2 * bsh_kt_stride;- bssh_ptr += 256;- }- if constexpr ((CKT % 2) == 1) {- DO_MFMA(0, 0, 0, ks_start + CKT - 1)- }- #undef DO_MFMA+ offs_k = tl.arange(0, BLOCK_SIZE_K // 2)+ offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k+ 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[None, :] * stride_ak- const int out_col = tile_n + lrow;- const int out_row_base = tile_m + kgrp * 4;+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)+ offs_k_shuffle = pid_k * (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- if constexpr (!B_VALID) { if (out_col >= N) return; }+ offs_ks_a = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE)+ a_scale_ptrs = a_scales_ptr + offs_am[:, None] * stride_asm + offs_ks_a[None, :] * stride_ask- constexpr bool out_rows_always_valid = A_VALID && (BM >= 16);+ offs_asn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N+ offs_ks_b = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)+ b_scale_ptrs = b_scales_ptr + offs_asn[:, None] * stride_bsn + offs_ks_b[None, :] * stride_bsk- if constexpr (SPLITK) {- auto c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;- #pragma unroll- for (int i = 0; i < 4; ++i) {- if constexpr (out_rows_always_valid) c_out[i * N] = acc[i];- else if (out_row_base + i < M) c_out[i * N] = acc[i];- }- } else {- auto c_out = C_final + (long)out_row_base * N + out_col;- #pragma unroll- for (int i = 0; i < 4; ++i) {- if constexpr (out_rows_always_valid) c_out[i * N] = float_to_bf16(acc[i]);- else if (out_row_base + i < M) c_out[i * N] = float_to_bf16(acc[i]);- }+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)++ for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):+ a_scales = tl.load(a_scale_ptrs)+ 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 = tl.load(a_ptrs)+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)+ else:+ a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0)+ b = tl.load(b_ptrs, mask=offs_k_shuffle_arr[None, :] < (K - k * (BLOCK_SIZE_K // 2)) * 16, other=0, 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)+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)++ a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk+ a_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * stride_ask+ 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+ c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)+ tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")+++ @triton.jit+ def _mxfp4_reduce_kernel(+ c_in_ptr, c_out_ptr, M, N,+ stride_c_in_k, stride_c_in_m, stride_c_in_n,+ stride_c_out_m, stride_c_out_n,+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,+ ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr,+ ):+ pid_m = tl.program_id(axis=0)+ pid_n = tl.program_id(axis=1)+ offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M+ offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N+ offs_k = tl.arange(0, MAX_KSPLIT)+ c_in_ptrs = c_in_ptr + offs_k[:, None, None] * stride_c_in_k + offs_m[None, :, None] * stride_c_in_m + offs_n[None, None, :] * stride_c_in_n+ if ACTUAL_KSPLIT == MAX_KSPLIT:+ c = tl.load(c_in_ptrs)+ else:+ c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)+ c = tl.sum(c, axis=0).to(c_out_ptr.type.element_ty)+ c_out_ptrs = c_out_ptr + offs_m[:, None] * stride_c_out_m + offs_n[None, :] * stride_c_out_n+ tl.store(c_out_ptrs, c)+++ # ---------------------------------------------------------------------------+ # AITER default configs (fallback / seeds)+ # ---------------------------------------------------------------------------++ def _cfg(bm, bn, bk, gsm, nw, ns, wpe, cm, nks):+ return {+ "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,+ "GROUP_SIZE_M": gsm, "num_warps": nw, "num_stages": ns,+ "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,+ "cache_modifier": cm, "NUM_KSPLIT": nks,}++ _DEFAULTS = {+ (4, 2880, 512): _cfg(8, 64, 512, 1, 2, 1, 1, None, 1),+ (16, 2112, 7168): _cfg(16, 32, 512, 1, 4, 1, 4, ".cg", 14),+ (32, 4096, 512): _cfg(32, 64, 512, 1, 4, 1, 1, None, 1),+ (32, 2880, 512): _cfg(32, 64, 512, 1, 4, 1, 1, None, 1),+ (64, 7168, 2048): _cfg(32, 64, 512, 4, 2, 2, 1, None, 1),+ (256,3072, 1536): _cfg(128,32, 512, 1, 4, 2, 1, None, 1),+ # test shapes+ (8, 2112, 7168): _cfg(8, 32, 512, 1, 2, 2, 1, ".cg", 7),+ (16, 3072, 1536): _cfg(8, 32, 512, 1, 4, 2, 1, None, 1),+ (64, 3072, 1536): _cfg(64, 32, 512, 1, 2, 2, 1, ".cg", 1),+ (256,2880, 512): _cfg(32, 64, 512, 1, 2, 2, 1, None, 1),}- template<int C_N>- __global__ void mxfp4_reduce(- const float* __restrict__ C_partial,- uint16_t* __restrict__ C_out,- int M, int NUM_KSPLIT)- {- constexpr int N = C_N;- const int col = blockIdx.x * 32 + threadIdx.x;- const int row = blockIdx.y * 16 + threadIdx.y;- if (row >= M || col >= N) return;- float sum = 0.f;- const long mn = (long)row * N + col;- const long mn_stride = (long)M * N;- for (int k = 0; k < NUM_KSPLIT; ++k)- sum += C_partial[k * mn_stride + mn];+ # ---------------------------------------------------------------------------+ # splitK helper (from AITER)+ # ---------------------------------------------------------------------------- bf16x2 v;- v[0] = static_cast<__bf16>(sum);- uint16_t r;- __builtin_memcpy(&r, &v, sizeof(r));- C_out[mn] = r;- }+ 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):+ 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+ else:+ break+ SPLITK_BLOCK_SIZE = triton.cdiv(2 * triton.cdiv(K, NUM_KSPLIT), BLOCK_SIZE_K) * BLOCK_SIZE_K+ NUM_KSPLIT = triton.cdiv(K, SPLITK_BLOCK_SIZE // 2)+ return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT- extern "C" void launch_quant(- const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)- {- const int KS = K / 32;- const int n_groups = M * KS;- const dim3 block{128};- const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};- mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);- }- template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>- void launch_gemm_nk(- const uint8_t* A, const uint8_t* As,- const uint8_t* Bsh, const uint8_t* Bssh,- float* C_partial, uint16_t* C_final, int M)- {- constexpr bool do_splitk = C_NUM_KSPLIT > 1;+ # ---------------------------------------------------------------------------+ # Kernel runner: launches GEMM (+ reduce) with a given config+ # ---------------------------------------------------------------------------- auto launch = [&]<int BM, int BN, int NWARPS>() {- static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);+ def _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp):+ cfg = dict(config)+ BK = cfg["BLOCK_SIZE_K"]+ NUM_KSPLIT = cfg["NUM_KSPLIT"]- const int full_m = (BM >= 16) ? M / BM : 0;- constexpr int full_n = C_N / BN;- const int total_m = (M + BM - 1) / BM;- constexpr int total_n = (C_N + BN - 1) / BN;- const int edge_m = total_m - full_m;- constexpr int edge_n = total_n - full_n;+ if BK >= 2 * K_packed:+ BK = triton.next_power_of_2(2 * K_packed)+ cfg["BLOCK_SIZE_K"] = BK+ cfg["NUM_KSPLIT"] = 1+ NUM_KSPLIT = 1- const dim3 block{static_cast<uint32_t>(NWARPS * 64)};+ cfg["BLOCK_SIZE_K"] = max(cfg["BLOCK_SIZE_K"], 256)+ BK = cfg["BLOCK_SIZE_K"]- auto sub = [&]<bool AV, bool BV>(int gx, int gy, int ox, int oy) {- if (gx <= 0 || gy <= 0) return;- const dim3 grid{- static_cast<uint32_t>(gx),- static_cast<uint32_t>(gy),- static_cast<uint32_t>(C_NUM_KSPLIT)- };- if constexpr (do_splitk)- mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>- <<<grid,block>>>(A,As,Bsh,Bssh,C_partial,nullptr,M,ox,oy);- else- mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>- <<<grid,block>>>(A,As,Bsh,Bssh,nullptr,C_final,M,ox,oy);- };+ if NUM_KSPLIT > 1:+ SPLITK_BS, BK, NUM_KSPLIT = get_splitk(K_packed, BK, NUM_KSPLIT)+ cfg["BLOCK_SIZE_K"] = BK+ cfg["NUM_KSPLIT"] = NUM_KSPLIT+ cfg["SPLITK_BLOCK_SIZE"] = SPLITK_BS+ else:+ cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed- sub.template operator()<true, true >(full_n, full_m, 0, 0);- sub.template operator()<true, false>(edge_n, full_m, full_n, 0);- sub.template operator()<false, true >(full_n, edge_m, 0, full_m);- sub.template operator()<false, false>(edge_n, edge_m, full_n, full_m);+ c_out = y_pp if NUM_KSPLIT > 1 else y- if constexpr (do_splitk) {- const dim3 rblock{32, 16};- const dim3 rgrid{- static_cast<uint32_t>((C_N + 31) / 32),- static_cast<uint32_t>((M + 15) / 16)- };- mxfp4_reduce<C_N><<<rgrid, rblock>>>(C_partial, C_final, M, C_NUM_KSPLIT);- }- };+ KS = K_elem // 32+ scaleN = ((KS + 7) // 8) * 8- if (M <= 8) launch.template operator()< 8, 32, 2>();- else if (M <= 16) launch.template operator()< 16, 32, 2>();- else if (M <= 32) launch.template operator()< 16, 32, 2>();- else if (M <= 64) launch.template operator()< 32, 32, 4>();- else if (M <=128) launch.template operator()< 32, 32, 4>();- else launch.template operator()< 64, 32, 8>();- }+ grid = lambda META: (+ META["NUM_KSPLIT"]+ * triton.cdiv(M, META["BLOCK_SIZE_M"])+ * triton.cdiv(N, META["BLOCK_SIZE_N"]),+ )- template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT>- void launch_gemm_nk_7168(- const uint8_t* A, const uint8_t* As,- const uint8_t* Bsh, const uint8_t* Bssh,- float* C_partial, uint16_t* C_final, int M)- {- if (M <= 8)- launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 8, 7>(A, As, Bsh, Bssh, C_partial, C_final, M);- else if (M <= 16)- launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 4, 14>(A, As, Bsh, Bssh, C_partial, C_final, M);- else- launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 56, 1>(A, As, Bsh, Bssh, C_partial, C_final, M);- }+ _mxfp4_gemm_kernel[grid](+ A_fp4, w, c_out, A_scale, b_scales,+ M, N, K_packed,+ A_fp4.stride(0), A_fp4.stride(1),+ w.stride(0), w.stride(1),+ 0 if NUM_KSPLIT == 1 else y_pp.stride(0),+ c_out.stride(-2), c_out.stride(-1),+ A_scale.stride(0), A_scale.stride(1),+ 32 * scaleN, 1,+ **cfg,+ )- // Precomputed raw-pointer fast path — avoids ALL tensor ops in hot path- extern "C" void launch_gemm_raw(- const uint8_t* A_fp4, const uint8_t* A_scale,- const uint8_t* Bsh, const uint8_t* Bssh,- float* C_partial, uint16_t* C_final,- int M, int N, int K)- {- if (N == 2880 && K == 512)- launch_gemm_nk<2880, 512, 16, 4, 4, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);- else if (N == 2112 && K == 7168)- launch_gemm_nk_7168<2112, 7168, 224, 56>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);- else if (N == 4096 && K == 512)- launch_gemm_nk<4096, 512, 16, 4, 4, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);- else if (N == 7168 && K == 2048)- launch_gemm_nk<7168, 2048, 64, 16, 16, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);- else if (N == 3072 && K == 1536)- launch_gemm_nk<3072, 1536, 48, 12, 12, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);- }- """+ if NUM_KSPLIT > 1:+ ACTUAL_KSPLIT = triton.cdiv(K_packed, cfg["SPLITK_BLOCK_SIZE"] // 2)+ grid_r = (triton.cdiv(M, 16), triton.cdiv(N, 64))+ _mxfp4_reduce_kernel[grid_r](+ y_pp, y, M, N,+ y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),+ y.stride(0), y.stride(1),+ 16, 64, ACTUAL_KSPLIT, triton.next_power_of_2(NUM_KSPLIT),+ )+ return y- CPP = r"""- #include <torch/extension.h>- #include <c10/core/DeviceGuard.h>- extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);- extern "C" void launch_gemm_raw(const uint8_t*, const uint8_t*, const uint8_t*, const uint8_t*,- float*, uint16_t*, int, int, int);+ # ---------------------------------------------------------------------------+ # Two-phase exhaustive config search+ # ---------------------------------------------------------------------------- static int get_num_ksplit(int M, int K) {- int total_ktiles = K / 128;- if (total_ktiles < 28) return 1;- if (M <= 8) return 7;- if (M <= 16) return 14;- return 1;- }+ _CONFIG_CACHE = {}- // Precomputed workspace — keyed by (M, N, K) tuple- struct ShapeWorkspace {- at::Tensor A_fp4;- at::Tensor A_scale;- at::Tensor C_partial;- at::Tensor C;- uint8_t* a_fp4_ptr = nullptr;- uint8_t* a_scale_ptr = nullptr;- float* c_partial_ptr = nullptr;- uint16_t* c_final_ptr = nullptr;- int M = 0, N = 0, K = 0;- int num_ksplit = 0;- };- // Cache for B tensor pointers (B doesn't change between calls for same N,K)- struct BCache {- const uint8_t* bsh_ptr = nullptr;- const uint8_t* bssh_ptr = nullptr;- int64_t bsh_data_ptr = 0; // for staleness check- int64_t bssh_data_ptr = 0;- };+ def _validate_config(config, K_packed):+ BN = config["BLOCK_SIZE_N"]+ BK = config["BLOCK_SIZE_K"]+ if BN < 32:+ return False+ if BK < 64:+ return False+ if K_packed % (BK // 2) != 0:+ return False+ return True- // Up to 10 different (M,N,K) combos (4 test + 6 bench)- static ShapeWorkspace g_ws[10];- static int g_ws_count = 0;- static BCache g_bcache;- static ShapeWorkspace* find_or_create_ws(int M, int N, int K, int num_ksplit,- const at::TensorOptions& opts) {- // Search existing- for (int i = 0; i < g_ws_count; ++i) {- if (g_ws[i].M == M && g_ws[i].N == N && g_ws[i].K == K)- return &g_ws[i];- }- // Create new- auto& ws = g_ws[g_ws_count++];- ws.M = M; ws.N = N; ws.K = K;- ws.num_ksplit = num_ksplit;- int64_t KS = K / 32;- ws.A_fp4 = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));- ws.A_scale = at::empty({(int64_t)M, KS}, opts.dtype(at::kByte));- ws.C = at::empty({(int64_t)M, (int64_t)N}, opts.dtype(at::kBFloat16));- if (num_ksplit > 1)- ws.C_partial = at::empty({(int64_t)num_ksplit, (int64_t)M, (int64_t)N}, opts.dtype(at::kFloat));- // Cache raw pointers- ws.a_fp4_ptr = ws.A_fp4.data_ptr<uint8_t>();- ws.a_scale_ptr = ws.A_scale.data_ptr<uint8_t>();- ws.c_partial_ptr = (num_ksplit > 1) ? ws.C_partial.data_ptr<float>() : nullptr;- ws.c_final_ptr = reinterpret_cast<uint16_t*>(ws.C.data_ptr<at::BFloat16>());- return &ws;- }+ def _time_config(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed,+ n_warmup=2, n_iter=8):+ NUM_KSPLIT = config.get("NUM_KSPLIT", 1)+ y = torch.empty((M, N), dtype=torch.bfloat16, device=A_fp4.device)+ y_pp = None+ if NUM_KSPLIT > 1:+ y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A_fp4.device)- at::Tensor fwd(const at::Tensor& A,- const at::Tensor& B_q,- const at::Tensor& B_shuffle,- const at::Tensor& B_scale_sh) {- auto guard = at::DeviceGuard(A.device());+ for _ in range(n_warmup):+ _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp)- const int M = A.size(0);- const int K = A.size(1);- const int N = B_q.size(0);+ torch.cuda.synchronize()+ start_evt = torch.cuda.Event(enable_timing=True)+ end_evt = torch.cuda.Event(enable_timing=True)+ start_evt.record()+ for _ in range(n_iter):+ _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp)+ end_evt.record()+ end_evt.synchronize()- // Fast path: check if A is already bf16 contiguous- const __bf16* a_bf16_ptr;- at::Tensor A_bf16;- if (A.scalar_type() == at::kBFloat16 && A.is_contiguous()) {- a_bf16_ptr = reinterpret_cast<const __bf16*>(A.data_ptr<at::BFloat16>());- } else {- A_bf16 = A.to(A.device(), at::kBFloat16, false, false, at::MemoryFormat::Contiguous);- a_bf16_ptr = reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>());- }+ return start_evt.elapsed_time(end_evt) / n_iter * 1000 # ms -> us- // Cache B pointers — B tensors don't change between benchmark iterations- auto bsh_dp = reinterpret_cast<int64_t>(B_shuffle.data_ptr());- auto bssh_dp = reinterpret_cast<int64_t>(B_scale_sh.data_ptr());- if (bsh_dp != g_bcache.bsh_data_ptr || bssh_dp != g_bcache.bssh_data_ptr) {- // First call or B changed — resolve views once- at::Tensor Bsh = B_shuffle.view(at::kByte);- if (!Bsh.is_contiguous()) Bsh = Bsh.contiguous();- at::Tensor Bssh = B_scale_sh.view(at::kByte);- if (!Bssh.is_contiguous()) Bssh = Bssh.contiguous();- g_bcache.bsh_ptr = Bsh.data_ptr<uint8_t>();- g_bcache.bssh_ptr = Bssh.data_ptr<uint8_t>();- g_bcache.bsh_data_ptr = bsh_dp;- g_bcache.bssh_data_ptr = bssh_dp;- }- const int num_ksplit = get_num_ksplit(M, K);- auto* ws = find_or_create_ws(M, N, K, num_ksplit, A.options());+ def _get_param_choices(M, N, K_elem):+ K_packed = K_elem // 2- // Quant: A_bf16 -> A_fp4 + A_scale- launch_quant(a_bf16_ptr, ws->a_fp4_ptr, ws->a_scale_ptr, M, K);+ bm_choices = [b for b in [8, 16, 32, 64, 128, 256] if b <= max(M * 2, 8)]+ bn_choices = [b for b in [32, 64, 128, 256] if b <= N]+ bk_choices = [b for b in [256, 512, 1024] if K_packed % (b // 2) == 0]+ if not bk_choices:+ bk_choices = [256]- // GEMM: all raw pointers, no tensor ops- launch_gemm_raw(- ws->a_fp4_ptr, ws->a_scale_ptr,- g_bcache.bsh_ptr, g_bcache.bssh_ptr,- ws->c_partial_ptr, ws->c_final_ptr,- M, N, K);+ max_splits = K_packed // (min(bk_choices) // 2)+ sk_choices = [1]+ for s in [2, 3, 4, 7, 14]:+ if s <= max_splits and K_packed % s == 0:+ sk_choices.append(s)- return ws->C;- }- """+ return bm_choices, bn_choices, bk_choices, sk_choices- _ext = load_inline(- name=f"g_{uuid.uuid4().hex[:8]}",- cpp_sources=[CPP],- cuda_sources=[HIP_KERNEL],- functions=["fwd"],- with_cuda=True,- extra_cflags=["-O3", "-std=c++20"],- extra_cuda_cflags=[- "-O3",- "--offload-arch=gfx950",- "-ffast-math",- "-ffinite-math-only",- "-munsafe-fp-atomics",- "-std=c++20",- "-mllvm", "-amdgpu-early-inline-all=true",- "-mllvm", "-amdgpu-function-calls=false",- "-mwavefrontsize64",- "-mcumode",- "-mllvm", "--amdgpu-kernarg-preload-count=16",- "-mllvm", "-enable-post-misched=0",- "-mllvm", "--lsr-drop-solution=1",- "-mllvm", "-amdgpu-coerce-illegal-types=1",- "-fgpu-flush-denormals-to-zero",- "-fno-offload-uniform-block",- # New aggressive flags- "-mllvm", "-amdgpu-loop-prefetch=true",- "-mllvm", "-enable-unroll-and-jam=true",- "-mllvm", "-amdgpu-set-wave-priority=true",- "-mllvm", "-unroll-threshold=1000",- "-mllvm", "-amdgpu-internalize-symbols=true",- ],- extra_ldflags=["-lamdhip64"],- )+ def _search_config(M, N, K_elem, A_fp4, A_scale, w, b_scales):+ K_packed = K_elem // 2+ bm_choices, bn_choices, bk_choices, sk_choices = _get_param_choices(M, N, K_elem)+ # Phase 1: sweep tile params (BM, BN, BK, NUM_KSPLIT) with sensible defaults+ tile_combos = list(itertools.product(bm_choices, bn_choices, bk_choices, sk_choices))+ p1_total = len(tile_combos)+ print(f" Phase 1: {p1_total} tile combos (BM x BN x BK x SK = "+ f"{len(bm_choices)}x{len(bn_choices)}x{len(bk_choices)}x{len(sk_choices)})", flush=True)++ best_time = float("inf")+ best_tile = None+ best_config = _DEFAULTS.get((M, N, K_elem), _cfg(16, 64, 256, 4, 2, 2, 1, None, 1))++ for i, (bm, bn, bk, sk) in enumerate(tile_combos):+ cfg = {+ "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": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": None, "NUM_KSPLIT": sk,+ }+ if not _validate_config(cfg, K_packed):+ continue+ try:+ t = _time_config(cfg, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed)+ except Exception:+ continue+ if t < best_time:+ best_time = t+ best_tile = (bm, bn, bk, sk)+ best_config = cfg+ print(f" [P1 {i+1}/{p1_total}] best={t:.1f}us BM={bm} BN={bn} BK={bk} SK={sk}", flush=True)++ if best_tile is None:+ print(" Phase 1: no valid tile found, using default", flush=True)+ return best_config++ bm, bn, bk, sk = best_tile+ print(f" Phase 1 winner: BM={bm} BN={bn} BK={bk} SK={sk} = {best_time:.1f}us", flush=True)++ # Phase 2: sweep (num_warps, num_stages, waves_per_eu, GROUP_SIZE_M, cache_modifier)+ tune_combos = list(itertools.product(+ [2, 4], # num_warps+ [1, 2], # num_stages+ [1, 2, 4], # waves_per_eu+ [1, 4, 8], # GROUP_SIZE_M+ [None, ".cg"], # cache_modifier+ ))+ p2_total = len(tune_combos)+ print(f" Phase 2: {p2_total} tune combos around best tile", flush=True)++ for i, (nw, ns, wpe, gsm, cm) in enumerate(tune_combos):+ cfg = {+ "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,+ "GROUP_SIZE_M": gsm, "num_warps": nw, "num_stages": ns,+ "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,+ "cache_modifier": cm, "NUM_KSPLIT": sk,+ }+ try:+ t = _time_config(cfg, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed)+ except Exception:+ continue+ if t < best_time:+ best_time = t+ best_config = cfg+ print(f" [P2 {i+1}/{p2_total}] best={t:.1f}us nw={nw} ns={ns} wpe={wpe} gsm={gsm} cm={cm}", flush=True)++ print(f" Phase 2 winner: {best_time:.1f}us", flush=True)+ return best_config+++ # ---------------------------------------------------------------------------+ # Workspace caching+ # ---------------------------------------------------------------------------++ class _Workspace:+ __slots__ = ["y", "y_pp", "_key"]+ def __init__(self):+ self.y = None; self.y_pp = None; self._key = None+ def ensure(self, M, N, num_ksplit, device):+ key = (M, N, num_ksplit)+ if key == self._key: return+ self._key = key+ self.y = torch.empty((M, N), dtype=torch.bfloat16, device=device)+ if num_ksplit > 1:+ self.y_pp = torch.empty((num_ksplit, M, N), dtype=torch.float32, device=device)+ else:+ self.y_pp = None++ _ws = _Workspace()+ _b_cache = {"bsh_dp": 0, "bssh_dp": 0, "w": None, "b_scales": None}+++ # ---------------------------------------------------------------------------+ # Main entry point+ # ---------------------------------------------------------------------------+def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:- """MXFP4 GEMM v25d: v25b + aggressive compiler flags."""+ """MXFP4 GEMM v25t_1: two-phase exhaustive config search."""A, _, B_q, B_shuffle, B_scale_sh = data- return _ext.fwd(A.cuda(), B_q, B_shuffle, B_scale_sh)+ A = A.contiguous().cuda()++ M, K_elem = A.shape+ N = B_q.shape[0]+ K_packed = K_elem // 2++ # Quant+ quant_result = _ext.do_quant(A)+ A_fp4 = quant_result[0]+ A_scale = quant_result[1]++ # Cache B views+ bsh_dp = B_shuffle.data_ptr()+ bssh_dp = B_scale_sh.data_ptr()+ if bsh_dp != _b_cache["bsh_dp"] or bssh_dp != _b_cache["bssh_dp"]:+ w_raw = B_shuffle.view(torch.uint8) if B_shuffle.dtype != torch.uint8 else B_shuffle+ _b_cache["w"] = w_raw.reshape(N // 16, K_packed * 16).contiguous()+ bs_raw = B_scale_sh.view(torch.uint8) if B_scale_sh.dtype != torch.uint8 else B_scale_sh+ _b_cache["b_scales"] = bs_raw.contiguous()+ _b_cache["bsh_dp"] = bsh_dp+ _b_cache["bssh_dp"] = bssh_dp++ w = _b_cache["w"]+ b_scales = _b_cache["b_scales"]++ # --- Search only for shapes listed in SEARCH_SHAPES; defaults for rest ---+ shape_key = (M, N, K_elem)+ if shape_key not in _CONFIG_CACHE:+ if shape_key in SEARCH_SHAPES:+ print(f"[v25t_1] exhaustive search for M={M} N={N} K={K_elem} ...", flush=True)+ try:+ best = _search_config(M, N, K_elem, A_fp4, A_scale, w, b_scales)+ _CONFIG_CACHE[shape_key] = best+ cm = best.get("cache_modifier", None)+ print(f"[v25t_1] BEST M={M} N={N} K={K_elem}: "+ f"BM={best['BLOCK_SIZE_M']} BN={best['BLOCK_SIZE_N']} "+ f"BK={best['BLOCK_SIZE_K']} warps={best['num_warps']} "+ f"stages={best['num_stages']} wpe={best['waves_per_eu']} "+ f"GSM={best['GROUP_SIZE_M']} splitK={best['NUM_KSPLIT']} "+ f"cache={cm}", flush=True)+ except Exception as e:+ print(f"[v25t_1] search failed: {e}, using default", flush=True)+ _CONFIG_CACHE[shape_key] = _DEFAULTS.get(+ shape_key, _cfg(16, 64, 256, 4, 2, 2, 1, None, 1)+ )+ else:+ # Test shape — use defaults, no search+ _CONFIG_CACHE[shape_key] = _DEFAULTS.get(+ shape_key, _cfg(16, 64, 256, 4, 2, 2, 1, None, 1)+ )+ print(f"[v25t_1] test shape M={M} N={N} K={K_elem}, using default", flush=True)++ config = _CONFIG_CACHE[shape_key]+ NUM_KSPLIT = config.get("NUM_KSPLIT", 1)++ _ws.ensure(M, N, NUM_KSPLIT, A.device)++ return _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed,+ _ws.y, _ws.y_pp)
scrolls · 1044 diff lines total
Best evidence level for this revision: reported
JSON