submission 721183
zainhaider5020 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 736 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-721183?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:7f6e57a166ef2fcd418ed74d332ae61229e76870a1319e3f1f663b277b2f76cc
license declaredunknown
license concludedunknown
authorszainhaider5020
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""Optimized MXFP4 quant: replaces log2/exp2 with bitwise float32 exponent ops."""num-warps = 1
NUM_WARPS = 1split-k
get_splitk,stages = 1
NUM_STAGES = 1tile-k = 512
_BK = 512tile-m = 64
BLOCK_M = 64tile-n = 32
BLOCK_N = 32Kernel source
submission.py736 lines
"""v87: waves_per_eu=4 for s1/s3/s4 to improve occupancy and memory hiding."""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
get_splitk,
)
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
@triton.jit
def _mxfp4_quant_op_fast(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE):
"""Optimized MXFP4 quant: replaces log2/exp2 with bitwise float32 exponent ops."""
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)
# Calculate scale using bitwise ops (no log2/exp2 transcendental functions)
amax_f32 = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax_i32 = amax_f32.to(tl.int32, bitcast=True)
amax_u32 = (amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
# Extract biased exponent from float32 bits (bits 30..23)
biased_exp = tl.cast((amax_u32 >> 23) & 0xFF, tl.int32)
# bs_e8m0 = clamp(biased_exp - 2, 0, 254) equivalent to original formula
bs_e8m0_int = tl.maximum(tl.minimum(biased_exp - 2, 254), 0)
bs_e8m0 = bs_e8m0_int.to(tl.uint8)
# quant_scale = 2^(127 - bs_e8m0): float32 with biased_exp=(254-bs_e8m0), mantissa=0
quant_scale_u32 = tl.cast(254 - bs_e8m0_int, tl.uint32) << 23
quant_scale = quant_scale_u32.to(tl.float32, bitcast=True)
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)
_BK = 512
_BUF: dict = {}
_PARAMS: dict = {}
_PARTIAL: dict = {}
_RUNNER: dict = {}
_RUNNER_ARGS: dict = {}
_REDUCE_RUNNER: dict = {}
_REDUCE_RUNNER_ARGS: dict = {}
_STATES: dict = {}
_BK_CACHE: dict = {}
_BS_CACHE: dict = {}
_DISPATCH_CACHE: dict = {}
def _use_preshuffle(M: int, N: int, K: int) -> bool:
return K <= 512 or (K == 7168 and M <= 16)
def _unwrap_jit_fn(jit_fn):
while not hasattr(jit_fn, "device_caches"):
jit_fn = jit_fn.fn
return jit_fn
def _get_kernel_cache(jit_fn):
try:
inner = _unwrap_jit_fn(jit_fn)
device = torch.cuda.current_device()
kernel_cache = inner.device_caches[device][0]
return kernel_cache, device
except Exception:
return None, None
def _extract_runner(kernel_cache, keys_before, grid):
try:
new_keys = set(kernel_cache.keys()) - keys_before
if not new_keys:
return None
compiled_kernel = kernel_cache[next(iter(new_keys))]
compiled_kernel._init_handles()
if len(grid) == 1:
grid = (grid[0], 1, 1)
elif len(grid) == 2:
grid = (grid[0], grid[1], 1)
return compiled_kernel[grid]
except Exception:
return None
def _check_even_k(K_int: int, bk: int, spk_sz: int) -> bool:
return K_int % (bk // 2) == 0 and spk_sz % bk == 0 and K_int % (spk_sz // 2) == 0
def _shape_config(M: int, N: int = 0, K_bf16: int = 0) -> dict:
# Shape2: bm=16, nsplit=8, 264 blocks; 2 K-iters per block → enable num_stages=2 pipelining!
# spk_sz=1024, bk=512 → num_k_iter=2 per block; pipeline overlaps K-iter1 load with K-iter0 compute
if N == 2112 and K_bf16 == 7168 and M <= 16:
return {
"bm": 16,
"bn": 64,
"num_warps": 4,
"waves_per_eu": 2,
"cache_modifier": ".cg",
"num_stages": 2,
}
# Shape1: m=4, k=512 → bm=4, bn=32 (90 blocks, 30% CU) + .cg + waves=4
if M <= 8 and K_bf16 <= 512:
return {
"bm": 4,
"bn": 32,
"num_warps": 2,
"waves_per_eu": 4,
"cache_modifier": ".cg",
}
# Shapes 3,4: m=32, k=512 → bm=16, bn=32 (s3: 256 blocks, s4: 180 blocks) + .cg + waves=4
if M <= 32 and K_bf16 <= 512:
return {
"bm": 16,
"bn": 32,
"num_warps": 4,
"waves_per_eu": 4,
"cache_modifier": ".cg",
}
return {
"bm": 16,
"bn": 64,
"num_warps": 4,
"waves_per_eu": 2,
"cache_modifier": None,
}
def _splitk_target(M: int, N: int, K_bf16: int) -> int:
# Shape2: use 14 splits → 462 blocks (2 partial waves at 152% CU util)
# vs 8 splits (231 useful blocks, 76% util). 14-split fills all 304 CUs in wave1!
# Each block: 1 K-iter (vs 2 with 8-split), so time_per_block is halved.
# Wave1: 304 blocks × T/2 → fills 100% CUs; Wave2: 158 blocks × T/2
# Total ≈ T/2 × (1 + 158/304) = 0.76T vs 7-split's T (full T with 76% util)
if N == 2112 and K_bf16 == 7168 and M <= 16:
return 14
return 1
def _build_params(M: int, N: int, K_bf16: int) -> dict:
cfg = _shape_config(M, N, K_bf16)
K_int = K_bf16 // 2
num_ksplit = _splitk_target(M, N, K_bf16)
if num_ksplit > 1:
spk_sz, bk, num_ksplit = get_splitk(K_int, _BK, num_ksplit)
else:
spk_sz = 2 * K_int
bk = triton.next_power_of_2(2 * K_int) if _BK >= 2 * K_int else _BK
g0 = triton.cdiv(M, cfg["bm"]) * triton.cdiv(N, cfg["bn"])
grid = (g0 * num_ksplit,)
actual_ksplit = triton.cdiv(K_int, spk_sz // 2)
max_ksplit = triton.next_power_of_2(num_ksplit)
reduce_grid = (triton.cdiv(M, 16), triton.cdiv(N, 16)) if num_ksplit > 1 else None
return {
"bm": cfg["bm"],
"bn": cfg["bn"],
"num_warps": cfg["num_warps"],
"waves_per_eu": cfg["waves_per_eu"],
"cache_modifier": cfg["cache_modifier"],
"bk": bk,
"spk_sz": spk_sz,
"grid": grid,
"g0": g0,
"K_int": K_int,
"num_ksplit": num_ksplit,
"actual_ksplit": actual_ksplit,
"max_ksplit": max_ksplit,
"reduce_grid": reduce_grid,
}
def _build_main_args(p, A, B_k, C_out, B_s):
K_int = p["K_int"]
EVEN_K = _check_even_k(K_int, p["bk"], p["spk_sz"])
GRID_MN = p["g0"]
if p["num_ksplit"] == 1:
stride_ck = 0
stride_cm = C_out.stride(0)
stride_cn = C_out.stride(1)
else:
stride_ck = C_out.stride(0)
stride_cm = C_out.stride(1)
stride_cn = C_out.stride(2)
return [
A,
B_k,
C_out,
B_s,
p["M_val"],
p["N_val"],
K_int,
A.stride(0),
A.stride(1),
B_k.stride(0),
B_k.stride(1),
stride_ck,
stride_cm,
stride_cn,
B_s.stride(0),
B_s.stride(1),
p["bm"],
p["bn"],
p["bk"],
1,
p["num_ksplit"],
p["spk_sz"],
EVEN_K,
p["num_warps"],
p.get("num_stages", 1),
p["waves_per_eu"],
16,
GRID_MN,
True,
p["cache_modifier"],
]
def _build_reduce_args(p, y_pp, C):
return [
y_pp,
C,
p["M_val"],
p["N_val"],
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
C.stride(0),
C.stride(1),
16,
16,
p["actual_ksplit"],
p["max_ksplit"],
]
@triton.jit
def _mxfp4_quant_shuffled_kernel(
x_ptr,
x_fp4_ptr,
bs_shuffled_ptr,
stride_x_m,
stride_x_n,
stride_fp4_m,
stride_fp4_n,
M,
N,
SN,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
EVEN_M_N: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
stride_x_m = tl.cast(stride_x_m, tl.int64)
stride_x_n = tl.cast(stride_x_n, tl.int64)
stride_fp4_m = tl.cast(stride_fp4_m, tl.int64)
stride_fp4_n = tl.cast(stride_fp4_n, tl.int64)
SN64 = tl.cast(SN, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, start_n + NUM_ITER, num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(
x_ptr + x_offs, mask=x_mask, other=0.0, cache_modifier=".cg"
).to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op_fast(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
fp4_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
fp4_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
fp4_offs = (
fp4_offs_m[:, None] * stride_fp4_m + fp4_offs_n[None, :] * stride_fp4_n
)
if EVEN_M_N:
tl.store(x_fp4_ptr + fp4_offs, out_tensor)
else:
fp4_mask = (fp4_offs_m < M)[:, None] & (fp4_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + fp4_offs, out_tensor, mask=fp4_mask)
x_idx = tl.cast(pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M), tl.int64)
y_idx = tl.cast(
pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS), tl.int64
)
xr = x_idx[:, None]
yr = y_idx[None, :]
sh_idx = (
(xr // 32 * SN64) * 32
+ (yr // 8) * 256
+ (yr % 4) * 64
+ (xr % 16) * 4
+ (yr % 8) // 4 * 2
+ (xr % 32) // 16
)
if EVEN_M_N:
tl.store(bs_shuffled_ptr + sh_idx, bs_e8m0)
else:
bs_mask = (x_idx[:, None] < M) & (
y_idx[None, :]
< (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
)
tl.store(bs_shuffled_ptr + sh_idx, bs_e8m0, mask=bs_mask)
def _quant_config(M: int, K: int) -> dict:
if M <= 32:
BLOCK_M = triton.next_power_of_2(M)
BLOCK_N = 32
NUM_ITER = 1
NUM_STAGES = 1
NUM_WARPS = 1
else:
BLOCK_M = 64
BLOCK_N = 64
NUM_ITER = 4
NUM_STAGES = 2
NUM_WARPS = 4
if K <= 16384:
# s5 (M=64, K=2048): BLOCK_M=8 → 256 blocks (84%, 1-wave) vs 128 blocks (42%)
# s6 (M=256, K=1536): BLOCK_M=8 → 768 blocks (2.5 waves), keep 16
BLOCK_M = 8 if M <= 64 else 16
BLOCK_N = 64
NUM_ITER = 1
if K <= 1024:
BLOCK_N = min(256, triton.next_power_of_2(K))
BLOCK_N = max(32, BLOCK_N)
BLOCK_M = min(8, triton.next_power_of_2(M))
NUM_ITER = 1
NUM_STAGES = 1
NUM_WARPS = 4
EVEN_M_N = (M % BLOCK_M == 0) and (K % (BLOCK_N * NUM_ITER) == 0)
grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, BLOCK_N * NUM_ITER))
return {
"BLOCK_M": BLOCK_M,
"BLOCK_N": BLOCK_N,
"NUM_ITER": NUM_ITER,
"NUM_STAGES": NUM_STAGES,
"NUM_WARPS": NUM_WARPS,
"EVEN_M_N": EVEN_M_N,
"grid": grid,
}
def _select_gemm_kernel(M: int, N: int, K: int) -> str:
"""Select optimal ASM GEMM kernel name for register efficiency.
The 32x128 tile gives better occupancy than 192x128 when M is not a multiple
of 192. All benchmark shapes (M=64, M=256) benefit from 32x128 tile:
- M=64: 32x128 → 2x56=112 blocks vs 1x56=56 blocks (2x better occupancy)
- M=256: 32x128 → 8x24=192 blocks vs 2x24=48 blocks (4x better occupancy)
Use 32x128 for all shapes in quant+ASM path since register utilization is better.
"""
return "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
def _init_state(M: int, N: int, K: int, device) -> dict:
SM = (M + 255) // 256 * 256
SN = (K // 32 + 7) // 8 * 8
padded_M = (M + 31) // 32 * 32
A_q_buf = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
A_scale_sh_buf = torch.zeros((SM, SN), dtype=torch.uint8, device=device)
out_buf = torch.empty((padded_M, N), dtype=torch.bfloat16, device=device)
A_q_fp4 = A_q_buf.view(dtypes.fp4x2)
A_scale_fp8 = A_scale_sh_buf.view(dtypes.fp8_e8m0)
cfg = _quant_config(M, K)
gemm_kernel_name = _select_gemm_kernel(M, N, K)
return {
"A_q_buf": A_q_buf,
"A_scale_sh_buf": A_scale_sh_buf,
"out_buf": out_buf,
"A_q_fp4": A_q_fp4,
"A_scale_fp8": A_scale_fp8,
"SM": SM,
"SN": SN,
"padded_M": padded_M,
"cfg": cfg,
"gemm_kernel_name": gemm_kernel_name,
"quant_runner": None,
"quant_ra": None,
}
def _build_quant_args(cfg, SN, A, A_q_buf, A_scale_sh_buf, M, K) -> list:
return [
A,
A_q_buf,
A_scale_sh_buf,
A.stride(0),
A.stride(1),
A_q_buf.stride(0),
A_q_buf.stride(1),
M,
K,
SN,
cfg["BLOCK_M"],
cfg["BLOCK_N"],
cfg["NUM_ITER"],
cfg["NUM_STAGES"],
32,
cfg["EVEN_M_N"],
]
def _custom_kernel_preshuffle(
A: torch.Tensor,
B_shuffle: torch.Tensor,
B_scale_sh: torch.Tensor,
M: int,
N: int,
K_bf16: int,
) -> torch.Tensor:
shape_key = (M, N, K_bf16)
K_packed = B_shuffle.shape[1]
SN_B = B_scale_sh.shape[1]
C = _BUF.get(shape_key)
if C is None:
C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
_BUF[shape_key] = C
p = _PARAMS.get(shape_key)
if p is None:
p = _build_params(M, N, K_bf16)
p["M_val"] = M
p["N_val"] = N
_PARAMS[shape_key] = p
y_pp = _PARTIAL.get(shape_key)
if p["num_ksplit"] > 1 and y_pp is None:
y_pp = torch.empty(
(p["max_ksplit"], M, N), dtype=torch.float32, device=A.device
)
_PARTIAL[shape_key] = y_pp
# Always recompute B_k/B_s views from current B tensors (B changes with each seed)
B_k = torch.as_strided(
B_shuffle.view(torch.uint8),
(N // 16, K_packed * 16),
(K_packed * 16, 1),
)
B_s = torch.as_strided(
B_scale_sh.view(torch.uint8),
(N // 32, K_bf16),
(SN_B * 32, 1),
)
runner = _RUNNER.get(shape_key)
reduce_runner = _REDUCE_RUNNER.get(shape_key)
C_out = C if p["num_ksplit"] == 1 else y_pp
if runner is None:
main_kc, _ = _get_kernel_cache(_gemm_a16wfp4_preshuffle_kernel)
main_keys_before = set(main_kc.keys()) if main_kc is not None else None
_gemm_a16wfp4_preshuffle_kernel[p["grid"]](
A,
B_k,
C_out,
B_s,
M,
N,
p["K_int"],
A.stride(0),
A.stride(1),
B_k.stride(0),
B_k.stride(1),
0 if p["num_ksplit"] == 1 else y_pp.stride(0),
C_out.stride(0) if p["num_ksplit"] == 1 else y_pp.stride(1),
C_out.stride(1) if p["num_ksplit"] == 1 else y_pp.stride(2),
B_s.stride(0),
B_s.stride(1),
PREQUANT=True,
BLOCK_SIZE_M=p["bm"],
BLOCK_SIZE_N=p["bn"],
BLOCK_SIZE_K=p["bk"],
GROUP_SIZE_M=1,
NUM_KSPLIT=p["num_ksplit"],
SPLITK_BLOCK_SIZE=p["spk_sz"],
num_warps=p["num_warps"],
num_stages=p.get("num_stages", 1),
waves_per_eu=p["waves_per_eu"],
matrix_instr_nonkdim=16,
cache_modifier=p["cache_modifier"],
)
if main_keys_before is not None:
r = _extract_runner(main_kc, main_keys_before, p["grid"])
if r is not None:
_RUNNER[shape_key] = r
_RUNNER_ARGS[shape_key] = _build_main_args(p, A, B_k, C_out, B_s)
ra = _RUNNER_ARGS.get(shape_key)
if ra is not None:
ra[1] = B_k
ra[3] = B_s
if p["num_ksplit"] > 1:
red_kc, _ = _get_kernel_cache(_gemm_afp4wfp4_reduce_kernel)
red_keys_before = set(red_kc.keys()) if red_kc is not None else None
_gemm_afp4wfp4_reduce_kernel[p["reduce_grid"]](
y_pp,
C,
M,
N,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
C.stride(0),
C.stride(1),
16,
16,
p["actual_ksplit"],
p["max_ksplit"],
)
if red_keys_before is not None:
rr = _extract_runner(red_kc, red_keys_before, p["reduce_grid"])
if rr is not None:
_REDUCE_RUNNER[shape_key] = rr
_REDUCE_RUNNER_ARGS[shape_key] = _build_reduce_args(p, y_pp, C)
return C
ra = _RUNNER_ARGS[shape_key]
ra[0] = A
ra[1] = B_k # B changes each seed - always refresh
ra[3] = B_s # B changes each seed - always refresh
runner(*ra)
if p["num_ksplit"] > 1:
if reduce_runner is None:
_gemm_afp4wfp4_reduce_kernel[p["reduce_grid"]](
y_pp,
C,
M,
N,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
C.stride(0),
C.stride(1),
16,
16,
p["actual_ksplit"],
p["max_ksplit"],
)
red_kc, _ = _get_kernel_cache(_gemm_afp4wfp4_reduce_kernel)
if red_kc is not None:
rr = _extract_runner(red_kc, set(), p["reduce_grid"])
if rr is not None:
_REDUCE_RUNNER[shape_key] = rr
_REDUCE_RUNNER_ARGS[shape_key] = _build_reduce_args(p, y_pp, C)
else:
reduce_runner(*_REDUCE_RUNNER_ARGS[shape_key])
return C
def _custom_kernel_quant_asm(
A: torch.Tensor,
B_shuffle: torch.Tensor,
B_scale_sh: torch.Tensor,
M: int,
N: int,
K_bf16: int,
) -> torch.Tensor:
shape_key = (M, N, K_bf16)
state = _STATES.get(shape_key)
if state is None:
state = _init_state(M, N, K_bf16, A.device)
_STATES[shape_key] = state
A_q_buf = state["A_q_buf"]
A_scale_sh_buf = state["A_scale_sh_buf"]
out_buf = state["out_buf"]
A_q_fp4 = state["A_q_fp4"]
A_scale_fp8 = state["A_scale_fp8"]
cfg = state["cfg"]
quant_runner = state["quant_runner"]
if quant_runner is None:
kc, _ = _get_kernel_cache(_mxfp4_quant_shuffled_kernel)
kb = set(kc.keys()) if kc is not None else None
_mxfp4_quant_shuffled_kernel[cfg["grid"]](
A,
A_q_buf,
A_scale_sh_buf,
A.stride(0),
A.stride(1),
A_q_buf.stride(0),
A_q_buf.stride(1),
M,
K_bf16,
state["SN"],
BLOCK_SIZE_M=cfg["BLOCK_M"],
BLOCK_SIZE_N=cfg["BLOCK_N"],
NUM_ITER=cfg["NUM_ITER"],
NUM_STAGES=cfg["NUM_STAGES"],
MXFP4_QUANT_BLOCK_SIZE=32,
EVEN_M_N=cfg["EVEN_M_N"],
num_warps=cfg["NUM_WARPS"],
waves_per_eu=2,
num_stages=1,
)
if kb is not None:
r = _extract_runner(kc, kb, cfg["grid"])
if r is not None:
state["quant_runner"] = r
state["quant_ra"] = _build_quant_args(
cfg,
state["SN"],
A,
A_q_buf,
A_scale_sh_buf,
M,
K_bf16,
)
else:
state["quant_ra"][0] = A
quant_runner(*state["quant_ra"])
gemm_a4w4_asm(
A_q_fp4,
B_shuffle,
A_scale_fp8,
B_scale_sh,
out_buf,
state["gemm_kernel_name"],
None,
1.0,
0.0,
True,
0,
)
return out_buf[:M]
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
M, K_bf16 = A.shape
N = B_shuffle.shape[0]
shape_key = (M, N, K_bf16)
dispatch = _DISPATCH_CACHE.get(shape_key)
if dispatch is not None:
return dispatch(A, B_shuffle, B_scale_sh)
fn = (
_custom_kernel_preshuffle
if _use_preshuffle(M, N, K_bf16)
else _custom_kernel_quant_asm
)
result = fn(A, B_shuffle, B_scale_sh, M, N, K_bf16)
def _bound(_A, _b, _bs, _fn=fn, _M=M, _N=N, _K=K_bf16):
return _fn(_A, _b, _bs, _M, _N, _K)
_DISPATCH_CACHE[shape_key] = _bound
return result
scrolls · 736 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