submission 732588
rosehulman. · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 675 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-732588?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:f16ec430bbac57383f998a450cf27ae528d2ded11911a55924a86a21b4d12f22
license declaredunknown
license concludedunknown
authorsrosehulman.
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Fused BF16->MXFP4 quant + FP4 GEMM kernel for AMD MI355X.split-k
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)tile-k = 256
QUANT_BK = 256tile-m = 16
REDUCE_BLOCK_SIZE_M = 16tile-n = 64
REDUCE_BLOCK_SIZE_N = 64Kernel source
submission.py675 lines
"""
Fused BF16->MXFP4 quant + FP4 GEMM kernel for AMD MI355X.
Uses pre-shuffled B and shuffled B_scale with corrected scale indexing.
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
@triton.jit
def _mxfp4_quant_in_reg(
x_bf16,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
):
"""Quantize BF16 block to MXFP4 using HW v_cvt_scalef32_pk_fp4_bf16."""
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr = 32
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_K // MXFP4_QUANT_BLOCK_SIZE
# Compute scales from FP32 values
x_fp32 = x_bf16.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x_fp32), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
log2_amax = ((amax >> 23) & 0xFF).to(tl.int32) - 127
scale_e8m0_unbiased_i = log2_amax - 2
scale_e8m0_unbiased_i = tl.minimum(tl.maximum(scale_e8m0_unbiased_i, -127), 127)
bs_e8m0 = scale_e8m0_unbiased_i.to(tl.uint8) + 127
# HW instruction divides by scale: fp4 = convert(bf16 / hw_scale)
# hw_scale = 2^unbiased (reciprocal of SW quant_scale which is 2^(-unbiased))
hw_scale_bits = (scale_e8m0_unbiased_i.to(tl.int32) + 127).to(tl.uint32) << 23
hw_scale = hw_scale_bits.to(tl.float32, bitcast=True) # [M, NUM_QB, 1]
# Broadcast scale to per-pair granularity
hw_scale_flat = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE))
hw_scale_flat = hw_scale_flat.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K)
# Take scale for even element of each pair (both share same scale within 32-group)
hw_scale_pairs = hw_scale_flat.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2, 2)
hw_scale_even, _ = tl.split(hw_scale_pairs)
hw_scale_pair = hw_scale_even.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2)
# Pack BF16 pairs into uint32 for HW instruction
x_u16 = x_bf16.to(tl.uint16, bitcast=True).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2, 2)
lo_u16, hi_u16 = tl.split(x_u16)
x_u32 = lo_u16.to(tl.uint32) | (hi_u16.to(tl.uint32) << 16)
x_u32 = x_u32.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2)
# HW FP4 conversion
fp4_u32 = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
"=v, v, v",
[x_u32, hw_scale_pair],
dtype=tl.uint32,
is_pure=True,
pack=1,
)
x_fp4 = (fp4_u32 & 0xFF).to(tl.uint8)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _standalone_quant_kernel(
a_ptr, a_fp4_ptr, a_scale_ptr,
M, K,
stride_am, stride_ak,
stride_qm, stride_qk,
stride_sm, stride_sk,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_k = tl.program_id(1)
offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_k = pid_k * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
m_mask = offs_m[:, None] < M
k_mask = offs_k[None, :] < K
a_bf16 = tl.load(a_ptrs, mask=m_mask & k_mask, other=0.0)
a_fp4, a_scales = _mxfp4_quant_in_reg(a_bf16, BLOCK_SIZE_M, BLOCK_SIZE_K)
HALF_K: tl.constexpr = BLOCK_SIZE_K // 2
offs_qk = pid_k * HALF_K + tl.arange(0, HALF_K)
q_ptrs = a_fp4_ptr + offs_m[:, None] * stride_qm + offs_qk[None, :] * stride_qk
tl.store(q_ptrs, a_fp4, mask=m_mask & (offs_qk[None, :] < (K // 2)))
SCALE_K: tl.constexpr = BLOCK_SIZE_K // 32
offs_sk = pid_k * SCALE_K + tl.arange(0, SCALE_K)
s_ptrs = a_scale_ptr + offs_m[:, None] * stride_sm + offs_sk[None, :] * stride_sk
tl.store(s_ptrs, a_scales, mask=m_mask & (offs_sk[None, :] < (K // 32)))
@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),
}
)
@triton.jit
def _fused_quant_gemm_preshuffle_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
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,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsn > 0)
tl.assume(stride_bsk > 0)
SCALE_GROUP_SIZE: tl.constexpr = 32
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
if NUM_KSPLIT == 1:
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
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(pid_k >= 0)
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
# A: BF16 [M, 2*K]
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_ak = pid_k * SPLITK_BLOCK_SIZE + tl.arange(0, BLOCK_SIZE_K)
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_ak[None, :] * stride_ak)
# B: pre-shuffled MXFP4 [N//16, K_packed*16]
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)
# B_scale: shuffled E8M0 [N_pad, K_scale_pad]
# Each group of 32 N values occupies 32 consecutive rows.
# Row index = pid_n * BLOCK_SIZE_N + group_offset * 32
offs_bsn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32)
offs_ks = (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[None, :] * stride_bsk
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k_iter in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
# Load and quantize A in registers
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
else:
k_offset = (k_iter - pid_k * num_k_iter) * BLOCK_SIZE_K
a_bf16 = tl.load(
a_ptrs,
mask=tl.arange(0, BLOCK_SIZE_K)[None, :] < (2 * K - pid_k * SPLITK_BLOCK_SIZE - k_offset),
other=0.0,
)
a_fp4, a_scales = _mxfp4_quant_in_reg(a_bf16, BLOCK_SIZE_M, BLOCK_SIZE_K)
# Load and unshuffle B scales
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.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)
)
# Load and unshuffle B data
if EVEN_K:
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
b = tl.load(
b_ptrs,
mask=offs_k_shuffle_arr[None, :] < ((K - (pid_k * (SPLITK_BLOCK_SIZE // 2) + (k_iter - pid_k * num_k_iter) * (BLOCK_SIZE_K // 2))) * 16),
other=0,
cache_modifier=cache_modifier,
)
b = (
b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.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) * 16 * stride_bk
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 _reduce_kernel(
c_in_ptr, c_out_ptr, M, N,
stride_c_in_k, stride_c_in_m, stride_c_in_n,
stride_c_out_m, stride_c_out_n,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
# Sequential accumulation: load one partial at a time (fewer registers)
base_ptrs = (
c_in_ptr
+ (offs_m[:, None] * stride_c_in_m)
+ (offs_n[None, :] * stride_c_in_n)
)
acc = tl.load(base_ptrs).to(tl.float32)
for ks in tl.static_range(1, MAX_KSPLIT):
if ks < ACTUAL_KSPLIT:
acc += tl.load(base_ptrs + ks * stride_c_in_k).to(tl.float32)
c = acc.to(c_out_ptr.type.element_ty)
c_out_ptrs = (
c_out_ptr
+ (offs_m[:, None] * stride_c_out_m)
+ (offs_n[None, :] * stride_c_out_n)
)
tl.store(c_out_ptrs, c)
@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),
}
)
@triton.jit
def _gemm_only_preshuffle_kernel(
a_fp4_ptr, a_scale_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_qm, stride_qk,
stride_sm, stride_sk,
stride_bn, stride_bk,
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,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
tl.assume(stride_qm > 0)
tl.assume(stride_qk > 0)
tl.assume(stride_sm > 0)
tl.assume(stride_sk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsn > 0)
tl.assume(stride_bsk > 0)
SCALE_GROUP_SIZE: tl.constexpr = 32
HALF_BK: tl.constexpr = BLOCK_SIZE_K // 2
SCALE_BK: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
if NUM_KSPLIT == 1:
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
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(pid_k >= 0)
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, HALF_BK)
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_aqk = pid_k * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, HALF_BK)
a_fp4_ptrs = a_fp4_ptr + (offs_am[:, None] * stride_qm + offs_aqk[None, :] * stride_qk)
offs_ask = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + tl.arange(0, SCALE_BK)
a_scale_ptrs = a_scale_ptr + (offs_am[:, None] * stride_sm + offs_ask[None, :] * stride_sk)
offs_k_shuffle_arr = tl.arange(0, HALF_BK * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)
offs_bsn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32)
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
0, SCALE_BK * 32
)
b_scale_ptrs = (
b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k_iter in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
if EVEN_K:
a_fp4 = tl.load(a_fp4_ptrs, cache_modifier=cache_modifier)
a_scales = tl.load(a_scale_ptrs, cache_modifier=cache_modifier)
else:
k_off = (k_iter - pid_k * num_k_iter) * HALF_BK
k_remain = K - (pid_k * (SPLITK_BLOCK_SIZE // 2) + k_off)
a_fp4 = tl.load(
a_fp4_ptrs,
mask=tl.arange(0, HALF_BK)[None, :] < k_remain,
other=0,
cache_modifier=cache_modifier,
)
s_remain = (2 * K) // SCALE_GROUP_SIZE - (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + (k_iter - pid_k * num_k_iter) * SCALE_BK)
a_scales = tl.load(
a_scale_ptrs,
mask=tl.arange(0, SCALE_BK)[None, :] < s_remain,
other=0,
cache_modifier=cache_modifier,
)
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(
BLOCK_SIZE_N // 32,
SCALE_BK // 8,
4, 16, 2, 2, 1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, SCALE_BK)
)
if EVEN_K:
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
b = tl.load(
b_ptrs,
mask=offs_k_shuffle_arr[None, :] < ((K - (pid_k * (SPLITK_BLOCK_SIZE // 2) + (k_iter - pid_k * num_k_iter) * HALF_BK)) * 16),
other=0,
cache_modifier=cache_modifier,
)
b = (
b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, HALF_BK)
.trans(1, 0)
)
accumulator = tl.dot_scaled(
a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", accumulator
)
a_fp4_ptrs += HALF_BK * stride_qk
a_scale_ptrs += SCALE_BK * stride_sk
b_ptrs += HALF_BK * 16 * stride_bk
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)
def get_splitk(K, BLOCK_SIZE_K, NUM_KSPLIT):
SPLITK_BLOCK_SIZE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
if (
K % (SPLITK_BLOCK_SIZE // 2) == 0
and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0
and K % (BLOCK_SIZE_K // 2) == 0
):
break
elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:
if NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
elif BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // 2
elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // 2
else:
break
SPLITK_BLOCK_SIZE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
NUM_KSPLIT = triton.cdiv(K, (SPLITK_BLOCK_SIZE // 2))
return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT
CONFIGS = {
(4, 2880, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(16, 2112, 7168): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 7},
(32, 4096, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(32, 2880, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(256, 3072, 1536): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None, "NUM_KSPLIT": 1},
}
DEFAULT_CONFIG = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}
GEMM_CONFIGS = {
(32, 4096, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(32, 2880, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 2},
}
GEMM_DEFAULT_CONFIG = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}
_buf_cache = {}
_config_cache = {}
def _get_buffers(m, n, num_ksplit, device):
key = (m, n, num_ksplit)
if key not in _buf_cache:
y = torch.empty((m, n), dtype=torch.bfloat16, device=device)
y_pp = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=device) if num_ksplit > 1 else None
_buf_cache[key] = (y, y_pp)
return _buf_cache[key]
def _get_config(m, n, k):
key = (m, n, k)
if key not in _config_cache:
config = CONFIGS.get(key, DEFAULT_CONFIG).copy()
K_packed = k // 2
if config["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(
K_packed, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
)
config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
config["BLOCK_SIZE_K"] = BLOCK_SIZE_K
config["NUM_KSPLIT"] = NUM_KSPLIT
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_N"] = max(config["BLOCK_SIZE_N"], 32)
_config_cache[key] = config
return _config_cache[key]
def fused_quant_gemm(A_bf16, B_shuffle, B_scale_sh, m, n, k):
config = _get_config(m, n, k)
K_packed = k // 2
y, y_pp = _get_buffers(m, n, config["NUM_KSPLIT"], A_bf16.device)
# Reshape B from [N, K_packed] to [N//16, K_packed*16] for preshuffle layout
b_uint8 = B_shuffle.view(torch.uint8)
b_reshaped = b_uint8.reshape(n // 16, K_packed * 16)
b_scale_uint8 = B_scale_sh.view(torch.uint8)
grid = lambda META: (
META["NUM_KSPLIT"]
* triton.cdiv(m, META["BLOCK_SIZE_M"])
* triton.cdiv(n, META["BLOCK_SIZE_N"]),
)
_fused_quant_gemm_preshuffle_kernel[grid](
A_bf16, b_reshaped,
y if config["NUM_KSPLIT"] == 1 else y_pp,
b_scale_uint8,
m, n, K_packed,
A_bf16.stride(0), A_bf16.stride(1),
b_reshaped.stride(0), b_reshaped.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),
b_scale_uint8.stride(0), b_scale_uint8.stride(1),
**config,
)
if config["NUM_KSPLIT"] > 1:
REDUCE_BLOCK_SIZE_M = 16
REDUCE_BLOCK_SIZE_N = 64
ACTUAL_KSPLIT = triton.cdiv(K_packed, (config["SPLITK_BLOCK_SIZE"] // 2))
grid_reduce = (
triton.cdiv(m, REDUCE_BLOCK_SIZE_M),
triton.cdiv(n, REDUCE_BLOCK_SIZE_N),
)
_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),
REDUCE_BLOCK_SIZE_M, REDUCE_BLOCK_SIZE_N,
ACTUAL_KSPLIT, triton.next_power_of_2(config["NUM_KSPLIT"]),
)
return y
def separate_quant_gemm(A_bf16, B_shuffle, B_scale_sh, m, n, k):
K_packed = k // 2
K_bf16 = k
QUANT_BM = 16
QUANT_BK = 256
A_fp4 = torch.empty((m, K_packed), dtype=torch.uint8, device=A_bf16.device)
A_scale = torch.empty((m, K_bf16 // 32), dtype=torch.uint8, device=A_bf16.device)
grid_quant = (triton.cdiv(m, QUANT_BM), triton.cdiv(K_bf16, QUANT_BK))
_standalone_quant_kernel[grid_quant](
A_bf16, A_fp4, A_scale,
m, K_bf16,
A_bf16.stride(0), A_bf16.stride(1),
A_fp4.stride(0), A_fp4.stride(1),
A_scale.stride(0), A_scale.stride(1),
QUANT_BM, QUANT_BK,
)
config = GEMM_CONFIGS.get((m, n, k), GEMM_DEFAULT_CONFIG).copy()
if config["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(
K_packed, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
)
config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
config["BLOCK_SIZE_K"] = BLOCK_SIZE_K
config["NUM_KSPLIT"] = NUM_KSPLIT
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_N"] = max(config["BLOCK_SIZE_N"], 32)
y = torch.empty((m, n), dtype=torch.bfloat16, device=A_bf16.device)
if config["NUM_KSPLIT"] > 1:
y_pp = torch.empty(
(config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=A_bf16.device
)
else:
y_pp = None
b_uint8 = B_shuffle.view(torch.uint8)
b_reshaped = b_uint8.reshape(n // 16, K_packed * 16)
b_scale_uint8 = B_scale_sh.view(torch.uint8)
grid = lambda META: (
META["NUM_KSPLIT"]
* triton.cdiv(m, META["BLOCK_SIZE_M"])
* triton.cdiv(n, META["BLOCK_SIZE_N"]),
)
_gemm_only_preshuffle_kernel[grid](
A_fp4, A_scale,
b_reshaped,
y if config["NUM_KSPLIT"] == 1 else y_pp,
b_scale_uint8,
m, n, K_packed,
A_fp4.stride(0), A_fp4.stride(1),
A_scale.stride(0), A_scale.stride(1),
b_reshaped.stride(0), b_reshaped.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),
b_scale_uint8.stride(0), b_scale_uint8.stride(1),
**config,
)
if config["NUM_KSPLIT"] > 1:
REDUCE_BLOCK_SIZE_M = 16
REDUCE_BLOCK_SIZE_N = 64
ACTUAL_KSPLIT = triton.cdiv(K_packed, (config["SPLITK_BLOCK_SIZE"] // 2))
grid_reduce = (
triton.cdiv(m, REDUCE_BLOCK_SIZE_M),
triton.cdiv(n, REDUCE_BLOCK_SIZE_N),
)
_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),
REDUCE_BLOCK_SIZE_M, REDUCE_BLOCK_SIZE_N,
ACTUAL_KSPLIT, triton.next_power_of_2(config["NUM_KSPLIT"]),
)
return y
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n, _ = B.shape
return fused_quant_gemm(A, B_shuffle, B_scale_sh, m, n, k)
scrolls · 675 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 650646.
⋯ 8 unchanged lines@triton.jit- def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):- pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS- tall_xcds = GRID_MN % NUM_XCDS- tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds- xcd = pid % NUM_XCDS- local_pid = pid // NUM_XCDS- if xcd < tall_xcds:- pid = xcd * pids_per_xcd + local_pid- else:- pid = (tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid)- return pid--- @triton.jitdef _mxfp4_quant_in_reg(x_bf16,BLOCK_SIZE_M: tl.constexpr,BLOCK_SIZE_K: tl.constexpr,):- """Hardware-accelerated BF16 -> MXFP4 quantization using v_cvt_scalef32_pk_fp4_bf16."""+ """Quantize BF16 block to MXFP4 using HW v_cvt_scalef32_pk_fp4_bf16."""MXFP4_QUANT_BLOCK_SIZE: tl.constexpr = 32NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_K // MXFP4_QUANT_BLOCK_SIZE- HALF_QUANT_BLOCK: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE // 2- # Compute amax and scale in fp32 (same logic as before)- x_fp32 = x_bf16.to(tl.float32)- x_fp32 = x_fp32.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)+ # Compute scales from FP32 values+ x_fp32 = x_bf16.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)amax = tl.max(tl.abs(x_fp32), axis=-1, keep_dims=True)amax = amax.to(tl.int32, bitcast=True)amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000- exp_biased = (amax >> 23).to(tl.int32)- scale_e8m0_unbiased = exp_biased - 129- scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased < -127, -127, scale_e8m0_unbiased)- bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127+ log2_amax = ((amax >> 23) & 0xFF).to(tl.int32) - 127+ scale_e8m0_unbiased_i = log2_amax - 2+ scale_e8m0_unbiased_i = tl.minimum(tl.maximum(scale_e8m0_unbiased_i, -127), 127)+ bs_e8m0 = scale_e8m0_unbiased_i.to(tl.uint8) + 127- # block_scale = 2^(scale_e8m0_unbiased): HW instruction divides by this- block_scale = ((127 + scale_e8m0_unbiased) << 23).to(tl.float32, bitcast=True)+ # HW instruction divides by scale: fp4 = convert(bf16 / hw_scale)+ # hw_scale = 2^unbiased (reciprocal of SW quant_scale which is 2^(-unbiased))+ hw_scale_bits = (scale_e8m0_unbiased_i.to(tl.int32) + 127).to(tl.uint32) << 23+ hw_scale = hw_scale_bits.to(tl.float32, bitcast=True) # [M, NUM_QB, 1]- # Pack bf16 pairs into uint32 for the HW instruction- x_bf16_3d = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)- x_bf16_pairs = x_bf16_3d.reshape(- BLOCK_SIZE_M * NUM_QUANT_BLOCKS * HALF_QUANT_BLOCK, 2- )- lo_bf16, hi_bf16 = tl.split(x_bf16_pairs)- lo_u32 = lo_bf16.reshape(BLOCK_SIZE_M * NUM_QUANT_BLOCKS * HALF_QUANT_BLOCK).to(tl.uint16, bitcast=True).to(tl.uint32)- hi_u32 = hi_bf16.reshape(BLOCK_SIZE_M * NUM_QUANT_BLOCKS * HALF_QUANT_BLOCK).to(tl.uint16, bitcast=True).to(tl.uint32)- bf16x2_packed = lo_u32 | (hi_u32 << 16)+ # Broadcast scale to per-pair granularity+ hw_scale_flat = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE))+ hw_scale_flat = hw_scale_flat.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K)+ # Take scale for even element of each pair (both share same scale within 32-group)+ hw_scale_pairs = hw_scale_flat.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2, 2)+ hw_scale_even, _ = tl.split(hw_scale_pairs)+ hw_scale_pair = hw_scale_even.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2)- # Broadcast scale to match flattened pairs- scale_bc = tl.broadcast_to(block_scale, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HALF_QUANT_BLOCK])- scale_flat = scale_bc.reshape(BLOCK_SIZE_M * NUM_QUANT_BLOCKS * HALF_QUANT_BLOCK)+ # Pack BF16 pairs into uint32 for HW instruction+ x_u16 = x_bf16.to(tl.uint16, bitcast=True).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2, 2)+ lo_u16, hi_u16 = tl.split(x_u16)+ x_u32 = lo_u16.to(tl.uint32) | (hi_u16.to(tl.uint32) << 16)+ x_u32 = x_u32.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2)- # v_cvt_scalef32_pk_fp4_bf16: converts 2 bf16 to 2 packed fp4 with scale- result_u32 = tl.inline_asm_elementwise(+ # HW FP4 conversion+ fp4_u32 = tl.inline_asm_elementwise("v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",- "=v,v,v",- [bf16x2_packed, scale_flat],+ "=v, v, v",+ [x_u32, hw_scale_pair],dtype=tl.uint32,is_pure=True,pack=1,)-- x_fp4 = (result_u32 & 0xFF).to(tl.uint8)+ x_fp4 = (fp4_u32 & 0xFF).to(tl.uint8)x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2)return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)⋯ 55 unchanged lineswaves_per_eu: tl.constexpr,matrix_instr_nonkdim: tl.constexpr,cache_modifier: tl.constexpr,- XCD_REMAP: tl.constexpr = False,):tl.assume(stride_am > 0)tl.assume(stride_ak > 0)⋯ 9 unchanged linesnum_pid_n = tl.cdiv(N, BLOCK_SIZE_N)pid_unified = tl.program_id(axis=0)- if XCD_REMAP:- pid_unified = remap_xcd(pid_unified, num_pid_m * num_pid_n * NUM_KSPLIT)pid_k = pid_unified % NUM_KSPLITpid = pid_unified // NUM_KSPLIT⋯ 90 unchanged linesb_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bkb_scale_ptrs += BLOCK_SIZE_K * stride_bsk- c = accumulator.to(tl.float32)+ 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_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)-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)⋯ 9 unchanged linespid_n = tl.program_id(axis=1)offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % Moffs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N- offs_k = tl.arange(0, MAX_KSPLIT)- c_in_ptrs = (++ # Sequential accumulation: load one partial at a time (fewer registers)+ base_ptrs = (c_in_ptr- + (offs_k[:, None, None] * stride_c_in_k)- + (offs_m[None, :, None] * stride_c_in_m)- + (offs_n[None, None, :] * stride_c_in_n)+ + (offs_m[:, None] * stride_c_in_m)+ + (offs_n[None, :] * stride_c_in_n))- if ACTUAL_KSPLIT == MAX_KSPLIT:- c = tl.load(c_in_ptrs)- else:- c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)- c = tl.sum(c, axis=0)- c = c.to(c_out_ptr.type.element_ty)+ acc = tl.load(base_ptrs).to(tl.float32)+ for ks in tl.static_range(1, MAX_KSPLIT):+ if ks < ACTUAL_KSPLIT:+ acc += tl.load(base_ptrs + ks * stride_c_in_k).to(tl.float32)++ c = acc.to(c_out_ptr.type.element_ty)c_out_ptrs = (c_out_ptr+ (offs_m[:, None] * stride_c_out_m)⋯ 49 unchanged linesnum_pid_n = tl.cdiv(N, BLOCK_SIZE_N)pid_unified = tl.program_id(axis=0)- pid_unified = remap_xcd(pid_unified, num_pid_m * num_pid_n * NUM_KSPLIT)pid_k = pid_unified % NUM_KSPLITpid = pid_unified // NUM_KSPLIT⋯ 39 unchanged linesfor k_iter 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)+ a_fp4 = tl.load(a_fp4_ptrs, cache_modifier=cache_modifier)+ a_scales = tl.load(a_scale_ptrs, cache_modifier=cache_modifier)else:k_off = (k_iter - pid_k * num_k_iter) * HALF_BKk_remain = K - (pid_k * (SPLITK_BLOCK_SIZE // 2) + k_off)⋯ 1 unchanged linesa_fp4_ptrs,mask=tl.arange(0, HALF_BK)[None, :] < k_remain,other=0,+ cache_modifier=cache_modifier,)s_remain = (2 * K) // SCALE_GROUP_SIZE - (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + (k_iter - pid_k * num_k_iter) * SCALE_BK)a_scales = tl.load(a_scale_ptrs,mask=tl.arange(0, SCALE_BK)[None, :] < s_remain,other=0,+ cache_modifier=cache_modifier,)b_scales = (⋯ 77 unchanged linesCONFIGS = {- (4, 2880, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},- (16, 2112, 7168): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 7},- (32, 4096, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},- (32, 2880, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},- (64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 2},- (256, 3072, 1536): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 2, "num_warps": 8, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (4, 2880, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (16, 2112, 7168): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 7},+ (32, 4096, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (32, 2880, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (256, 3072, 1536): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None, "NUM_KSPLIT": 1},}- DEFAULT_CONFIG = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}+ DEFAULT_CONFIG = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}GEMM_CONFIGS = {- (64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 1024, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (32, 4096, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (32, 2880, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 2},}GEMM_DEFAULT_CONFIG = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}+ _buf_cache = {}+ _config_cache = {}- def fused_quant_gemm(A_bf16, B_shuffle, B_scale_sh, m, n, k):- config = CONFIGS.get((m, n, k), DEFAULT_CONFIG).copy()- K_packed = k // 2+ def _get_buffers(m, n, num_ksplit, device):+ key = (m, n, num_ksplit)+ if key not in _buf_cache:+ y = torch.empty((m, n), dtype=torch.bfloat16, device=device)+ y_pp = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=device) if num_ksplit > 1 else None+ _buf_cache[key] = (y, y_pp)+ return _buf_cache[key]- if config["NUM_KSPLIT"] > 1:- SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(- K_packed, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]- )- config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE- config["BLOCK_SIZE_K"] = BLOCK_SIZE_K- config["NUM_KSPLIT"] = NUM_KSPLIT- else:- config["SPLITK_BLOCK_SIZE"] = 2 * K_packed- config["NUM_KSPLIT"] = 1+ def _get_config(m, n, k):+ key = (m, n, k)+ if key not in _config_cache:+ config = CONFIGS.get(key, DEFAULT_CONFIG).copy()+ K_packed = k // 2+ if config["NUM_KSPLIT"] > 1:+ SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(+ K_packed, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]+ )+ config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE+ config["BLOCK_SIZE_K"] = BLOCK_SIZE_K+ config["NUM_KSPLIT"] = NUM_KSPLIT+ 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_N"] = max(config["BLOCK_SIZE_N"], 32)+ _config_cache[key] = config+ return _config_cache[key]- 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_N"] = max(config["BLOCK_SIZE_N"], 32)+ def fused_quant_gemm(A_bf16, B_shuffle, B_scale_sh, m, n, k):+ config = _get_config(m, n, k)+ K_packed = k // 2- y = torch.empty((m, n), dtype=torch.bfloat16, device=A_bf16.device)+ y, y_pp = _get_buffers(m, n, config["NUM_KSPLIT"], A_bf16.device)- if config["NUM_KSPLIT"] > 1:- y_pp = torch.empty((config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=A_bf16.device)- else:- y_pp = None-+ # Reshape B from [N, K_packed] to [N//16, K_packed*16] for preshuffle layoutb_uint8 = B_shuffle.view(torch.uint8)b_reshaped = b_uint8.reshape(n // 16, K_packed * 16)b_scale_uint8 = B_scale_sh.view(torch.uint8)⋯ 20 unchanged linesif config["NUM_KSPLIT"] > 1:REDUCE_BLOCK_SIZE_M = 16- REDUCE_BLOCK_SIZE_N = 16+ REDUCE_BLOCK_SIZE_N = 64ACTUAL_KSPLIT = triton.cdiv(K_packed, (config["SPLITK_BLOCK_SIZE"] // 2))grid_reduce = (triton.cdiv(m, REDUCE_BLOCK_SIZE_M),⋯ 52 unchanged linesy = torch.empty((m, n), dtype=torch.bfloat16, device=A_bf16.device)if config["NUM_KSPLIT"] > 1:- y_pp = torch.empty((config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=A_bf16.device)+ y_pp = torch.empty(+ (config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=A_bf16.device+ )else:y_pp = None⋯ 25 unchanged linesif config["NUM_KSPLIT"] > 1:REDUCE_BLOCK_SIZE_M = 16- REDUCE_BLOCK_SIZE_N = 16+ REDUCE_BLOCK_SIZE_N = 64ACTUAL_KSPLIT = triton.cdiv(K_packed, (config["SPLITK_BLOCK_SIZE"] // 2))grid_reduce = (triton.cdiv(m, REDUCE_BLOCK_SIZE_M),⋯ 15 unchanged linesm, k = A.shapen, _ = B.shape- if m > 32 and n > 4096:- return separate_quant_gemm(A, B_shuffle, B_scale_sh, m, n, k)return fused_quant_gemm(A, B_shuffle, B_scale_sh, m, n, k)
scrolls · 338 diff lines total
Best evidence level for this revision: reported
JSON