submission 670608
CaymanYang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 820 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-670608?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:e5617be8ab35e8d790f2fd1971e0794500e055d5583cc1d37705ec8198defc06
license declaredunknown
license concludedunknown
authorsCaymanYang
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Two-stage FP4 path:num-warps = 2
num_warps = 2 if m <= 8 else (4 if m <= 64 else 8)stages = 3
num_stages = 3 if k >= 2048 else 2Kernel source
submission.py820 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Two-stage FP4 path:
1) Quantize bf16 A -> (MXFP4 A_q + E8M0 A_scales).
2) GEMM with tl.dot_scaled using pre-quantized A and shuffled B/B-scales.
"""
from __future__ import annotations
import torch
import triton
import triton.language as tl
from task import input_t, output_t
@triton.jit
def _mxfp4_quant_op(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
EXP_BIAS_FP32: tl.constexpr = 127
EXP_BIAS_FP4: tl.constexpr = 1
EBITS_F32: tl.constexpr = 8
EBITS_FP4: tl.constexpr = 2
MBITS_F32: tl.constexpr = 23
MBITS_FP4: tl.constexpr = 1
max_normal: tl.constexpr = 6
min_normal: tl.constexpr = 1
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)
'''
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
'''#有精度问题
# 位操作: 向上舍入到 2 的幂
amax_bits = amax.to(tl.uint32, bitcast=True) #有精度问题
# 简化: 直接提取指数 + 1 (向上取整效果)
amax_exp = ((amax_bits >> 23) & 0xFF).to(tl.int32)
# 如果尾数非零,指数+1 (向上舍入)
has_mant = (amax_bits & 0x7FFFFF) != 0
amax_exp = amax_exp + has_mant.to(tl.int32)
# E8M0 指数: exp - 2 (让 max 映射到 4.0)
scale_exp = amax_exp - 2 - 127 # 无偏指数
'''
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
#scale_exp = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
#scale_exp = scale_exp.to(tl.float32)
#scale_exp = tl.clamp(scale_exp, -127, 127)
#bs_e8m0 = (scale_exp + 127).to(tl.uint8)
# 扩展 scale 到每个元素
scale = tl.exp2(scale_e8m0_unbiased.to(tl.float32))
scale = tl.broadcast_to(
scale, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE]
)
# ===== 步骤3: 按原始位级流程量化,保证与基准语义一致 =====
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
qx_bits = qx.to(tl.uint32, bitcast=True)
s = qx_bits & 0x80000000
qx_mag = qx_bits ^ s
qx_fp32 = qx_mag.to(tl.float32, bitcast=True)
saturate_mask = qx_fp32 >= max_normal
denormal_mask = (~saturate_mask) & (qx_fp32 < min_normal)
normal_mask = ~(saturate_mask | denormal_mask)
denorm_exp: tl.constexpr = (
(EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
)
denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
denormal_x = qx_fp32 + denorm_mask_float
denormal_x = denormal_x.to(tl.uint32, bitcast=True)
denormal_x -= denorm_mask_int
denormal_x = denormal_x.to(tl.uint8)
normal_x = qx_mag
mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
normal_x += val_to_add
normal_x += mant_odd
normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
normal_x = normal_x.to(tl.uint8)
e2m1_value = tl.full(
[BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE], 0x7, dtype=tl.uint8
)
e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
sign_lp = sign_lp.to(tl.uint8)
e2m1_value = e2m1_value | sign_lp
'''
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
qx = qx ^ s
qx_fp32 = qx.to(tl.float32, bitcast=True)
saturate_mask = qx_fp32 >= max_normal
denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
normal_mask = not (saturate_mask | denormal_mask)
denorm_exp: tl.constexpr = (
(EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
)
denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
denormal_x = qx_fp32 + denorm_mask_float
denormal_x = denormal_x.to(tl.uint32, bitcast=True)
denormal_x -= denorm_mask_int
denormal_x = denormal_x.to(tl.uint8)
normal_x = qx
mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
normal_x += val_to_add
normal_x += mant_odd
normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
normal_x = normal_x.to(tl.uint8)
e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
sign_lp = sign_lp.to(tl.uint8)
e2m1_value = e2m1_value | sign_lp
'''
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)
def _quant_config(m: int, k: int) -> tuple[int, int, int, int, int]:
"""
Return (BLOCK_M, BLOCK_K, NUM_WARPS, NUM_STAGES, K_CHUNK) for
_amd_mxfp4_quant_a_kernel.
Quant kernel parallelizes M and K-chunk in 2D grid, and each CTA walks its
K-chunk with inner tl.range. num_stages participates in software pipelining.
"""
bench: dict[tuple[int, int], tuple[int, int, int, int, int]] = {
# Short-K: larger K tile with moderate pipeline depth.
(4, 512): (4, 512, 2, 2, 512),
(32, 512): (32, 512, 4, 2, 512),
# Long-K: reduce BLOCK_K for register pressure; increase stages to hide latency.
(16, 7168): (16, 128, 4, 3, 256),
(64, 2048): (32, 128, 4, 3, 256),
(256, 1536): (32, 64, 4, 3, 256),
}
if (m, k) in bench:
return bench[(m, k)]
block_m = 16 if m <= 16 else (32 if m <= 64 else 64)
block_k = 512 if k <= 1024 else 256
num_warps = 2 if m <= 8 else (4 if m <= 64 else 8)
num_stages = 3 if k >= 2048 else 2
# Keep each K-chunk large enough for overlap but bounded for occupancy.
k_chunk = block_k * (4 if k >= 2048 else 2)
return (block_m, block_k, num_warps, num_stages, k_chunk)
def _gemm_config(m: int, n: int, k: int) -> tuple[int, int, int, int, int]:
"""
Return (BLOCK_M, BLOCK_N, BLOCK_K, NUM_WARPS, NUM_STAGES).
"""
bench: dict[tuple[int, int, int], tuple[int, int, int, int, int]] = {
#(4, 2880, 512): (16, 128, 64, 4, 2),
(4, 2880, 512): (16, 16, 128, 4, 3),#fix config
# Long-K: deeper pipeline to hide global-memory latency.
#(16, 2112, 7168): (16, 128, 128, 4, 3),
(16, 2112, 7168): (16, 16, 512, 4, 3),
#(8, 2112, 7168): (16, 16, 128, 4, 2),
(32, 4096, 512): (16, 16, 128, 4, 3), #fix config
(32, 2880, 512): (16, 16, 128, 4, 3), #fix config
# Case5: recover opt_3 behavior (more CTAs, lower per-CTA pressure).
(64, 7168, 2048): (64, 16, 256, 4, 3),
(256, 3072, 1536): (32, 16, 256, 4, 3),
}
if (m, n, k) in bench:
return bench[(m, n, k)]
# Unknown benchmark shapes: avoid KeyError on CI / extra tests.
if m <= 16:
if k >= 4096:
return (16, 128, 256, 8, 2)
return (16, 128, 128, 4, 2)
if m <= 64:
if m <= 32:
return (32, 128, 256, 8, 2)
return (64, 128, 256, 8, 2)
if k >= 2048:
return (64, 128, 256, 8, 2)
return (64, 128, 256, 8, 2)
@triton.jit
def _amd_mxfp4_quant_a_kernel(
a_bf16_ptr,
a_q_ptr,
a_sc_ptr,
stride_a_m,
stride_a_k,
stride_aq_m,
stride_aq_kh,
stride_asc_m,
stride_asc_kg,
M,
K,
BLOCK_M: tl.constexpr,
BLOCK_K: tl.constexpr,
NUM_STAGES: tl.constexpr,
K_CHUNK: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_k_chunk = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
m_mask = offs_m < M
quant_block_size: tl.constexpr = 32
k_half = K // 2
k_scale = K // 32
k_begin = pid_k_chunk * K_CHUNK
for k_iter in tl.range(0, K_CHUNK, BLOCK_K, num_stages=NUM_STAGES):
k0 = k_begin + k_iter
offs_kh = (k0 // 2) + tl.arange(0, BLOCK_K // 2)
offs_kg = (k0 // 32) + tl.arange(0, BLOCK_K // 32)
a_block_ptr = tl.make_block_ptr(
base=a_bf16_ptr,
shape=(M, K),
strides=(stride_a_m, stride_a_k),
offsets=(pid_m * BLOCK_M, k0),
block_shape=(BLOCK_M, BLOCK_K),
order=(1, 0),
)
a_f32 = tl.load(a_block_ptr, boundary_check=(0, 1), padding_option="zero").to(
tl.float32
)
a_q, a_scales = _mxfp4_quant_op(
a_f32,
BLOCK_SIZE_N=BLOCK_K,
BLOCK_SIZE_M=BLOCK_M,
MXFP4_QUANT_BLOCK_SIZE=quant_block_size,
)
tl.store(
a_q_ptr + offs_m[:, None] * stride_aq_m + offs_kh[None, :] * stride_aq_kh,
a_q,
mask=m_mask[:, None] & (offs_kh[None, :] < k_half),
)
tl.store(
a_sc_ptr + offs_m[:, None] * stride_asc_m + offs_kg[None, :] * stride_asc_kg,
a_scales,
mask=m_mask[:, None] & (offs_kg[None, :] < k_scale),
)
@triton.jit
def _amd_mxfp4_qs_gemm_kernel(
a_q_ptr,
a_sc_ptr,
b_sh_ptr,
b_sc_sh_ptr,
c_ptr,
stride_aq_m,
stride_aq_kh,
stride_asc_m,
stride_asc_kg,
stride_b_m,
stride_b_kh,
stride_bs_m,
stride_bs_n,
stride_c_m,
stride_c_n,
M,
N,
K,
BS_SN,
N_TILES: tl.constexpr,
K_TILES: tl.constexpr,
USE_B_SHUFFLE_TILE_LOAD: tl.constexpr,
USE_BS_SHUFFLE_TILE_LOAD: tl.constexpr,
BS_BLOCK_K_TILES: tl.constexpr,
BS_BLOCK4_TILES: tl.constexpr,
BLOCK_K_TILES: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
NUM_STAGES: 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)
offs_k = tl.arange(0, BLOCK_K)
offs_kh = tl.arange(0, BLOCK_K // 2)
offs_kg = tl.arange(0, BLOCK_K // 32)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k_half = K // 2
k_blocks_half = k_half // 32
k_scale_valid = K // 32
for k0 in tl.range(0, K, BLOCK_K, num_stages=NUM_STAGES):
kh0 = k0 // 2
kg0 = k0 // 32
a_kh = kh0 + offs_kh
a_q = tl.load(
a_q_ptr + offs_m[:, None] * stride_aq_m + a_kh[None, :] * stride_aq_kh,
mask=(offs_m[:, None] < M) & (a_kh[None, :] < k_half),
other=0,
cache_modifier=".ca",
)
a_kg = kg0 + offs_kg
a_scales = tl.load(
a_sc_ptr + offs_m[:, None] * stride_asc_m + a_kg[None, :] * stride_asc_kg,
mask=(offs_m[:, None] < M) & (a_kg[None, :] < k_scale_valid),
other=127,
cache_modifier=".ca",
)
#b_kh = kh0 + offs_kh
#bn = offs_n[None, :] // 16
#bi = offs_n[None, :] % 16
#kb = b_kh[:, None] // 32
#kk = b_kh[:, None] % 32
#k4 = kk // 16
#k5 = kk % 16
# BLOCK_N % 16 == 0, BLOCK_K multiple of 64: load BLOCK_N//16 shuffle n-tiles along dim 1.
# Raw (Nn,Kt,2,16,16) -> permute (1,2,4,0,3) -> (k_tile,sub,k5,n_tile,bi) -> row-major (kh,n) for dot_scaled.
'''
if (
tl.constexpr(USE_B_SHUFFLE_TILE_LOAD)
and tl.constexpr(BLOCK_N % 16 == 0)
and tl.constexpr(BLOCK_K_TILES >= 1)
):
'''
kb0 = kh0 // 32
b_block_ptr = tl.make_block_ptr(
base=b_sh_ptr,
shape=(1, N_TILES, K_TILES, 2, 16, 16),
strides=(
N_TILES * K_TILES * 512,
K_TILES * 512,
512,
256,
16,
1,
),
offsets=(0, pid_n * (BLOCK_N // 16), kb0, 0, 0, 0),
block_shape=(1, BLOCK_N // 16, BLOCK_K_TILES, 2, 16, 16),
order=(5, 4, 3, 2, 1, 0),
)
b_raw = tl.load(b_block_ptr)
b4 = tl.reshape(b_raw, (BLOCK_N // 16, BLOCK_K_TILES, 2, 16, 16))
b4 = tl.permute(b4, (1, 2, 4, 0, 3))
b = tl.reshape(b4, (BLOCK_K // 2, BLOCK_N))
'''
else:
b_lin = (((((bn * k_blocks_half + kb) * 2 + k4) * 16 + bi) * 16) + k5)
b_row = b_lin // k_half
b_col = b_lin % k_half
b = tl.load(
b_sh_ptr + b_row * stride_b_m + b_col * stride_b_kh,
mask=(offs_n[None, :] < N) & (b_kh[:, None] < k_half),
other=0,
cache_modifier=".cg",
)
'''
b_kg = kg0 + offs_kg
b_sc_mask = (offs_n[:, None] < N) & (b_kg[None, :] < k_scale_valid)
if tl.constexpr(USE_BS_SHUFFLE_TILE_LOAD):
# Fast path for B_scale_sh shuffle layout:
# load (b3, b4, b5, b2) tile and remap to (n, kg).
n_tile16 = pid_n
n_block32 = n_tile16 // 2
n_half16 = n_tile16 % 2
kg_block8 = kg0 // 8
kg_half4 = (kg0 % 8) // 4
bs_base = (
b_sc_sh_ptr
+ n_block32 * 32 * stride_bs_m
+ (kg_block8 * 256 + kg_half4 * 2 + n_half16) * stride_bs_n
)
bs_block_ptr = tl.make_block_ptr(
base=bs_base,
shape=(BS_BLOCK_K_TILES, BS_BLOCK4_TILES, 4, 16),
strides=(256, 2, 64, 4),
offsets=(0, 0, 0, 0),
block_shape=(BS_BLOCK_K_TILES, BS_BLOCK4_TILES, 4, 16),
order=(3, 2, 1, 0),
)
bs_raw = tl.load(bs_block_ptr)
bs4 = tl.permute(bs_raw, (3, 0, 1, 2))
b_scales = tl.reshape(bs4, (BLOCK_N, BLOCK_K // 32))
b_scales = tl.where(b_sc_mask, b_scales, 127)
else:
b0 = offs_n[:, None] // 32
b1 = offs_n[:, None] % 32
b2 = b1 % 16
b1 = b1 // 16
b3 = b_kg[None, :] // 8
b4 = b_kg[None, :] % 8
b5 = b4 % 4
b4 = b4 // 4
b_lin_sc = b1 + b4 * 2 + b2 * 4 + b5 * 64 + b3 * 256 + b0 * 32 * BS_SN
b_sm = b_lin_sc // BS_SN
b_sn = b_lin_sc % BS_SN
b_scales = tl.load(
b_sc_sh_ptr + b_sm * stride_bs_m + b_sn * stride_bs_n,
mask=b_sc_mask,
other=127,
cache_modifier=".cg",
)
acc = tl.dot_scaled(a_q, a_scales, "e2m1", b, b_scales, "e2m1", acc)
c = acc.to(tl.bfloat16)
c_ptrs = c_ptr + offs_m[:, None] * stride_c_m + offs_n[None, :] * stride_c_n
c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
@triton.jit
def _amd_fused_mxfp4_qs_gemm_kernel(
a_bf16_ptr,
b_sh_ptr,
b_sc_sh_ptr,
c_ptr,
stride_a_m,
stride_a_k,
stride_b_m,
stride_b_kh,
stride_bs_m,
stride_bs_n,
stride_c_m,
stride_c_n,
M,
N,
K,
BS_SN,
N_TILES: tl.constexpr,
K_TILES: tl.constexpr,
USE_B_SHUFFLE_TILE_LOAD: tl.constexpr,
USE_BS_SHUFFLE_TILE_LOAD: tl.constexpr,
BS_BLOCK_K_TILES: tl.constexpr,
BS_BLOCK4_TILES: tl.constexpr,
BLOCK_K_TILES: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
NUM_STAGES: 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)
offs_k = tl.arange(0, BLOCK_K)
offs_kh = tl.arange(0, BLOCK_K // 2)
offs_kg = tl.arange(0, BLOCK_K // 32)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k_half = K // 2
k_blocks_half = k_half // 32
k_scale_valid = K // 32
quant_block_size: tl.constexpr = 32
for k0 in tl.range(0, K, BLOCK_K, num_stages=NUM_STAGES):
kh0 = k0 // 2
kg0 = k0 // 32
a_block_ptr = tl.make_block_ptr(
base=a_bf16_ptr,
shape=(M, K),
strides=(stride_a_m, stride_a_k),
offsets=(pid_m * BLOCK_M, k0),
block_shape=(BLOCK_M, BLOCK_K),
order=(1, 0),
)
a_f32 = tl.load(
a_block_ptr, boundary_check=(0, 1), padding_option="zero"
).to(tl.float32)
a_q, a_scales = _mxfp4_quant_op(
a_f32,
BLOCK_SIZE_N=BLOCK_K,
BLOCK_SIZE_M=BLOCK_M,
MXFP4_QUANT_BLOCK_SIZE=quant_block_size,
)
#b_kh = kh0 + offs_kh
#bn = offs_n[None, :] // 16
#bi = offs_n[None, :] % 16
#kb = b_kh[:, None] // 32
#kk = b_kh[:, None] % 32
#k4 = kk // 16
#k5 = kk % 16
# Same B layout as _amd_mxfp4_qs_gemm_kernel; BLOCK_K_TILES = BLOCK_K//64.
'''
if (
tl.constexpr(USE_B_SHUFFLE_TILE_LOAD)
and tl.constexpr(BLOCK_N % 16 == 0)
and tl.constexpr(BLOCK_K_TILES >= 1)
):
'''
kb0 = kh0 // 32
b_block_ptr = tl.make_block_ptr(
base=b_sh_ptr,
shape=(1, N_TILES, K_TILES, 2, 16, 16),
strides=(
N_TILES * K_TILES * 512,
K_TILES * 512,
512,
256,
16,
1,
),
offsets=(0, pid_n * (BLOCK_N // 16), kb0, 0, 0, 0),
block_shape=(1, BLOCK_N // 16, BLOCK_K_TILES, 2, 16, 16),
order=(5, 4, 3, 2, 1, 0),
)
b_raw = tl.load(b_block_ptr)
b4 = tl.reshape(b_raw, (BLOCK_N // 16, BLOCK_K_TILES, 2, 16, 16))
b4 = tl.permute(b4, (1, 2, 4, 0, 3))
b = tl.reshape(b4, (BLOCK_K // 2, BLOCK_N))
'''
else:
b_lin = (((((bn * k_blocks_half + kb) * 2 + k4) * 16 + bi) * 16) + k5)
b_row = b_lin // k_half
b_col = b_lin % k_half
b = tl.load(
b_sh_ptr + b_row * stride_b_m + b_col * stride_b_kh,
mask=(offs_n[None, :] < N) & (b_kh[:, None] < k_half),
other=0,
cache_modifier=".cg",
)
'''
b_kg = kg0 + offs_kg
b_sc_mask = (offs_n[:, None] < N) & (b_kg[None, :] < k_scale_valid)
if tl.constexpr(USE_BS_SHUFFLE_TILE_LOAD):
n_tile16 = pid_n
n_block32 = n_tile16 // 2
n_half16 = n_tile16 % 2
kg_block8 = kg0 // 8
kg_half4 = (kg0 % 8) // 4
bs_base = (
b_sc_sh_ptr
+ n_block32 * 32 * stride_bs_m
+ (kg_block8 * 256 + kg_half4 * 2 + n_half16) * stride_bs_n
)
bs_block_ptr = tl.make_block_ptr(
base=bs_base,
shape=(BS_BLOCK_K_TILES, BS_BLOCK4_TILES, 4, 16),
strides=(256, 2, 64, 4),
offsets=(0, 0, 0, 0),
block_shape=(BS_BLOCK_K_TILES, BS_BLOCK4_TILES, 4, 16),
order=(3, 2, 1, 0),
)
bs_raw = tl.load(bs_block_ptr)
bs4 = tl.permute(bs_raw, (3, 0, 1, 2))
b_scales = tl.reshape(bs4, (BLOCK_N, BLOCK_K // 32))
b_scales = tl.where(b_sc_mask, b_scales, 127)
else:
b0 = offs_n[:, None] // 32
b1 = offs_n[:, None] % 32
b2 = b1 % 16
b1 = b1 // 16
b3 = b_kg[None, :] // 8
b4 = b_kg[None, :] % 8
b5 = b4 % 4
b4 = b4 // 4
b_lin_sc = b1 + b4 * 2 + b2 * 4 + b5 * 64 + b3 * 256 + b0 * 32 * BS_SN
b_sm = b_lin_sc // BS_SN
b_sn = b_lin_sc % BS_SN
b_scales = tl.load(
b_sc_sh_ptr + b_sm * stride_bs_m + b_sn * stride_bs_n,
mask=b_sc_mask,
other=127,
cache_modifier=".cg",
)
acc = tl.dot_scaled(a_q, a_scales, "e2m1", b, b_scales, "e2m1", acc)
c = acc.to(tl.bfloat16)
c_ptrs = c_ptr + offs_m[:, None] * stride_c_m + offs_n[None, :] * stride_c_n
c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
#mxfp4_gemm_shuffle_direct = _amd_fused_mxfp4_qs_gemm_kernel
def _b_shuffle_tile_load_ok(B_shuffle: torch.Tensor, block_n: int, block_k: int) -> bool:
"""
Block contiguous load matches shuffle_weight [N, K/2] layout when:
- BLOCK_N is a multiple of 16 (covers BLOCK_N//16 shuffle n-tiles per program),
- K-tile count BLOCK_K//64 is integer (BLOCK_K divisible by 64),
- tensor is contiguous row-major (stride (kh, 1)).
"""
if block_n % 16 != 0:
return False
if block_k % 64 != 0:
return False
if not B_shuffle.is_contiguous():
return False
n, kh = B_shuffle.shape
return B_shuffle.stride(0) == kh and B_shuffle.stride(1) == 1
def _b_scale_shuffle_tile_load_ok(
B_scale_sh: torch.Tensor, block_n: int, block_k: int
) -> bool:
"""
Fast path for B_scale_sh shuffle layout.
Conditions match current kernel mapping:
- BLOCK_N == 16 (single half-32 n tile per program),
- BLOCK_K in {128, 256, 512} (and 128-aligned),
- B_scale_sh is contiguous row-major with stride (BS_SN, 1).
"""
if block_n != 16:
return False
if block_k % 128 != 0:
return False
if block_k > 512:
return False
if not B_scale_sh.is_contiguous():
return False
bs_m, bs_n = B_scale_sh.shape
if bs_n % 8 != 0:
return False
return B_scale_sh.stride(0) == bs_n and B_scale_sh.stride(1) == 1
def amd_fused_mxfp4_qs_gemm(
A: torch.Tensor,
B_shuffle: torch.Tensor,
B_scale_sh: torch.Tensor,
*,
dtype: torch.dtype,
) -> torch.Tensor:
m, k = A.shape
n, kh_b = B_shuffle.shape
assert kh_b * 2 == k, "A and B_shuffle K mismatch"
assert k % 64 == 0, "K must be divisible by 64"
if B_shuffle.dtype != torch.uint8:
B_shuffle = B_shuffle.view(torch.uint8)
if B_scale_sh.dtype != torch.uint8:
B_scale_sh = B_scale_sh.view(torch.uint8)
out = torch.empty((m, n), device=A.device, dtype=dtype)
block_m, block_n, block_k, num_warps, num_stages = _gemm_config(m, n, k)
grid = (triton.cdiv(m, block_m), triton.cdiv(n, block_n))
n_tiles = n // 16
k_tiles = k // 64
use_b_tile = _b_shuffle_tile_load_ok(B_shuffle, block_n, block_k)
use_bs_tile = _b_scale_shuffle_tile_load_ok(B_scale_sh, block_n, block_k)
block_k_tiles = block_k // 64
bs_block_k_tiles = max(1, block_k // 256)
bs_block4_tiles = 1 if block_k == 128 else 2
_amd_fused_mxfp4_qs_gemm_kernel[grid](
A,
B_shuffle,
B_scale_sh,
out,
*A.stride(),
*B_shuffle.stride(),
*B_scale_sh.stride(),
*out.stride(),
m,
n,
k,
B_scale_sh.shape[1],
N_TILES=n_tiles,
K_TILES=k_tiles,
USE_B_SHUFFLE_TILE_LOAD=use_b_tile,
USE_BS_SHUFFLE_TILE_LOAD=use_bs_tile,
BS_BLOCK_K_TILES=bs_block_k_tiles,
BS_BLOCK4_TILES=bs_block4_tiles,
BLOCK_K_TILES=block_k_tiles,
BLOCK_M=block_m,
BLOCK_N=block_n,
BLOCK_K=block_k,
NUM_STAGES=num_stages,
num_warps=num_warps,
num_stages=num_stages,
)
return out
def amd_mxfp4_qs_gemm_two_stage(
A: torch.Tensor,
B_shuffle: torch.Tensor,
B_scale_sh: torch.Tensor,
*,
dtype: torch.dtype,
) -> torch.Tensor:
m, k = A.shape
n, kh_b = B_shuffle.shape
assert kh_b * 2 == k, "A and B_shuffle K mismatch"
assert k % 64 == 0, "K must be divisible by 64"
if B_shuffle.dtype != torch.uint8:
B_shuffle = B_shuffle.view(torch.uint8)
if B_scale_sh.dtype != torch.uint8:
B_scale_sh = B_scale_sh.view(torch.uint8)
out = torch.empty((m, n), device=A.device, dtype=dtype)
q_block_m, q_block_k, q_num_warps, q_num_stages, q_k_chunk = _quant_config(m, k)
g_block_m, g_block_n, g_block_k, g_num_warps, g_num_stages = _gemm_config(m, n, k)
a_q = torch.empty((m, k // 2), device=A.device, dtype=torch.uint8)
a_scales = torch.empty((m, k // 32), device=A.device, dtype=torch.uint8)
grid_quant = (triton.cdiv(m, q_block_m), triton.cdiv(k, q_k_chunk))
_amd_mxfp4_quant_a_kernel[grid_quant](
A,
a_q,
a_scales,
*A.stride(),
*a_q.stride(),
*a_scales.stride(),
m,
k,
BLOCK_M=q_block_m,
BLOCK_K=q_block_k,
NUM_STAGES=q_num_stages,
K_CHUNK=q_k_chunk,
num_warps=q_num_warps,
num_stages=q_num_stages,
)
grid = (triton.cdiv(m, g_block_m), triton.cdiv(n, g_block_n))
n_tiles = n // 16
k_tiles = k // 64
use_b_tile = _b_shuffle_tile_load_ok(B_shuffle, g_block_n, g_block_k)
use_bs_tile = _b_scale_shuffle_tile_load_ok(B_scale_sh, g_block_n, g_block_k)
block_k_tiles = g_block_k // 64
bs_block_k_tiles = max(1, g_block_k // 256)
bs_block4_tiles = 1 if g_block_k == 128 else 2
_amd_mxfp4_qs_gemm_kernel[grid](
a_q,
a_scales,
B_shuffle,
B_scale_sh,
out,
*a_q.stride(),
*a_scales.stride(),
*B_shuffle.stride(),
*B_scale_sh.stride(),
*out.stride(),
m,
n,
k,
B_scale_sh.shape[1],
N_TILES=n_tiles,
K_TILES=k_tiles,
USE_B_SHUFFLE_TILE_LOAD=use_b_tile,
USE_BS_SHUFFLE_TILE_LOAD=use_bs_tile,
BS_BLOCK_K_TILES=bs_block_k_tiles,
BS_BLOCK4_TILES=bs_block4_tiles,
BLOCK_K_TILES=block_k_tiles,
BLOCK_M=g_block_m,
BLOCK_N=g_block_n,
BLOCK_K=g_block_k,
NUM_STAGES=g_num_stages,
num_warps=g_num_warps,
num_stages=g_num_stages,
)
return out
def custom_kernel(data: input_t) -> output_t:
A, _B, _B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
# First threshold pass: small-M/small-K tends to favor fused path.
use_fused = (k <1536 )
if use_fused:
return amd_fused_mxfp4_qs_gemm(
A,
B_shuffle,
B_scale_sh,
dtype=torch.bfloat16,
)
return amd_mxfp4_qs_gemm_two_stage(
#return amd_fused_mxfp4_qs_gemm(
A,
B_shuffle,
B_scale_sh,
dtype=torch.bfloat16,
)
scrolls · 820 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