submission 754587
ooousay · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1020 lines, June 9 Researcher Reciprocity License v1.0.
kernel.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754587?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:5ee9f51ffa298859b904a97ff234a1dabb7831a48ad90e943bf5534a9fd197fa
license declaredunknown
license concludedunknown
authorsooousay
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4, num_stages=2,split-k
"""M=16, N=2112, K=7168 ? custom hardcoded split-K kernel with native hw quant."""stages = 2
num_warps=4, num_stages=2,tile-k = 256
BLOCK_SIZE_K=256,tile-m = 4
BLOCK_SIZE_M=4, BLOCK_SIZE_N=128,tile-n = 128
BLOCK_SIZE_M=4, BLOCK_SIZE_N=128,Kernel source
kernel.py1020 lines
#!POPCORN leaderboard amd-mxfp4-mm
"""Auto-generated by build.py"""
# ============================================================
# m=4, n=2880, k=512
# ============================================================
"""M=4, N=2880, K=512 ? constexpr shapes (no binary patching)."""
import torch
import triton
import triton.language as tl
@triton.jit
def _mxfp4_quant_op_native(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
"""Native hw quant using v_cvt_scalef32_pk_fp4_bf16 -- shared by all shapes."""
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
# Scale computation
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
# Native convert+round+pack via hardware instruction
hw_scale = tl.exp2(scale_e8m0_unbiased)
x_bf16 = x.to(tl.bfloat16)
x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
x_lo, x_hi = tl.split(x_pairs)
x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF
x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16
packed_bf16x2 = x_lo_i32 | x_hi_i32
hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2))
hw_result = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
"=v,v,v",
[packed_bf16x2, hw_scale_bc],
dtype=tl.int32,
is_pure=True,
pack=1,
)
x_fp4 = (hw_result & 0xFF).to(tl.uint8)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _hardcoded_m4_kernel_4_2880_512(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
stride_am, stride_ak, stride_bn, stride_bk,
stride_cm, stride_cn, stride_bsn, stride_bsk,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_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,
):
M: tl.constexpr = 4
N: tl.constexpr = 2880
K: tl.constexpr = 256 # k // 2
SCALE_GROUP_SIZE: tl.constexpr = 32
NUM_PID_N: tl.constexpr = 23 # cdiv(2880, 128)
NUM_K_ITER: tl.constexpr = 2 # cdiv(512 // 2, 256 // 2)
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsk > 0)
tl.assume(stride_bsn > 0)
pid = tl.program_id(axis=0)
pid_m = pid // NUM_PID_N
pid_n = pid % NUM_PID_N
tl.assume(pid_m >= 0)
tl.assume(pid_n >= 0)
# A offsets (constant across k)
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
# B offsets (constant across k)
offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
# B scale offsets (constant across k)
offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(NUM_K_ITER):
# Recompute pointers each iteration from function args
# (enables ConvertToBufferOps ? buffer_load)
a_offs = (
offs_am[:, None] * stride_am
+ (k * BLOCK_SIZE_K + offs_k_bf16)[None, :] * stride_ak
)
b_offs = (
offs_bn[:, None] * stride_bn
+ (k * (BLOCK_SIZE_K // 2) * 16 + offs_k_shuffle_arr)[None, :] * stride_bk
)
bs_offs = (
offs_bsn[:, None] * stride_bsn
+ (k * BLOCK_SIZE_K + offs_ks)[None, :] * stride_bsk
)
b_scales = (
tl.load(b_scales_ptr + bs_offs, 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)
)
a_bf16 = tl.load(a_ptr + a_offs)
b = tl.load(b_ptr + b_offs, 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)
)
a, a_scales = _mxfp4_quant_op_native(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
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, :]
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
_state_4_2880_512 = None
def _run_4_2880_512(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
global _state_4_2880_512
if _state_4_2880_512 is None:
_state_4_2880_512 = {
'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
'B_w': None, 'B_sc': None, '_b_ptr': None,
}
b_ptr = B_shuffle.data_ptr()
if _state_4_2880_512['_b_ptr'] != b_ptr:
_state_4_2880_512['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
_state_4_2880_512['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
_state_4_2880_512['_b_ptr'] = b_ptr
_hardcoded_m4_kernel_4_2880_512[(23,)](
A, _state_4_2880_512['B_w'], _state_4_2880_512['out'], _state_4_2880_512['B_sc'],
A.stride(0), A.stride(1),
_state_4_2880_512['B_w'].stride(0), _state_4_2880_512['B_w'].stride(1),
_state_4_2880_512['out'].stride(0), _state_4_2880_512['out'].stride(1),
_state_4_2880_512['B_sc'].stride(0), _state_4_2880_512['B_sc'].stride(1),
BLOCK_SIZE_M=4, BLOCK_SIZE_N=128,
BLOCK_SIZE_K=256,
num_warps=4, num_stages=2,
waves_per_eu=0, matrix_instr_nonkdim=16,
cache_modifier=".cg",
)
return _state_4_2880_512['out']
# ============================================================
# m=16, n=2112, k=7168
# ============================================================
"""M=16, N=2112, K=7168 ? custom hardcoded split-K kernel with native hw quant."""
import torch
import triton
import triton.language as tl
from aiter.ops.triton.gluon.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel,
)
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
@triton.jit
def _mxfp4_quant_op_native_16_2112_7168(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
# Scale computation
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
# Native convert+round+pack via hardware instruction
hw_scale = tl.exp2(scale_e8m0_unbiased)
x_bf16 = x.to(tl.bfloat16)
x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
x_lo, x_hi = tl.split(x_pairs)
x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF
x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16
packed_bf16x2 = x_lo_i32 | x_hi_i32
hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2))
hw_result = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
"=v,v,v",
[packed_bf16x2, hw_scale_bc],
dtype=tl.int32,
is_pure=True,
pack=1,
)
x_fp4 = (hw_result & 0xFF).to(tl.uint8)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _hardcoded_preshuffle_splitk_kernel_16_2112_7168(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
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,
NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
# Hardcoded shape constants
M: tl.constexpr = 16
N: tl.constexpr = 2112
K: tl.constexpr = 3584 # k // 2
SCALE_GROUP_SIZE: tl.constexpr = 32
NUM_PID_N: tl.constexpr = 17 # cdiv(2112, 128)
NUM_K_ITER: tl.constexpr = 2 # cdiv(SPLITK_BLOCK_SIZE//2, BLOCK_SIZE_K//2)
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_ck > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsk > 0)
tl.assume(stride_bsn > 0)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
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:
# A offsets (constant across k)
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
# B offsets (constant across k)
offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
# B scale offsets (constant across k)
offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(pid_k * NUM_K_ITER, (pid_k + 1) * NUM_K_ITER):
# Recompute pointers each iteration from function args
# (enables ConvertToBufferOps)
a_offs = (
offs_am[:, None] * stride_am
+ (k * BLOCK_SIZE_K + offs_k_bf16)[None, :] * stride_ak
)
b_offs = (
offs_bn[:, None] * stride_bn
+ (k * (BLOCK_SIZE_K // 2) * 16 + offs_k_shuffle_arr)[None, :] * stride_bk
)
bs_offs = (
offs_bsn[:, None] * stride_bsn
+ (k * BLOCK_SIZE_K + offs_ks)[None, :] * stride_bsk
)
b_scales = (
tl.load(b_scales_ptr + bs_offs, 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)
)
a_bf16 = tl.load(a_ptr + a_offs)
b = tl.load(b_ptr + b_offs, 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)
)
a, a_scales = _mxfp4_quant_op_native_16_2112_7168(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
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)
_state_16_2112_7168 = None
def _run_16_2112_7168(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
global _state_16_2112_7168
if _state_16_2112_7168 is None:
K_kernel = k // 2
SPLITK_BLOCK_SIZE, BSK, actual_ksplit = get_splitk(K_kernel, 256, 14)
grid_size = actual_ksplit * triton.cdiv(m, 16) * triton.cdiv(n, 128)
y_pp = torch.empty((actual_ksplit, m, n), dtype=torch.float32, device=A.device)
_state_16_2112_7168 = {
'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
'B_w': None, 'B_sc': None, '_b_ptr': None,
'y_pp': y_pp,
'grid_size': grid_size,
'K_kernel': K_kernel,
'BSK': BSK,
'SPLITK_BLOCK_SIZE': SPLITK_BLOCK_SIZE,
'NUM_KSPLIT': actual_ksplit,
'ACTUAL_KSPLIT': triton.cdiv(K_kernel, (SPLITK_BLOCK_SIZE // 2)),
'MAX_KSPLIT': triton.next_power_of_2(actual_ksplit),
'reduce_grid': (triton.cdiv(m, 16), triton.cdiv(n, 64)),
}
s = _state_16_2112_7168
b_ptr = B_shuffle.data_ptr()
if s['_b_ptr'] != b_ptr:
s['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
s['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
s['_b_ptr'] = b_ptr
_hardcoded_preshuffle_splitk_kernel_16_2112_7168[(s['grid_size'],)](
A, s['B_w'], s['y_pp'], s['B_sc'],
A.stride(0), A.stride(1),
s['B_w'].stride(0), s['B_w'].stride(1),
s['y_pp'].stride(0), s['y_pp'].stride(1), s['y_pp'].stride(2),
s['B_sc'].stride(0), s['B_sc'].stride(1),
BLOCK_SIZE_M=16, BLOCK_SIZE_N=128,
BLOCK_SIZE_K=s['BSK'],
NUM_KSPLIT=s['NUM_KSPLIT'], SPLITK_BLOCK_SIZE=s['SPLITK_BLOCK_SIZE'],
num_warps=4, num_stages=2,
waves_per_eu=2, matrix_instr_nonkdim=16,
cache_modifier=".cg",
)
_gluon_reduce_kernel[s['reduce_grid']](
s['y_pp'], s['out'],
m, n,
s['y_pp'].stride(0), s['y_pp'].stride(1), s['y_pp'].stride(2),
s['out'].stride(0), s['out'].stride(1),
16, 64,
s['ACTUAL_KSPLIT'], s['MAX_KSPLIT'],
)
return s['out']
# ============================================================
# m=32, n=4096, k=512
# ============================================================
"""M=32, N=4096, K=512 ? v13: custom hardcoded preshuffle, all constexpr."""
import torch
import triton
import triton.language as tl
from aiter import dtypes
@triton.jit
def _hardcoded_preshuffle_kernel_32_4096_512(
a_ptr,
b_ptr,
c_ptr,
b_scales_ptr,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_cm,
stride_cn,
stride_bsn,
stride_bsk,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_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,
):
# Shape constants ? M=32, N=4096, K=512
M: tl.constexpr = 32
N: tl.constexpr = 4096
K: tl.constexpr = 256 # K_kernel = 512 // 2
SCALE_GROUP_SIZE: tl.constexpr = 32
# Grid: 4 x 32 = 128 WGs
NUM_PID_M: tl.constexpr = 4 # cdiv(32, 8)
NUM_PID_N: tl.constexpr = 32 # cdiv(4096, 128)
NUM_K_ITER: tl.constexpr = 2 # (512//2) / (256//2)
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsk > 0)
tl.assume(stride_bsn > 0)
pid = tl.program_id(axis=0)
pid_m = pid // NUM_PID_N
pid_n = pid % NUM_PID_N
tl.assume(pid_m >= 0)
tl.assume(pid_n >= 0)
# A pointers ? no mask (32%8=0)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
)
# B pointers (preshuffled) ? no mask (4096%128=0)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
b_ptrs = b_ptr + (
offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk
)
# B scale pointers
offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
offs_ks = 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 in range(NUM_K_ITER):
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)
)
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, 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)
)
a, a_scales = _mxfp4_quant_op_native(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
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, :]
)
tl.store(c_ptrs, c)
_state_32_4096_512 = None
def _run_32_4096_512(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
global _state_32_4096_512
if _state_32_4096_512 is None:
_state_32_4096_512 = {
'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
'B_w': None, 'B_sc': None, '_b_ptr': None,
}
b_ptr = B_shuffle.data_ptr()
if _state_32_4096_512['_b_ptr'] != b_ptr:
_state_32_4096_512['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
_state_32_4096_512['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
_state_32_4096_512['_b_ptr'] = b_ptr
_hardcoded_preshuffle_kernel_32_4096_512[(128,)](
A, _state_32_4096_512['B_w'], _state_32_4096_512['out'], _state_32_4096_512['B_sc'],
A.stride(0), A.stride(1),
_state_32_4096_512['B_w'].stride(0), _state_32_4096_512['B_w'].stride(1),
_state_32_4096_512['out'].stride(0), _state_32_4096_512['out'].stride(1),
_state_32_4096_512['B_sc'].stride(0), _state_32_4096_512['B_sc'].stride(1),
BLOCK_SIZE_M=8, BLOCK_SIZE_N=128,
BLOCK_SIZE_K=256,
num_warps=4, num_stages=2,
waves_per_eu=2, matrix_instr_nonkdim=16,
cache_modifier=".cg",
)
return _state_32_4096_512['out']
# ============================================================
# m=32, n=2880, k=512
# ============================================================
"""M=32, N=2880, K=512 ? fused_direct path. BSM=8 BSN=128 BSK=256."""
import torch
import triton
from aiter import dtypes
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
)
_state_32_2880_512 = None
def _run_32_2880_512(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
global _state_32_2880_512
if _state_32_2880_512 is None:
_state_32_2880_512 = {
'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
'B_w': None, 'B_sc': None, '_b_ptr': None,
'grid_size': triton.cdiv(m, 8) * triton.cdiv(n, 128),
'K_kernel': k // 2,
'SPLITK_BLOCK_SIZE': k,
}
b_ptr = B_shuffle.data_ptr()
if _state_32_2880_512['_b_ptr'] != b_ptr:
_state_32_2880_512['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
_state_32_2880_512['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
_state_32_2880_512['_b_ptr'] = b_ptr
_gemm_a16wfp4_preshuffle_kernel[(_state_32_2880_512['grid_size'],)](
A, _state_32_2880_512['B_w'], _state_32_2880_512['out'], _state_32_2880_512['B_sc'],
m, n, _state_32_2880_512['K_kernel'],
A.stride(0), A.stride(1),
_state_32_2880_512['B_w'].stride(0), _state_32_2880_512['B_w'].stride(1),
0, _state_32_2880_512['out'].stride(0), _state_32_2880_512['out'].stride(1),
_state_32_2880_512['B_sc'].stride(0), _state_32_2880_512['B_sc'].stride(1),
BLOCK_SIZE_M=8, BLOCK_SIZE_N=128,
BLOCK_SIZE_K=256, GROUP_SIZE_M=1,
NUM_KSPLIT=1, SPLITK_BLOCK_SIZE=_state_32_2880_512['SPLITK_BLOCK_SIZE'],
num_warps=4, num_stages=2,
waves_per_eu=2, matrix_instr_nonkdim=16,
PREQUANT=True, cache_modifier=None,
)
return _state_32_2880_512['out']
# ============================================================
# m=64, n=7168, k=2048
# ============================================================
"""M=64, N=7168, K=2048 ? v14 with native v_cvt_scalef32_pk_fp4_bf16 quant."""
import torch
import triton
import triton.language as tl
from aiter import dtypes
@triton.jit
def _mxfp4_quant_op_native_64_7168_2048(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
# Scale computation ? identical to aiter's _mxfp4_quant_op
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
# Native convert+round+pack
hw_scale = tl.exp2(scale_e8m0_unbiased)
x_bf16 = x.to(tl.bfloat16)
x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
x_lo, x_hi = tl.split(x_pairs)
x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF
x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16
packed_bf16x2 = x_lo_i32 | x_hi_i32
hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2))
hw_result = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
"=v,v,v",
[packed_bf16x2, hw_scale_bc],
dtype=tl.int32,
is_pure=True,
pack=1,
)
x_fp4 = (hw_result & 0xFF).to(tl.uint8)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _hardcoded_preshuffle_kernel_64_7168_2048(
a_ptr,
b_ptr,
c_ptr,
b_scales_ptr,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_cm,
stride_cn,
stride_bsn,
stride_bsk,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_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,
):
M: tl.constexpr = 64
N: tl.constexpr = 7168
K: tl.constexpr = 1024
SCALE_GROUP_SIZE: tl.constexpr = 32
SPLITK_BLOCK_SIZE: tl.constexpr = 2048
NUM_PID_M: tl.constexpr = 4
NUM_PID_N: tl.constexpr = 56
NUM_K_ITER: tl.constexpr = 8
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsk > 0)
tl.assume(stride_bsn > 0)
pid = tl.program_id(axis=0)
pid_m = pid // NUM_PID_N
pid_n = pid % NUM_PID_N
tl.assume(pid_m >= 0)
tl.assume(pid_n >= 0)
# A: row offsets (constant across k)
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
# B: n-group offsets (constant across k)
offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
# B scales: n-group offsets (constant across k)
offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(NUM_K_ITER):
# Recompute pointers each iteration: splat(func_arg) + tensor_offset
# Enables ConvertToBufferOps to match and emit buffer_load
a_offs = (
offs_am[:, None] * stride_am
+ (k * BLOCK_SIZE_K + offs_k_bf16)[None, :] * stride_ak
)
b_offs = (
offs_bn[:, None] * stride_bn
+ (k * (BLOCK_SIZE_K // 2) * 16 + offs_k_shuffle_arr)[None, :] * stride_bk
)
bs_offs = (
offs_bsn[:, None] * stride_bsn
+ (k * BLOCK_SIZE_K + offs_ks)[None, :] * stride_bsk
)
b_scales = (
tl.load(b_scales_ptr + bs_offs, 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)
)
a_bf16 = tl.load(a_ptr + a_offs)
b = tl.load(b_ptr + b_offs, 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)
)
a, a_scales = _mxfp4_quant_op_native_64_7168_2048(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
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, :]
)
tl.store(c_ptrs, c)
_state_64_7168_2048 = None
def _run_64_7168_2048(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
global _state_64_7168_2048
if _state_64_7168_2048 is None:
_state_64_7168_2048 = {
'out': torch.empty((m, n), dtype=torch.bfloat16, device=A.device),
'B_w': None, 'B_sc': None, '_b_ptr': None,
}
b_ptr = B_shuffle.data_ptr()
if _state_64_7168_2048['_b_ptr'] != b_ptr:
_state_64_7168_2048['B_w'] = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
_state_64_7168_2048['B_sc'] = B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
_state_64_7168_2048['_b_ptr'] = b_ptr
_hardcoded_preshuffle_kernel_64_7168_2048[(224,)](
A, _state_64_7168_2048['B_w'], _state_64_7168_2048['out'], _state_64_7168_2048['B_sc'],
A.stride(0), A.stride(1),
_state_64_7168_2048['B_w'].stride(0), _state_64_7168_2048['B_w'].stride(1),
_state_64_7168_2048['out'].stride(0), _state_64_7168_2048['out'].stride(1),
_state_64_7168_2048['B_sc'].stride(0), _state_64_7168_2048['B_sc'].stride(1),
BLOCK_SIZE_M=16, BLOCK_SIZE_N=128,
BLOCK_SIZE_K=256,
num_warps=4, num_stages=2,
waves_per_eu=2, matrix_instr_nonkdim=16,
cache_modifier=".cg",
)
return _state_64_7168_2048['out']
# ============================================================
# m=256, n=3072, k=1536
# ============================================================
"""M=256, N=3072, K=1536 ? native v_cvt_scalef32_pk_fp4_bf16 quant."""
import torch
import triton
import triton.language as tl
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
ASM_KERNEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
MXFP4_QUANT_BLOCK_SIZE = 32
_state_256_3072_1536 = None
@triton.jit
def _mxfp4_quant_op_native_256_3072_1536(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
hw_scale = tl.exp2(scale_e8m0_unbiased)
x_bf16 = x.to(tl.bfloat16)
x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
x_lo, x_hi = tl.split(x_pairs)
x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF
x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16
packed_bf16x2 = x_lo_i32 | x_hi_i32
hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2))
hw_result = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
"=v,v,v",
[packed_bf16x2, hw_scale_bc],
dtype=tl.int32,
is_pure=True,
pack=1,
)
x_fp4 = (hw_result & 0xFF).to(tl.uint8)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.heuristics({
"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
})
@triton.jit
def _fused_mxfp4_quant_shuffle_kernel_256_3072_1536(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m_in, stride_x_n_in, stride_x_fp4_m_in, stride_x_fp4_n_in,
M, N,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr, EVEN_M_N: tl.constexpr,
SCALING_MODE: tl.constexpr, SCALE_N_PAD: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op_native_256_3072_1536(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
if EVEN_M_N:
tl.store(x_fp4_ptr + out_offs, out_tensor, cache_modifier=".cg")
else:
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, cache_modifier=".cg")
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
num_bs_cols = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
bs_offs_0 = bs_offs_m[:, None] // 32
bs_offs_1 = bs_offs_m[:, None] % 32
bs_offs_2 = bs_offs_1 % 16
bs_offs_1 = bs_offs_1 // 16
bs_offs_3 = bs_offs_n[None, :] // 8
bs_offs_4 = bs_offs_n[None, :] % 8
bs_offs_5 = bs_offs_4 % 4
bs_offs_4 = bs_offs_4 // 4
bs_offs = (bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 2 * 2
+ bs_offs_5 * 2 * 2 * 16 + bs_offs_3 * 2 * 2 * 16 * 4
+ bs_offs_0 * 2 * 16 * SCALE_N_PAD)
bs_mask_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < num_bs_cols)[None, :]
bs_e8m0 = tl.where(bs_mask_valid, bs_e8m0, 127)
SCALE_M_PAD = (M + 255) // 256 * 256
bs_mask = (bs_offs_m < SCALE_M_PAD)[:, None] & (bs_offs_n < SCALE_N_PAD)[None, :]
tl.store(bs_ptr + bs_offs, bs_e8m0.to(tl.uint8), mask=bs_mask, cache_modifier=".cg")
def _run_256_3072_1536(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
global _state_256_3072_1536
if _state_256_3072_1536 is None:
SCALE_N_valid = triton.cdiv(k, MXFP4_QUANT_BLOCK_SIZE)
SCALE_M = triton.cdiv(m, 256) * 256
SCALE_N = triton.cdiv(SCALE_N_valid, 8) * 8
BLOCK_SIZE_M = triton.cdiv(min(32, triton.next_power_of_2(m)), 32) * 32
BLOCK_SIZE_N = 64
grid = (triton.cdiv(m, BLOCK_SIZE_M), triton.cdiv(k, BLOCK_SIZE_N * 1))
padded_M = (m + 31) // 32 * 32
_state_256_3072_1536 = {
'x_fp4': torch.empty((m, k // 2), dtype=torch.uint8, device=A.device),
'blockscale': torch.empty((SCALE_M, SCALE_N), dtype=torch.uint8, device=A.device),
'gemm_out': torch.empty((padded_M, n), dtype=torch.bfloat16, device=A.device),
'SCALE_N': SCALE_N,
'BLOCK_SIZE_M': BLOCK_SIZE_M,
'BLOCK_SIZE_N': BLOCK_SIZE_N,
'grid': grid,
}
s = _state_256_3072_1536
_fused_mxfp4_quant_shuffle_kernel_256_3072_1536[s['grid']](
A, s['x_fp4'], s['blockscale'],
*A.stride(), *s['x_fp4'].stride(),
M=m, N=k,
BLOCK_SIZE_M=s['BLOCK_SIZE_M'], BLOCK_SIZE_N=s['BLOCK_SIZE_N'],
NUM_ITER=1, NUM_STAGES=1,
MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0, SCALE_N_PAD=s['SCALE_N'],
num_warps=2, waves_per_eu=0, num_stages=1,
)
gemm_a4w4_asm(
s['x_fp4'].view(dtypes.fp4x2), B_shuffle,
s['blockscale'].view(dtypes.fp8_e8m0), B_scale_sh,
s['gemm_out'], ASM_KERNEL, None, 1.0, 0.0, True, log2_k_split=0,
)
return s['gemm_out'][:m]
import aiter as _aiter
from aiter import dtypes as _dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant as _dq
from aiter.utility.fp4_utils import e8m0_shuffle as _es
def _run_default(A, B, B_q, B_shuffle, B_scale_sh, m, n, k):
A = A.contiguous()
x_fp4, bs = _dq(A)
bs = _es(bs)
return _aiter.gemm_a4w4(x_fp4.view(_dtypes.fp4x2), B_shuffle, bs.view(_dtypes.fp8_e8m0), B_scale_sh, dtype=_dtypes.bf16, bpreshuffle=True)
from task import input_t, output_t
_DISPATCH = {
(4, 2880, 512): _run_4_2880_512,
(16, 2112, 7168): _run_16_2112_7168,
(32, 4096, 512): _run_32_4096_512,
(32, 2880, 512): _run_32_2880_512,
(64, 7168, 2048): _run_64_7168_2048,
(256, 3072, 1536): _run_256_3072_1536,
}
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
fn = _DISPATCH.get((m, n, k), _run_default)
return fn(A, B, B_q, B_shuffle, B_scale_sh, m, n, k)
scrolls · 1020 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON