submission 561029
kkosey · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 544 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-561029?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:abb28e16426ebc4332d35e0d1a9ae9dd8b67d787a573e8e138e431ed72ab5bb8
license declaredunknown
license concludedunknown
authorskkosey
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Optimized MXFP4 GEMM: Custom Triton kernel for small-K + ASM for large-K.num-warps = 1
NUM_WARPS = 1split-k
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitkstages = 1
num_stages = 1tile-m = 32
BLOCK_SIZE_M = 32tile-n = 32
BLOCK_SIZE_N = 32Kernel source
submission.py544 lines
"""
Optimized MXFP4 GEMM: Custom Triton kernel for small-K + ASM for large-K.
The Triton kernel fuses bf16→MXFP4 quant + GEMM + shuffled B_scale read.
MI355X: 256 CUs, 8 XCDs, gfx950, native FP4 MFMA 16x16.
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _get_config
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
try:
from aiter.ops.gemm_op_a4w4 import get_padded_m as _get_padded_m
except ImportError:
_get_padded_m = None
# Raw ASM dispatch: bypass torch.ops.aiter wrapper overhead (~0.4-0.6µs savings)
_raw_gemm_fn = None
def _init_raw_gemm():
global _raw_gemm_fn
if _raw_gemm_fn is not None:
return
try:
from aiter.jit.core import get_module
mod = get_module("module_gemm_a4w4_asm")
_raw_gemm_fn = mod.gemm_a4w4_asm
except Exception:
pass # .so not built yet, will use wrapper
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_bf16 = dtypes.bf16
def _asm_kernel_name(tile_m, tile_n):
name = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}"
return f"_ZN5aiter{len(name)}{name}E"
_DEFAULT_32x128 = _asm_kernel_name(32, 128)
# ---------------------------------------------------------------------------
# Custom GEMM kernel: bf16 A × fp4 B → bf16 C, reading SHUFFLED B_scale
# Based on aiter's _gemm_a16wfp4_kernel but with inline e8m0 unshuffle.
# Eliminates the need for a separate unshuffle kernel launch.
# ---------------------------------------------------------------------------
@triton.jit
def _gemm_a16wfp4_shuffled_scale_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K, # K = K_half (packed)
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
stride_ck, # 0 for no splitK, M*N for splitK (fp32 output)
# constexpr meta-parameters
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
SN_SCALE: tl.constexpr, # padded scale columns (sn from e8m0_shuffle)
EVEN_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,
):
"""GEMM C = A @ B^T with inline MXFP4 quant of A and shuffled B_scale read."""
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)
SCALE_GROUP_SIZE: tl.constexpr = 32
GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
pid_unified = tl.program_id(axis=0)
if NUM_KSPLIT > 1:
pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
else:
pid_k = 0
pid = pid_unified
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
tl.assume(pid_m >= 0)
tl.assume(pid_n >= 0)
tl.assume(pid_k >= 0)
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
# A pointers (bf16) — offset by splitK range
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
k_offset_bf16 = pid_k * SPLITK_BLOCK_SIZE # bf16 element offset
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + offs_am[:, None] * stride_am + (k_offset_bf16 + offs_k_bf16[None, :]) * stride_ak
# B pointers (fp4 packed as uint8) — offset by splitK range
offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
k_offset_packed = pid_k * (SPLITK_BLOCK_SIZE // 2) # packed byte offset
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
b_ptrs = b_ptr + (k_offset_packed + offs_k[:, None]) * stride_bk + offs_bn[None, :] * stride_bn
# Pre-compute row decomposition for shuffled B_scale access
bs_d0 = (offs_bn // 32)[:, None]
bs_d1 = ((offs_bn % 32) // 16)[:, None]
bs_d2 = (offs_bn % 16)[:, None]
# Scale column tracking — start at splitK offset
SCALES_PER_BLOCK: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
cur_scale_col = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)
scale_k_range = tl.arange(0, SCALES_PER_BLOCK)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(0, num_k_iter):
# Load B_scale from SHUFFLED layout (inline unshuffle)
cols = cur_scale_col + scale_k_range
bs_d3 = (cols // 8)[None, :]
bs_d4 = ((cols % 8) // 4)[None, :]
bs_d5 = (cols % 4)[None, :]
shuffled_idx = (bs_d0 * (32 * SN_SCALE) + bs_d3 * 256
+ bs_d5 * 64 + bs_d2 * 4 + bs_d4 * 2 + bs_d1)
b_scales = tl.load(b_scales_ptr + shuffled_idx)
# Load A (bf16) and B (fp4 packed)
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
a_bf16 = tl.load(
a_ptrs,
mask=offs_k_bf16[None, :] < 2 * K - (pid_k * num_k_iter + k) * BLOCK_SIZE_K,
other=0,
)
b = tl.load(
b_ptrs,
mask=offs_k[:, None] < K - (pid_k * num_k_iter + k) * (BLOCK_SIZE_K // 2),
other=0,
cache_modifier=cache_modifier,
)
# In-register quant: bf16 A → mxfp4 + e8m0 scale
a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
# Scaled dot product
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
# Advance pointers
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk
cur_scale_col += SCALES_PER_BLOCK
# Store output
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)
if NUM_KSPLIT == 1:
c = accumulator.to(c_ptr.type.element_ty)
else:
c = accumulator # keep fp32 for splitK
tl.store(c_ptrs, c, mask=c_mask)
@triton.jit
def _splitk_reduce_kernel(
y_pp_ptr, y_ptr,
M, N, NUM_KSPLIT: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
num_warps: tl.constexpr,
):
"""Reduce splitK partial results: y = sum(y_pp[k]) converted to bf16."""
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)
mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
base = offs_m[:, None] * N + offs_n[None, :]
for k in range(NUM_KSPLIT):
val = tl.load(y_pp_ptr + k * M * N + base, mask=mask, other=0.0)
acc += val
tl.store(y_ptr + base, acc.to(tl.bfloat16), mask=mask)
def _use_triton(M, K):
return (K <= 1024 and M <= 64) or M <= 16
_buf = {}
@triton.jit
def _fused_quant_shuffle_kernel(
x_ptr,
x_fp4_ptr,
shuffled_bs_ptr,
stride_x_m_in,
stride_x_n_in,
stride_x_fp4_m_in,
stride_x_fp4_n_in,
M: tl.constexpr,
N: tl.constexpr,
sn: tl.constexpr,
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,
):
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).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).to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(
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)
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)
row = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
col = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
d0 = (row // 32)[:, None]
d1 = ((row % 32) // 16)[:, None]
d2 = (row % 16)[:, None]
d3 = (col // 8)[None, :]
d4 = ((col % 8) // 4)[None, :]
d5 = (col % 4)[None, :]
flat_out = d0 * (32 * sn) + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1
if EVEN_M_N:
tl.store(shuffled_bs_ptr + flat_out, bs_e8m0)
else:
SCALE_N: tl.constexpr = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
bs_mask = (row[:, None] < M) & (col[None, :] < SCALE_N)
tl.store(shuffled_bs_ptr + flat_out, bs_e8m0, mask=bs_mask)
def _prepare_triton(M, N, K, device):
K_half = K // 2
SCALE_K = K // 32
sn = (SCALE_K + 7) // 8 * 8 # padded scale cols (for shuffle formula)
# Get config (same as wrapper: _get_config(M, N, K_half))
raw_config, _ = _get_config(M, N, K_half)
BSK = raw_config["BLOCK_SIZE_K"]
if BSK >= 2 * K_half:
BSK = triton.next_power_of_2(2 * K_half)
BSK = max(BSK, 64)
BSM = raw_config["BLOCK_SIZE_M"]
BSN = raw_config["BLOCK_SIZE_N"]
NW = raw_config["num_warps"]
# Determine splitK based on shape
grid_mn = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
if K > 1024 and grid_mn < 128:
# Large K, few MN blocks → use splitK
# Use BSM that matches M for maximum MFMA utilization
BSM = triton.next_power_of_2(min(M, 16))
BSN = 64 # larger N-tiles for better data reuse
BSK = 512 # larger K-tiles: fewer iterations, better register reuse
NW = 4
num_stages = 1
grid_mn = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
# Target ~2 waves (fewer tiles = larger work per tile = better efficiency)
target_blocks = 256 * 2
NUM_KSPLIT = max(1, min(target_blocks // max(grid_mn, 1), K_half // (BSK // 2)))
# Use get_splitk to ensure EVEN_K
SPLITK_BS, BSK, NUM_KSPLIT = get_splitk(K_half, BSK, NUM_KSPLIT)
EVEN_K = (K_half % (BSK // 2) == 0) and (SPLITK_BS % BSK == 0) and (K_half % (SPLITK_BS // 2) == 0)
else:
NUM_KSPLIT = 1
# Override: use BSN=64 for more grid parallelism
if BSN > 64:
BSN = 64
NW = max(2, NW // 2)
SPLITK_BS = 2 * K_half
EVEN_K = (K_half % (BSK // 2) == 0) and (SPLITK_BS % BSK == 0) and (K_half % (SPLITK_BS // 2) == 0)
num_stages = raw_config["num_stages"]
# Filter config to only keys our custom kernel accepts
config = {
"BLOCK_SIZE_M": BSM,
"BLOCK_SIZE_N": BSN,
"BLOCK_SIZE_K": BSK,
"GROUP_SIZE_M": raw_config["GROUP_SIZE_M"],
"num_warps": NW,
"num_stages": num_stages,
"waves_per_eu": 0 if NUM_KSPLIT > 1 else raw_config["waves_per_eu"],
"matrix_instr_nonkdim": raw_config["matrix_instr_nonkdim"],
"cache_modifier": raw_config.get("cache_modifier", ".cg"),
"NUM_KSPLIT": NUM_KSPLIT,
"SPLITK_BLOCK_SIZE": SPLITK_BS,
}
gemm_out = torch.empty((M, N), dtype=_bf16, device=device)
grid_size = NUM_KSPLIT * triton.cdiv(M, config["BLOCK_SIZE_M"]) * triton.cdiv(N, config["BLOCK_SIZE_N"])
# Allocate splitK intermediate buffer if needed
if NUM_KSPLIT > 1:
y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=device)
reduce_bm = min(M, 16)
reduce_bn = 64
reduce_grid = (triton.cdiv(M, reduce_bm), triton.cdiv(N, reduce_bn))
else:
y_pp = None
reduce_bm = 0
reduce_bn = 0
reduce_grid = None
# Pre-compute all stride/scalar values as Python ints for fast hot path
stride_am = K # bf16 elements per row
stride_ak = 1
stride_bk = 1 # B_q is K-contiguous (after .T)
stride_bn = K_half # B_q row stride (after .T)
stride_cm = N
stride_cn = 1
stride_ck = M * N if NUM_KSPLIT > 1 else 0
return (gemm_out, config, sn, EVEN_K, grid_size, K_half, y_pp,
NUM_KSPLIT, reduce_grid, reduce_bm, reduce_bn,
stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, stride_ck)
def _prepare_asm(M, N, K, device):
# Cache raw ASM function on first ASM call
_init_raw_gemm()
MXFP4_QUANT_BLOCK_SIZE = 32
SCALE_N = (K + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
sm = (M + 255) // 256 * 256
sn = (SCALE_N + 7) // 8 * 8
tile_m = 32
if _get_padded_m is not None:
padded_m = _get_padded_m(M, N, K, tile_m)
else:
padded_m = ((M + tile_m - 1) // tile_m) * tile_m
x_fp4 = torch.empty((padded_m, K // 2), dtype=torch.uint8, device=device)
shuffled_scale = torch.zeros(sm * sn, dtype=torch.uint8, device=device)
if M <= 32:
NUM_ITER = 1
BLOCK_SIZE_M = triton.next_power_of_2(M)
BLOCK_SIZE_N = 32
NUM_WARPS = 1
NUM_STAGES = 1
else:
NUM_ITER = 1
BLOCK_SIZE_M = 32
BLOCK_SIZE_N = 64 # reduced from 128 for better CU utilization
NUM_WARPS = 2
NUM_STAGES = 1
if K <= 1024:
NUM_ITER = 1
NUM_STAGES = 1
BLOCK_SIZE_N = 32
NUM_WARPS = 1
BLOCK_SIZE_M = min(32, triton.next_power_of_2(M))
EVEN_M_N = (M % BLOCK_SIZE_M == 0) and (K % (BLOCK_SIZE_N * NUM_ITER) == 0)
grid = (
triton.cdiv(M, BLOCK_SIZE_M),
triton.cdiv(K, BLOCK_SIZE_N * NUM_ITER),
)
gemm_out = torch.empty((padded_m, N), dtype=_bf16, device=device)
fp4_stride = x_fp4.stride()
kernel_name = _DEFAULT_32x128
# Determine splitK based on tile count (32×128 tiles)
tiles_m = (padded_m + 31) // 32
tiles_n = (N + 127) // 128
grid_mn = tiles_m * tiles_n
# Avoid log2_k_split=1 (known broken on MI355X)
# Only use splitK for large K with few tiles
if grid_mn < 32 and K >= 4096:
log2_k_split = 4 # 16-way K-split for very few tiles + large K
elif grid_mn < 128 and K >= 1024:
log2_k_split = 3 # 8-way K-split for low tile counts
else:
log2_k_split = 0 # no split
return (
x_fp4, shuffled_scale, grid,
BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_ITER, NUM_WARPS, NUM_STAGES,
EVEN_M_N, sn,
gemm_out, kernel_name, log2_k_split,
x_fp4.view(_fp4x2), shuffled_scale.view(sm, sn).view(_fp8_e8m0),
gemm_out[:M], fp4_stride[0], fp4_stride[1], K, M,
)
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[0]
key = (M, K, N)
use_t = _use_triton(M, K)
if key not in _buf:
if use_t:
_buf[key] = ('t',) + _prepare_triton(M, N, K, A.device)
else:
_buf[key] = ('a',) + _prepare_asm(M, N, K, A.device)
buf = _buf[key]
if buf[0] == 't':
(_, gemm_out, config, sn, EVEN_K, grid_size, K_half, y_pp,
NUM_KSPLIT, reduce_grid, reduce_bm, reduce_bn,
stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, stride_ck) = buf
# Recompute B views each call (cheap views, avoids stale B caching)
# Skip .T — strides are pre-cached, data_ptr is the same
b_uint8 = B_q.view(torch.uint8)
b_scale_raw = B_scale_sh.view(torch.uint8)
if NUM_KSPLIT > 1:
# SplitK: write fp32 partials to y_pp, then reduce
_gemm_a16wfp4_shuffled_scale_kernel[(grid_size,)](
A, b_uint8, y_pp, b_scale_raw,
M, N, K_half,
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
stride_ck,
SN_SCALE=sn,
EVEN_K=EVEN_K,
**config,
)
# Reduce: sum splitK partials → bf16 output
_splitk_reduce_kernel[reduce_grid](
y_pp, gemm_out, M, N,
NUM_KSPLIT=NUM_KSPLIT,
BLOCK_M=reduce_bm, BLOCK_N=reduce_bn,
num_warps=4,
)
else:
# No splitK: direct bf16 output
_gemm_a16wfp4_shuffled_scale_kernel[(grid_size,)](
A, b_uint8, gemm_out, b_scale_raw,
M, N, K_half,
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
0,
SN_SCALE=sn,
EVEN_K=EVEN_K,
**config,
)
return gemm_out
else:
(_, x_fp4, shuffled_scale, grid,
BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_ITER, NUM_WARPS, NUM_STAGES,
EVEN_M_N, sn,
gemm_out, kernel_name, log2_k_split,
A_q, A_scale, out_view, fp4_s0, fp4_s1, K_val, M_val) = buf
_fused_quant_shuffle_kernel[grid](
A, x_fp4, shuffled_scale,
K_val, 1, fp4_s0, fp4_s1,
M=M_val, N=K_val, sn=sn,
MXFP4_QUANT_BLOCK_SIZE=32,
NUM_ITER=NUM_ITER,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
NUM_STAGES=NUM_STAGES,
EVEN_M_N=EVEN_M_N,
num_warps=NUM_WARPS,
waves_per_eu=0, num_stages=1,
)
if _raw_gemm_fn is not None:
_raw_gemm_fn(
A_q, B_shuffle, A_scale, B_scale_sh,
gemm_out, kernel_name,
None, 1.0, 0.0, True, log2_k_split,
)
else:
gemm_a4w4_asm(
A_q, B_shuffle, A_scale, B_scale_sh,
gemm_out, kernel_name,
bpreshuffle=True, log2_k_split=log2_k_split,
)
_init_raw_gemm() # Cache for subsequent calls
return out_viewscrolls · 544 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