submission 752215
nataliakokoromyti · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1113 lines, June 9 Researcher Reciprocity License v1.0.
aiter_patched_quant_direct_dual_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-752215?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:98e9e14cb3da19901b8915c8ee93a91b9259f4d3be1efd891ac30365da653cee
license declaredunknown
license concludedunknown
authorsnataliakokoromyti
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 1
num_warps = 1split-k
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitkstages = 1
num_stages=1,tile-m = 32
BLOCK_M=32,tile-n = 8
BLOCK_N=8,Kernel source
aiter_patched_quant_direct_dual_submission.py1113 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from task import input_t, output_t
import triton
import triton.language as tl
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.quant.quant import _dynamic_mxfp4_quant_kernel
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid
_ASM_OUT_BUFFER_CACHE = {}
_RAW_OUT_BUFFER_CACHE = {}
_DIRECT_QUANT_CACHE = {}
_ASM_KERNEL_OVERRIDES = {
(64, 7168, 2048): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
(256, 3072, 1536): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
}
def _as_u8_tensor(x):
import torch
if x.dtype == torch.uint8:
return x.contiguous()
return x.view(torch.uint8).contiguous()
def _get_cached_asm_out(torch_mod, dtypes, device, m: int, n: int):
key = (device.type, device.index, m, n)
cached = _ASM_OUT_BUFFER_CACHE.get(key)
if cached is None:
padded_m = ((m + 31) // 32) * 32
cached = torch_mod.empty((padded_m, n), dtype=dtypes.bf16, device=device)
_ASM_OUT_BUFFER_CACHE[key] = cached
return cached
def _get_cached_raw_out(torch_mod, device, m: int, n: int):
key = (device.type, device.index, m, n)
cached = _RAW_OUT_BUFFER_CACHE.get(key)
if cached is None:
cached = torch_mod.empty((m, n), dtype=torch_mod.bfloat16, device=device)
_RAW_OUT_BUFFER_CACHE[key] = cached
return cached
def _get_cached_direct_quant(torch_mod, device, m: int, k: int):
scale_n = (k + 31) // 32
scale_n_pad = ((scale_n + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
key = (device.type, device.index, m, k)
cached = _DIRECT_QUANT_CACHE.get(key)
if cached is None:
a_q = torch_mod.empty((m, k // 2), dtype=torch_mod.uint8, device=device)
a_scale_raw = torch_mod.empty((m, scale_n), dtype=torch_mod.uint8, device=device)
a_scale_sh = torch_mod.empty(
(scale_m_pad, scale_n_pad), dtype=torch_mod.uint8, device=device
)
cached = (a_q, a_scale_raw, a_scale_sh)
_DIRECT_QUANT_CACHE[key] = cached
return cached
def _get_cached_direct_quant_exact(torch_mod, device):
key = (device.type, device.index, 256, 1536, "exact")
cached = _DIRECT_QUANT_CACHE.get(key)
if cached is None:
a_q = torch_mod.empty((256, 768), dtype=torch_mod.uint8, device=device)
a_scale_sh = torch_mod.empty((256, 48), dtype=torch_mod.uint8, device=device)
cached = (a_q, a_scale_sh)
_DIRECT_QUANT_CACHE[key] = cached
return cached
def _get_cached_direct_quant_exact_64x2048(torch_mod, device):
key = (device.type, device.index, 64, 2048, "exact")
cached = _DIRECT_QUANT_CACHE.get(key)
if cached is None:
a_q = torch_mod.empty((64, 1024), dtype=torch_mod.uint8, device=device)
a_scale_sh = torch_mod.empty((256, 64), dtype=torch_mod.uint8, device=device)
cached = (a_q, a_scale_sh)
_DIRECT_QUANT_CACHE[key] = cached
return cached
def _pick_config(m: int, n: int, k: int) -> dict | None:
# AITER ships only one specialized gfx950 A16WFP4_PRESHUFFLED config for the
# public leaderboard shapes. Everything else falls back to the generic family.
if (n, k) == (2112, 7168):
if m <= 8:
return {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 14,
}
if m <= 16:
return {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 14,
}
if m <= 64:
return {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 14,
}
if m <= 256:
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 14,
}
# Borrow known-good preshuffled FP4 configs for public shapes that do not
# have an explicit A16WFP4 specialization in AITER.
if (n, k) == (4096, 512):
if m <= 31:
return {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
if m <= 64:
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
if m <= 256:
return {
"BLOCK_SIZE_M": 256,
"BLOCK_SIZE_N": 256,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
return {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 256,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
if (n, k) == (3072, 1536):
if m <= 31:
return {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
if m <= 64:
return {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
if m <= 256:
return {
"BLOCK_SIZE_M": 128,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
return None
def _default_config() -> dict:
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
def _pick_raw_config(m: int, n: int, k: int) -> dict:
if (n, k) == (2112, 7168):
if m <= 8:
return {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 16,
}
if m <= 16:
return {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 6,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 8,
}
if m <= 32:
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 1024,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 8,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 8,
}
if m <= 64:
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 1024,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
if m <= 128:
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 1024,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 6,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
if m <= 256:
return {
"BLOCK_SIZE_M": 128,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 2,
"num_warps": 4,
"num_stages": 3,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
return {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 1024,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
if (n, k) == (3072, 1536):
if m <= 16:
return {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 16,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"waves_per_eu": 6,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
if m <= 32:
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 16,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
if m <= 128:
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 2,
"num_warps": 4,
"num_stages": 3,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
if m <= 256:
return {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 2,
"num_warps": 4,
"num_stages": 3,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
if m <= 2048:
return {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 2,
"num_warps": 4,
"num_stages": 3,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
return {
"BLOCK_SIZE_M": 128,
"BLOCK_SIZE_N": 256,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 16,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 32,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
if (n, k) == (7168, 2048):
if m <= 8:
return {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 4,
}
if m <= 16:
return {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 4,
}
if m <= 32:
return {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 4,
}
if m <= 64:
return {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
if m <= 128:
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
return {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 256,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
}
# Generic gfx950-GEMM-A16WFP4 config.
if m <= 16:
return {
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
return {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg" if m <= 256 else None,
"NUM_KSPLIT": 1,
}
def _e8m0_unshuffle(scale_sh, *, rows: int, cols: int):
restored = scale_sh.view(scale_sh.shape[0] // 32, scale_sh.shape[1] // 8, 4, 16, 2, 2)
restored = restored.permute(0, 5, 3, 1, 4, 2).contiguous()
restored = restored.view(scale_sh.shape[0], scale_sh.shape[1])
return restored[:rows, :cols].contiguous()
@triton.jit
def _e8m0_shuffle_filled_kernel(
raw_ptr,
sh_ptr,
raw_stride_m,
raw_stride_n,
sh_stride_m,
sh_stride_n,
M,
M_PAD,
N_VALID,
N_PAD,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
linear = offs_m[:, None] * N_PAD + offs_n[None, :]
out_k = linear & 1
linear = linear // 2
out_j = linear & 1
linear = linear // 2
out_i = linear % 16
linear = linear // 16
out_f = linear % 4
linear = linear // 4
out_c = linear % (N_PAD // 8)
src_m = (offs_m[:, None] // 32) * 32 + out_k * 16 + out_i
src_n = out_c * 8 + out_j * 4 + out_f
src_mask = (src_m < M) & (src_n < N_VALID)
raw = tl.load(
raw_ptr + src_m * raw_stride_m + src_n * raw_stride_n,
mask=src_mask,
other=127,
)
dst_mask = (offs_m[:, None] < M_PAD) & (offs_n[None, :] < N_PAD)
tl.store(
sh_ptr + offs_m[:, None] * sh_stride_m + offs_n[None, :] * sh_stride_n,
raw,
mask=dst_mask,
)
@triton.jit
def _patched_quant_asm_exact_256x1536_kernel(
x_ptr,
x_fp4_ptr,
bs_sh_ptr,
stride_x_m,
stride_x_n,
stride_x_fp4_m,
stride_x_fp4_n,
BLOCK_M: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * 32 + tl.arange(0, 32)
x_offs = offs_m[:, None] * stride_x_m + offs_n[None, :] * stride_x_n
x = tl.load(x_ptr + x_offs).to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(x, 32, BLOCK_M, 32)
out_offs_n = pid_n * 16 + tl.arange(0, 16)
out_offs = offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
tl.store(x_fp4_ptr + out_offs, out_tensor)
# Equivalent to fp4_utils.e8m0_shuffle for an exact [256, 48] scale tensor.
row = offs_m
col = pid_n
block_row = row // 32
rem = row % 32
b = rem // 16
c = rem % 16
d = col // 8
remc = col % 8
e = remc // 4
f = remc % 4
linear = (((((block_row * 6 + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b)
tl.store(bs_sh_ptr + linear, bs_e8m0.reshape(BLOCK_M))
@triton.jit
def _patched_quant_asm_exact_64x2048_kernel(
x_ptr,
x_fp4_ptr,
bs_sh_ptr,
stride_x_m,
stride_x_n,
stride_x_fp4_m,
stride_x_fp4_n,
BLOCK_M: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * 32 + tl.arange(0, 32)
x_mask = offs_m[:, None] < 64
x_offs = offs_m[:, None] * stride_x_m + offs_n[None, :] * stride_x_n
x = tl.load(x_ptr + x_offs, mask=x_mask, other=0).to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(x, 32, BLOCK_M, 32)
bs_e8m0 = tl.where((offs_m < 64)[:, None], bs_e8m0, 127)
out_offs_n = pid_n * 16 + tl.arange(0, 16)
out_offs = offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
out_mask = offs_m[:, None] < 64
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
row = offs_m
col = pid_n
block_row = row // 32
rem = row % 32
b = rem // 16
c = rem % 16
d = col // 8
remc = col % 8
e = remc // 4
f = remc % 4
linear = (((((block_row * 8 + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b)
tl.store(bs_sh_ptr + linear, bs_e8m0.reshape(BLOCK_M))
def _quant_mxfp4_direct(torch_mod, x, a_q, a_scale_raw, a_scale_sh):
m, k = x.shape
scale_n = (k + 31) // 32
scale_n_pad = ((scale_n + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
if m <= 32:
num_iter = 1
block_size_m = triton.next_power_of_2(m)
block_size_n = 32
num_warps = 1
num_stages_quant = 1
else:
num_iter = 4
block_size_m = 64
block_size_n = 64
num_warps = 4
num_stages_quant = 2
if k <= 16384:
block_size_m = 32
block_size_n = 128
if k <= 1024:
num_iter = 1
num_warps = 4
num_stages_quant = 1
block_size_n = min(256, triton.next_power_of_2(k))
block_size_n = max(32, block_size_n)
block_size_m = min(8, triton.next_power_of_2(m))
grid = (
triton.cdiv(m, block_size_m),
triton.cdiv(k, block_size_n * num_iter),
)
_dynamic_mxfp4_quant_kernel[grid](
x,
a_q,
a_scale_raw,
*x.stride(),
*a_q.stride(),
*a_scale_raw.stride(),
M=m,
N=k,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
NUM_ITER=num_iter,
BLOCK_SIZE_M=block_size_m,
BLOCK_SIZE_N=block_size_n,
NUM_STAGES=num_stages_quant,
num_warps=num_warps,
waves_per_eu=0,
num_stages=1,
)
grid_shuffle = (
triton.cdiv(a_scale_sh.shape[0], 32),
triton.cdiv(a_scale_sh.shape[1], 8),
)
_e8m0_shuffle_filled_kernel[grid_shuffle](
a_scale_raw,
a_scale_sh,
a_scale_raw.stride(0),
a_scale_raw.stride(1),
a_scale_sh.stride(0),
a_scale_sh.stride(1),
m,
scale_m_pad,
scale_n,
scale_n_pad,
BLOCK_M=32,
BLOCK_N=8,
)
def _quant_mxfp4_direct_exact_256x1536(x, a_q, a_scale_sh):
grid = (2, 48)
_patched_quant_asm_exact_256x1536_kernel[grid](
x,
a_q,
a_scale_sh,
x.stride(0),
x.stride(1),
a_q.stride(0),
a_q.stride(1),
BLOCK_M=128,
num_warps=4,
num_stages=1,
)
def _quant_mxfp4_direct_exact_64x2048(x, a_q, a_scale_sh):
grid = (1, 64)
_patched_quant_asm_exact_64x2048_kernel[grid](
x,
a_q,
a_scale_sh,
x.stride(0),
x.stride(1),
a_q.stride(0),
a_q.stride(1),
BLOCK_M=64,
num_warps=4,
num_stages=2,
)
@triton.heuristics(
{
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
* triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
}
)
@triton.jit
def _fixed_a16wfp4_preshuffle_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,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr,
PREQUANT: tl.constexpr,
cache_modifier: tl.constexpr,
):
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)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
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(0, 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)
)
a_bf16 = tl.load(a_ptrs)
b_raw = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (
b_raw.reshape(
1,
BLOCK_SIZE_N // 16,
BLOCK_SIZE_K // 64,
2,
16,
16,
)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
if PREQUANT:
a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
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)
def _run_fixed_a16wfp4_preshuffle(x, w, w_scales, *, dtype, config):
import torch
M, _ = x.shape
n_outer, k_outer = w.shape
N = n_outer * 16
K = k_outer // 16
cfg = dict(config)
if cfg["NUM_KSPLIT"] > 1:
splitk_block_size, block_size_k, num_ksplit = get_splitk(
K, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
)
cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size
cfg["BLOCK_SIZE_K"] = block_size_k
cfg["NUM_KSPLIT"] = num_ksplit
if cfg["BLOCK_SIZE_K"] >= 2 * K:
cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
cfg["SPLITK_BLOCK_SIZE"] = 2 * K
cfg["NUM_KSPLIT"] = 1
else:
cfg.setdefault("SPLITK_BLOCK_SIZE", 2 * K)
cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
if cfg["NUM_KSPLIT"] > 1:
y_pp = torch.empty((cfg["NUM_KSPLIT"], M, N), dtype=torch.float32, device=x.device)
y = torch.empty((M, N), dtype=dtype, device=x.device)
else:
y_pp = None
y = torch.empty((M, N), dtype=dtype, device=x.device)
grid = lambda META: ( # noqa: E731
(
META["NUM_KSPLIT"]
* triton.cdiv(M, META["BLOCK_SIZE_M"])
* triton.cdiv(N, META["BLOCK_SIZE_N"])
),
)
_fixed_a16wfp4_preshuffle_kernel[grid](
x,
w,
y if y_pp is None else y_pp,
w_scales,
M,
N,
K,
x.stride(0),
x.stride(1),
w.stride(0),
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),
w_scales.stride(0),
w_scales.stride(1),
PREQUANT=True,
**cfg,
)
if cfg["NUM_KSPLIT"] > 1:
reduce_block_size_m = 16
reduce_block_size_n = 64
actual_ksplit = triton.cdiv(K, (cfg["SPLITK_BLOCK_SIZE"] // 2))
grid_reduce = (
triton.cdiv(M, reduce_block_size_m),
triton.cdiv(N, reduce_block_size_n),
)
_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,
triton.next_power_of_2(cfg["NUM_KSPLIT"]),
)
return y
def custom_kernel(data: input_t) -> output_t:
import os
os.environ.setdefault("AITER_LOG_LEVEL", "ERROR")
import torch
import aiter
from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
def _quant_mxfp4(x):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
bs_e8m0 = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
a, _b, b_q, b_shuffle, b_scale_sh = data
a = a.contiguous()
b_shuffle = _as_u8_tensor(b_shuffle)
b_scale_sh = _as_u8_tensor(b_scale_sh)
m, k_real = a.shape
n = b_shuffle.shape[0]
packed_k = k_real // 2
override = _ASM_KERNEL_OVERRIDES.get((m, n, k_real))
if override is not None:
if (m, n, k_real) == (64, 7168, 2048):
a_q_u8, a_scale_sh_u8 = _get_cached_direct_quant_exact_64x2048(
torch, a.device
)
_quant_mxfp4_direct_exact_64x2048(a, a_q_u8, a_scale_sh_u8)
a_q = a_q_u8.view(dtypes.fp4x2)
a_scale_sh = a_scale_sh_u8.view(dtypes.fp8_e8m0)
elif (m, n, k_real) == (256, 3072, 1536):
a_q_u8, a_scale_sh_u8 = _get_cached_direct_quant_exact(torch, a.device)
_quant_mxfp4_direct_exact_256x1536(a, a_q_u8, a_scale_sh_u8)
a_q = a_q_u8.view(dtypes.fp4x2)
a_scale_sh = a_scale_sh_u8.view(dtypes.fp8_e8m0)
else:
a_q, a_scale_sh = _quant_mxfp4(a)
kernel_name, split_k = override
out = _get_cached_asm_out(torch, aiter.dtypes, a.device, m, n)
aiter.gemm_a4w4_asm(
a_q.view(m, packed_k),
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
kernel_name,
None,
1.0,
0.0,
True,
split_k,
)
return out[:m]
# The long/wide 7168x2048 family still prefers the older fused-raw path on MI355X.
if (n, k_real) == (7168, 2048):
b_q_u8 = _as_u8_tensor(b_q)
b_scale_raw = _e8m0_unshuffle(b_scale_sh, rows=b_q_u8.shape[0], cols=k_real // 32)
out = _get_cached_raw_out(torch, a.device, m, b_q_u8.shape[0])
return gemm_a16wfp4(
a,
b_q_u8,
b_scale_raw,
atomic_add=False,
dtype=torch.bfloat16,
y=out,
config=_pick_raw_config(m, b_q_u8.shape[0], k_real),
)
# GPU MODE's task tensors use the logical task shapes:
# B_shuffle : (N, K/2)
# B_scale_sh : (*, K/32)
# AITER's preshuffled kernel expects the exact same bytes reinterpreted as:
# B : (N/16, (K/2) * 16)
# B_scales : (* / 32, K)
# The byte-level round-trip matches the task's public shuffle operators.
b_phys = b_shuffle.contiguous().view(b_shuffle.numel() // (packed_k * 16), packed_k * 16)
b_scale_phys = b_scale_sh.contiguous().view(b_scale_sh.numel() // k_real, k_real)
config = _pick_config(m, n, k_real) or _default_config()
return _run_fixed_a16wfp4_preshuffle(
a,
b_phys,
b_scale_phys,
dtype=torch.bfloat16,
config=config,
)
scrolls · 1113 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