submission 563736
Danishlynx · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 274 lines, June 9 Researcher Reciprocity License v1.0.
submission_gemm_v58.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-563736?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:56f7e3108f5c22312c88f002dc4396248cdeccf23595ada55fb3500a86478d12
license declaredunknown
license concludedunknown
authorsDanishlynx
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
fp4 = evens | (odds << 4)num-warps = 4
num_warps=4, num_stages=1,stages = 1
num_warps=4, num_stages=1,tile-n = 128
BLOCK_N = 128Kernel source
submission_gemm_v58.py274 lines
# /// script
# requires-python = ">=3.9"
# dependencies = []
# ///
# leaderboard = "amd-mxfp4-mm"
"""
v58: Hybrid, tuned block sizes for both paths.
- K<=1024: Fused quant+GEMM, load B_q directly + shuffled scale indexing
- K>1024: Fused quant+shuffle + gemm_a4w4
No B transpose, no inverse shuffle. Everything computed in-kernel.
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes as _dt
_FP4X2 = _dt.fp4x2
_E8M0 = _dt.fp8_e8m0
_BF16 = _dt.bf16
_cache = {}
@triton.jit
def _mxfp4_quant_op(x, BLOCK_K: tl.constexpr, BLOCK_M: tl.constexpr):
EXP_BIAS_FP32: tl.constexpr = 127
EXP_BIAS_FP4: tl.constexpr = 1
MBITS_F32: tl.constexpr = 23
MBITS_FP4: tl.constexpr = 1
EBITS_F32: tl.constexpr = 8
EBITS_FP4: tl.constexpr = 2
max_normal: tl.constexpr = 6
min_normal: tl.constexpr = 1
QUANT: tl.constexpr = 32
NUM_QB: tl.constexpr = BLOCK_K // QUANT
x = x.reshape(BLOCK_M, NUM_QB, QUANT)
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_unb = tl.log2(amax).floor() - 2
scale_unb = tl.clamp(scale_unb, min=-127, max=127)
bs = scale_unb.to(tl.uint8) + 127
qscale = tl.exp2(-scale_unb)
qx = x * qscale
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.to(tl.int32)
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
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 = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e2m1 = tl.where(normal_mask, normal_x, e2m1)
e2m1 = tl.where(denormal_mask, denormal_x, e2m1)
sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
sign_lp = sign_lp.to(tl.uint8)
e2m1 = e2m1 | sign_lp
e2m1 = tl.reshape(e2m1, [BLOCK_M, NUM_QB, QUANT // 2, 2])
evens, odds = tl.split(e2m1)
fp4 = evens | (odds << 4)
fp4 = fp4.reshape(BLOCK_M, BLOCK_K // 2)
return fp4, bs.reshape(BLOCK_M, NUM_QB)
@triton.jit
def _fused_quant_gemm_kernel(
A_ptr, Bq_ptr, Bscale_sh_ptr, C_ptr,
M, N, K: tl.constexpr,
stride_a_m, stride_a_k,
stride_bq_n, stride_bq_k,
SN_DIV8_MUL256: tl.constexpr,
stride_c_m, stride_c_n,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
QUANT: tl.constexpr = 32
NSK: tl.constexpr = BLOCK_K // QUANT
for ki in tl.range(0, K, BLOCK_K):
# Load A tile (bf16) and quantize to FP4
a_offs_k = ki + tl.arange(0, BLOCK_K)
a_ptrs = A_ptr + offs_m[:, None] * stride_a_m + a_offs_k[None, :] * stride_a_k
a_mask = (offs_m < M)[:, None] & (a_offs_k < K)[None, :]
a_tile = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)
a_fp4, a_scale = _mxfp4_quant_op(a_tile, BLOCK_K, BLOCK_M)
# Load B tile directly from B_q (N, K//2) — no transpose needed
# We need (BLOCK_K//2, BLOCK_N) for rhs of dot_scaled
# B_q[n, k_half] is the FP4 packed pair at position (n, 2*k_half)
b_offs_k = ki // 2 + tl.arange(0, BLOCK_K // 2)
# Load with transposed access: iterate K dim (inner) x N dim (outer)
b_ptrs = Bq_ptr + offs_n[None, :] * stride_bq_n + b_offs_k[:, None] * stride_bq_k
b_mask = (offs_n < N)[None, :] & (b_offs_k < K // 2)[:, None]
b_tile = tl.load(b_ptrs, mask=b_mask, other=0) # (BLOCK_K//2, BLOCK_N)
# Load B scales from SHUFFLED tensor using inverse shuffle indexing
# We need (BLOCK_N, NSK) scale values
# Raw scale at (row=n, col=ki//32+j) maps to shuffled index:
bs_row = offs_n # N dimension
bs_col_base = ki // QUANT
bs_col_offs = tl.arange(0, NSK)
row = bs_row[:, None]
col = (bs_col_base + bs_col_offs)[None, :]
shuf_idx = (row // 32) * SN_DIV8_MUL256 + (col // 8) * 256 + (col % 4) * 64 + (row % 16) * 4 + ((col % 8) // 4) * 2 + ((row % 32) // 16)
bs_mask = (offs_n[:, None] < N) & (bs_col_offs[None, :] < (K // QUANT - bs_col_base))
b_scale = tl.load(Bscale_sh_ptr + shuf_idx, mask=bs_mask, other=0)
acc = tl.dot_scaled(a_fp4, a_scale, "e2m1", b_tile, b_scale, "e2m1", acc)
c_ptrs = C_ptr + offs_m[:, None] * stride_c_m + offs_n[None, :] * stride_c_n
c_mask = (offs_m < M)[:, None] & (offs_n < N)[None, :]
tl.store(c_ptrs, acc.to(tl.bfloat16), mask=c_mask)
@triton.jit
def _fused_quant_shuffle_kernel(
x_ptr, x_fp4_ptr, bs_shuf_ptr,
stride_x_m, stride_x_n,
stride_fp4_m, stride_fp4_n,
M, N,
SN_DIV8_MUL256, SCALE_COLS,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
QUANT: tl.constexpr = 32
NUM_QB: tl.constexpr = BLOCK_SIZE_N // QUANT
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0, cache_modifier=".cg").to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = out_offs_m[:, None] * stride_fp4_m + out_offs_n[None, :] * stride_fp4_n
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
bs_row = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_col = pid_n * NUM_QB + tl.arange(0, NUM_QB)
row = bs_row[:, None]
col = bs_col[None, :]
shuf_idx = (row // 32) * SN_DIV8_MUL256 + (col // 8) * 256 + (col % 4) * 64 + (row % 16) * 4 + ((col % 8) // 4) * 2 + ((row % 32) // 16)
bs_mask = (bs_row[:, None] < M) & (bs_col[None, :] < SCALE_COLS)
tl.store(bs_shuf_ptr + shuf_idx, bs_e8m0, mask=bs_mask)
def _init(m, k, n, device):
QUANT = 32
scale_cols = (k + QUANT - 1) // QUANT
if k <= 1024:
# PATH 1: Fused quant+GEMM — NO B preprocessing
BLOCK_K = max(32, triton.next_power_of_2(k))
BLOCK_M = max(16, min(32, triton.next_power_of_2(m)))
BLOCK_N = 128
grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N))
out = torch.empty(m, n, dtype=torch.bfloat16, device=device)
# Compute SN_DIV8_MUL256 for shuffle indexing of B scales
sn = ((scale_cols + 7) // 8) * 8
sn_div8_mul256 = (sn // 8) * 256
return {
'mode': 'fused', 'out': out,
'grid': grid, 'BM': BLOCK_M, 'BN': BLOCK_N, 'BK': BLOCK_K,
'sn_div8_mul256': sn_div8_mul256,
}
else:
# PATH 2: Fused quant+shuffle + gemm_a4w4
sm = ((m + 255) // 256) * 256
sn = ((scale_cols + 7) // 8) * 8
x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
bs_shuffled = torch.empty(sm, sn, dtype=torch.uint8, device=device)
if m <= 32:
BSM = triton.next_power_of_2(m)
if k <= 2048:
# For moderate K, use larger BSN for fewer blocks but more work/block
NUM_ITER, BSN, NW, NS = 1, 128, 4, 1
else:
NUM_ITER, BSN, NW, NS = 1, 32, 1, 1
elif m <= 64:
# M=64: use smaller blocks for more parallelism
NUM_ITER, BSM, BSN, NW, NS = 2, 32, 64, 4, 2
else:
# M=256
NUM_ITER, BSM, BSN, NW, NS = 2, 32, 64, 4, 2
if k > 16384:
BSM, BSN = 64, 64
grid = (triton.cdiv(m, BSM), triton.cdiv(k, BSN * NUM_ITER))
return {
'mode': 'separate',
'x_fp4': x_fp4, 'bs_shuffled': bs_shuffled,
'sc': scale_cols, 'sn_div8_mul256': (sn // 8) * 256,
'grid': grid, 'BSM': BSM, 'BSN': BSN,
'NW': NW, 'NS': NS, 'NI': NUM_ITER,
}
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m = A.shape[0]
k = A.shape[1]
n = B_q.shape[0]
key = (m, k, n)
if key not in _cache:
_cache[key] = _init(m, k, n, A.device)
c = _cache[key]
if c['mode'] == 'fused':
# Pass B_q and B_scale_sh directly — no preprocessing!
Bq_uint8 = B_q.view(torch.uint8)
Bscale_uint8 = B_scale_sh.view(torch.uint8)
_fused_quant_gemm_kernel[c['grid']](
A, Bq_uint8, Bscale_uint8, c['out'],
m, n, k,
A.stride(0), A.stride(1),
Bq_uint8.stride(0), Bq_uint8.stride(1),
c['sn_div8_mul256'],
c['out'].stride(0), c['out'].stride(1),
BLOCK_M=c['BM'], BLOCK_N=c['BN'], BLOCK_K=c['BK'],
num_warps=4, num_stages=1,
)
return c['out']
else:
x_fp4 = c['x_fp4']
bs_shuf = c['bs_shuffled']
_fused_quant_shuffle_kernel[c['grid']](
A, x_fp4, bs_shuf,
A.stride(0), A.stride(1),
x_fp4.stride(0), x_fp4.stride(1),
m, k,
c['sn_div8_mul256'], c['sc'],
BLOCK_SIZE_M=c['BSM'], BLOCK_SIZE_N=c['BSN'],
NUM_ITER=c['NI'], NUM_STAGES=c['NS'],
num_warps=c['NW'], waves_per_eu=0, num_stages=1,
)
return aiter.gemm_a4w4(
x_fp4.view(_FP4X2), B_shuffle,
bs_shuf.view(_E8M0), B_scale_sh,
dtype=_BF16, bpreshuffle=True,
)
scrolls · 274 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 518127.
+ # /// script+ # requires-python = ">=3.9"+ # dependencies = []+ # ///+ # leaderboard = "amd-mxfp4-mm"+"""- MXFP4 GEMM submission v1 - AITER a4w4 baseline.- Flow: bf16 A -> MXFP4 quant A -> gemm_a4w4(A_q, B_shuffled) -> bf16 C+ v58: Hybrid, tuned block sizes for both paths.+ - K<=1024: Fused quant+GEMM, load B_q directly + shuffled scale indexing+ - K>1024: Fused quant+shuffle + gemm_a4w4+ No B transpose, no inverse shuffle. Everything computed in-kernel."""from task import input_t, output_t+ import torch+ import triton+ import triton.language as tl+ import aiter+ from aiter import dtypes as _dt+ _FP4X2 = _dt.fp4x2+ _E8M0 = _dt.fp8_e8m0+ _BF16 = _dt.bf16+ _cache = {}- def custom_kernel(data: input_t) -> output_t:- import aiter- from aiter import QuantType, dtypes+ @triton.jit+ def _mxfp4_quant_op(x, BLOCK_K: tl.constexpr, BLOCK_M: tl.constexpr):+ EXP_BIAS_FP32: tl.constexpr = 127+ EXP_BIAS_FP4: tl.constexpr = 1+ MBITS_F32: tl.constexpr = 23+ MBITS_FP4: tl.constexpr = 1+ EBITS_F32: tl.constexpr = 8+ EBITS_FP4: tl.constexpr = 2+ max_normal: tl.constexpr = 6+ min_normal: tl.constexpr = 1+ QUANT: tl.constexpr = 32+ NUM_QB: tl.constexpr = BLOCK_K // QUANT+ x = x.reshape(BLOCK_M, NUM_QB, QUANT)+ 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_unb = tl.log2(amax).floor() - 2+ scale_unb = tl.clamp(scale_unb, min=-127, max=127)+ bs = scale_unb.to(tl.uint8) + 127+ qscale = tl.exp2(-scale_unb)+ qx = x * qscale+ 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.to(tl.int32)+ 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+ 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 = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)+ e2m1 = tl.where(normal_mask, normal_x, e2m1)+ e2m1 = tl.where(denormal_mask, denormal_x, e2m1)+ sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)+ sign_lp = sign_lp.to(tl.uint8)+ e2m1 = e2m1 | sign_lp+ e2m1 = tl.reshape(e2m1, [BLOCK_M, NUM_QB, QUANT // 2, 2])+ evens, odds = tl.split(e2m1)+ fp4 = evens | (odds << 4)+ fp4 = fp4.reshape(BLOCK_M, BLOCK_K // 2)+ return fp4, bs.reshape(BLOCK_M, NUM_QB)+++ @triton.jit+ def _fused_quant_gemm_kernel(+ A_ptr, Bq_ptr, Bscale_sh_ptr, C_ptr,+ M, N, K: tl.constexpr,+ stride_a_m, stride_a_k,+ stride_bq_n, stride_bq_k,+ SN_DIV8_MUL256: tl.constexpr,+ stride_c_m, stride_c_n,+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,+ ):+ pid_m = tl.program_id(0)+ pid_n = tl.program_id(1)+ offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)+ offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)+ acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)+ QUANT: tl.constexpr = 32+ NSK: tl.constexpr = BLOCK_K // QUANT++ for ki in tl.range(0, K, BLOCK_K):+ # Load A tile (bf16) and quantize to FP4+ a_offs_k = ki + tl.arange(0, BLOCK_K)+ a_ptrs = A_ptr + offs_m[:, None] * stride_a_m + a_offs_k[None, :] * stride_a_k+ a_mask = (offs_m < M)[:, None] & (a_offs_k < K)[None, :]+ a_tile = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)+ a_fp4, a_scale = _mxfp4_quant_op(a_tile, BLOCK_K, BLOCK_M)++ # Load B tile directly from B_q (N, K//2) — no transpose needed+ # We need (BLOCK_K//2, BLOCK_N) for rhs of dot_scaled+ # B_q[n, k_half] is the FP4 packed pair at position (n, 2*k_half)+ b_offs_k = ki // 2 + tl.arange(0, BLOCK_K // 2)+ # Load with transposed access: iterate K dim (inner) x N dim (outer)+ b_ptrs = Bq_ptr + offs_n[None, :] * stride_bq_n + b_offs_k[:, None] * stride_bq_k+ b_mask = (offs_n < N)[None, :] & (b_offs_k < K // 2)[:, None]+ b_tile = tl.load(b_ptrs, mask=b_mask, other=0) # (BLOCK_K//2, BLOCK_N)++ # Load B scales from SHUFFLED tensor using inverse shuffle indexing+ # We need (BLOCK_N, NSK) scale values+ # Raw scale at (row=n, col=ki//32+j) maps to shuffled index:+ bs_row = offs_n # N dimension+ bs_col_base = ki // QUANT+ bs_col_offs = tl.arange(0, NSK)+ row = bs_row[:, None]+ col = (bs_col_base + bs_col_offs)[None, :]+ shuf_idx = (row // 32) * SN_DIV8_MUL256 + (col // 8) * 256 + (col % 4) * 64 + (row % 16) * 4 + ((col % 8) // 4) * 2 + ((row % 32) // 16)+ bs_mask = (offs_n[:, None] < N) & (bs_col_offs[None, :] < (K // QUANT - bs_col_base))+ b_scale = tl.load(Bscale_sh_ptr + shuf_idx, mask=bs_mask, other=0)++ acc = tl.dot_scaled(a_fp4, a_scale, "e2m1", b_tile, b_scale, "e2m1", acc)++ c_ptrs = C_ptr + offs_m[:, None] * stride_c_m + offs_n[None, :] * stride_c_n+ c_mask = (offs_m < M)[:, None] & (offs_n < N)[None, :]+ tl.store(c_ptrs, acc.to(tl.bfloat16), mask=c_mask)+++ @triton.jit+ def _fused_quant_shuffle_kernel(+ x_ptr, x_fp4_ptr, bs_shuf_ptr,+ stride_x_m, stride_x_n,+ stride_fp4_m, stride_fp4_n,+ M, N,+ SN_DIV8_MUL256, SCALE_COLS,+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,+ NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,+ ):+ pid_m = tl.program_id(0)+ start_n = tl.program_id(1) * NUM_ITER+ QUANT: tl.constexpr = 32+ NUM_QB: tl.constexpr = BLOCK_SIZE_N // QUANT++ for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):+ x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)+ x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)+ x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n+ x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]+ x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0, cache_modifier=".cg").to(tl.float32)++ out_tensor, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M)++ out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)+ out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)+ out_offs = out_offs_m[:, None] * stride_fp4_m + out_offs_n[None, :] * stride_fp4_n+ out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]+ tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)++ bs_row = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)+ bs_col = pid_n * NUM_QB + tl.arange(0, NUM_QB)+ row = bs_row[:, None]+ col = bs_col[None, :]+ shuf_idx = (row // 32) * SN_DIV8_MUL256 + (col // 8) * 256 + (col % 4) * 64 + (row % 16) * 4 + ((col % 8) // 4) * 2 + ((row % 32) // 16)+ bs_mask = (bs_row[:, None] < M) & (bs_col[None, :] < SCALE_COLS)+ tl.store(bs_shuf_ptr + shuf_idx, bs_e8m0, mask=bs_mask)+++ def _init(m, k, n, device):+ QUANT = 32+ scale_cols = (k + QUANT - 1) // QUANT++ if k <= 1024:+ # PATH 1: Fused quant+GEMM — NO B preprocessing+ BLOCK_K = max(32, triton.next_power_of_2(k))+ BLOCK_M = max(16, min(32, triton.next_power_of_2(m)))+ BLOCK_N = 128+ grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N))+ out = torch.empty(m, n, dtype=torch.bfloat16, device=device)++ # Compute SN_DIV8_MUL256 for shuffle indexing of B scales+ sn = ((scale_cols + 7) // 8) * 8+ sn_div8_mul256 = (sn // 8) * 256++ return {+ 'mode': 'fused', 'out': out,+ 'grid': grid, 'BM': BLOCK_M, 'BN': BLOCK_N, 'BK': BLOCK_K,+ 'sn_div8_mul256': sn_div8_mul256,+ }+ else:+ # PATH 2: Fused quant+shuffle + gemm_a4w4+ sm = ((m + 255) // 256) * 256+ sn = ((scale_cols + 7) // 8) * 8+ x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=device)+ bs_shuffled = torch.empty(sm, sn, dtype=torch.uint8, device=device)++ if m <= 32:+ BSM = triton.next_power_of_2(m)+ if k <= 2048:+ # For moderate K, use larger BSN for fewer blocks but more work/block+ NUM_ITER, BSN, NW, NS = 1, 128, 4, 1+ else:+ NUM_ITER, BSN, NW, NS = 1, 32, 1, 1+ elif m <= 64:+ # M=64: use smaller blocks for more parallelism+ NUM_ITER, BSM, BSN, NW, NS = 2, 32, 64, 4, 2+ else:+ # M=256+ NUM_ITER, BSM, BSN, NW, NS = 2, 32, 64, 4, 2+ if k > 16384:+ BSM, BSN = 64, 64++ grid = (triton.cdiv(m, BSM), triton.cdiv(k, BSN * NUM_ITER))+ return {+ 'mode': 'separate',+ 'x_fp4': x_fp4, 'bs_shuffled': bs_shuffled,+ 'sc': scale_cols, 'sn_div8_mul256': (sn // 8) * 256,+ 'grid': grid, 'BSM': BSM, 'BSN': BSN,+ 'NW': NW, 'NS': NS, 'NI': NUM_ITER,+ }+++ def custom_kernel(data: input_t) -> output_t:A, B, B_q, B_shuffle, B_scale_sh = data- A = A.contiguous()- m, k = A.shape- n = B.shape[0]+ m = A.shape[0]+ k = A.shape[1]+ n = B_q.shape[0]- # Quantize A to MXFP4 with shuffle for CK kernel- quant_func = aiter.get_triton_quant(QuantType.per_1x32)- A_q, A_scale_sh = quant_func(A, shuffle=True)+ key = (m, k, n)+ if key not in _cache:+ _cache[key] = _init(m, k, n, A.device)+ c = _cache[key]- # GEMM via AITER CK kernel (pre-shuffled weights)- out_gemm = aiter.gemm_a4w4(- A_q,- B_shuffle,- A_scale_sh,- B_scale_sh,- dtype=dtypes.bf16,- bpreshuffle=True,- )- return out_gemm+ if c['mode'] == 'fused':+ # Pass B_q and B_scale_sh directly — no preprocessing!+ Bq_uint8 = B_q.view(torch.uint8)+ Bscale_uint8 = B_scale_sh.view(torch.uint8)++ _fused_quant_gemm_kernel[c['grid']](+ A, Bq_uint8, Bscale_uint8, c['out'],+ m, n, k,+ A.stride(0), A.stride(1),+ Bq_uint8.stride(0), Bq_uint8.stride(1),+ c['sn_div8_mul256'],+ c['out'].stride(0), c['out'].stride(1),+ BLOCK_M=c['BM'], BLOCK_N=c['BN'], BLOCK_K=c['BK'],+ num_warps=4, num_stages=1,+ )+ return c['out']+ else:+ x_fp4 = c['x_fp4']+ bs_shuf = c['bs_shuffled']+ _fused_quant_shuffle_kernel[c['grid']](+ A, x_fp4, bs_shuf,+ A.stride(0), A.stride(1),+ x_fp4.stride(0), x_fp4.stride(1),+ m, k,+ c['sn_div8_mul256'], c['sc'],+ BLOCK_SIZE_M=c['BSM'], BLOCK_SIZE_N=c['BSN'],+ NUM_ITER=c['NI'], NUM_STAGES=c['NS'],+ num_warps=c['NW'], waves_per_eu=0, num_stages=1,+ )+ return aiter.gemm_a4w4(+ x_fp4.view(_FP4X2), B_shuffle,+ bs_shuf.view(_E8M0), B_scale_sh,+ dtype=_BF16, bpreshuffle=True,+ )
scrolls · 294 diff lines total
Best evidence level for this revision: reported
JSON