submission 539978
johnny.t.shi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 396 lines, June 9 Researcher Reciprocity License v1.0.
submission_v49_fused_all_small.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-539978?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:e758ae7a2b9e310d6d43b19678a8368aef0798a7ff73e9225d3825c81fda6cf7
license declaredunknown
license concludedunknown
authorsjohnny.t.shi
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 1
num_warps=1, waves_per_eu=0, num_stages=1,split-k
_get_splitk_fn = _gemm_mod.get_splitkstages = 1
NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,tile-k = 256
BLOCK_SIZE_K = 256tile-n = 32
SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,Kernel source
submission_v49_fused_all_small.py396 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v49: Fused quant+GEMM for ALL M<=16 shapes (not just K>=2048).
Saves 1 kernel launch for M=4/K=512 shapes.
ASM GEMM for M>=32 (where ASM is faster).
XCD remap bug fix: pad grid to multiple of 8.
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import _mxfp4_quant_op
from aiter.ops.gemm_op_a4w4 import get_GEMM_config
from aiter.ops.gemm_op_common import get_padded_m
import aiter.ops.triton.gemm_afp4wfp4 as _gemm_mod
_reduce_kernel = _gemm_mod._gemm_afp4wfp4_reduce_kernel
_get_splitk_fn = _gemm_mod.get_splitk
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_bf16 = dtypes.bf16
@triton.jit
def _remap_xcd(pid, num_pids, NUM_XCDS: tl.constexpr):
chunk_size = tl.cdiv(num_pids, NUM_XCDS)
xcd = pid % NUM_XCDS
pid_in_xcd = pid // NUM_XCDS
return xcd * chunk_size + pid_in_xcd
@triton.jit
def _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr):
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + (pid % num_pid_in_group) % group_size_m
pid_n = (pid % num_pid_in_group) // group_size_m
return pid_m, pid_n
@triton.jit
def _fused_quant_gemm_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K_real,
stride_am, stride_ak,
stride_bk, stride_bn,
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,
QUANT_BLOCK: tl.constexpr,
):
SCALE_GROUP_SIZE: tl.constexpr = 32
K_packed = K_real // 2
GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
total_pids = GRID_MN * NUM_KSPLIT
# Pad to multiple of 8 for correct XCD remapping
total_pids_padded = ((total_pids + 7) // 8) * 8
pid_unified = tl.program_id(axis=0)
pid_unified = _remap_xcd(pid_unified, total_pids_padded, NUM_XCDS=8)
if pid_unified < total_pids:
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)
if (pid_k * SPLITK_BLOCK_SIZE) < K_real:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE, BLOCK_SIZE_K)
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_k = pid_k * SPLITK_BLOCK_SIZE + tl.arange(0, BLOCK_SIZE_K)
a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak
offs_k_packed = pid_k * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, BLOCK_SIZE_K // 2)
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
b_ptrs = b_ptr + offs_k_packed[:, None] * stride_bk + offs_bn[None, :] * stride_bn
offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
offs_ks_scale = (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_scale[None, :] * stride_bsk
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in tl.range(0, num_k_iter):
a_bf16 = tl.load(a_ptrs, mask=offs_k[None, :] < K_real, other=0.0).to(tl.float32)
a_fp4, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, QUANT_BLOCK)
b_fp4 = tl.load(b_ptrs, mask=offs_k_packed[:, None] < K_packed, other=0)
b_scales = (
tl.load(b_scale_ptrs)
.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)
)
accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b_fp4, b_scales, "e2m1", accumulator)
a_ptrs += BLOCK_SIZE_K * stride_ak
offs_k += BLOCK_SIZE_K
b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk
offs_k_packed += BLOCK_SIZE_K // 2
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)
@triton.jit
def _fused_quant_shuffle_kernel(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m, stride_x_n,
stride_x_fp4_m, stride_x_fp4_n,
M, N, scale_n_valid,
SCALE_N: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), 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
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).to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
m_idx = bs_offs_m[:, None]
n_idx = bs_offs_n[None, :]
i0 = m_idx // 32
i1 = (m_idx // 16) % 2
i2 = m_idx % 16
i3 = n_idx // 8
i4 = (n_idx // 4) % 2
i5 = n_idx % 4
shuffled_offset = (i0 * (SCALE_N * 32) + i3 * 256 + i5 * 64 + i2 * 4 + i4 * 2 + i1)
bs_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < scale_n_valid)[None, :]
bs_e8m0 = tl.where(bs_valid, bs_e8m0, 127)
bs_store_mask = (m_idx < (M + 255) // 256 * 256) & (n_idx < SCALE_N)
tl.store(bs_ptr + shuffled_offset, bs_e8m0, mask=bs_store_mask)
_cache_asm = {}
_cache_fused = {}
_gemm_asm = None
_warmup_done = False
def custom_kernel(data: input_t) -> output_t:
global _gemm_asm, _warmup_done
A, B, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape
N = B_shuffle.shape[0]
# Fused for ALL M<=16 shapes (saves 1 kernel launch)
use_fused = (M <= 16)
# Warmup: use ASM path to init aiter module
if not _warmup_done:
scale_n_valid = (K + 31) // 32
SCALE_M = ((M + 255) // 256) * 256
SCALE_N = ((scale_n_valid + 7) // 8) * 8
BSM = triton.next_power_of_2(M) if M <= 32 else 16
grid = (triton.cdiv(M, BSM), triton.cdiv(K, 32))
x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)
_fused_quant_shuffle_kernel[grid](
A, x_fp4, bs_sh,
A.stride(0), A.stride(1),
x_fp4.stride(0), x_fp4.stride(1),
M, K, scale_n_valid,
SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,
NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=1, waves_per_eu=0, num_stages=1,
)
result = aiter.gemm_a4w4(
x_fp4.view(_fp4x2), B_shuffle,
bs_sh.view(_fp8_e8m0), B_scale_sh,
dtype=_bf16, bpreshuffle=True,
)
_warmup_done = True
try:
_gemm_asm = torch.ops.aiter.gemm_a4w4_asm
except Exception:
try:
import aiter.jit.core as _jc
_gemm_asm = getattr(_jc, 'gemm_a4w4_asm', None)
except Exception:
pass
return result
if use_fused:
# --- Fused quant+GEMM: single kernel launch ---
key = (M, K, N)
c = _cache_fused.get(key)
if c is None:
K_packed = K // 2
scale_n = (K + 31) // 32
SCALE_N_B = ((scale_n + 7) // 8) * 8
BLOCK_SIZE_M = max(16, triton.next_power_of_2(M))
BLOCK_SIZE_N = 128
BLOCK_SIZE_K = 256
base_blocks = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
target_ksplit = max(1, 128 // max(1, base_blocks))
if target_ksplit > 1:
SPLITK_BLOCK_SIZE, BLOCK_SIZE_K_adj, NUM_KSPLIT = _get_splitk_fn(
K_packed, BLOCK_SIZE_K, target_ksplit
)
if BLOCK_SIZE_K_adj < 256:
BLOCK_SIZE_K_adj = 256
SPLITK_BLOCK_SIZE = 2 * K_packed
NUM_KSPLIT = 1
else:
BLOCK_SIZE_K = BLOCK_SIZE_K_adj
else:
NUM_KSPLIT = 1
SPLITK_BLOCK_SIZE = 2 * K_packed
if NUM_KSPLIT > 1:
y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
else:
y_pp = None
SPLITK_BLOCK_SIZE = 2 * K_packed
y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
total_blocks_raw = NUM_KSPLIT * triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
# Pad to multiple of 8 for correct XCD remapping
total_blocks = ((total_blocks_raw + 7) // 8) * 8
bs_stride_n = 32 * SCALE_N_B
bs_stride_k = 1
c = (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
NUM_KSPLIT, SPLITK_BLOCK_SIZE,
y, y_pp, total_blocks, bs_stride_n, bs_stride_k)
_cache_fused[key] = c
(K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
NUM_KSPLIT, SPLITK_BLOCK_SIZE,
y, y_pp, total_blocks, bs_stride_n, bs_stride_k) = c
B_q_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
B_q_T = B_q_u8.T
B_scale_u8 = B_scale_sh.view(torch.uint8)
out_tensor = y if NUM_KSPLIT == 1 else y_pp
_fused_quant_gemm_kernel[(total_blocks,)](
A, B_q_T, out_tensor, B_scale_u8,
M, N, K,
A.stride(0), A.stride(1),
B_q_T.stride(0), B_q_T.stride(1),
0 if NUM_KSPLIT == 1 else y_pp.stride(0),
y.stride(0) if NUM_KSPLIT == 1 else y_pp.stride(1),
y.stride(1) if NUM_KSPLIT == 1 else y_pp.stride(2),
bs_stride_n, bs_stride_k,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
BLOCK_SIZE_K=BLOCK_SIZE_K,
GROUP_SIZE_M=8,
NUM_KSPLIT=NUM_KSPLIT,
SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
QUANT_BLOCK=32,
num_warps=4,
num_stages=2,
waves_per_eu=0,
)
if NUM_KSPLIT > 1:
ACTUAL_KSPLIT = triton.cdiv(K_packed, (SPLITK_BLOCK_SIZE // 2))
grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 64))
_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),
16, 64, ACTUAL_KSPLIT,
triton.next_power_of_2(NUM_KSPLIT),
)
return y
else:
# --- ASM GEMM path ---
key = (M, K, N)
c = _cache_asm.get(key)
if c is None:
scale_n_valid = (K + 31) // 32
SCALE_M = ((M + 255) // 256) * 256
SCALE_N = ((scale_n_valid + 7) // 8) * 8
padded_m = get_padded_m(M, N, K, 0)
BSM = triton.next_power_of_2(M) if M <= 32 else 16
NW = 1
BSN = 32
grid = (triton.cdiv(M, BSM), triton.cdiv(K, BSN))
ck_config = get_GEMM_config(M, N, K)
kernel_name = ""
split_k = 0
if ck_config is not None:
split_k = ck_config.get("splitK", 0) or 0
kernel_name = ck_config["kernelName"]
x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)
out = torch.empty((padded_m, N), dtype=torch.bfloat16, device=A.device)
x_fp4_view = x_fp4.view(_fp4x2)
bs_sh_view = bs_sh.view(_fp8_e8m0)
out_view = out[:M] if M < padded_m else out
c = (scale_n_valid, SCALE_N, BSM, BSN, grid,
x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
kernel_name, split_k,
A.stride(0), A.stride(1), x_fp4.stride(0), x_fp4.stride(1))
_cache_asm[key] = c
(scale_n_valid, SCALE_N, BSM, BSN, grid,
x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
kernel_name, split_k,
stride_a0, stride_a1, stride_fp4_0, stride_fp4_1) = c
_fused_quant_shuffle_kernel[grid](
A, x_fp4, bs_sh,
stride_a0, stride_a1,
stride_fp4_0, stride_fp4_1,
M, K, scale_n_valid,
SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=1, waves_per_eu=0, num_stages=1,
)
if _gemm_asm is not None:
_gemm_asm(x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
out, kernel_name, None, 1.0, 0.0, True, split_k)
return out_view
return aiter.gemm_a4w4(
x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
dtype=_bf16, bpreshuffle=True,
)
scrolls · 396 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 539788.
⋯ 1 unchanged lines#!POPCORN gpu MI355X"""- v46: Preshuffle-scales GEMM — reads B_scale_sh directly (zero unshuffle overhead).- For M<=16, K>=2048: call _gemm_afp4wfp4_kernel_preshuffle_scales directly,- bypassing the M>=32 wrapper assertion. The kernel handles M<32 with raw A scales- and does in-register B scale unshuffle via reshape+permute (free, hidden by mem latency).- This eliminates the ~5µs unshuffle that killed v45 ranked performance.- For all others: ASM GEMM (v35 approach).+ v49: Fused quant+GEMM for ALL M<=16 shapes (not just K>=2048).+ Saves 1 kernel launch for M=4/K=512 shapes.+ ASM GEMM for M>=32 (where ASM is faster).+ XCD remap bug fix: pad grid to multiple of 8."""from task import input_t, output_t⋯ 6 unchanged linesfrom aiter.ops.gemm_op_a4w4 import get_GEMM_configfrom aiter.ops.gemm_op_common import get_padded_m- # Import preshuffle-scales kernel directly (bypasses M>=32 wrapper assert)- try:- import aiter.ops.triton.gemm_afp4wfp4 as _gemm_mod- _ps_kernel = _gemm_mod._gemm_afp4wfp4_kernel_preshuffle_scales- _reduce_kernel = _gemm_mod._gemm_afp4wfp4_reduce_kernel- _get_splitk_fn = _gemm_mod.get_splitk- _HAS_PS = True- except (ImportError, AttributeError):- _HAS_PS = False+ import aiter.ops.triton.gemm_afp4wfp4 as _gemm_mod+ _reduce_kernel = _gemm_mod._gemm_afp4wfp4_reduce_kernel+ _get_splitk_fn = _gemm_mod.get_splitk_fp4x2 = dtypes.fp4x2_fp8_e8m0 = dtypes.fp8_e8m0⋯ 1 unchanged lines@triton.jit- def _quant_raw_kernel(- x_ptr, x_fp4_ptr, scale_ptr,- stride_x_m, stride_x_n,- stride_fp4_m, stride_fp4_n,- stride_sc_m, stride_sc_n,- M, K,- BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr,+ def _remap_xcd(pid, num_pids, NUM_XCDS: tl.constexpr):+ chunk_size = tl.cdiv(num_pids, NUM_XCDS)+ xcd = pid % NUM_XCDS+ pid_in_xcd = pid // NUM_XCDS+ return xcd * chunk_size + pid_in_xcd+++ @triton.jit+ def _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr):+ num_pid_in_group = GROUP_SIZE_M * num_pid_n+ group_id = pid // num_pid_in_group+ first_pid_m = group_id * GROUP_SIZE_M+ group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)+ pid_m = first_pid_m + (pid % num_pid_in_group) % group_size_m+ pid_n = (pid % num_pid_in_group) // group_size_m+ return pid_m, pid_n+++ @triton.jit+ def _fused_quant_gemm_kernel(+ a_ptr, b_ptr, c_ptr, b_scales_ptr,+ M, N, K_real,+ stride_am, stride_ak,+ stride_bk, stride_bn,+ 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,QUANT_BLOCK: tl.constexpr,):- pid_m = tl.program_id(0)- pid_k = tl.program_id(1)- offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)- offs_k = pid_k * BLOCK_K + tl.arange(0, BLOCK_K)- mask = (offs_m[:, None] < M) & (offs_k[None, :] < K)- x = tl.load(x_ptr + offs_m[:, None] * stride_x_m + offs_k[None, :] * stride_x_n,- mask=mask, other=0.0).to(tl.float32)+ SCALE_GROUP_SIZE: tl.constexpr = 32+ K_packed = K_real // 2+ GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)+ total_pids = GRID_MN * NUM_KSPLIT+ # Pad to multiple of 8 for correct XCD remapping+ total_pids_padded = ((total_pids + 7) // 8) * 8- out_fp4, scales_e8m0 = _mxfp4_quant_op(x, BLOCK_K, BLOCK_M, QUANT_BLOCK)+ pid_unified = tl.program_id(axis=0)+ pid_unified = _remap_xcd(pid_unified, total_pids_padded, NUM_XCDS=8)- fp4_offs_k = pid_k * BLOCK_K // 2 + tl.arange(0, BLOCK_K // 2)- fp4_mask = (offs_m[:, None] < M) & (fp4_offs_k[None, :] < K // 2)- tl.store(x_fp4_ptr + offs_m[:, None] * stride_fp4_m + fp4_offs_k[None, :] * stride_fp4_n,- out_fp4, mask=fp4_mask)+ if pid_unified < total_pids:+ 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)- NUM_SC: tl.constexpr = BLOCK_K // QUANT_BLOCK- sc_offs_k = pid_k * NUM_SC + tl.arange(0, NUM_SC)- sc_mask = (offs_m[:, None] < M) & (sc_offs_k[None, :] < (K + QUANT_BLOCK - 1) // QUANT_BLOCK)- tl.store(scale_ptr + offs_m[:, None] * stride_sc_m + sc_offs_k[None, :] * stride_sc_n,- scales_e8m0, mask=sc_mask)+ 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)+ if (pid_k * SPLITK_BLOCK_SIZE) < K_real:+ num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE, BLOCK_SIZE_K)++ offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M+ offs_k = pid_k * SPLITK_BLOCK_SIZE + tl.arange(0, BLOCK_SIZE_K)+ a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak++ offs_k_packed = pid_k * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, BLOCK_SIZE_K // 2)+ offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N+ b_ptrs = b_ptr + offs_k_packed[:, None] * stride_bk + offs_bn[None, :] * stride_bn++ offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N+ offs_ks_scale = (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_scale[None, :] * stride_bsk++ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)++ for k in tl.range(0, num_k_iter):+ a_bf16 = tl.load(a_ptrs, mask=offs_k[None, :] < K_real, other=0.0).to(tl.float32)+ a_fp4, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, QUANT_BLOCK)++ b_fp4 = tl.load(b_ptrs, mask=offs_k_packed[:, None] < K_packed, other=0)++ b_scales = (+ tl.load(b_scale_ptrs)+ .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)+ )++ accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b_fp4, b_scales, "e2m1", accumulator)++ a_ptrs += BLOCK_SIZE_K * stride_ak+ offs_k += BLOCK_SIZE_K+ b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk+ offs_k_packed += BLOCK_SIZE_K // 2+ 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)++@triton.jitdef _fused_quant_shuffle_kernel(x_ptr, x_fp4_ptr, bs_ptr,⋯ 46 unchanged lines_cache_asm = {}- _cache_triton = {}+ _cache_fused = {}_gemm_asm = None_warmup_done = False⋯ 5 unchanged linesM, K = A.shapeN = B_shuffle.shape[0]- use_triton = _HAS_PS and (M <= 16) and (K >= 2048)+ # Fused for ALL M<=16 shapes (saves 1 kernel launch)+ use_fused = (M <= 16)- # Warmup: use ASM path to initialize module+ # Warmup: use ASM path to init aiter moduleif not _warmup_done:scale_n_valid = (K + 31) // 32SCALE_M = ((M + 255) // 256) * 256⋯ 30 unchanged linespassreturn result- if use_triton:- # --- Preshuffle-scales GEMM: reads B_scale_sh directly ---+ if use_fused:+ # --- Fused quant+GEMM: single kernel launch ---key = (M, K, N)- c = _cache_triton.get(key)+ c = _cache_fused.get(key)if c is None:K_packed = K // 2scale_n = (K + 31) // 32SCALE_N_B = ((scale_n + 7) // 8) * 8- # Quant config- BSM_q = triton.next_power_of_2(M)- BSK_q = 32- grid_q = (triton.cdiv(M, BSM_q), triton.cdiv(K, BSK_q))-- x_fp4 = torch.empty((M, K_packed), dtype=torch.uint8, device=A.device)- x_scales = torch.empty((M, scale_n), dtype=torch.uint8, device=A.device)-- # GEMM configBLOCK_SIZE_M = max(16, triton.next_power_of_2(M))BLOCK_SIZE_N = 128- BLOCK_SIZE_K = 256 # min for preshuffle_scales reshape+ BLOCK_SIZE_K = 256- # SplitK for occupancybase_blocks = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)target_ksplit = max(1, 128 // max(1, base_blocks))- SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = _get_splitk_fn(- K_packed, BLOCK_SIZE_K, target_ksplit- )-- # Ensure BLOCK_SIZE_K >= 256 for reshape- if BLOCK_SIZE_K < 256:- BLOCK_SIZE_K = 256- SPLITK_BLOCK_SIZE = 2 * K_packed+ if target_ksplit > 1:+ SPLITK_BLOCK_SIZE, BLOCK_SIZE_K_adj, NUM_KSPLIT = _get_splitk_fn(+ K_packed, BLOCK_SIZE_K, target_ksplit+ )+ if BLOCK_SIZE_K_adj < 256:+ BLOCK_SIZE_K_adj = 256+ SPLITK_BLOCK_SIZE = 2 * K_packed+ NUM_KSPLIT = 1+ else:+ BLOCK_SIZE_K = BLOCK_SIZE_K_adj+ else:NUM_KSPLIT = 1+ SPLITK_BLOCK_SIZE = 2 * K_packedif NUM_KSPLIT > 1:y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)⋯ 3 unchanged linesy = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)- config = {- "BLOCK_SIZE_M": BLOCK_SIZE_M,- "BLOCK_SIZE_N": BLOCK_SIZE_N,- "BLOCK_SIZE_K": BLOCK_SIZE_K,- "GROUP_SIZE_M": 8,- "NUM_KSPLIT": NUM_KSPLIT,- "SPLITK_BLOCK_SIZE": SPLITK_BLOCK_SIZE,- "num_warps": 4,- "num_stages": 2,- "waves_per_eu": 0,- "matrix_instr_nonkdim": 32,- "cache_modifier": ".ca",- }+ total_blocks_raw = NUM_KSPLIT * triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)+ # Pad to multiple of 8 for correct XCD remapping+ total_blocks = ((total_blocks_raw + 7) // 8) * 8- total_blocks = NUM_KSPLIT * triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)-- # B_scale strides for shuffled layout (32 rows per N-group)bs_stride_n = 32 * SCALE_N_Bbs_stride_k = 1- c = (K_packed, scale_n, SCALE_N_B, BSM_q, BSK_q, grid_q,- x_fp4, x_scales, y, y_pp, config, total_blocks,- bs_stride_n, bs_stride_k,- A.stride(0), A.stride(1),- x_fp4.stride(0), x_fp4.stride(1),- x_scales.stride(0), x_scales.stride(1))- _cache_triton[key] = c+ c = (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,+ NUM_KSPLIT, SPLITK_BLOCK_SIZE,+ y, y_pp, total_blocks, bs_stride_n, bs_stride_k)+ _cache_fused[key] = c- (K_packed, scale_n, SCALE_N_B, BSM_q, BSK_q, grid_q,- x_fp4, x_scales, y, y_pp, config, total_blocks,- bs_stride_n, bs_stride_k,- sa0, sa1, sf0, sf1, ss0, ss1) = c+ (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,+ NUM_KSPLIT, SPLITK_BLOCK_SIZE,+ y, y_pp, total_blocks, bs_stride_n, bs_stride_k) = c- # 1. Raw quant (produces x_fp4 + raw x_scales)- _quant_raw_kernel[grid_q](- A, x_fp4, x_scales,- sa0, sa1, sf0, sf1, ss0, ss1,- M, K,- BLOCK_M=BSM_q, BLOCK_K=BSK_q,- QUANT_BLOCK=32,- num_warps=1, waves_per_eu=0, num_stages=1,- )-- # 2. Preshuffle-scales GEMM — reads B_scale_sh directly, no unshuffle neededB_q_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q- B_q_T = B_q_u8.T # (K//2, N) non-contiguous view- # View B_scale_sh as uint8 to avoid Triton float8_e8m0fnu type error+ B_q_T = B_q_u8.TB_scale_u8 = B_scale_sh.view(torch.uint8)- out_tensor = y if config["NUM_KSPLIT"] == 1 else y_pp+ out_tensor = y if NUM_KSPLIT == 1 else y_pp- _ps_kernel[(total_blocks,)](- x_fp4,- B_q_T,- out_tensor,- x_scales,- B_scale_u8,- M, N, K_packed,- x_fp4.stride(0), x_fp4.stride(1),+ _fused_quant_gemm_kernel[(total_blocks,)](+ A, B_q_T, out_tensor, B_scale_u8,+ M, N, K,+ A.stride(0), A.stride(1),B_q_T.stride(0), B_q_T.stride(1),- 0 if config["NUM_KSPLIT"] == 1 else y_pp.stride(0),- y.stride(0) if config["NUM_KSPLIT"] == 1 else y_pp.stride(1),- y.stride(1) if config["NUM_KSPLIT"] == 1 else y_pp.stride(2),- x_scales.stride(0), x_scales.stride(1),+ 0 if NUM_KSPLIT == 1 else y_pp.stride(0),+ y.stride(0) if NUM_KSPLIT == 1 else y_pp.stride(1),+ y.stride(1) if NUM_KSPLIT == 1 else y_pp.stride(2),bs_stride_n, bs_stride_k,- **config,+ BLOCK_SIZE_M=BLOCK_SIZE_M,+ BLOCK_SIZE_N=BLOCK_SIZE_N,+ BLOCK_SIZE_K=BLOCK_SIZE_K,+ GROUP_SIZE_M=8,+ NUM_KSPLIT=NUM_KSPLIT,+ SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,+ QUANT_BLOCK=32,+ num_warps=4,+ num_stages=2,+ waves_per_eu=0,)- # 3. Reduce if SplitK- if config["NUM_KSPLIT"] > 1:- ACTUAL_KSPLIT = triton.cdiv(K_packed, (config["SPLITK_BLOCK_SIZE"] // 2))- grid_reduce = (- triton.cdiv(M, 16),- triton.cdiv(N, 64),- )+ if NUM_KSPLIT > 1:+ ACTUAL_KSPLIT = triton.cdiv(K_packed, (SPLITK_BLOCK_SIZE // 2))+ grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 64))_reduce_kernel[grid_reduce](- y_pp, y,- M, N,+ y_pp, y, M, N,y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),y.stride(0), y.stride(1),- 16, 64,- ACTUAL_KSPLIT,- triton.next_power_of_2(config["NUM_KSPLIT"]),+ 16, 64, ACTUAL_KSPLIT,+ triton.next_power_of_2(NUM_KSPLIT),)return yelse:- # --- ASM GEMM path (v35) ---+ # --- ASM GEMM path ---key = (M, K, N)c = _cache_asm.get(key)if c is None:
scrolls · 386 diff lines total
Best evidence level for this revision: reported
JSON