submission 685211
jiahuizz · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 634 lines, June 9 Researcher Reciprocity License v1.0.
submit_cc_v7.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-685211?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:73ca812d1bc9d9df897741d293a6885bdf7bd7fd17bf5e7a4ea6331aca8695fd
license declaredunknown
license concludedunknown
authorsjiahuizz
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submit_cc_v7.py634 lines
"""
submit_cc_v7.py - Hybrid MXFP4 GEMM with shape-specific large-M K tiles.
Changes vs submit_cc_v4.py:
- Keep the low-overhead fast paths from v4.
- Use a 512-wide bf16 K tile only for the 64x7168x2048 benchmark shape, where
it helps materially.
- Keep the 256x3072x1536 shape on the safer 256-wide K tile.
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
_get_config,
)
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
# ---------------------------------------------------------------------------
# Inline MXFP4 quantization (shared by fused kernel and quant-only kernel)
# ---------------------------------------------------------------------------
@triton.jit
def _mxfp4_quant_op(x, BK: tl.constexpr, BM: tl.constexpr, QBS: tl.constexpr):
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
NQB: tl.constexpr = BK // QBS
x = x.reshape(BM, NQB, QBS)
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)
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: tl.constexpr = (((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1) & 0xFFFFFFFF
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, [BM, NQB, QBS // 2, 2])
evens, odds = tl.split(e2m1_value)
x_fp4 = evens | (odds << 4)
return x_fp4.reshape(BM, BK // 2), bs_e8m0.reshape(BM, NQB)
# ---------------------------------------------------------------------------
# Standalone A quantization kernel (for large M path)
# ---------------------------------------------------------------------------
@triton.heuristics({
"EVEN_M": lambda args: (args["M"] % args["BM_Q"]) == 0,
"EVEN_KQ": lambda args: (args["K_bf16"] % args["BK_Q"]) == 0,
})
@triton.jit
def _quant_a_kernel(
a_ptr, a_fp4_ptr, a_scale_ptr,
M, K_bf16,
stride_am, stride_ak,
stride_afm, stride_afk,
stride_asm, stride_ask,
BM_Q: tl.constexpr, BK_Q: tl.constexpr,
EVEN_M: tl.constexpr, EVEN_KQ: tl.constexpr,
):
SCALE_GROUP_SIZE: tl.constexpr = 32
NQB: tl.constexpr = BK_Q // SCALE_GROUP_SIZE
pid_m = tl.program_id(0)
pid_k = tl.program_id(1)
offs_m = pid_m * BM_Q + tl.arange(0, BM_Q)
offs_k = pid_k * BK_Q + tl.arange(0, BK_Q)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
if EVEN_M and EVEN_KQ:
a = tl.load(a_ptrs)
else:
mask = (offs_m[:, None] < M) & (offs_k[None, :] < K_bf16)
a = tl.load(a_ptrs, mask=mask, other=0.0)
a_fp4, a_scales = _mxfp4_quant_op(a.to(tl.float32), BK_Q, BM_Q, SCALE_GROUP_SIZE)
offs_kp = pid_k * (BK_Q // 2) + tl.arange(0, BK_Q // 2)
a_fp4_ptrs = a_fp4_ptr + offs_m[:, None] * stride_afm + offs_kp[None, :] * stride_afk
if EVEN_M and EVEN_KQ:
tl.store(a_fp4_ptrs, a_fp4)
else:
fp4_mask = (offs_m[:, None] < M) & (offs_kp[None, :] < K_bf16 // 2)
tl.store(a_fp4_ptrs, a_fp4, mask=fp4_mask)
offs_ks = pid_k * NQB + tl.arange(0, NQB)
a_scale_ptrs = a_scale_ptr + offs_m[:, None] * stride_asm + offs_ks[None, :] * stride_ask
if EVEN_M and EVEN_KQ:
tl.store(a_scale_ptrs, a_scales.reshape(BM_Q, NQB))
else:
scale_mask = (offs_m[:, None] < M) & (offs_ks[None, :] < K_bf16 // SCALE_GROUP_SIZE)
tl.store(a_scale_ptrs, a_scales.reshape(BM_Q, NQB), mask=scale_mask)
@triton.jit
def _load_b_scales(
b_scales_ptr,
stride_bsn,
stride_bsk,
n_group,
n2,
n16,
k_iter,
BK_SCALES: tl.constexpr,
K_SCALE_PAD: tl.constexpr,
):
if BK_SCALES == 8:
k_local = tl.arange(0, 8)
k4 = k_local % 4
k2 = k_local // 4
base = n16[:, None] * 4 + k2[None, :] * 2 + n2[:, None]
if K_SCALE_PAD == 64:
bs_row = n_group[:, None] * 32 + k_iter * 4 + k4[None, :]
bs_col = base
return tl.load(b_scales_ptr + bs_row * stride_bsn + bs_col * stride_bsk)
if K_SCALE_PAD == 48:
k_mod3 = k_iter % 3
k_div3 = k_iter // 3
mixed = base + k_mod3 * 16 + k4[None, :] * 16
carry = mixed >= 48
bs_row = n_group[:, None] * 32 + k_iter * 5 + k_div3 + k4[None, :] + carry
bs_col = mixed - carry * 48
return tl.load(b_scales_ptr + bs_row * stride_bsn + bs_col * stride_bsk)
if K_SCALE_PAD == 16:
bs_row = n_group[:, None] * 32 + k_iter * 16 + k4[None, :] * 4 + (base >> 4)
bs_col = base & 0xF
return tl.load(b_scales_ptr + bs_row * stride_bsn + bs_col * stride_bsk)
flat_local = k4 * 64
flat_bs = k_iter * 256 + base + flat_local[None, :]
else:
offs_ks = k_iter * BK_SCALES + tl.arange(0, BK_SCALES)
ks8 = offs_ks // 8
k2 = (offs_ks // 4) % 2
k4 = offs_ks % 4
flat_bs = ks8[None, :] * 256 + k4[None, :] * 64 + n16[:, None] * 4 + k2[None, :] * 2 + n2[:, None]
bs_row = n_group[:, None] * 32 + flat_bs // K_SCALE_PAD
bs_col = flat_bs % K_SCALE_PAD
return tl.load(b_scales_ptr + bs_row * stride_bsn + bs_col * stride_bsk)
# ---------------------------------------------------------------------------
# GEMM kernel without inline quant (reads pre-quantized A fp4 + A scales)
# ---------------------------------------------------------------------------
@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),
"EVEN_M": lambda args: (args["M"] % args["BLOCK_SIZE_M"]) == 0,
"EVEN_N": lambda args: (args["N"] % args["BLOCK_SIZE_N"]) == 0,
})
@triton.jit
def _gemm_noquant_kernel(
a_fp4_ptr, a_scale_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_afm, stride_afk,
stride_asm, stride_ask,
stride_bk, stride_bn,
stride_cm, stride_cn, stride_bsn, stride_bsk,
K_SCALE_PAD: tl.constexpr,
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,
EVEN_M: tl.constexpr, EVEN_N: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
SCALE_GROUP_SIZE: tl.constexpr = 32
BK_PACKED: tl.constexpr = BLOCK_SIZE_K // 2
BK_SCALES: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
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(stride_afm > 0)
tl.assume(stride_afk > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsn > 0)
tl.assume(stride_bsk > 0)
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BK_PACKED)
if EVEN_M:
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
else:
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
# A fp4 pointers [M, K_packed]
offs_afk = tl.arange(0, BK_PACKED)
offs_afk_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_afk
a_fp4_ptrs = a_fp4_ptr + offs_am[:, None] * stride_afm + offs_afk_split[None, :] * stride_afk
# A scale pointers [M, K_bf16//32]
offs_ask = tl.arange(0, BK_SCALES)
offs_ask_split = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + offs_ask
a_scale_ptrs = a_scale_ptr + offs_am[:, None] * stride_asm + offs_ask_split[None, :] * stride_ask
# B pointers [N, K_packed] loaded as [BN, BK_packed] coalesced
offs_k = tl.arange(0, BK_PACKED)
offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
if EVEN_N:
offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
else:
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_split[None, :] * stride_bk
# B scale N-side indices (invariant across K iterations)
n_group = offs_bn // 32
n2 = (offs_bn % 32) // 16
n16 = offs_bn % 16
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 EVEN_K:
a_fp4 = tl.load(a_fp4_ptrs)
a_scales = tl.load(a_scale_ptrs)
else:
a_fp4 = tl.load(a_fp4_ptrs, mask=offs_afk[None, :] < K - k * BK_PACKED, other=0)
a_scales = tl.load(a_scale_ptrs, mask=offs_ask[None, :] < tl.cdiv(2 * K, SCALE_GROUP_SIZE) - k * BK_SCALES, other=0)
b_scales = _load_b_scales(
b_scales_ptr,
stride_bsn,
stride_bsk,
n_group,
n2,
n16,
k,
BK_SCALES,
K_SCALE_PAD,
)
if EVEN_K:
b = tl.load(b_ptrs, cache_modifier=cache_modifier).trans(1, 0)
else:
b = tl.load(b_ptrs, mask=offs_k[None, :] < K - k * BK_PACKED, other=0, cache_modifier=cache_modifier).trans(1, 0)
accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
a_fp4_ptrs += BK_PACKED * stride_afk
a_scale_ptrs += BK_SCALES * stride_ask
b_ptrs += BK_PACKED * stride_bk
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, :]
if EVEN_M and EVEN_N:
tl.store(c_ptrs, c)
else:
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
# ---------------------------------------------------------------------------
# Fused kernel (inline A quant + GEMM, for small M shapes)
# ---------------------------------------------------------------------------
@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),
"EVEN_M": lambda args: (args["M"] % args["BLOCK_SIZE_M"]) == 0,
"EVEN_N": lambda args: (args["N"] % args["BLOCK_SIZE_N"]) == 0,
})
@triton.jit
def _fused_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
K_SCALE_PAD: tl.constexpr,
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,
EVEN_M: tl.constexpr, EVEN_N: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
SCALE_GROUP_SIZE: tl.constexpr = 32
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(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_bsn > 0)
tl.assume(stride_bsk > 0)
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
if EVEN_M:
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
else:
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_ak = tl.arange(0, BLOCK_SIZE_K)
offs_ak_split = pid_k * SPLITK_BLOCK_SIZE + offs_ak
a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_ak_split[None, :] * stride_ak
offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
if EVEN_N:
offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
else:
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_split[None, :] * stride_bk
n_group = offs_bn // 32
n2 = (offs_bn % 32) // 16
n16 = offs_bn % 16
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 EVEN_K:
a_bf16 = tl.load(a_ptrs)
else:
a_bf16 = tl.load(a_ptrs, mask=offs_ak[None, :] < 2 * K - k * BLOCK_SIZE_K, other=0.0)
a_fp4, a_scales = _mxfp4_quant_op(a_bf16.to(tl.float32), BLOCK_SIZE_K, BLOCK_SIZE_M, SCALE_GROUP_SIZE)
b_scales = _load_b_scales(
b_scales_ptr,
stride_bsn,
stride_bsk,
n_group,
n2,
n16,
k,
BLOCK_SIZE_K // SCALE_GROUP_SIZE,
K_SCALE_PAD,
)
if EVEN_K:
b = tl.load(b_ptrs, cache_modifier=cache_modifier).trans(1, 0)
else:
b = tl.load(b_ptrs, mask=offs_k[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0, cache_modifier=cache_modifier).trans(1, 0)
accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk
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
if EVEN_M and EVEN_N:
tl.store(c_ptrs, c)
else:
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
# ---------------------------------------------------------------------------
# Per-shape config
# ---------------------------------------------------------------------------
def _make_cfg(bm, bn, bk, gm=4, nks=1, nw=4, ns=2, wpe=0, mind=16, cm=".cg"):
return {
"BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
"GROUP_SIZE_M": gm, "NUM_KSPLIT": nks, "SPLITK_BLOCK_SIZE": bk,
"num_warps": nw, "num_stages": ns, "waves_per_eu": wpe,
"matrix_instr_nonkdim": mind, "cache_modifier": cm,
}
# Fused kernel configs (small M shapes)
_SHAPE_CONFIGS = {
(4, 2880, 256): _make_cfg(16, 64, 256, nw=4, ns=2),
(16, 2112, 3584): _make_cfg(16, 128, 256, nks=7, nw=4, ns=2),
(32, 4096, 256): _make_cfg(32, 64, 256, nw=4, ns=2),
(32, 2880, 256): _make_cfg(32, 64, 256, nw=4, ns=2),
# Fallback for large M (used if noquant path disabled)
(64, 7168, 1024): _make_cfg(16, 64, 256, nw=4, ns=2),
(256, 3072, 768): _make_cfg(64, 64, 256, nw=8, ns=2, wpe=2),
# Test shapes
(8, 2112, 3584): _make_cfg(16, 128, 256, nw=4, ns=2),
(16, 3072, 768): _make_cfg(16, 128, 256, nw=4, ns=2),
(64, 3072, 768): _make_cfg(64, 64, 256, nw=8, ns=2),
(256, 2880, 256): _make_cfg(64, 64, 256, nw=8, ns=2),
}
# Noquant GEMM configs (large M shapes, lower VGPR → higher occupancy)
_NOQUANT_GEMM_CONFIGS = {
(64, 7168, 1024): _make_cfg(32, 64, 512, gm=4, nw=8, ns=2),
(256, 3072, 768): _make_cfg(64, 64, 256, nw=8, ns=2, wpe=2),
(64, 3072, 768): _make_cfg(64, 64, 256, nw=8, ns=2),
(256, 2880, 256): _make_cfg(64, 64, 256, nw=8, ns=2),
}
_NOQUANT_M_THRESHOLD = 64
_config_cache = {}
_noquant_config_cache = {}
_ksplit_cache = {}
_out_cache = {}
_a_fp4_cache = {}
_a_scale_cache = {}
def _get_quant_launch_params(m):
if m == 64:
return 32, 512, 8, 1
return 32, 256, 4, 1
def _get_shape_config(m, n, k_packed):
key = (m, n, k_packed)
cached = _config_cache.get(key)
if cached is not None:
return dict(cached)
cfg = _SHAPE_CONFIGS.get(key)
if cfg is None:
cfg, _ = _get_config(m, n, k_packed)
_config_cache[key] = cfg
return dict(cfg)
def _get_noquant_config(m, n, k_packed):
key = (m, n, k_packed)
cached = _noquant_config_cache.get(key)
if cached is not None:
return dict(cached)
cfg = _NOQUANT_GEMM_CONFIGS.get(key)
if cfg is None:
cfg, _ = _get_config(m, n, k_packed)
_noquant_config_cache[key] = cfg
return dict(cfg)
# ---------------------------------------------------------------------------
# Wrapper
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
A, _, B_q, _, B_scale_sh = data
m, k_bf16 = A.shape
n = B_q.shape[0]
k_packed = k_bf16 // 2
if m >= _NOQUANT_M_THRESHOLD:
return _noquant_path(A, B_q, B_scale_sh, m, n, k_bf16, k_packed)
return _fused_path(A, B_q, B_scale_sh, m, n, k_packed)
def _fused_path(A, B_q, B_scale_sh, m, n, k_packed):
w = B_q.view(torch.uint8)
b_scale = B_scale_sh.view(torch.uint8)
k_scale_pad = b_scale.shape[1]
config = _get_shape_config(m, n, k_packed)
key = (m, n, k_packed)
sk_cached = _ksplit_cache.get(key)
if sk_cached is not None:
config["SPLITK_BLOCK_SIZE"], config["BLOCK_SIZE_K"], config["NUM_KSPLIT"] = sk_cached
else:
if config["NUM_KSPLIT"] > 1:
sbs, bk, nks = get_splitk(k_packed, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"])
config["SPLITK_BLOCK_SIZE"] = sbs
config["BLOCK_SIZE_K"] = bk
config["NUM_KSPLIT"] = nks
else:
config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
config["NUM_KSPLIT"] = 1
if config["BLOCK_SIZE_K"] >= 2 * k_packed:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k_packed)
config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
config["NUM_KSPLIT"] = 1
config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 128)
_ksplit_cache[key] = (config["SPLITK_BLOCK_SIZE"], config["BLOCK_SIZE_K"], config["NUM_KSPLIT"])
NUM_KSPLIT = config["NUM_KSPLIT"]
config["K_SCALE_PAD"] = k_scale_pad
y_key = (m, n, A.device)
y = _out_cache.get(y_key)
if y is None:
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_out_cache[y_key] = y
if NUM_KSPLIT > 1:
y_pp = torch.empty((NUM_KSPLIT, m, n), dtype=torch.float32, device=A.device)
else:
y_pp = None
grid = lambda META: (
META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),
)
_fused_kernel[grid](
A, w, y if NUM_KSPLIT == 1 else y_pp, b_scale,
m, n, k_packed,
A.stride(0), A.stride(1), w.stride(1), w.stride(0),
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),
b_scale.stride(0), b_scale.stride(1),
**config,
)
if NUM_KSPLIT > 1:
ACTUAL_KSPLIT = triton.cdiv(k_packed, (config["SPLITK_BLOCK_SIZE"] // 2))
grid_r = (triton.cdiv(m, 16), triton.cdiv(n, 64))
_gemm_afp4wfp4_reduce_kernel[grid_r](
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
def _noquant_path(A, B_q, B_scale_sh, m, n, k_bf16, k_packed):
w = B_q.view(torch.uint8)
b_scale = B_scale_sh.view(torch.uint8)
k_scale_pad = b_scale.shape[1]
k_scales = k_bf16 // 32
# Cached intermediate buffers
buf_key = (m, k_packed, A.device)
a_fp4 = _a_fp4_cache.get(buf_key)
if a_fp4 is None:
a_fp4 = torch.empty((m, k_packed), dtype=torch.uint8, device=A.device)
_a_fp4_cache[buf_key] = a_fp4
scale_key = (m, k_scales, A.device)
a_scale = _a_scale_cache.get(scale_key)
if a_scale is None:
a_scale = torch.empty((m, k_scales), dtype=torch.uint8, device=A.device)
_a_scale_cache[scale_key] = a_scale
# Step 1: Quantize A (separate kernel)
BM_Q, BK_Q, Q_NUM_WARPS, Q_NUM_STAGES = _get_quant_launch_params(m)
grid_q = (triton.cdiv(m, BM_Q), triton.cdiv(k_bf16, BK_Q))
_quant_a_kernel[grid_q](
A, a_fp4, a_scale,
m, k_bf16,
A.stride(0), A.stride(1),
a_fp4.stride(0), a_fp4.stride(1),
a_scale.stride(0), a_scale.stride(1),
BM_Q=BM_Q, BK_Q=BK_Q,
num_warps=Q_NUM_WARPS, num_stages=Q_NUM_STAGES,
)
# Step 2: GEMM with pre-quantized A (no inline quant → lower VGPR → higher occupancy)
config = _get_noquant_config(m, n, k_packed)
config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
config["NUM_KSPLIT"] = 1
if config["BLOCK_SIZE_K"] >= 2 * k_packed:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k_packed)
config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 128)
config["K_SCALE_PAD"] = k_scale_pad
y_key = (m, n, A.device)
y = _out_cache.get(y_key)
if y is None:
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_out_cache[y_key] = y
grid = lambda META: (
triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),
)
_gemm_noquant_kernel[grid](
a_fp4, a_scale, w, y, b_scale,
m, n, k_packed,
a_fp4.stride(0), a_fp4.stride(1),
a_scale.stride(0), a_scale.stride(1),
w.stride(1), w.stride(0),
y.stride(0), y.stride(1),
b_scale.stride(0), b_scale.stride(1),
**config,
)
return y
scrolls · 634 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