submission 752099
bigmodel_wuzhigang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 704 lines, June 9 Researcher Reciprocity License v1.0.
submission_v10.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-752099?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:6dd10e6046fc763af9dd909003dc1c109cb64dbcd456369ff028f9c6b7d51495
license declaredunknown
license concludedunknown
authorsbigmodel_wuzhigang
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
4. Split-K support with reduction kerneltile-k = 256
QUANT_BK = 256tile-m = 16
QUANT_BM = 16tile-n = 64
REDUCE_TILE_N = 64Kernel source
submission_v10.py704 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Fused BF16->MXFP4 quant + FP4 GEMM kernel for AMD MI355X.
Optimized Triton kernel using:
1. v_cvt_scalef32_pk_fp4_bf16 hardware instruction for FP4 conversion
2. tl.dot_scaled for FP4 matrix multiplication
3. Per-shape tuned configurations
4. Split-K support with reduction kernel
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
@triton.jit
def _bf16_to_fp4_hw(
x_bf16,
TILE_M: tl.constexpr,
TILE_K: tl.constexpr,
):
"""Quantize BF16 block to MXFP4 using HW v_cvt_scalef32_pk_fp4_bf16."""
FP4_GRP_SZ: tl.constexpr = 32
NUM_QBLK: tl.constexpr = TILE_K // FP4_GRP_SZ
# Compute scales from FP32 values
x_fp32 = x_bf16.to(tl.float32).reshape(TILE_M, NUM_QBLK, FP4_GRP_SZ)
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_raw_i = log2_amax - 2
scale_e8m0_raw_i = tl.minimum(tl.maximum(scale_e8m0_raw_i, -127), 127)
bs_e8m0 = scale_e8m0_raw_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_raw_i.to(tl.int32) + 127).to(tl.uint32) << 23
hw_scale = hw_scale_bits.to(tl.float32, bitcast=True) # [M, NUM_QBLK, 1]
# Broadcast scale to per-pair granularity
hw_scale_flat = tl.broadcast_to(hw_scale, (TILE_M, NUM_QBLK, FP4_GRP_SZ))
hw_scale_flat = hw_scale_flat.reshape(TILE_M, TILE_K)
# Take scale for even element of each pair (both share same scale within 32-group)
hw_scale_pairs = hw_scale_flat.reshape(TILE_M, TILE_K // 2, 2)
hw_scale_even, _ = tl.split(hw_scale_pairs)
hw_scale_pair = hw_scale_even.reshape(TILE_M, TILE_K // 2)
# Pack BF16 pairs into uint32 for HW instruction
x_u16 = x_bf16.to(tl.uint16, bitcast=True).reshape(TILE_M, TILE_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(TILE_M, TILE_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(TILE_M, TILE_K // 2)
return x_fp4, bs_e8m0.reshape(TILE_M, NUM_QBLK)
@triton.jit
def _quant_only_launcher(
a_ptr, a_fp4_ptr, a_scale_ptr,
M, K,
stride_am, stride_ak,
stride_qm, stride_qk,
stride_sm, stride_sk,
TILE_M: tl.constexpr,
TILE_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_k = tl.program_id(1)
offs_m = pid_m * TILE_M + tl.arange(0, TILE_M)
offs_k = pid_k * TILE_K + tl.arange(0, TILE_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 = _bf16_to_fp4_hw(a_bf16, TILE_M, TILE_K)
HALF_K: tl.constexpr = TILE_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 = TILE_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["TILE_K"] // 2) == 0)
and (args["KSPLIT_TILE"] % args["TILE_K"] == 0)
and (args["K"] % (args["KSPLIT_TILE"] // 2) == 0),
}
)
@triton.jit
def _quant_gemm_fused_launcher(
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,
TILE_M: tl.constexpr,
TILE_N: tl.constexpr,
TILE_K: tl.constexpr,
M_CLUSTER: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
KSPLIT_TILE: 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_GRP_SZ: tl.constexpr = 32
num_pid_m = tl.cdiv(M, TILE_M)
num_pid_n = tl.cdiv(N, TILE_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 = M_CLUSTER * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * M_CLUSTER
group_size_m = min(num_pid_m - first_pid_m, M_CLUSTER)
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 * KSPLIT_TILE // 2) < K:
num_k_iter = tl.cdiv(KSPLIT_TILE // 2, TILE_K // 2)
# A: BF16 [M, 2*K]
offs_am = (pid_m * TILE_M + tl.arange(0, TILE_M)) % M
offs_ak = pid_k * KSPLIT_TILE + tl.arange(0, TILE_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, (TILE_K // 2) * 16)
offs_k_shuffle = pid_k * (KSPLIT_TILE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (TILE_N // 16) + tl.arange(0, TILE_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 * TILE_N + group_offset * 32
offs_bsn = (pid_n * TILE_N + tl.arange(0, TILE_N // 32) * 32)
offs_ks = (pid_k * (KSPLIT_TILE // SCALE_GRP_SZ) * 32) + tl.arange(
0, TILE_K // SCALE_GRP_SZ * 32
)
b_scale_ptrs = (
b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
)
acc = tl.zeros((TILE_M, TILE_N), dtype=tl.float32)
for k_iter in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
# Fire all loads first for better memory-level parallelism
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b_scales_raw = tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
b_raw = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
k_offset = (k_iter - pid_k * num_k_iter) * TILE_K
a_bf16 = tl.load(
a_ptrs,
mask=tl.arange(0, TILE_K)[None, :] < (2 * K - pid_k * KSPLIT_TILE - k_offset),
other=0.0,
)
b_scales_raw = tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
b_raw = tl.load(
b_ptrs,
mask=offs_k_shuffle_arr[None, :] < ((K - (pid_k * (KSPLIT_TILE // 2) + (k_iter - pid_k * num_k_iter) * (TILE_K // 2))) * 16),
other=0,
cache_modifier=cache_modifier,
)
# Quantize A in registers
a_fp4, a_scales = _bf16_to_fp4_hw(a_bf16, TILE_M, TILE_K)
# Unshuffle B scales
b_scales = (
b_scales_raw
.reshape(
TILE_N // 32,
TILE_K // SCALE_GRP_SZ // 8,
4, 16, 2, 2, 1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(TILE_N, TILE_K // SCALE_GRP_SZ)
)
# Unshuffle B data
b = (
b_raw.reshape(1, TILE_N // 16, TILE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(TILE_N, TILE_K // 2)
.trans(1, 0)
)
acc = tl.dot_scaled(
a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc
)
a_ptrs += TILE_K * stride_ak
b_ptrs += (TILE_K // 2) * 16 * stride_bk
b_scale_ptrs += TILE_K * stride_bsk
c = acc.to(c_ptr.type.element_ty)
offs_cm = pid_m * TILE_M + tl.arange(0, TILE_M).to(tl.int64)
offs_cn = pid_n * TILE_N + tl.arange(0, TILE_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 _splitk_sum_launcher(
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,
TILE_M: tl.constexpr, TILE_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 * TILE_M + tl.arange(0, TILE_M)) % M
offs_n = (pid_n * TILE_N + tl.arange(0, TILE_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["TILE_K"] // 2) == 0)
and (args["KSPLIT_TILE"] % args["TILE_K"] == 0)
and (args["K"] % (args["KSPLIT_TILE"] // 2) == 0),
}
)
@triton.jit
def _fp4_gemm_only_launcher(
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,
TILE_M: tl.constexpr,
TILE_N: tl.constexpr,
TILE_K: tl.constexpr,
M_CLUSTER: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
KSPLIT_TILE: 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_GRP_SZ: tl.constexpr = 32
HALF_TK: tl.constexpr = TILE_K // 2
SCALE_TK: tl.constexpr = TILE_K // SCALE_GRP_SZ
num_pid_m = tl.cdiv(M, TILE_M)
num_pid_n = tl.cdiv(N, TILE_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 = M_CLUSTER * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * M_CLUSTER
group_size_m = min(num_pid_m - first_pid_m, M_CLUSTER)
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 * KSPLIT_TILE // 2) < K:
num_k_iter = tl.cdiv(KSPLIT_TILE // 2, HALF_TK)
offs_am = (pid_m * TILE_M + tl.arange(0, TILE_M)) % M
offs_aqk = pid_k * (KSPLIT_TILE // 2) + tl.arange(0, HALF_TK)
a_fp4_ptrs = a_fp4_ptr + (offs_am[:, None] * stride_qm + offs_aqk[None, :] * stride_qk)
offs_ask = pid_k * (KSPLIT_TILE // SCALE_GRP_SZ) + tl.arange(0, SCALE_TK)
a_scale_ptrs = a_scale_ptr + (offs_am[:, None] * stride_sm + offs_ask[None, :] * stride_sk)
offs_k_shuffle_arr = tl.arange(0, HALF_TK * 16)
offs_k_shuffle = pid_k * (KSPLIT_TILE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (TILE_N // 16) + tl.arange(0, TILE_N // 16)) % (N // 16)
b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)
offs_bsn = (pid_n * TILE_N + tl.arange(0, TILE_N // 32) * 32)
offs_ks = (pid_k * (KSPLIT_TILE // SCALE_GRP_SZ) * 32) + tl.arange(
0, SCALE_TK * 32
)
b_scale_ptrs = (
b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
)
acc = tl.zeros((TILE_M, TILE_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_TK
k_remain = K - (pid_k * (KSPLIT_TILE // 2) + k_off)
a_fp4 = tl.load(
a_fp4_ptrs,
mask=tl.arange(0, HALF_TK)[None, :] < k_remain,
other=0,
cache_modifier=cache_modifier,
)
s_remain = (2 * K) // SCALE_GRP_SZ - (pid_k * (KSPLIT_TILE // SCALE_GRP_SZ) + (k_iter - pid_k * num_k_iter) * SCALE_TK)
a_scales = tl.load(
a_scale_ptrs,
mask=tl.arange(0, SCALE_TK)[None, :] < s_remain,
other=0,
cache_modifier=cache_modifier,
)
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(
TILE_N // 32,
SCALE_TK // 8,
4, 16, 2, 2, 1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(TILE_N, SCALE_TK)
)
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 * (KSPLIT_TILE // 2) + (k_iter - pid_k * num_k_iter) * HALF_TK)) * 16),
other=0,
cache_modifier=cache_modifier,
)
b = (
b.reshape(1, TILE_N // 16, TILE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(TILE_N, HALF_TK)
.trans(1, 0)
)
acc = tl.dot_scaled(
a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc
)
a_fp4_ptrs += HALF_TK * stride_qk
a_scale_ptrs += SCALE_TK * stride_sk
b_ptrs += HALF_TK * 16 * stride_bk
b_scale_ptrs += TILE_K * stride_bsk
c = acc.to(c_ptr.type.element_ty)
offs_cm = pid_m * TILE_M + tl.arange(0, TILE_M).to(tl.int64)
offs_cn = pid_n * TILE_N + tl.arange(0, TILE_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 _calc_splitk_params(K, TILE_K, NUM_KSPLIT):
KSPLIT_TILE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), TILE_K) * TILE_K
)
while NUM_KSPLIT > 1 and TILE_K > 16:
if (
K % (KSPLIT_TILE // 2) == 0
and KSPLIT_TILE % TILE_K == 0
and K % (TILE_K // 2) == 0
):
break
elif K % (KSPLIT_TILE // 2) != 0 and NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
elif KSPLIT_TILE % TILE_K != 0:
if NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
elif TILE_K > 16:
TILE_K = TILE_K // 2
elif K % (TILE_K // 2) != 0 and TILE_K > 16:
TILE_K = TILE_K // 2
else:
break
KSPLIT_TILE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), TILE_K) * TILE_K
)
NUM_KSPLIT = triton.cdiv(K, (KSPLIT_TILE // 2))
return KSPLIT_TILE, TILE_K, NUM_KSPLIT
_TUNE_PARAMS = {
# V10: Conservative tuning based on V8 baseline
# M=4: Try larger TILE_N for better N parallelism
(4, 2880, 512): {"TILE_M": 16, "TILE_N": 64, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 2, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
# M=16, K=7168: Keep Split-K=7, try num_stages=3
(16, 2112, 7168): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 512, "M_CLUSTER": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 7},
# M=32: Keep V8 config
(32, 4096, 512): {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(32, 2880, 512): {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
# M=64: Try waves_per_eu=2
(64, 7168, 2048): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
# M=256: Try waves_per_eu=3
(256, 3072, 1536): {"TILE_M": 16, "TILE_N": 256, "TILE_K": 512, "M_CLUSTER": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
}
_DEF_PARAMS = {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}
_FP4_GEMM_PARAMS = {
(32, 4096, 512): {"TILE_M": 32, "TILE_N": 128, "TILE_K": 256, "M_CLUSTER": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(32, 2880, 512): {"TILE_M": 32, "TILE_N": 64, "TILE_K": 256, "M_CLUSTER": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(64, 7168, 2048): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 512, "M_CLUSTER": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 2},
}
_FP4_GEMM_DEF = {"TILE_M": 16, "TILE_N": 64, "TILE_K": 256, "M_CLUSTER": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}
_mem_pool = {}
_param_pool = {}
def _acquire_buffers(m, n, num_ksplit, device):
key = (m, n, num_ksplit)
if key not in _mem_pool:
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
_mem_pool[key] = (y, y_pp)
return _mem_pool[key]
def _acquire_config(m, n, k):
key = (m, n, k)
if key not in _param_pool:
config = _TUNE_PARAMS.get(key, _DEF_PARAMS).copy()
K_packed = k // 2
if config["NUM_KSPLIT"] > 1:
KSPLIT_TILE, TILE_K, NUM_KSPLIT = _calc_splitk_params(
K_packed, config["TILE_K"], config["NUM_KSPLIT"]
)
config["KSPLIT_TILE"] = KSPLIT_TILE
config["TILE_K"] = TILE_K
config["NUM_KSPLIT"] = NUM_KSPLIT
else:
config["KSPLIT_TILE"] = 2 * K_packed
config["NUM_KSPLIT"] = 1
if config["TILE_K"] >= 2 * K_packed:
config["TILE_K"] = triton.next_power_of_2(2 * K_packed)
config["KSPLIT_TILE"] = 2 * K_packed
config["NUM_KSPLIT"] = 1
config["TILE_N"] = max(config["TILE_N"], 32)
_param_pool[key] = config
return _param_pool[key]
_grid_pool = {}
def _prepare_launch_meta(m, n, k, device):
"""Precompute ALL launch parameters once per shape."""
config = _acquire_config(m, n, k)
K_packed = k // 2
ks = config["NUM_KSPLIT"]
y, y_pp = _acquire_buffers(m, n, ks, device)
grid = (ks * triton.cdiv(m, config["TILE_M"]) * triton.cdiv(n, config["TILE_N"]),)
# Pre-store strides for y/y_pp
if ks == 1:
c_stride_k, c_stride_m, c_stride_n = 0, y.stride(0), y.stride(1)
else:
c_stride_k, c_stride_m, c_stride_n = y_pp.stride(0), y_pp.stride(1), y_pp.stride(2)
meta = {
'config': config,
'K_packed': K_packed,
'grid': grid,
'ks': ks,
'c_stride_k': c_stride_k,
'c_stride_m': c_stride_m,
'c_stride_n': c_stride_n,
}
if ks > 1:
meta['reduce_grid'] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
meta['actual_ksplit'] = triton.cdiv(K_packed, (config["KSPLIT_TILE"] // 2))
meta['max_ksplit'] = triton.next_power_of_2(ks)
return meta
def fused_quant_gemm(A_bf16, B_shuffle, B_scale_sh, m, n, k):
key = (m, n, k)
if key not in _grid_pool:
_grid_pool[key] = _prepare_launch_meta(m, n, k, A_bf16.device)
p = _grid_pool[key]
y, y_pp = _acquire_buffers(m, n, p['ks'], A_bf16.device)
b_reshaped = B_shuffle.view(torch.uint8).reshape(n // 16, p['K_packed'] * 16)
b_scale_uint8 = B_scale_sh.view(torch.uint8)
_quant_gemm_fused_launcher[p['grid']](
A_bf16, b_reshaped,
y if p['ks'] == 1 else y_pp,
b_scale_uint8,
m, n, p['K_packed'],
A_bf16.stride(0), A_bf16.stride(1),
b_reshaped.stride(0), b_reshaped.stride(1),
p['c_stride_k'], p['c_stride_m'], p['c_stride_n'],
b_scale_uint8.stride(0), b_scale_uint8.stride(1),
**p['config'],
)
if p['ks'] > 1:
_splitk_sum_launcher[p['reduce_grid']](
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,
p['actual_ksplit'], p['max_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))
_quant_only_launcher[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 = _FP4_GEMM_PARAMS.get((m, n, k), _FP4_GEMM_DEF).copy()
if config["NUM_KSPLIT"] > 1:
KSPLIT_TILE, TILE_K, NUM_KSPLIT = _calc_splitk_params(
K_packed, config["TILE_K"], config["NUM_KSPLIT"]
)
config["KSPLIT_TILE"] = KSPLIT_TILE
config["TILE_K"] = TILE_K
config["NUM_KSPLIT"] = NUM_KSPLIT
else:
config["KSPLIT_TILE"] = 2 * K_packed
config["NUM_KSPLIT"] = 1
if config["TILE_K"] >= 2 * K_packed:
config["TILE_K"] = triton.next_power_of_2(2 * K_packed)
config["KSPLIT_TILE"] = 2 * K_packed
config["NUM_KSPLIT"] = 1
config["TILE_N"] = max(config["TILE_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["TILE_M"])
* triton.cdiv(n, META["TILE_N"]),
)
_fp4_gemm_only_launcher[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_TILE_M = 16
REDUCE_TILE_N = 64
ACTUAL_KSPLIT = triton.cdiv(K_packed, (config["KSPLIT_TILE"] // 2))
grid_reduce = (
triton.cdiv(m, REDUCE_TILE_M),
triton.cdiv(n, REDUCE_TILE_N),
)
_splitk_sum_launcher[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_TILE_M, REDUCE_TILE_N,
ACTUAL_KSPLIT, triton.next_power_of_2(config["NUM_KSPLIT"]),
)
return y
def custom_kernel(data: input_t) -> output_t:
A = data[0]
return fused_quant_gemm(A, data[3], data[4], A.shape[0], data[1].shape[0], A.shape[1])scrolls · 704 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 746805.
- #!POPCORN leaderboard amd-mxfp4-mm- #!POPCORN gpu MI355X- """- MXFP4 GEMM Optimization V6 - Custom Triton kernel with fused quantization.-- Based on research:- 1. AITER Triton path: aiter.ops.triton.gemm.basic.gemm_afp4wfp4- 2. Hardware instruction: v_cvt_scalef32_pk_fp4_bf16- 3. Per-shape tuned tile configurations- 4. Split-K for large K dimensions-- The benchmark shapes are:- - (m=4, n=2880, k=512) - Small M, large N- - (m=16, n=2112, k=7168) - Medium M, very large K- - (m=32, n=4096, k=512) - Medium M, small K- - (m=32, n=2880, k=512) - Medium M, small K- - (m=64, n=7168, k=2048) - Medium M, medium K- - (m=256, n=3072, k=1536) - Large M, medium K-- Optimization strategy:- 1. For small M (<=16): Use smaller tiles, more K-split- 2. For medium M (32-64): Use medium tiles- 3. For large M (>=256): Use larger tiles- 4. For large K (>=7168): Use split-K for parallelism- """- from task import input_t, output_t-- import torch- import triton- import triton.language as tl--- # Per-shape tuned configurations based on AITER reference performance- # These are tuned for MI355X architecture- CONFIGS = {- # Small M, small K- (4, 2880, 512): {"BLOCK_M": 16, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},- (32, 4096, 512): {"BLOCK_M": 32, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},- (32, 2880, 512): {"BLOCK_M": 32, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},- # Medium M, large K- (16, 2112, 7168): {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 128, "split_k": 4},- # Medium M, medium K- (64, 7168, 2048): {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 2},- # Large M, medium K- (256, 3072, 1536): {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},- # Test shapes- (8, 2112, 7168): {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 128, "split_k": 4},- (16, 3072, 1536): {"BLOCK_M": 64, "BLOCK_N": 64, "BLOCK_K": 64, "split_k": 2},- (64, 3072, 1536): {"BLOCK_M": 128, "BLOCK_N": 64, "BLOCK_K": 64, "split_k": 1},- (256, 2880, 512): {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},- }--- def get_config(m, n, k):- """Get optimized configuration for given shape."""- key = (m, n, k)- if key in CONFIGS:- return CONFIGS[key]-- # Default heuristic for unknown shapes- if m <= 16:- return {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 128, "split_k": max(1, k // 2048)}- elif m <= 64:- return {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": max(1, k // 2048)}- else:- return {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1}--- def custom_kernel(data: input_t) -> output_t:- """- Optimized MXFP4 GEMM with fused quantization and tuned configurations.- """- import aiter- from aiter import QuantType, dtypes- from aiter.ops.triton.quant import dynamic_mxfp4_quant- from aiter.utility.fp4_utils import e8m0_shuffle-- A, B, B_q, B_shuffle, B_scale_sh = data- A = A.contiguous()- B = B.contiguous()- m, k = A.shape- n, _ = B.shape-- # Quantize A to MXFP4 with shuffling- def _quant_mxfp4(x, shuffle=True):- x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)- if shuffle:- bs_e8m0 = e8m0_shuffle(bs_e8m0)- return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)-- A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)-- # Get tuned configuration- config = get_config(m, n, k)-- # Use AITER's asm kernel with bpreshuffle - it's the fastest path- # The Triton path may be slower for these shapes- out_gemm = aiter.gemm_a4w4(- A_q,- B_shuffle,- A_scale_sh,- B_scale_sh,- dtype=dtypes.bf16,- bpreshuffle=True,- )-- return out_gemmNo newline at end of file+ #!POPCORN leaderboard amd-mxfp4-mm+ #!POPCORN gpu MI355X+ """+ Fused BF16->MXFP4 quant + FP4 GEMM kernel for AMD MI355X.+ Optimized Triton kernel using:+ 1. v_cvt_scalef32_pk_fp4_bf16 hardware instruction for FP4 conversion+ 2. tl.dot_scaled for FP4 matrix multiplication+ 3. Per-shape tuned configurations+ 4. Split-K support with reduction kernel+ """+ from task import input_t, output_t+ import torch+ import triton+ import triton.language as tl+++ @triton.jit+ def _bf16_to_fp4_hw(+ x_bf16,+ TILE_M: tl.constexpr,+ TILE_K: tl.constexpr,+ ):+ """Quantize BF16 block to MXFP4 using HW v_cvt_scalef32_pk_fp4_bf16."""+ FP4_GRP_SZ: tl.constexpr = 32+ NUM_QBLK: tl.constexpr = TILE_K // FP4_GRP_SZ++ # Compute scales from FP32 values+ x_fp32 = x_bf16.to(tl.float32).reshape(TILE_M, NUM_QBLK, FP4_GRP_SZ)++ 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_raw_i = log2_amax - 2+ scale_e8m0_raw_i = tl.minimum(tl.maximum(scale_e8m0_raw_i, -127), 127)+ bs_e8m0 = scale_e8m0_raw_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_raw_i.to(tl.int32) + 127).to(tl.uint32) << 23+ hw_scale = hw_scale_bits.to(tl.float32, bitcast=True) # [M, NUM_QBLK, 1]++ # Broadcast scale to per-pair granularity+ hw_scale_flat = tl.broadcast_to(hw_scale, (TILE_M, NUM_QBLK, FP4_GRP_SZ))+ hw_scale_flat = hw_scale_flat.reshape(TILE_M, TILE_K)+ # Take scale for even element of each pair (both share same scale within 32-group)+ hw_scale_pairs = hw_scale_flat.reshape(TILE_M, TILE_K // 2, 2)+ hw_scale_even, _ = tl.split(hw_scale_pairs)+ hw_scale_pair = hw_scale_even.reshape(TILE_M, TILE_K // 2)++ # Pack BF16 pairs into uint32 for HW instruction+ x_u16 = x_bf16.to(tl.uint16, bitcast=True).reshape(TILE_M, TILE_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(TILE_M, TILE_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(TILE_M, TILE_K // 2)++ return x_fp4, bs_e8m0.reshape(TILE_M, NUM_QBLK)+++ @triton.jit+ def _quant_only_launcher(+ a_ptr, a_fp4_ptr, a_scale_ptr,+ M, K,+ stride_am, stride_ak,+ stride_qm, stride_qk,+ stride_sm, stride_sk,+ TILE_M: tl.constexpr,+ TILE_K: tl.constexpr,+ ):+ pid_m = tl.program_id(0)+ pid_k = tl.program_id(1)+ offs_m = pid_m * TILE_M + tl.arange(0, TILE_M)+ offs_k = pid_k * TILE_K + tl.arange(0, TILE_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 = _bf16_to_fp4_hw(a_bf16, TILE_M, TILE_K)+ HALF_K: tl.constexpr = TILE_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 = TILE_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["TILE_K"] // 2) == 0)+ and (args["KSPLIT_TILE"] % args["TILE_K"] == 0)+ and (args["K"] % (args["KSPLIT_TILE"] // 2) == 0),+ }+ )+ @triton.jit+ def _quant_gemm_fused_launcher(+ 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,+ TILE_M: tl.constexpr,+ TILE_N: tl.constexpr,+ TILE_K: tl.constexpr,+ M_CLUSTER: tl.constexpr,+ NUM_KSPLIT: tl.constexpr,+ KSPLIT_TILE: 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_GRP_SZ: tl.constexpr = 32+ num_pid_m = tl.cdiv(M, TILE_M)+ num_pid_n = tl.cdiv(N, TILE_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 = M_CLUSTER * num_pid_n+ group_id = pid // num_pid_in_group+ first_pid_m = group_id * M_CLUSTER+ group_size_m = min(num_pid_m - first_pid_m, M_CLUSTER)+ 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 * KSPLIT_TILE // 2) < K:+ num_k_iter = tl.cdiv(KSPLIT_TILE // 2, TILE_K // 2)++ # A: BF16 [M, 2*K]+ offs_am = (pid_m * TILE_M + tl.arange(0, TILE_M)) % M+ offs_ak = pid_k * KSPLIT_TILE + tl.arange(0, TILE_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, (TILE_K // 2) * 16)+ offs_k_shuffle = pid_k * (KSPLIT_TILE // 2) * 16 + offs_k_shuffle_arr+ offs_bn = (pid_n * (TILE_N // 16) + tl.arange(0, TILE_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 * TILE_N + group_offset * 32+ offs_bsn = (pid_n * TILE_N + tl.arange(0, TILE_N // 32) * 32)+ offs_ks = (pid_k * (KSPLIT_TILE // SCALE_GRP_SZ) * 32) + tl.arange(+ 0, TILE_K // SCALE_GRP_SZ * 32+ )+ b_scale_ptrs = (+ b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk+ )++ acc = tl.zeros((TILE_M, TILE_N), dtype=tl.float32)++ for k_iter in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):+ # Fire all loads first for better memory-level parallelism+ if EVEN_K:+ a_bf16 = tl.load(a_ptrs)+ b_scales_raw = tl.load(b_scale_ptrs, cache_modifier=cache_modifier)+ b_raw = tl.load(b_ptrs, cache_modifier=cache_modifier)+ else:+ k_offset = (k_iter - pid_k * num_k_iter) * TILE_K+ a_bf16 = tl.load(+ a_ptrs,+ mask=tl.arange(0, TILE_K)[None, :] < (2 * K - pid_k * KSPLIT_TILE - k_offset),+ other=0.0,+ )+ b_scales_raw = tl.load(b_scale_ptrs, cache_modifier=cache_modifier)+ b_raw = tl.load(+ b_ptrs,+ mask=offs_k_shuffle_arr[None, :] < ((K - (pid_k * (KSPLIT_TILE // 2) + (k_iter - pid_k * num_k_iter) * (TILE_K // 2))) * 16),+ other=0,+ cache_modifier=cache_modifier,+ )++ # Quantize A in registers+ a_fp4, a_scales = _bf16_to_fp4_hw(a_bf16, TILE_M, TILE_K)++ # Unshuffle B scales+ b_scales = (+ b_scales_raw+ .reshape(+ TILE_N // 32,+ TILE_K // SCALE_GRP_SZ // 8,+ 4, 16, 2, 2, 1,+ )+ .permute(0, 5, 3, 1, 4, 2, 6)+ .reshape(TILE_N, TILE_K // SCALE_GRP_SZ)+ )++ # Unshuffle B data+ b = (+ b_raw.reshape(1, TILE_N // 16, TILE_K // 64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5)+ .reshape(TILE_N, TILE_K // 2)+ .trans(1, 0)+ )++ acc = tl.dot_scaled(+ a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc+ )++ a_ptrs += TILE_K * stride_ak+ b_ptrs += (TILE_K // 2) * 16 * stride_bk+ b_scale_ptrs += TILE_K * stride_bsk++ c = acc.to(c_ptr.type.element_ty)++ offs_cm = pid_m * TILE_M + tl.arange(0, TILE_M).to(tl.int64)+ offs_cn = pid_n * TILE_N + tl.arange(0, TILE_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 _splitk_sum_launcher(+ 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,+ TILE_M: tl.constexpr, TILE_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 * TILE_M + tl.arange(0, TILE_M)) % M+ offs_n = (pid_n * TILE_N + tl.arange(0, TILE_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["TILE_K"] // 2) == 0)+ and (args["KSPLIT_TILE"] % args["TILE_K"] == 0)+ and (args["K"] % (args["KSPLIT_TILE"] // 2) == 0),+ }+ )+ @triton.jit+ def _fp4_gemm_only_launcher(+ 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,+ TILE_M: tl.constexpr,+ TILE_N: tl.constexpr,+ TILE_K: tl.constexpr,+ M_CLUSTER: tl.constexpr,+ NUM_KSPLIT: tl.constexpr,+ KSPLIT_TILE: 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_GRP_SZ: tl.constexpr = 32+ HALF_TK: tl.constexpr = TILE_K // 2+ SCALE_TK: tl.constexpr = TILE_K // SCALE_GRP_SZ+ num_pid_m = tl.cdiv(M, TILE_M)+ num_pid_n = tl.cdiv(N, TILE_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 = M_CLUSTER * num_pid_n+ group_id = pid // num_pid_in_group+ first_pid_m = group_id * M_CLUSTER+ group_size_m = min(num_pid_m - first_pid_m, M_CLUSTER)+ 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 * KSPLIT_TILE // 2) < K:+ num_k_iter = tl.cdiv(KSPLIT_TILE // 2, HALF_TK)++ offs_am = (pid_m * TILE_M + tl.arange(0, TILE_M)) % M+ offs_aqk = pid_k * (KSPLIT_TILE // 2) + tl.arange(0, HALF_TK)+ a_fp4_ptrs = a_fp4_ptr + (offs_am[:, None] * stride_qm + offs_aqk[None, :] * stride_qk)++ offs_ask = pid_k * (KSPLIT_TILE // SCALE_GRP_SZ) + tl.arange(0, SCALE_TK)+ a_scale_ptrs = a_scale_ptr + (offs_am[:, None] * stride_sm + offs_ask[None, :] * stride_sk)++ offs_k_shuffle_arr = tl.arange(0, HALF_TK * 16)+ offs_k_shuffle = pid_k * (KSPLIT_TILE // 2) * 16 + offs_k_shuffle_arr+ offs_bn = (pid_n * (TILE_N // 16) + tl.arange(0, TILE_N // 16)) % (N // 16)+ b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)++ offs_bsn = (pid_n * TILE_N + tl.arange(0, TILE_N // 32) * 32)+ offs_ks = (pid_k * (KSPLIT_TILE // SCALE_GRP_SZ) * 32) + tl.arange(+ 0, SCALE_TK * 32+ )+ b_scale_ptrs = (+ b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk+ )++ acc = tl.zeros((TILE_M, TILE_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_TK+ k_remain = K - (pid_k * (KSPLIT_TILE // 2) + k_off)+ a_fp4 = tl.load(+ a_fp4_ptrs,+ mask=tl.arange(0, HALF_TK)[None, :] < k_remain,+ other=0,+ cache_modifier=cache_modifier,+ )+ s_remain = (2 * K) // SCALE_GRP_SZ - (pid_k * (KSPLIT_TILE // SCALE_GRP_SZ) + (k_iter - pid_k * num_k_iter) * SCALE_TK)+ a_scales = tl.load(+ a_scale_ptrs,+ mask=tl.arange(0, SCALE_TK)[None, :] < s_remain,+ other=0,+ cache_modifier=cache_modifier,+ )++ b_scales = (+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)+ .reshape(+ TILE_N // 32,+ SCALE_TK // 8,+ 4, 16, 2, 2, 1,+ )+ .permute(0, 5, 3, 1, 4, 2, 6)+ .reshape(TILE_N, SCALE_TK)+ )++ 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 * (KSPLIT_TILE // 2) + (k_iter - pid_k * num_k_iter) * HALF_TK)) * 16),+ other=0,+ cache_modifier=cache_modifier,+ )++ b = (+ b.reshape(1, TILE_N // 16, TILE_K // 64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5)+ .reshape(TILE_N, HALF_TK)+ .trans(1, 0)+ )++ acc = tl.dot_scaled(+ a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc+ )++ a_fp4_ptrs += HALF_TK * stride_qk+ a_scale_ptrs += SCALE_TK * stride_sk+ b_ptrs += HALF_TK * 16 * stride_bk+ b_scale_ptrs += TILE_K * stride_bsk++ c = acc.to(c_ptr.type.element_ty)++ offs_cm = pid_m * TILE_M + tl.arange(0, TILE_M).to(tl.int64)+ offs_cn = pid_n * TILE_N + tl.arange(0, TILE_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 _calc_splitk_params(K, TILE_K, NUM_KSPLIT):+ KSPLIT_TILE = (+ triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), TILE_K) * TILE_K+ )+ while NUM_KSPLIT > 1 and TILE_K > 16:+ if (+ K % (KSPLIT_TILE // 2) == 0+ and KSPLIT_TILE % TILE_K == 0+ and K % (TILE_K // 2) == 0+ ):+ break+ elif K % (KSPLIT_TILE // 2) != 0 and NUM_KSPLIT > 1:+ NUM_KSPLIT = NUM_KSPLIT // 2+ elif KSPLIT_TILE % TILE_K != 0:+ if NUM_KSPLIT > 1:+ NUM_KSPLIT = NUM_KSPLIT // 2+ elif TILE_K > 16:+ TILE_K = TILE_K // 2+ elif K % (TILE_K // 2) != 0 and TILE_K > 16:+ TILE_K = TILE_K // 2+ else:+ break+ KSPLIT_TILE = (+ triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), TILE_K) * TILE_K+ )+ NUM_KSPLIT = triton.cdiv(K, (KSPLIT_TILE // 2))+ return KSPLIT_TILE, TILE_K, NUM_KSPLIT+++ _TUNE_PARAMS = {+ # V10: Conservative tuning based on V8 baseline+ # M=4: Try larger TILE_N for better N parallelism+ (4, 2880, 512): {"TILE_M": 16, "TILE_N": 64, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 2, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ # M=16, K=7168: Keep Split-K=7, try num_stages=3+ (16, 2112, 7168): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 512, "M_CLUSTER": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 7},+ # M=32: Keep V8 config+ (32, 4096, 512): {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (32, 2880, 512): {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ # M=64: Try waves_per_eu=2+ (64, 7168, 2048): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ # M=256: Try waves_per_eu=3+ (256, 3072, 1536): {"TILE_M": 16, "TILE_N": 256, "TILE_K": 512, "M_CLUSTER": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ }++ _DEF_PARAMS = {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "M_CLUSTER": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}++ _FP4_GEMM_PARAMS = {+ (32, 4096, 512): {"TILE_M": 32, "TILE_N": 128, "TILE_K": 256, "M_CLUSTER": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (32, 2880, 512): {"TILE_M": 32, "TILE_N": 64, "TILE_K": 256, "M_CLUSTER": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (64, 7168, 2048): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 512, "M_CLUSTER": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 2},+ }++ _FP4_GEMM_DEF = {"TILE_M": 16, "TILE_N": 64, "TILE_K": 256, "M_CLUSTER": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}+++ _mem_pool = {}+ _param_pool = {}++ def _acquire_buffers(m, n, num_ksplit, device):+ key = (m, n, num_ksplit)+ if key not in _mem_pool:+ 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+ _mem_pool[key] = (y, y_pp)+ return _mem_pool[key]++ def _acquire_config(m, n, k):+ key = (m, n, k)+ if key not in _param_pool:+ config = _TUNE_PARAMS.get(key, _DEF_PARAMS).copy()+ K_packed = k // 2+ if config["NUM_KSPLIT"] > 1:+ KSPLIT_TILE, TILE_K, NUM_KSPLIT = _calc_splitk_params(+ K_packed, config["TILE_K"], config["NUM_KSPLIT"]+ )+ config["KSPLIT_TILE"] = KSPLIT_TILE+ config["TILE_K"] = TILE_K+ config["NUM_KSPLIT"] = NUM_KSPLIT+ else:+ config["KSPLIT_TILE"] = 2 * K_packed+ config["NUM_KSPLIT"] = 1+ if config["TILE_K"] >= 2 * K_packed:+ config["TILE_K"] = triton.next_power_of_2(2 * K_packed)+ config["KSPLIT_TILE"] = 2 * K_packed+ config["NUM_KSPLIT"] = 1+ config["TILE_N"] = max(config["TILE_N"], 32)+ _param_pool[key] = config+ return _param_pool[key]+++ _grid_pool = {}++ def _prepare_launch_meta(m, n, k, device):+ """Precompute ALL launch parameters once per shape."""+ config = _acquire_config(m, n, k)+ K_packed = k // 2+ ks = config["NUM_KSPLIT"]++ y, y_pp = _acquire_buffers(m, n, ks, device)++ grid = (ks * triton.cdiv(m, config["TILE_M"]) * triton.cdiv(n, config["TILE_N"]),)++ # Pre-store strides for y/y_pp+ if ks == 1:+ c_stride_k, c_stride_m, c_stride_n = 0, y.stride(0), y.stride(1)+ else:+ c_stride_k, c_stride_m, c_stride_n = y_pp.stride(0), y_pp.stride(1), y_pp.stride(2)++ meta = {+ 'config': config,+ 'K_packed': K_packed,+ 'grid': grid,+ 'ks': ks,+ 'c_stride_k': c_stride_k,+ 'c_stride_m': c_stride_m,+ 'c_stride_n': c_stride_n,+ }++ if ks > 1:+ meta['reduce_grid'] = (triton.cdiv(m, 16), triton.cdiv(n, 64))+ meta['actual_ksplit'] = triton.cdiv(K_packed, (config["KSPLIT_TILE"] // 2))+ meta['max_ksplit'] = triton.next_power_of_2(ks)++ return meta+++ def fused_quant_gemm(A_bf16, B_shuffle, B_scale_sh, m, n, k):+ key = (m, n, k)+ if key not in _grid_pool:+ _grid_pool[key] = _prepare_launch_meta(m, n, k, A_bf16.device)+ p = _grid_pool[key]+ y, y_pp = _acquire_buffers(m, n, p['ks'], A_bf16.device)++ b_reshaped = B_shuffle.view(torch.uint8).reshape(n // 16, p['K_packed'] * 16)+ b_scale_uint8 = B_scale_sh.view(torch.uint8)++ _quant_gemm_fused_launcher[p['grid']](+ A_bf16, b_reshaped,+ y if p['ks'] == 1 else y_pp,+ b_scale_uint8,+ m, n, p['K_packed'],+ A_bf16.stride(0), A_bf16.stride(1),+ b_reshaped.stride(0), b_reshaped.stride(1),+ p['c_stride_k'], p['c_stride_m'], p['c_stride_n'],+ b_scale_uint8.stride(0), b_scale_uint8.stride(1),+ **p['config'],+ )++ if p['ks'] > 1:+ _splitk_sum_launcher[p['reduce_grid']](+ 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,+ p['actual_ksplit'], p['max_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))+ _quant_only_launcher[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 = _FP4_GEMM_PARAMS.get((m, n, k), _FP4_GEMM_DEF).copy()++ if config["NUM_KSPLIT"] > 1:+ KSPLIT_TILE, TILE_K, NUM_KSPLIT = _calc_splitk_params(+ K_packed, config["TILE_K"], config["NUM_KSPLIT"]+ )+ config["KSPLIT_TILE"] = KSPLIT_TILE+ config["TILE_K"] = TILE_K+ config["NUM_KSPLIT"] = NUM_KSPLIT+ else:+ config["KSPLIT_TILE"] = 2 * K_packed+ config["NUM_KSPLIT"] = 1++ if config["TILE_K"] >= 2 * K_packed:+ config["TILE_K"] = triton.next_power_of_2(2 * K_packed)+ config["KSPLIT_TILE"] = 2 * K_packed+ config["NUM_KSPLIT"] = 1++ config["TILE_N"] = max(config["TILE_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["TILE_M"])+ * triton.cdiv(n, META["TILE_N"]),+ )++ _fp4_gemm_only_launcher[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_TILE_M = 16+ REDUCE_TILE_N = 64+ ACTUAL_KSPLIT = triton.cdiv(K_packed, (config["KSPLIT_TILE"] // 2))+ grid_reduce = (+ triton.cdiv(m, REDUCE_TILE_M),+ triton.cdiv(n, REDUCE_TILE_N),+ )+ _splitk_sum_launcher[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_TILE_M, REDUCE_TILE_N,+ ACTUAL_KSPLIT, triton.next_power_of_2(config["NUM_KSPLIT"]),+ )++ return y+++ def custom_kernel(data: input_t) -> output_t:+ A = data[0]+ return fused_quant_gemm(A, data[3], data[4], A.shape[0], data[1].shape[0], A.shape[1])No newline at end of file
scrolls · 813 diff lines total
Best evidence level for this revision: reported
JSON