submission 539664
johnny.t.shi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 336 lines, June 9 Researcher Reciprocity License v1.0.
submission_v45_hybrid_smart_cache.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-539664?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:23bf9a834f3277764129cd7b1b880e15fc63478d1cc62cd4d534564d4dff14f7
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
- M<=16, K>=2048: Triton GEMM with SplitK (fixes 6.6% CU utilization for M=16/K=7168)stages = 1
NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,tile-k = 256
BLOCK_K = 256tile-n = 32
SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,Kernel source
submission_v45_hybrid_smart_cache.py336 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v45: Hybrid GEMM with smart B_scale caching.
- M<=16, K>=2048: Triton GEMM with SplitK (fixes 6.6% CU utilization for M=16/K=7168)
- All others: ASM GEMM (v35 approach)
Smart cache: uses Python `is` identity to detect when B_scale_sh changes,
avoiding both stale cache bugs and per-call unshuffle overhead (~5µs).
"""
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
from aiter.ops.triton.gemm_afp4wfp4 import gemm_afp4wfp4
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_bf16 = dtypes.bf16
@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,
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)
out_fp4, scales_e8m0 = _mxfp4_quant_op(x, BLOCK_K, BLOCK_M, QUANT_BLOCK)
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)
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)
@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)
def _unshuffle_b_scale(B_scale_sh, n, k):
sn = k // 32
b_u8 = B_scale_sh.contiguous().view(torch.uint8)
total = b_u8.numel()
SN = ((sn + 7) // 8) * 8
padded_n = total // SN
if padded_n < 32 or SN < 8:
return None
try:
raw = b_u8.reshape(padded_n // 32, SN // 8, 4, 16, 2, 2)
raw = raw.permute(0, 5, 3, 1, 4, 2).contiguous().view(padded_n, SN)
return raw[:n, :sn].contiguous()
except Exception:
return None
_cache_asm = {}
_cache_triton = {}
_b_scale_cache = {} # key: (N, K), value: (B_scale_sh_ref, B_scale_raw)
_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]
# Use Triton GEMM only for small M with large K (where ASM has terrible occupancy)
use_triton = (M <= 16) and (K >= 2048)
# Warmup: always use ASM path to initialize the 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_triton:
# --- Triton GEMM path with SplitK ---
key = (M, K, N)
c = _cache_triton.get(key)
if c is None:
scale_n = (K + 31) // 32
BSM_q = triton.next_power_of_2(M)
BSK_q = 32
NW_q = 1
grid_q = (triton.cdiv(M, BSM_q), triton.cdiv(K, BSK_q))
x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
x_scales = torch.empty((M, scale_n), dtype=torch.uint8, device=A.device)
out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
# Compute SplitK for better CU occupancy
BLOCK_M = max(16, triton.next_power_of_2(M))
BLOCK_N = 128
BLOCK_K = 256
base_blocks = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N)
k_iters = max(1, K // BLOCK_K)
target_ksplit = max(1, 128 // max(1, base_blocks))
target_ksplit = min(target_ksplit, k_iters)
NUM_KSPLIT = 1
if target_ksplit > 1:
for ks in range(target_ksplit, k_iters + 1):
if k_iters % ks == 0:
NUM_KSPLIT = ks
break
if NUM_KSPLIT == 1:
NUM_KSPLIT = target_ksplit
config = {
"BLOCK_SIZE_M": BLOCK_M,
"BLOCK_SIZE_N": BLOCK_N,
"BLOCK_SIZE_K": BLOCK_K,
"GROUP_SIZE_M": 8,
"NUM_KSPLIT": NUM_KSPLIT,
"SPLITK_BLOCK_SIZE": K,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 0,
"matrix_instr_nonkdim": 32,
"cache_modifier": ".ca",
}
c = (scale_n, BSM_q, BSK_q, NW_q, grid_q,
x_fp4, x_scales, out, config,
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
(scale_n, BSM_q, BSK_q, NW_q, grid_q,
x_fp4, x_scales, out, config,
sa0, sa1, sf0, sf1, ss0, ss1) = c
# 1. Raw quant
_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=NW_q, waves_per_eu=0, num_stages=1,
)
# 2. Smart B_scale cache: use Python `is` identity to detect changes
bkey = (N, K)
cached = _b_scale_cache.get(bkey)
if cached is not None:
old_ref, B_scale_raw = cached
if old_ref is not B_scale_sh:
# Different tensor object → recompute
B_scale_raw = _unshuffle_b_scale(B_scale_sh, N, K)
_b_scale_cache[bkey] = (B_scale_sh, B_scale_raw)
else:
B_scale_raw = _unshuffle_b_scale(B_scale_sh, N, K)
_b_scale_cache[bkey] = (B_scale_sh, B_scale_raw)
if B_scale_raw is None:
return aiter.gemm_a4w4(
x_fp4.view(_fp4x2), B_shuffle,
torch.empty(0, dtype=torch.uint8, device=A.device).view(_fp8_e8m0),
B_scale_sh, dtype=_bf16, bpreshuffle=True,
)
# 3. Triton GEMM with SplitK
B_q_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
return gemm_afp4wfp4(
x_fp4, B_q_u8,
x_scales, B_scale_raw,
dtype=_bf16, y=out, config=config,
)
else:
# --- ASM GEMM path (v35) ---
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 · 336 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 528902.
⋯ 1 unchanged lines#!POPCORN gpu MI355X"""- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).+ v45: Hybrid GEMM with smart B_scale caching.+ - M<=16, K>=2048: Triton GEMM with SplitK (fixes 6.6% CU utilization for M=16/K=7168)+ - All others: ASM GEMM (v35 approach)+ Smart cache: uses Python `is` identity to detect when B_scale_sh changes,+ avoiding both stale cache bugs and per-call unshuffle overhead (~5µs)."""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+ from aiter.ops.triton.gemm_afp4wfp4 import gemm_afp4wfp4+ _fp4x2 = dtypes.fp4x2+ _fp8_e8m0 = dtypes.fp8_e8m0+ _bf16 = dtypes.bf16+++ @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,+ 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)++ out_fp4, scales_e8m0 = _mxfp4_quant_op(x, BLOCK_K, BLOCK_M, QUANT_BLOCK)++ 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)++ 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)+++ @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)+++ def _unshuffle_b_scale(B_scale_sh, n, k):+ sn = k // 32+ b_u8 = B_scale_sh.contiguous().view(torch.uint8)+ total = b_u8.numel()+ SN = ((sn + 7) // 8) * 8+ padded_n = total // SN+ if padded_n < 32 or SN < 8:+ return None+ try:+ raw = b_u8.reshape(padded_n // 32, SN // 8, 4, 16, 2, 2)+ raw = raw.permute(0, 5, 3, 1, 4, 2).contiguous().view(padded_n, SN)+ return raw[:n, :sn].contiguous()+ except Exception:+ return None+++ _cache_asm = {}+ _cache_triton = {}+ _b_scale_cache = {} # key: (N, K), value: (B_scale_sh_ref, B_scale_raw)+ _gemm_asm = None+ _warmup_done = False++def custom_kernel(data: input_t) -> output_t:- """- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.- gemm_a4w4 with bpreshuffle=True.- """- import aiter- from aiter import QuantType, dtypes+ global _gemm_asm, _warmup_doneA, B, B_q, B_shuffle, B_scale_sh = data- A = A.contiguous()- B = B.contiguous()- m, k = A.shape- n, _ = B.shape+ M, K = A.shape+ N = B_shuffle.shape[0]- quant_func = aiter.get_triton_quant(QuantType.per_1x32)- A_q, A_scale_sh = quant_func(A, shuffle=True)- out_gemm = aiter.gemm_a4w4(- A_q,- B_shuffle,- A_scale_sh,- B_scale_sh,- dtype=dtypes.bf16,- bpreshuffle=True,- )- return out_gemm+ # Use Triton GEMM only for small M with large K (where ASM has terrible occupancy)+ use_triton = (M <= 16) and (K >= 2048)++ # Warmup: always use ASM path to initialize the 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_triton:+ # --- Triton GEMM path with SplitK ---+ key = (M, K, N)+ c = _cache_triton.get(key)+ if c is None:+ scale_n = (K + 31) // 32+ BSM_q = triton.next_power_of_2(M)+ BSK_q = 32+ NW_q = 1+ grid_q = (triton.cdiv(M, BSM_q), triton.cdiv(K, BSK_q))++ x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)+ x_scales = torch.empty((M, scale_n), dtype=torch.uint8, device=A.device)+ out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)++ # Compute SplitK for better CU occupancy+ BLOCK_M = max(16, triton.next_power_of_2(M))+ BLOCK_N = 128+ BLOCK_K = 256+ base_blocks = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N)+ k_iters = max(1, K // BLOCK_K)+ target_ksplit = max(1, 128 // max(1, base_blocks))+ target_ksplit = min(target_ksplit, k_iters)+ NUM_KSPLIT = 1+ if target_ksplit > 1:+ for ks in range(target_ksplit, k_iters + 1):+ if k_iters % ks == 0:+ NUM_KSPLIT = ks+ break+ if NUM_KSPLIT == 1:+ NUM_KSPLIT = target_ksplit++ config = {+ "BLOCK_SIZE_M": BLOCK_M,+ "BLOCK_SIZE_N": BLOCK_N,+ "BLOCK_SIZE_K": BLOCK_K,+ "GROUP_SIZE_M": 8,+ "NUM_KSPLIT": NUM_KSPLIT,+ "SPLITK_BLOCK_SIZE": K,+ "num_warps": 4,+ "num_stages": 2,+ "waves_per_eu": 0,+ "matrix_instr_nonkdim": 32,+ "cache_modifier": ".ca",+ }++ c = (scale_n, BSM_q, BSK_q, NW_q, grid_q,+ x_fp4, x_scales, out, config,+ 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++ (scale_n, BSM_q, BSK_q, NW_q, grid_q,+ x_fp4, x_scales, out, config,+ sa0, sa1, sf0, sf1, ss0, ss1) = c++ # 1. Raw quant+ _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=NW_q, waves_per_eu=0, num_stages=1,+ )++ # 2. Smart B_scale cache: use Python `is` identity to detect changes+ bkey = (N, K)+ cached = _b_scale_cache.get(bkey)+ if cached is not None:+ old_ref, B_scale_raw = cached+ if old_ref is not B_scale_sh:+ # Different tensor object → recompute+ B_scale_raw = _unshuffle_b_scale(B_scale_sh, N, K)+ _b_scale_cache[bkey] = (B_scale_sh, B_scale_raw)+ else:+ B_scale_raw = _unshuffle_b_scale(B_scale_sh, N, K)+ _b_scale_cache[bkey] = (B_scale_sh, B_scale_raw)++ if B_scale_raw is None:+ return aiter.gemm_a4w4(+ x_fp4.view(_fp4x2), B_shuffle,+ torch.empty(0, dtype=torch.uint8, device=A.device).view(_fp8_e8m0),+ B_scale_sh, dtype=_bf16, bpreshuffle=True,+ )++ # 3. Triton GEMM with SplitK+ B_q_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q+ return gemm_afp4wfp4(+ x_fp4, B_q_u8,+ x_scales, B_scale_raw,+ dtype=_bf16, y=out, config=config,+ )++ else:+ # --- ASM GEMM path (v35) ---+ 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 · 358 diff lines total
Best evidence level for this revision: reported
JSON