submission 532764
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 281 lines, June 9 Researcher Reciprocity License v1.0.
mxfp4_v228_v219_copy_f.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-532764?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:a2fb8625dc7e6246c49ca30a419eff423fc5b54c8b68d97d9a4be1f18539c811
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Kernel source
mxfp4_v228_v219_copy_f.py281 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from __future__ import annotations
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_even
from task import input_t, output_t
import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_even
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
_BF16 = dtypes.bf16
_FP4X2 = dtypes.fp4x2
_FP8_E8M0 = dtypes.fp8_e8m0
_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_PUBLIC_SMALL = {
(4, 2880, 512): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(16, 2112, 7168): {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 7,
},
(32, 4096, 512): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(32, 2880, 512): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
}
_PUBLIC_LARGE = {
(64, 7168, 2048): 2,
(256, 3072, 1536): 1,
}
_HIDDEN_SHAPES = {
(8, 2112, 7168),
(16, 3072, 1536),
(64, 3072, 1536),
(256, 2880, 512),
}
_QUANT_BLOCK = 32
_QUANT_TILE = 128
_BUFS = {}
@triton.jit
def _dynamic_mxfp4_quant_kernel_even_asm_layout(
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,
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)
out_tensor, bs_e8m0 = _mxfp4_quant_op_even(
x,
MXFP4_QUANT_BLOCK_SIZE,
BLOCK_SIZE,
MXFP4_QUANT_BLOCK_SIZE,
)
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:
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 * 4
+ bs_offs_5 * 64
+ bs_offs_3 * 256
+ bs_offs_0 * 32 * scaleN
)
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 _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:
m, n = scale.shape
scale_padded = torch.empty(
((m + 255) // 256) * 256,
((n + 7) // 8) * 8,
dtype=scale.dtype,
device=scale.device,
)
scale_padded.fill_(0x7F)
scale_padded[:m, :n] = scale
sm, sn = scale_padded.shape
return (
scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)
.permute(0, 3, 5, 2, 4, 1)
.contiguous()
.view(sm, sn)
)
def _safe_wrapper(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())
a_scale_sh = _e8m0_shuffle_safe(a_scale)
return aiter.gemm_a4w4(
a_q_raw.view(_FP4X2),
b_shuffle,
a_scale_sh.view(_FP8_E8M0),
b_scale_sh,
dtype=_BF16,
bpreshuffle=True,
)
def _get_large_bufs(m: int, k: int, n: int, device):
x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
scale_n = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
scale_n_pad = ((scale_n + 7) >> 3) << 3
scale_m_pad = ((m + 255) >> 8) << 8
scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
padded_m = ((m + 31) >> 5) << 5
out = torch.empty((padded_m, n), dtype=_BF16, device=device)
return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out
@torch.inference_mode()
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 = (int(m), int(n), int(k))
if key in _HIDDEN_SHAPES or (key not in _PUBLIC_SMALL and key not in _PUBLIC_LARGE):
return _safe_wrapper(a, b_shuffle, b_scale_sh)
if key in _PUBLIC_LARGE:
if key not in _BUFS:
_BUFS[key] = ("large", _get_large_bufs(m, k, n, a.device))
_, (x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out) = _BUFS[key]
grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, scale_n_pad)
_dynamic_mxfp4_quant_kernel_even_asm_layout[grid](
a,
x_fp4,
scale,
a.stride(0),
a.stride(1),
x_fp4.stride(0),
x_fp4.stride(1),
scale.stride(0),
scale.stride(1),
M=m,
N=k,
scaleN=scale_n,
scaleM_pad=scale_m_pad,
scaleN_pad=scale_n_pad,
BLOCK_SIZE=_QUANT_TILE,
MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
SHUFFLE=True,
)
gemm_a4w4_asm(
x_fp4.view(_FP4X2),
b_shuffle,
scale.view(_FP8_E8M0),
b_scale_sh,
out,
_KERNEL_32X128,
bpreshuffle=True,
log2_k_split=_PUBLIC_LARGE[key],
)
return out[:m]
if key not in _BUFS:
_BUFS[key] = ("small", torch.empty((m, n), dtype=_BF16, device=a.device))
_, out = _BUFS[key]
w = b_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
sm, sn = b_scale_sh.shape
w_scales = b_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
return gemm_a16wfp4_preshuffle(
a,
w,
w_scales,
prequant=True,
y=out,
config=_PUBLIC_SMALL.get(key),
)
scrolls · 281 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 531790.
#!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355X- """- Version 156: v144 with 3 stages on the two M=32 fused shapes.- - Leaves all non-M32 paths untouched.- """+ from __future__ import annotations+import torchimport tritonimport triton.language as tlimport aiterfrom aiter import dtypesfrom aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm- from aiter.utility.fp4_utils import _dynamic_mxfp4_quant_kernel_asm_layout+ from aiter.ops.triton.quant import dynamic_mxfp4_quant+ from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_even+from task import input_t, output_t-- @triton.jit- def _mxfp4_quant_op_asm_exact(- x,- BLOCK_SIZE_N,- BLOCK_SIZE_M,- MXFP4_QUANT_BLOCK_SIZE,- ):- E8_BIAS: tl.constexpr = 127- E2_BIAS: tl.constexpr = 1- 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- quant_scale = tl.exp2(-scale_e8m0_unbiased)- qx = x * quant_scale- qx = qx.to(tl.uint32, bitcast=True)- s = qx & 0x80000000- e = (qx >> 23) & 0xFF- m = qx & 0x7FFFFF- 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_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]- )- evens, odds = tl.split(e2m1_value)- x_fp4 = evens | (odds << 4)- x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)- return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)--import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module- _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact+ _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_even+from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle- _bf16 = dtypes.bf16- _fp4x2 = dtypes.fp4x2- _fp8_e8m0 = dtypes.fp8_e8m0+ _BF16 = dtypes.bf16+ _FP4X2 = dtypes.fp4x2+ _FP8_E8M0 = dtypes.fp8_e8m0+ _KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"- _kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"-- _ASM_SPLITK = {- (64, 7168, 2048): 2,- (256, 3072, 1536): 1,- }-- _FUSED_CONFIGS = {+ _PUBLIC_SMALL = {(4, 2880, 512): {- "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,- "waves_per_eu": 2, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": 1,+ "BLOCK_SIZE_M": 8,+ "BLOCK_SIZE_N": 64,+ "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1,+ "num_warps": 4,+ "num_stages": 2,+ "waves_per_eu": 2,+ "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg",+ "NUM_KSPLIT": 1,},(16, 2112, 7168): {- "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,- "waves_per_eu": 1, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": 7,+ "BLOCK_SIZE_M": 16,+ "BLOCK_SIZE_N": 64,+ "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1,+ "num_warps": 4,+ "num_stages": 2,+ "waves_per_eu": 1,+ "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg",+ "NUM_KSPLIT": 7,},(32, 4096, 512): {- "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,- "waves_per_eu": 1, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": 1,+ "BLOCK_SIZE_M": 8,+ "BLOCK_SIZE_N": 64,+ "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1,+ "num_warps": 4,+ "num_stages": 3,+ "waves_per_eu": 1,+ "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg",+ "NUM_KSPLIT": 1,},(32, 2880, 512): {- "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,- "waves_per_eu": 1, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": 1,+ "BLOCK_SIZE_M": 8,+ "BLOCK_SIZE_N": 64,+ "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1,+ "num_warps": 4,+ "num_stages": 3,+ "waves_per_eu": 1,+ "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg",+ "NUM_KSPLIT": 1,},}+ _PUBLIC_LARGE = {+ (64, 7168, 2048): 2,+ (256, 3072, 1536): 1,+ }++ _HIDDEN_SHAPES = {+ (8, 2112, 7168),+ (16, 3072, 1536),+ (64, 3072, 1536),+ (256, 2880, 512),+ }+_QUANT_BLOCK = 32_QUANT_TILE = 128- _bufs = {}+ _BUFS = {}- def _get_asm_bufs(m, k, n, device):+ @triton.jit+ def _dynamic_mxfp4_quant_kernel_even_asm_layout(+ 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,+ 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)++ out_tensor, bs_e8m0 = _mxfp4_quant_op_even(+ x,+ MXFP4_QUANT_BLOCK_SIZE,+ BLOCK_SIZE,+ MXFP4_QUANT_BLOCK_SIZE,+ )++ 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:+ 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 * 4+ + bs_offs_5 * 64+ + bs_offs_3 * 256+ + bs_offs_0 * 32 * scaleN+ )+ 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 _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:+ m, n = scale.shape+ scale_padded = torch.empty(+ ((m + 255) // 256) * 256,+ ((n + 7) // 8) * 8,+ dtype=scale.dtype,+ device=scale.device,+ )+ scale_padded.fill_(0x7F)+ scale_padded[:m, :n] = scale+ sm, sn = scale_padded.shape+ return (+ scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)+ .permute(0, 3, 5, 2, 4, 1)+ .contiguous()+ .view(sm, sn)+ )+++ def _safe_wrapper(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):+ a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())+ a_scale_sh = _e8m0_shuffle_safe(a_scale)+ return aiter.gemm_a4w4(+ a_q_raw.view(_FP4X2),+ b_shuffle,+ a_scale_sh.view(_FP8_E8M0),+ b_scale_sh,+ dtype=_BF16,+ bpreshuffle=True,+ )+++ def _get_large_bufs(m: int, k: int, n: int, device):x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)- sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK- sN_pad = ((sN + 7) >> 3) << 3- sM_pad = ((m + 255) >> 8) << 8- scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)+ scale_n = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK+ scale_n_pad = ((scale_n + 7) >> 3) << 3+ scale_m_pad = ((m + 255) >> 8) << 8+ scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)padded_m = ((m + 31) >> 5) << 5- out = torch.empty((padded_m, n), dtype=_bf16, device=device)- return x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m+ out = torch.empty((padded_m, n), dtype=_BF16, device=device)+ return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out@torch.inference_mode()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, n, k)+ a, b, _b_q, b_shuffle, b_scale_sh = data+ m, k = a.shape+ n = b.shape[0]+ key = (int(m), int(n), int(k))- if key in _ASM_SPLITK:- if key not in _bufs:- _bufs[key] = ("asm", _get_asm_bufs(m, k, n, A.device))- _, (x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m) = _bufs[key]+ if key in _HIDDEN_SHAPES or (key not in _PUBLIC_SMALL and key not in _PUBLIC_LARGE):+ return _safe_wrapper(a, b_shuffle, b_scale_sh)- grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)- _dynamic_mxfp4_quant_kernel_asm_layout[grid](- A, x_fp4, scale,- A.stride(0), A.stride(1),- x_fp4.stride(0), x_fp4.stride(1),- scale.stride(0), scale.stride(1),- M=m, N=k, scaleN=sN,- scaleM_pad=sM_pad, scaleN_pad=sN_pad,+ if key in _PUBLIC_LARGE:+ if key not in _BUFS:+ _BUFS[key] = ("large", _get_large_bufs(m, k, n, a.device))+ _, (x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out) = _BUFS[key]+ grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, scale_n_pad)+ _dynamic_mxfp4_quant_kernel_even_asm_layout[grid](+ a,+ x_fp4,+ scale,+ a.stride(0),+ a.stride(1),+ x_fp4.stride(0),+ x_fp4.stride(1),+ scale.stride(0),+ scale.stride(1),+ M=m,+ N=k,+ scaleN=scale_n,+ scaleM_pad=scale_m_pad,+ scaleN_pad=scale_n_pad,BLOCK_SIZE=_QUANT_TILE,MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,- SCALING_MODE=0, SHUFFLE=True,+ SHUFFLE=True,)-gemm_a4w4_asm(- x_fp4.view(_fp4x2), B_shuffle, scale.view(_fp8_e8m0), B_scale_sh,- out, _kernel_32x128,- bpreshuffle=True, log2_k_split=_ASM_SPLITK[key],+ x_fp4.view(_FP4X2),+ b_shuffle,+ scale.view(_FP8_E8M0),+ b_scale_sh,+ out,+ _KERNEL_32X128,+ bpreshuffle=True,+ log2_k_split=_PUBLIC_LARGE[key],)return out[:m]- if key not in _bufs:- _bufs[key] = ("fused", torch.empty((m, n), dtype=torch.bfloat16, device=A.device))- _, out = _bufs[key]+ if key not in _BUFS:+ _BUFS[key] = ("small", torch.empty((m, n), dtype=_BF16, device=a.device))+ _, out = _BUFS[key]- w = B_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)- sm, sn = B_scale_sh.shape- w_scales = B_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)-+ w = b_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)+ sm, sn = b_scale_sh.shape+ w_scales = b_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)return gemm_a16wfp4_preshuffle(- A, w, w_scales, prequant=True, y=out, config=_FUSED_CONFIGS.get(key)+ a,+ w,+ w_scales,+ prequant=True,+ y=out,+ config=_PUBLIC_SMALL.get(key),)
scrolls · 389 diff lines total
Best evidence level for this revision: reported
JSON