submission 521876
rishi048401 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 177 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-521876?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:c7e3ca40db8bdcfea3a59749ef3e696d11c865fca1c2480bda4eb8d6705c778f
license declaredunknown
license concludedunknown
authorsrishi048401
imported2026-08-26
Kernel source
submission.py177 lines
try:
from task import input_t, output_t
except ImportError:
input_t = output_t = any
import torch
import triton
import triton.language as tl
from aiter.utility import dtypes
import aiter
@triton.jit
def my_quant_kernel(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m, stride_x_n, stride_x_fp4_m, stride_x_fp4_n,
stride_bs_m, stride_bs_n,
M: tl.constexpr, N: tl.constexpr,
scaleN: tl.constexpr, scaleM_pad: tl.constexpr, scaleN_pad: tl.constexpr,
BLOCK_SIZE: tl.constexpr, MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
SCALING_MODE: tl.constexpr, SHUFFLE: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_x_m = tl.cast(stride_x_m, tl.int64)
stride_x_n = tl.cast(stride_x_n, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
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).to(tl.float32)
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)
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
e = (qx >> 23) & 0xFF
m = qx & 0x7FFFFF
E8_BIAS: tl.constexpr = 127
E2_BIAS: tl.constexpr = 1
adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
e2m1_value = tl.reshape(e2m1_value, [BLOCK_SIZE, MXFP4_QUANT_BLOCK_SIZE // 2, 2])
evens, odds = tl.split(e2m1_value)
out_tensor = evens | (odds << 4)
out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 2)
out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_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_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
bs_offs_n = pid_n
if SHUFFLE:
# bitwise ALUs for massive optimization since AMD emulates integer division
bs_offs_0 = bs_offs_m[:, None] >> 5 # // 32
bs_offs_1 = bs_offs_m[:, None] & 31 # % 32
bs_offs_2 = bs_offs_1 & 15 # % 16
bs_offs_1 = bs_offs_1 >> 4 # // 16
bs_offs_3 = bs_offs_n[None, :] >> 3 # // 8
bs_offs_4 = bs_offs_n[None, :] & 7 # % 8
bs_offs_5 = bs_offs_4 & 3 # % 4
bs_offs_4 = bs_offs_4 >> 2 # // 4
bs_offs = (
bs_offs_1
+ (bs_offs_4 << 1)
+ (bs_offs_2 << 2)
+ (bs_offs_5 << 6)
+ (bs_offs_3 << 8)
+ bs_offs_0 * (32 * scaleN_pad)
)
bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
else:
bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N)[None, :]
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)
def custom_quant(x: torch.Tensor, config: dict):
M, N = x.shape
scaleN_valid = N // 32
scaleM_pad = (M + 31) // 32 * 32
scaleN_pad = (scaleN_valid + 7) // 8 * 8
scaleN = ((scaleN_valid + 7) // 8) * 8
# Use tuned config or golden defaults
b_size = config.get("bs", 16 if M <= 16 else 64)
n_warps = config.get("w", 1 if M <= 16 else 4)
n_stages = config.get("s", 1 if M <= 16 else 3)
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
bs = torch.empty((((M + 255) // 256) * 256, scaleN), dtype=torch.uint8, device=x.device)
grid = ((M + b_size - 1) // b_size, scaleN)
my_quant_kernel[grid](
x, x_fp4, bs,
N, 1,
N // 2, 1,
scaleN, 1,
M, N,
scaleN, scaleM_pad, scaleN_pad,
BLOCK_SIZE=b_size,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0, SHUFFLE=True,
num_warps=n_warps, num_stages=n_stages,
)
return x_fp4.view(dtypes.fp4x2), bs.view(dtypes.fp8_e8m0)
# ============================================================
# PHASE 5: SNIPER TUNING TABLE
# Config: {"bs": BLOCK_SIZE, "w": num_warps, "s": num_stages, "k": log2_k_split}
# ============================================================
CONFIG_TABLE = {
(4, 2880): {"bs": 16, "w": 1, "s": 1, "k": None},
(8, 2112): {"bs": 16, "w": 1, "s": 1, "k": None},
(16, 2112): {"bs": 16, "w": 1, "s": 1, "k": 1},
(16, 3072): {"bs": 16, "w": 1, "s": 1, "k": None},
(32, 2880): {"bs": 32, "w": 4, "s": 3, "k": None},
(32, 4096): {"bs": 32, "w": 4, "s": 3, "k": None},
(64, 3072): {"bs": 64, "w": 4, "s": 3, "k": None},
(64, 7168): {"bs": 64, "w": 4, "s": 3, "k": None},
(256, 2880): {"bs": 64, "w": 4, "s": 3, "k": None},
(256, 3072): {"bs": 64, "w": 4, "s": 3, "k": None},
}
_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_KERNEL_64x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"
_OUT_CACHE = {}
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]
# Shape-Specific Lookup
config = CONFIG_TABLE.get((M, N), {})
A_q, A_scale_sh = custom_quant(A, config)
key = (M, N)
if key not in _OUT_CACHE:
_OUT_CACHE[key] = torch.empty((M, N), dtype=dtypes.bf16, device=A.device)
out = _OUT_CACHE[key]
kernel_name = _KERNEL_32x128 if M <= 64 else _KERNEL_64x128
k_split = config.get("k", 1 if (M <= 16 and K >= 4096) else None)
return aiter.gemm_a4w4_asm(
A_q, B_shuffle, A_scale_sh, B_scale_sh, out,
kernel_name, bpreshuffle=True, log2_k_split=k_split,
)
scrolls · 177 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