submission 735661
lmw-perfxlab · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 755 lines, June 9 Researcher Reciprocity License v1.0.
triton_a4w4_merge.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-735661?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:fcaec78fa61367b92420af5f832ffbc4918292216adaba468214f5f0a577a34c
license declaredunknown
license concludedunknown
authorslmw-perfxlab
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.num-warps = 1
num_warps=1split-k
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)stages = 2
num_stages = 2tile-k = 512
BLOCK_SIZE_K = 512 #kwargs.get("BLOCK_SIZE_K", 256)tile-m = 16
BLOCK_SIZE_M = 16 #kwargs.get("BLOCK_SIZE_M", 32)tile-n = 32
BLOCK_SIZE_N = 32 #kwargs.get("BLOCK_SIZE_N", 32)Kernel source
triton_a4w4_merge.py755 lines
# SPDX-License-Identifier: MIT
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import aiter
from aiter import QuantType, dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from aiter.ops.triton.utils.gemm_config_utils import get_gemm_config
MXBLK = 32
GROUP_N = 64
@triton.jit
def mxfp4_quant_kernel_smallM(
x_ptr, out_ptr, scale_ptr,
stride_x_m, stride_x_n,
stride_out_m, stride_out_n,
stride_s_m, stride_s_n,
M: tl.constexpr,
N: tl.constexpr,
GROUP_N_: tl.constexpr,
MXBLK_: tl.constexpr,
):
m_id = tl.program_id(0)
g_id = tl.program_id(1)
# ---------- 子块0 (0..31) ----------
offs0 = g_id * GROUP_N_ + tl.arange(0, MXBLK_)
mask0 = offs0 < N
x0 = tl.load(
x_ptr + m_id * stride_x_m + offs0 * stride_x_n,
mask=mask0, other=0.0
).to(tl.float32)
# 计算 scale0
amax0 = tl.max(tl.abs(x0))
amax0 = tl.where(amax0 == 0, 1e-8, amax0)
amax_i32 = amax0.to(tl.int32, bitcast=True)
amax_i32 = (amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax0 = amax_i32.to(tl.float32, bitcast=True)
se0 = tl.log2(amax0).floor() - 2
se0 = tl.clamp(se0, -127, 127)
scale0 = tl.exp2(-se0)
bs0 = (se0.to(tl.uint8) + 127)
# 量化子块0
q0 = x0 * scale0
qbits0 = q0.to(tl.uint32, bitcast=True)
s0 = qbits0 & 0x80000000
qbits0 ^= s0
qf0 = qbits0.to(tl.float32, bitcast=True)
sat_mask0 = qf0 >= 6.0
den_mask0 = (qf0 < 1.0) & (~sat_mask0)
norm_mask0 = ~(sat_mask0 | den_mask0)
denorm_exp = ((127 - 1) + (23 - 1) + 1) << 23
denorm_float = tl.cast(denorm_exp, tl.float32, bitcast=True)
dval0 = (qf0 + denorm_float).to(tl.uint32, bitcast=True) - denorm_exp
dval0 = dval0.to(tl.uint8)
nval0 = qbits0
mant_odd0 = (nval0 >> (23 - 1)) & 1
b1 = ((1 - 127) << 23) + (1 << 21) - 1
nval0 = nval0 + b1 + mant_odd0
nval0 = (nval0 >> (23 - 1)).to(tl.uint8)
fp4_0 = tl.where(norm_mask0, nval0, tl.full([1], 0x7, tl.uint8))
fp4_0 = tl.where(den_mask0, dval0, fp4_0)
sign_lp0 = (s0 >> (23 + 8 - 1 - 2)).to(tl.uint8)
fp4_0 |= sign_lp0 # shape: [32]
# 打包子块0
fp4_reshaped0 = tl.reshape(fp4_0, (MXBLK_ // 2, 2)) # [16, 2]
ev0, od0 = tl.split(fp4_reshaped0)
packed0 = (ev0 | (od0 << 4)).to(tl.uint8) # [16]
# 存储子块0的 packed 结果
out_offs0 = g_id * (GROUP_N_ // 2) + tl.arange(0, MXBLK_ // 2)
mask_out0 = out_offs0 < (N // 2)
tl.store(
out_ptr + m_id * stride_out_m + out_offs0 * stride_out_n,
packed0,
mask=mask_out0
)
# 存储 scale0
s0_idx = g_id * 2
if s0_idx < (N // MXBLK_):
tl.store(scale_ptr + m_id * stride_s_m + s0_idx * stride_s_n, bs0)
# ---------- 子块1 (32..63) ----------
offs1 = g_id * GROUP_N_ + MXBLK_ + tl.arange(0, MXBLK_)
mask1 = offs1 < N
x1 = tl.load(
x_ptr + m_id * stride_x_m + offs1 * stride_x_n,
mask=mask1, other=0.0
).to(tl.float32)
amax1 = tl.max(tl.abs(x1))
amax1 = tl.where(amax1 == 0, 1e-8, amax1)
amax_i32 = amax1.to(tl.int32, bitcast=True)
amax_i32 = (amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax1 = amax_i32.to(tl.float32, bitcast=True)
se1 = tl.log2(amax1).floor() - 2
se1 = tl.clamp(se1, -127, 127)
scale1 = tl.exp2(-se1)
bs1 = (se1.to(tl.uint8) + 127)
q1 = x1 * scale1
qbits1 = q1.to(tl.uint32, bitcast=True)
s1 = qbits1 & 0x80000000
qbits1 ^= s1
qf1 = qbits1.to(tl.float32, bitcast=True)
sat_mask1 = qf1 >= 6.0
den_mask1 = (qf1 < 1.0) & (~sat_mask1)
norm_mask1 = ~(sat_mask1 | den_mask1)
dval1 = (qf1 + denorm_float).to(tl.uint32, bitcast=True) - denorm_exp
dval1 = dval1.to(tl.uint8)
nval1 = qbits1
mant_odd1 = (nval1 >> (23 - 1)) & 1
nval1 = nval1 + b1 + mant_odd1
nval1 = (nval1 >> (23 - 1)).to(tl.uint8)
fp4_1 = tl.where(norm_mask1, nval1, tl.full([1], 0x7, tl.uint8))
fp4_1 = tl.where(den_mask1, dval1, fp4_1)
sign_lp1 = (s1 >> (23 + 8 - 1 - 2)).to(tl.uint8)
fp4_1 |= sign_lp1
fp4_reshaped1 = tl.reshape(fp4_1, (MXBLK_ // 2, 2))
ev1, od1 = tl.split(fp4_reshaped1)
packed1 = (ev1 | (od1 << 4)).to(tl.uint8)
# 存储子块1的 packed 结果(紧接着子块0)
out_offs1 = g_id * (GROUP_N_ // 2) + (MXBLK_ // 2) + tl.arange(0, MXBLK_ // 2)
mask_out1 = out_offs1 < (N // 2)
tl.store(
out_ptr + m_id * stride_out_m + out_offs1 * stride_out_n,
packed1,
mask=mask_out1
)
# 存储 scale1
s1_idx = g_id * 2 + 1
if s1_idx < (N // MXBLK_):
tl.store(scale_ptr + m_id * stride_s_m + s1_idx * stride_s_n, bs1)
def dynamic_mxfp4_quant_smallM(A: torch.Tensor):
M, N = A.shape
assert N % MXBLK == 0
out = torch.empty((M, N // 2), dtype=torch.uint8, device=A.device)
scale = torch.empty((M, N // MXBLK), dtype=torch.uint8, device=A.device)
grid = (M, triton.cdiv(N, GROUP_N))
mxfp4_quant_kernel_smallM[grid](
A, out, scale,
*A.stride(), *out.stride(), *scale.stride(),
M=M, N=N, GROUP_N_=GROUP_N, MXBLK_=MXBLK,
num_warps=1
)
return out, scale
_gemm_afp4wfp4_repr = make_kernel_repr(
"_gemm_afp4wfp4_kernel",
[
"BLOCK_SIZE_M",
"BLOCK_SIZE_N",
"BLOCK_SIZE_K",
"GROUP_SIZE_M",
"num_warps",
"num_stages",
"waves_per_eu",
"matrix_instr_nonkdim",
"cache_modifier",
"NUM_KSPLIT",
],
)
@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(repr=_gemm_afp4wfp4_repr)
def _gemm_afp4wfp4_kernel(
a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
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)
tl.assume(pid_k >= 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_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak)
b_ptrs = b_ptr + (offs_k_split[:, None] * stride_bk + offs_bn[None, :] * stride_bn)
offs_ks = (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[None, :] * stride_ask
b_scale_ptrs = b_scales_ptr + offs_bn[:, None] * stride_bsn + offs_ks[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)
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[:, None] < K - k * (BLOCK_SIZE_K // 2), other=0, cache_modifier=cache_modifier)
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) * stride_bk
a_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * stride_ask
b_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * 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)
_gemm_afp4wfp4_preshuffle_scales_repr = make_kernel_repr(
"_gemm_afp4wfp4_preshuffle_kernel",
[
"BLOCK_SIZE_M",
"BLOCK_SIZE_N",
"BLOCK_SIZE_K",
"GROUP_SIZE_M",
"num_warps",
"num_stages",
"waves_per_eu",
"matrix_instr_nonkdim",
"cache_modifier",
"NUM_KSPLIT",
],
)
@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(repr=_gemm_afp4wfp4_preshuffle_scales_repr)
def _gemm_afp4wfp4_kernel_preshuffle_scales(
a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
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_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak)
b_ptrs = b_ptr + (offs_k_split[:, None] * stride_bk + offs_bn[None, :] * stride_bn)
offs_asn = (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_asn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
if BLOCK_SIZE_M < 32:
offs_ks_non_shufl = (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_shufl[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[:, None] < K - k * (BLOCK_SIZE_K // 2), other=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) * 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")
_gemm_afp4wfp4_preshuffle_repr = make_kernel_repr(
"_gemm_afp4wfp4_preshuffle_kernel",
[
"BLOCK_SIZE_M",
"BLOCK_SIZE_N",
"BLOCK_SIZE_K",
"GROUP_SIZE_M",
"num_warps",
"num_stages",
"waves_per_eu",
"matrix_instr_nonkdim",
"cache_modifier",
"NUM_KSPLIT",
],
)
@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(repr=_gemm_afp4wfp4_preshuffle_repr)
def _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_asn = (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_asn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
if BLOCK_SIZE_M < 32:
offs_ks_non_shufl = (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_shufl[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)
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")
_gemm_afp4wfp4_reduce_repr = make_kernel_repr(
"_gemm_afp4wfp4_reduce_kernel",
[
"BLOCK_SIZE_M",
"BLOCK_SIZE_N",
"ACTUAL_KSPLIT",
"MAX_KSPLIT",
],
)
@triton.heuristics({})
@triton.jit(repr=_gemm_afp4wfp4_reduce_repr)
def _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)
def _get_config(M: int, N: int, K: int, shuffle: bool = False):
K = 2 * K
if shuffle:
return get_gemm_config(
"GEMM-AFP4WFP4_PRESHUFFLED",
M, N, K,
bounds=(4, 8, 16, 31, 32, 64, 128, 256, 512, 1024, 2048, 4096, 8192),
)
else:
return get_gemm_config("GEMM-AFP4WFP4", M, N, K)
def custom_kernel(data: input_t) -> output_t:
"""
Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
Replaced Aiter generic gemm call with the direct Triton preshuffle kernel.
"""
def _quant_mxfp4(x, shuffle=False):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
if shuffle:
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 = B.contiguous()
# 保留张量的真实物理内存排布,规避前端类型检查
M, K = A.shape
N, _ = B.shape
A_q, A_scale_sh = dynamic_mxfp4_quant_smallM(A)
# A_q, A_scale_sh = _quant_mxfp4(A, False)
# K dimension handling (fp4 packs 2 elements per uint8, logical physical length K_fp4 = K // 2)
K_fp4 = K // 2
# Fetching configuration and resolving parameters
# config = _get_config(M, N, K_fp4, shuffle=True)
# if hasattr(config, "kwargs"):
# kwargs = config.kwargs
# num_warps = getattr(config, "num_warps", 4)
# num_stages = getattr(config, "num_stages", 4)
# elif isinstance(config, dict):
# kwargs = config
# num_warps = kwargs.get("num_warps", 4)
# num_stages = kwargs.get("num_stages", 4)
# else:
# kwargs = {}
num_warps = 2
num_stages = 2
# BLOCK_SIZE_M 8
# BLOCK_SIZE_N 64
# BLOCK_SIZE_K 512
# GROUP_SIZE_M 4
# num_warps 2
# num_stages 2
# waves_per_eu 4
# matrix_instr_nonkdim 16
# cache_modifier CG
# NUM_KSPLIT_1
BLOCK_SIZE_M = 16 #kwargs.get("BLOCK_SIZE_M", 32)
BLOCK_SIZE_N = 32 #kwargs.get("BLOCK_SIZE_N", 32)
BLOCK_SIZE_K = 512 #kwargs.get("BLOCK_SIZE_K", 256)
# if BLOCK_SIZE_M >= 32 and BLOCK_SIZE_M % 32 != 0:
# BLOCK_SIZE_M = (BLOCK_SIZE_M // 32 + 1) * 32
# if BLOCK_SIZE_N % 32 != 0:
# BLOCK_SIZE_N = max(32, (BLOCK_SIZE_N // 32 + 1) * 32)
# if BLOCK_SIZE_K % 256 != 0:
# BLOCK_SIZE_K = max(256, (BLOCK_SIZE_K // 256 + 1) * 256)
GROUP_SIZE_M = 4 #kwargs.get("GROUP_SIZE_M", 4)
NUM_KSPLIT = 1 #kwargs.get("NUM_KSPLIT", 1)
waves_per_eu = 4 #kwargs.get("waves_per_eu", 0)
matrix_instr_nonkdim = 16 #kwargs.get("matrix_instr_nonkdim", 16)
cache_modifier = "" #kwargs.get("cache_modifier", "")
# shuffle_a = BLOCK_SIZE_M >= 32
# A_q, A_scale_sh = _quant_mxfp4(A, False)
# Preparing splits and outputs
SPLITK_BLOCK_SIZE = max(1, K // NUM_KSPLIT)
SPLITK_BLOCK_SIZE = triton.cdiv(SPLITK_BLOCK_SIZE, BLOCK_SIZE_K) * BLOCK_SIZE_K
if NUM_KSPLIT == 1:
C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
stride_ck = 0
stride_cm, stride_cn = C.stride()
else:
C = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
stride_ck, stride_cm, stride_cn = C.stride()
# Extracting layout strides
stride_am, stride_ak = A_q.stride()
B_shuffle = B_shuffle.contiguous()
stride_bn, stride_bk = B_shuffle.stride()
# Forcing architectural scales strides
A_scale_sh = A_scale_sh.contiguous()
B_scale_sh = B_scale_sh.contiguous()
stride_asm, stride_ask = A_scale_sh.stride()
stride_bsn, stride_bsk = B_scale_sh.stride()
# print(A_q, A_scale_sh)
# 补偿宏块物理跨度,使前端步长与内核寻址游标对齐
stride_bn *= 16
stride_bsn *= 32
# if shuffle_a:
# stride_asm *= 32
grid = lambda META: (
triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]) * META["NUM_KSPLIT"],
)
A_q_uint8 = A_q.contiguous().view(torch.uint8)
A_scale_sh_uint8 = A_scale_sh.contiguous().view(torch.uint8)
B_shuffle_uint8 = B_shuffle.contiguous().view(torch.uint8)
B_scale_sh_uint8 = B_scale_sh.contiguous().view(torch.uint8)
# Executing Triton Kernel
_gemm_afp4wfp4_preshuffle_kernel[grid](
A_q_uint8,
B_shuffle_uint8,
C,
A_scale_sh_uint8,
B_scale_sh_uint8,
M, N, K_fp4,
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=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
BLOCK_SIZE_K=BLOCK_SIZE_K,
GROUP_SIZE_M=GROUP_SIZE_M,
NUM_KSPLIT=NUM_KSPLIT,
SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
num_warps=num_warps,
num_stages=num_stages,
waves_per_eu=waves_per_eu,
matrix_instr_nonkdim=matrix_instr_nonkdim,
cache_modifier=cache_modifier,
)
# Optional reduction when KSPLIT > 1
if NUM_KSPLIT > 1:
C_out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
grid_reduce = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), triton.cdiv(N, META["BLOCK_SIZE_N"]))
_gemm_afp4wfp4_reduce_kernel[grid_reduce](
C, C_out,
M, N,
stride_ck, stride_cm, stride_cn,
C_out.stride(0), C_out.stride(1),
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
ACTUAL_KSPLIT=NUM_KSPLIT,
MAX_KSPLIT=NUM_KSPLIT
)
return C_out
else:
return Cscrolls · 755 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