submission 527012
_radna · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1455 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-527012?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:8b3d7f7252d0d369d3b97d3a5d05bde337522edbaee1cfcd4542dfea1d0b8e5a
license declaredunknown
license concludedunknown
authors_radna
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission.py1455 lines
"""
Hybrid MXFP4 submission:
- Keep the proven explicit asm dispatch for five benchmark tuples and the general fallback path.
- Replace only the dominant `(16, 2112, 7168)` tuple with a vendored fix of AITER's
broken `gemm_afp4wfp4_preshuffle` Triton kernel.
"""
from functools import lru_cache
from task import input_t, output_t
try:
import triton
import triton.language as tl
except Exception:
triton = None
tl = None
_EXPLICIT_ASM_KERNELS = {
(4, 2880, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
(16, 2112, 7168): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
(32, 4096, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
(32, 2880, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
(64, 7168, 2048): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
(256, 3072, 1536): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
}
_VENDORED_TRITON_PRESHUFFLE_CONFIGS = {
(16, 2112, 7168): {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 1024,
"GROUP_SIZE_M": 1,
"NUM_KSPLIT": 7,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
},
}
_HOT_BIGK_STATE = {"enabled": True, "error": None}
_A16WFP4_PRESHUFFLE_STATE = {"enabled": True, "error": None}
_FUSED_A16WFP4_SHAPES = {
(4, 2880, 512),
(32, 4096, 512),
(32, 2880, 512),
}
_ASM_SPLIT_K = 0
_EXPLICIT_ASM_LOG2_K_SPLIT = {}
_WORKSPACE_CACHE = {}
_WORKSPACE_CACHE_LIMIT = 16
def _cdiv(x, y):
return (x + y - 1) // y
def _next_power_of_two(x):
return 1 if x <= 1 else 1 << (x - 1).bit_length()
def _get_workspace_tensor(tag, reference, shape, dtype):
import torch
device = reference.device
device_index = -1 if device.index is None else device.index
key = (tag, device.type, device_index, str(dtype), tuple(shape))
cached = _WORKSPACE_CACHE.get(key)
if cached is not None:
return cached
if len(_WORKSPACE_CACHE) >= _WORKSPACE_CACHE_LIMIT:
_WORKSPACE_CACHE.clear()
cached = torch.empty(shape, dtype=dtype, device=device)
_WORKSPACE_CACHE[key] = cached
return cached
def _get_splitk(K, block_size_k, num_ksplit):
num_ksplit_step = 2
block_size_k_step = 2
splitk_block_size = _cdiv(2 * _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
if K % (splitk_block_size // 2) != 0 and num_ksplit > 1:
num_ksplit //= num_ksplit_step
elif splitk_block_size % block_size_k != 0:
if num_ksplit > 1:
num_ksplit //= num_ksplit_step
elif block_size_k > 16:
block_size_k //= block_size_k_step
elif K % (block_size_k // 2) != 0 and block_size_k > 16:
block_size_k //= block_size_k_step
else:
break
splitk_block_size = _cdiv(2 * _cdiv(K, num_ksplit), block_size_k) * block_size_k
num_ksplit = _cdiv(K, splitk_block_size // 2)
return splitk_block_size, block_size_k, num_ksplit
if triton is not None:
@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: int, num_pid_m: int, num_pid_n: int, 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 = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
tl.assume(group_size_m >= 0)
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.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 _vendored_gemm_afp4wfp4_preshuffle_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)
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)
scale_group_size: tl.constexpr = 32
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_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak
)
b_ptrs = b_ptr + (
offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
)
offs_bsn = (
pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
) % N
offs_ks = (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_bsn[:, None] * stride_bsn
+ offs_ks[None, :] * stride_bsk
)
if BLOCK_SIZE_M < 32:
offs_ks_non_shuffle = (
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_non_shuffle[None, :] * stride_ask
)
else:
offs_asm = (
pid_m * (BLOCK_SIZE_M // 32) + tl.arange(0, (BLOCK_SIZE_M // 32))
) % M
a_scale_ptrs = (
a_scales_ptr
+ offs_asm[:, None] * stride_asm
+ offs_ks[None, :] * stride_ask
)
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):
if BLOCK_SIZE_M < 32:
a_scales = tl.load(a_scale_ptrs)
else:
a_scales = (
tl.load(a_scale_ptrs)
.reshape(
BLOCK_SIZE_M // 32,
BLOCK_SIZE_K // scale_group_size // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // scale_group_size)
)
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
if BLOCK_SIZE_M < 32:
a_scale_ptrs += (BLOCK_SIZE_K // scale_group_size) * stride_ask
else:
a_scale_ptrs += BLOCK_SIZE_K * 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 _vendored_gemm_afp4wfp4_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)
c = c.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)
@triton.jit
def _vendored_gemm_afp4wfp4_preshuffle_hot_exact_kernel(
a_ptr,
b_ptr,
c_ptr,
a_scales_ptr,
b_scales_ptr,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_asm,
stride_ask,
stride_bsn,
stride_bsk,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
pid_unified = tl.program_id(axis=0)
pid_unified = _remap_xcd(pid_unified, 66 * 7, NUM_XCDS=8)
pid_k = pid_unified % 7
pid_n = pid_unified // 7
offs_am = tl.arange(0, 16)
offs_bn = pid_n * 2 + tl.arange(0, 2)
offs_k = tl.arange(0, 256)
offs_k_shuffle_arr = tl.arange(0, 4096)
a_ptrs0 = a_ptr + offs_am[:, None] * stride_am + (pid_k * 512 + offs_k)[None, :] * stride_ak
a_ptrs1 = a_ptr + offs_am[:, None] * stride_am + (pid_k * 512 + 256 + offs_k)[None, :] * stride_ak
b_ptrs0 = b_ptr + offs_bn[:, None] * stride_bn + (pid_k * 8192 + offs_k_shuffle_arr)[None, :] * stride_bk
b_ptrs1 = b_ptr + offs_bn[:, None] * stride_bn + (pid_k * 8192 + 4096 + offs_k_shuffle_arr)[None, :] * stride_bk
offs_bsn = pid_n + tl.arange(0, 1)
b_scale_ptrs0 = (
b_scales_ptr
+ offs_bsn[:, None] * stride_bsn
+ (pid_k * 1024 + tl.arange(0, 512))[None, :] * stride_bsk
)
b_scale_ptrs1 = (
b_scales_ptr
+ offs_bsn[:, None] * stride_bsn
+ (pid_k * 1024 + 512 + tl.arange(0, 512))[None, :] * stride_bsk
)
a_scale_ptrs0 = (
a_scales_ptr
+ offs_am[:, None] * stride_asm
+ (pid_k * 32 + tl.arange(0, 16))[None, :] * stride_ask
)
a_scale_ptrs1 = (
a_scales_ptr
+ offs_am[:, None] * stride_asm
+ (pid_k * 32 + 16 + tl.arange(0, 16))[None, :] * stride_ask
)
accumulator = tl.zeros((16, 32), dtype=tl.float32)
a_scales0 = tl.load(a_scale_ptrs0)
b_scales0 = (
tl.load(b_scale_ptrs0, cache_modifier=cache_modifier)
.reshape(1, 2, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(32, 16)
)
a0 = tl.load(a_ptrs0)
b0 = tl.load(b_ptrs0, cache_modifier=cache_modifier)
b0 = (
b0.reshape(1, 2, 8, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(32, 256)
.trans(1, 0)
)
accumulator = tl.dot_scaled(
a0, a_scales0, "e2m1", b0, b_scales0, "e2m1", accumulator
)
a_scales1 = tl.load(a_scale_ptrs1)
b_scales1 = (
tl.load(b_scale_ptrs1, cache_modifier=cache_modifier)
.reshape(1, 2, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(32, 16)
)
a1 = tl.load(a_ptrs1)
b1 = tl.load(b_ptrs1, cache_modifier=cache_modifier)
b1 = (
b1.reshape(1, 2, 8, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(32, 256)
.trans(1, 0)
)
accumulator = tl.dot_scaled(
a1, a_scales1, "e2m1", b1, b_scales1, "e2m1", accumulator
)
offs_cn = pid_n * 32 + tl.arange(0, 32).to(tl.int64)
c_ptrs = (
c_ptr
+ pid_k * stride_ck
+ stride_cm * offs_am[:, None]
+ stride_cn * offs_cn[None, :]
)
tl.store(c_ptrs, accumulator.to(c_ptr.type.element_ty))
@triton.jit
def _vendored_gemm_afp4wfp4_preshuffle_hot_exact_k1024_kernel(
a_ptr,
b_ptr,
c_ptr,
a_scales_ptr,
b_scales_ptr,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_asm,
stride_ask,
stride_bsn,
stride_bsk,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
# Derive the 1024-element packing directly from the known-correct generic preshuffle kernel:
# - K is byte-addressed for fp4x2 tensors (two fp4 values per byte)
# - For the hot tuple, each split covers 1024 fp4 elements = 512 bytes.
pid_unified = tl.program_id(axis=0)
pid_unified = _remap_xcd(pid_unified, 66 * 7, NUM_XCDS=8)
pid_k = pid_unified % 7
pid_n = pid_unified // 7
offs_am = tl.arange(0, 16)
offs_bn = pid_n * 2 + tl.arange(0, 2)
offs_k = tl.arange(0, 512)
offs_k_shuffle_arr = tl.arange(0, 8192)
a_ptrs = (
a_ptr
+ offs_am[:, None] * stride_am
+ (pid_k * 512 + offs_k)[None, :] * stride_ak
)
b_ptrs = (
b_ptr
+ offs_bn[:, None] * stride_bn
+ (pid_k * 8192 + offs_k_shuffle_arr)[None, :] * stride_bk
)
offs_bsn = pid_n + tl.arange(0, 1)
b_scale_ptrs = (
b_scales_ptr
+ offs_bsn[:, None] * stride_bsn
+ (pid_k * 1024 + tl.arange(0, 1024))[None, :] * stride_bsk
)
a_scale_ptrs = (
a_scales_ptr
+ offs_am[:, None] * stride_asm
+ (pid_k * 32 + tl.arange(0, 32))[None, :] * stride_ask
)
accumulator = tl.zeros((16, 32), dtype=tl.float32)
a_scales = tl.load(a_scale_ptrs)
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(1, 4, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(32, 32)
)
a = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (
b.reshape(1, 2, 16, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(32, 512)
.trans(1, 0)
)
accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
offs_cn = pid_n * 32 + tl.arange(0, 32).to(tl.int64)
c_ptrs = (
c_ptr
+ pid_k * stride_ck
+ stride_cm * offs_am[:, None]
+ stride_cn * offs_cn[None, :]
)
tl.store(c_ptrs, accumulator.to(c_ptr.type.element_ty))
@triton.jit
def _vendored_gemm_afp4wfp4_reduce_hot_exact_kernel(
c_in_ptr,
c_out_ptr,
stride_c_in_k,
stride_c_in_m,
stride_c_in_n,
stride_c_out_m,
stride_c_out_n,
BLOCK_SIZE_N: tl.constexpr,
):
pid_n = tl.program_id(axis=0)
offs_m = tl.arange(0, 16)
offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
c0 = tl.load(
c_in_ptr
+ 0 * stride_c_in_k
+ offs_m[:, None] * stride_c_in_m
+ offs_n[None, :] * stride_c_in_n
)
c1 = tl.load(
c_in_ptr
+ 1 * stride_c_in_k
+ offs_m[:, None] * stride_c_in_m
+ offs_n[None, :] * stride_c_in_n
)
c2 = tl.load(
c_in_ptr
+ 2 * stride_c_in_k
+ offs_m[:, None] * stride_c_in_m
+ offs_n[None, :] * stride_c_in_n
)
c3 = tl.load(
c_in_ptr
+ 3 * stride_c_in_k
+ offs_m[:, None] * stride_c_in_m
+ offs_n[None, :] * stride_c_in_n
)
c4 = tl.load(
c_in_ptr
+ 4 * stride_c_in_k
+ offs_m[:, None] * stride_c_in_m
+ offs_n[None, :] * stride_c_in_n
)
c5 = tl.load(
c_in_ptr
+ 5 * stride_c_in_k
+ offs_m[:, None] * stride_c_in_m
+ offs_n[None, :] * stride_c_in_n
)
c6 = tl.load(
c_in_ptr
+ 6 * stride_c_in_k
+ offs_m[:, None] * stride_c_in_m
+ offs_n[None, :] * stride_c_in_n
)
c = (c0 + c1 + c2 + c3 + c4 + c5 + c6).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)
@triton.jit
def _mxfp4_quant_contract_op(
x,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
# Contract-aligned MXFP4 quantization op (matches fp4_utils conversion).
#
# Inputs:
# x: [BLOCK_SIZE_M, BLOCK_SIZE_N] bf16/fp16/fp32
# Returns:
# x_fp4: [BLOCK_SIZE_M, BLOCK_SIZE_N // 2] uint8 (packed e2m1 fp4x2)
# bs_e8m0: [BLOCK_SIZE_M, BLOCK_SIZE_N // 32] uint8 (e8m0, biased by 127)
x = x.to(tl.float32)
num_quant_blocks: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, num_quant_blocks, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
e = (qx >> 23) & 0xFF
m = qx & 0x7FFFFF
e8_bias: tl.constexpr = 127
e2_bias: tl.constexpr = 1
adjusted_exponents = tl.core.sub(e8_bias, e + 1, sanitize_overflow=False)
m = tl.where(e < e8_bias, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
e = tl.maximum(e, e8_bias - e2_bias) - (e8_bias - e2_bias)
e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
e2m1_value = tl.reshape(
e2m1_value,
[BLOCK_SIZE_M, num_quant_blocks, MXFP4_QUANT_BLOCK_SIZE // 2, 2],
)
evens, odds = tl.split(e2m1_value)
x_fp4 = evens | (odds << 4)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, num_quant_blocks)
@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 _vendored_gemm_a16wfp4_preshuffle_contract_kernel(
a_ptr,
b_ptr,
c_ptr,
b_scales_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
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_bsk > 0)
tl.assume(stride_bsn > 0)
grid_mn = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
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)
scale_group_size: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * 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_bsn = (
pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
) % N
offs_ks = (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_bsn[:, None] * stride_bsn
+ offs_ks[None, :] * stride_bsk
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for _ in range(num_k_iter):
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(
BLOCK_SIZE_N // 32,
BLOCK_SIZE_K // scale_group_size // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // scale_group_size)
)
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
a_bf16 = tl.load(
a_ptrs,
mask=offs_k_bf16[None, :] < 2 * K,
other=0,
)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (
b.reshape(
1,
BLOCK_SIZE_N // 16,
BLOCK_SIZE_K // 64,
2,
16,
16,
)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
a, a_scales = _mxfp4_quant_contract_op(
a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, scale_group_size
)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
@lru_cache(maxsize=1)
def _get_quant_func():
import aiter
from aiter import QuantType
return aiter.get_triton_quant(QuantType.per_1x32)
@lru_cache(maxsize=1)
def _get_fp4_utils():
from aiter.utility import fp4_utils
return fp4_utils
@lru_cache(maxsize=128)
def _get_a16wfp4_preshuffle_config(M, N, K_bytes):
# Reuse AITER's shape-based config selection (tuned when available).
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _get_config
config, _tuned = _get_config(M, N, K_bytes, True)
return config
def _quant_block_size_m(m):
if m <= 8:
return 8
if m <= 16:
return 16
if m <= 32:
return 32
if m <= 64:
return 64
return 128
def _shape_tuned_quant(A, shuffle):
import torch
from aiter import dtypes
if triton is None:
return _get_quant_func()(A, shuffle=shuffle)
fp4_utils = _get_fp4_utils()
M, N = A.shape
# Preserve AITER's output contract while shrinking row tiles for small-M cases.
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=A.device)
scale_n_valid = _cdiv(N, 32)
scale_n_pad = _cdiv(scale_n_valid, 8) * 8
scale_m_pad = _cdiv(M, 32) * 32
blockscale_e8m0 = torch.empty(
(_cdiv(M, 256) * 256, scale_n_pad),
dtype=torch.uint8,
device=A.device,
)
block_size_m = _quant_block_size_m(M)
grid = (_cdiv(M, block_size_m), scale_n_pad)
fp4_utils._dynamic_mxfp4_quant_kernel_asm_layout[grid](
A,
x_fp4,
blockscale_e8m0,
*A.stride(),
*x_fp4.stride(),
*blockscale_e8m0.stride(),
M=M,
N=N,
scaleN=scale_n_valid,
scaleM_pad=scale_m_pad,
scaleN_pad=scale_n_pad,
BLOCK_SIZE=block_size_m,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
SHUFFLE=shuffle,
)
return (x_fp4.view(dtypes.fp4x2), blockscale_e8m0.view(dtypes.fp8_e8m0))
def _shape_tuned_quant_hot_exact(A):
import torch
from aiter import dtypes
if triton is None:
return _shape_tuned_quant(A, shuffle=False)
fp4_utils = _get_fp4_utils()
M, N = A.shape
if (M, N) != (16, 7168):
raise RuntimeError(f"unexpected hot exact quant shape {(M, N)}")
x_fp4 = _get_workspace_tensor("hot_exact_x_fp4", A, (16, 7168 // 2), torch.uint8)
blockscale_e8m0 = _get_workspace_tensor(
"hot_exact_scale_e8m0", A, (16, 7168 // 32), torch.uint8
)
grid = (1, 7168 // 32)
fp4_utils._dynamic_mxfp4_quant_kernel_asm_layout[grid](
A,
x_fp4,
blockscale_e8m0,
*A.stride(),
*x_fp4.stride(),
*blockscale_e8m0.stride(),
M=16,
N=7168,
scaleN=7168 // 32,
scaleM_pad=32,
scaleN_pad=7168 // 32,
BLOCK_SIZE=16,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
SHUFFLE=False,
)
return (x_fp4.view(dtypes.fp4x2), blockscale_e8m0.view(dtypes.fp8_e8m0))
def _run_explicit_asm(
aiter,
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
m,
n,
kernel_name,
log2_k_split,
):
import torch
from aiter import dtypes
out = torch.empty(((m + 31) // 32 * 32, n), dtype=dtypes.bf16, device=A_q.device)
aiter.gemm_a4w4_asm(
A_q.view(m, A_q.shape[-1]),
B_shuffle,
A_scale_sh,
B_scale_sh,
out,
kernel_name,
None,
1.0,
0.0,
True,
log2_k_split=log2_k_split,
)
return out[:m]
def _reshape_preshuffle_weight(weight, n, k):
import torch
weight_u8 = weight.view(torch.uint8)
assert weight_u8.shape == (n, k // 2), (
f"expected preshuffled weight shape {(n, k // 2)}, got {tuple(weight_u8.shape)}"
)
return weight_u8.view(n // 16, weight_u8.shape[1] * 16)
def _reshape_preshuffle_scales(scales, k):
import torch
scales_u8 = scales.view(torch.uint8)
assert scales_u8.shape[0] % 32 == 0, (
f"expected padded scale rows divisible by 32, got {scales_u8.shape[0]}"
)
assert scales_u8.numel() % k == 0, (
f"expected scale bytes divisible by K={k}, got {scales_u8.numel()}"
)
return scales_u8.view(scales_u8.shape[0] // 32, scales_u8.shape[1] * 32)
def _run_vendored_triton_a16wfp4_preshuffle_contract(A, B_shuffle, B_scale_sh):
import torch
from aiter import dtypes
if triton is None:
raise RuntimeError("triton is required for the fused BF16xFP4 preshuffle path")
M, K_full = A.shape
N = B_shuffle.shape[0]
if K_full % 32 != 0:
raise RuntimeError(f"unsupported K={K_full} for MXFP4 group=32")
if K_full % 2 != 0:
raise RuntimeError(f"unsupported odd K={K_full} for fp4x2 packing")
K_bytes = K_full // 2
packed_w = _reshape_preshuffle_weight(B_shuffle, N, K_full)
packed_scales = _reshape_preshuffle_scales(B_scale_sh, K_full)
local_config = dict(_get_a16wfp4_preshuffle_config(M, N, K_bytes))
return_y_pp = local_config["NUM_KSPLIT"] > 1
splitk_block_size, block_size_k, num_ksplit = _get_splitk(
K_bytes,
local_config["BLOCK_SIZE_K"],
local_config["NUM_KSPLIT"],
)
local_config["SPLITK_BLOCK_SIZE"] = splitk_block_size
local_config["BLOCK_SIZE_K"] = block_size_k
local_config["NUM_KSPLIT"] = num_ksplit
if local_config["BLOCK_SIZE_K"] >= 2 * K_bytes:
local_config["BLOCK_SIZE_K"] = _next_power_of_two(2 * K_bytes)
local_config["SPLITK_BLOCK_SIZE"] = 2 * K_bytes
local_config["NUM_KSPLIT"] = 1
return_y_pp = False
local_config["BLOCK_SIZE_N"] = max(local_config["BLOCK_SIZE_N"], 32)
y = torch.empty((M, N), dtype=dtypes.bf16, device=A.device)
if return_y_pp:
y_pp = torch.empty(
(local_config["NUM_KSPLIT"], M, N),
dtype=torch.float32,
device=A.device,
)
else:
y_pp = None
grid = lambda META: ( # noqa: E731
(
META["NUM_KSPLIT"]
* triton.cdiv(M, META["BLOCK_SIZE_M"])
* triton.cdiv(N, META["BLOCK_SIZE_N"])
),
)
_vendored_gemm_a16wfp4_preshuffle_contract_kernel[grid](
A,
packed_w,
y if y_pp is None else y_pp,
packed_scales,
M,
N,
K_bytes,
A.stride(0),
A.stride(1),
packed_w.stride(0),
packed_w.stride(1),
0 if y_pp is None else y_pp.stride(0),
y.stride(0) if y_pp is None else y_pp.stride(1),
y.stride(1) if y_pp is None else y_pp.stride(2),
packed_scales.stride(0),
packed_scales.stride(1),
**local_config,
)
if y_pp is None:
return y
reduce_block_size_m = 16
reduce_block_size_n = 64
actual_ksplit = triton.cdiv(K_bytes, local_config["SPLITK_BLOCK_SIZE"] // 2)
grid_reduce = (
triton.cdiv(M, reduce_block_size_m),
triton.cdiv(N, reduce_block_size_n),
)
_vendored_gemm_afp4wfp4_reduce_kernel[grid_reduce](
y_pp,
y,
M,
N,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
y.stride(0),
y.stride(1),
reduce_block_size_m,
reduce_block_size_n,
actual_ksplit,
_next_power_of_two(local_config["NUM_KSPLIT"]),
)
return y
def _run_vendored_triton_preshuffle(A_q, B_shuffle, A_scale, B_scale_sh, config):
import torch
from aiter import dtypes
if triton is None:
raise RuntimeError("triton is required for the vendored preshuffle path")
M, K = A_q.shape
N = B_shuffle.shape[0]
K = K
packed_w = _reshape_preshuffle_weight(B_shuffle, N, K * 2)
packed_scales = _reshape_preshuffle_scales(B_scale_sh, K * 2)
local_config = dict(config)
return_y_pp = local_config["NUM_KSPLIT"] > 1
splitk_block_size, block_size_k, num_ksplit = _get_splitk(
K,
local_config["BLOCK_SIZE_K"],
local_config["NUM_KSPLIT"],
)
local_config["SPLITK_BLOCK_SIZE"] = splitk_block_size
local_config["BLOCK_SIZE_K"] = block_size_k
local_config["NUM_KSPLIT"] = num_ksplit
if local_config["BLOCK_SIZE_K"] >= 2 * K:
local_config["BLOCK_SIZE_K"] = _next_power_of_two(2 * K)
local_config["SPLITK_BLOCK_SIZE"] = 2 * K
local_config["NUM_KSPLIT"] = 1
return_y_pp = False
local_config["BLOCK_SIZE_N"] = max(local_config["BLOCK_SIZE_N"], 32)
y = torch.empty((M, N), dtype=dtypes.bf16, device=A_q.device)
if return_y_pp:
y_pp = torch.empty(
(local_config["NUM_KSPLIT"], M, N),
dtype=torch.float32,
device=A_q.device,
)
else:
y_pp = None
grid = lambda META: ( # noqa: E731
(
META["NUM_KSPLIT"]
* triton.cdiv(M, META["BLOCK_SIZE_M"])
* triton.cdiv(N, META["BLOCK_SIZE_N"])
),
)
_vendored_gemm_afp4wfp4_preshuffle_kernel[grid](
A_q.view(torch.uint8),
packed_w,
y if y_pp is None else y_pp,
A_scale.view(torch.uint8),
packed_scales,
M,
N,
K,
A_q.stride(0),
A_q.stride(1),
packed_w.stride(0),
packed_w.stride(1),
0 if y_pp is None else y_pp.stride(0),
y.stride(0) if y_pp is None else y_pp.stride(1),
y.stride(1) if y_pp is None else y_pp.stride(2),
A_scale.stride(0),
A_scale.stride(1),
packed_scales.stride(0),
packed_scales.stride(1),
**local_config,
)
if y_pp is None:
return y
reduce_block_size_m = 16
reduce_block_size_n = 64
actual_ksplit = triton.cdiv(K, local_config["SPLITK_BLOCK_SIZE"] // 2)
grid_reduce = (
triton.cdiv(M, reduce_block_size_m),
triton.cdiv(N, reduce_block_size_n),
)
_vendored_gemm_afp4wfp4_reduce_kernel[grid_reduce](
y_pp,
y,
M,
N,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
y.stride(0),
y.stride(1),
reduce_block_size_m,
reduce_block_size_n,
actual_ksplit,
_next_power_of_two(local_config["NUM_KSPLIT"]),
)
return y
def _run_vendored_triton_preshuffle_hot_exact(A_q, B_shuffle, A_scale, B_scale_sh, config):
import torch
from aiter import dtypes
if triton is None:
raise RuntimeError("triton is required for the vendored preshuffle hot path")
M, K = A_q.shape
N = B_shuffle.shape[0]
if (M, N, K) != (16, 2112, 3584):
raise RuntimeError(f"unexpected hot exact shape {(M, N, K)}")
packed_w = _reshape_preshuffle_weight(B_shuffle, N, K * 2)
packed_scales = _reshape_preshuffle_scales(B_scale_sh, K * 2)
local_config = dict(config)
splitk_block_size, block_size_k, num_ksplit = _get_splitk(
K,
local_config["BLOCK_SIZE_K"],
local_config["NUM_KSPLIT"],
)
local_config["SPLITK_BLOCK_SIZE"] = splitk_block_size
local_config["BLOCK_SIZE_K"] = block_size_k
local_config["NUM_KSPLIT"] = num_ksplit
if not (
local_config["BLOCK_SIZE_M"] == 16
and local_config["BLOCK_SIZE_N"] == 32
and local_config["BLOCK_SIZE_K"] in (512, 1024)
and local_config["NUM_KSPLIT"] == 7
):
return _run_vendored_triton_preshuffle(A_q, B_shuffle, A_scale, B_scale_sh, config)
y_pp = _get_workspace_tensor("hot_exact_y_pp", A_q, (7, 16, 2112), torch.float32)
y = torch.empty((16, 2112), dtype=dtypes.bf16, device=A_q.device)
grid = (66 * 7,)
if local_config["BLOCK_SIZE_K"] == 1024 and _HOT_BIGK_STATE["enabled"]:
try:
_vendored_gemm_afp4wfp4_preshuffle_hot_exact_k1024_kernel[grid](
A_q.view(torch.uint8),
packed_w,
y_pp,
A_scale.view(torch.uint8),
packed_scales,
A_q.stride(0),
A_q.stride(1),
packed_w.stride(0),
packed_w.stride(1),
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
A_scale.stride(0),
A_scale.stride(1),
packed_scales.stride(0),
packed_scales.stride(1),
num_warps=local_config["num_warps"],
num_stages=local_config["num_stages"],
waves_per_eu=local_config["waves_per_eu"],
matrix_instr_nonkdim=local_config["matrix_instr_nonkdim"],
cache_modifier=local_config["cache_modifier"],
)
except Exception as exc:
# One-shot disable: if the 1024-kernel compilation fails in the runner, keep the
# canonical 512-kernel behavior for the remainder of the process.
_HOT_BIGK_STATE["enabled"] = False
_HOT_BIGK_STATE["error"] = repr(exc)
_vendored_gemm_afp4wfp4_preshuffle_hot_exact_kernel[grid](
A_q.view(torch.uint8),
packed_w,
y_pp,
A_scale.view(torch.uint8),
packed_scales,
A_q.stride(0),
A_q.stride(1),
packed_w.stride(0),
packed_w.stride(1),
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
A_scale.stride(0),
A_scale.stride(1),
packed_scales.stride(0),
packed_scales.stride(1),
num_warps=local_config["num_warps"],
num_stages=local_config["num_stages"],
waves_per_eu=local_config["waves_per_eu"],
matrix_instr_nonkdim=local_config["matrix_instr_nonkdim"],
cache_modifier=local_config["cache_modifier"],
)
else:
_vendored_gemm_afp4wfp4_preshuffle_hot_exact_kernel[grid](
A_q.view(torch.uint8),
packed_w,
y_pp,
A_scale.view(torch.uint8),
packed_scales,
A_q.stride(0),
A_q.stride(1),
packed_w.stride(0),
packed_w.stride(1),
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
A_scale.stride(0),
A_scale.stride(1),
packed_scales.stride(0),
packed_scales.stride(1),
num_warps=local_config["num_warps"],
num_stages=local_config["num_stages"],
waves_per_eu=local_config["waves_per_eu"],
matrix_instr_nonkdim=local_config["matrix_instr_nonkdim"],
cache_modifier=local_config["cache_modifier"],
)
grid_reduce = (33,)
_vendored_gemm_afp4wfp4_reduce_hot_exact_kernel[grid_reduce](
y_pp,
y,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
y.stride(0),
y.stride(1),
BLOCK_SIZE_N=64,
)
return y
def custom_kernel(data: input_t) -> output_t:
"""
Hybrid path:
- exact `(16, 2112, 7168)` benchmark tuple uses the vendored preshuffled Triton path
- everything else keeps the explicit asm or reference a4w4 fallback path
"""
import aiter
from aiter import dtypes
A, _B, _B_q, B_shuffle, B_scale_sh = data
B_shuffle = B_shuffle.contiguous()
B_scale_sh = B_scale_sh.contiguous()
m, k = A.shape
n = B_shuffle.shape[0]
shape = (m, n, k)
if (
shape in _FUSED_A16WFP4_SHAPES
and triton is not None
and _A16WFP4_PRESHUFFLE_STATE["enabled"]
):
try:
return _run_vendored_triton_a16wfp4_preshuffle_contract(A, B_shuffle, B_scale_sh)
except Exception as exc:
# One-shot disable: if the fused preshuffle kernel fails to compile or run in the
# runner, keep the known-good paths for the remainder of the process.
_A16WFP4_PRESHUFFLE_STATE["enabled"] = False
_A16WFP4_PRESHUFFLE_STATE["error"] = repr(exc)
vendored_config = _VENDORED_TRITON_PRESHUFFLE_CONFIGS.get(shape)
if vendored_config is not None and triton is not None:
# The preshuffled AFP4/WFP4 path requires unshuffled A scales for M < 32.
if shape == (16, 2112, 7168):
A_q, A_scale = _shape_tuned_quant_hot_exact(A)
return _run_vendored_triton_preshuffle_hot_exact(
A_q.contiguous(),
B_shuffle,
A_scale.contiguous(),
B_scale_sh,
vendored_config,
)
A_q, A_scale = _shape_tuned_quant(A, shuffle=False)
return _run_vendored_triton_preshuffle(
A_q.contiguous(),
B_shuffle,
A_scale.contiguous(),
B_scale_sh,
vendored_config,
)
A_q, A_scale_sh = _shape_tuned_quant(A, shuffle=True)
kernel_name = _EXPLICIT_ASM_KERNELS.get(shape)
if kernel_name is not None:
log2_k_split = _EXPLICIT_ASM_LOG2_K_SPLIT.get(shape, _ASM_SPLIT_K)
return _run_explicit_asm(
aiter,
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
m,
n,
kernel_name,
log2_k_split,
)
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
scrolls · 1455 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON