submission 527938
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 302 lines, June 9 Researcher Reciprocity License v1.0.
submission-rent.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-527938?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:8bb1b43b1f9aa75532d568f93ff89c580051c5f4ba856814244176b0cf65ce85
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Kernel source
submission-rent.py302 lines
from __future__ import annotations
import csv
from pathlib import Path
from task import input_t, output_t
_AITER = None
_CLEAR_L2_CACHE = None
_DTYPES = None
_FP4_UTILS = None
_TRITON_QUANT_FUNC = None
_ASM_KERNELS = None
_USE_CACHED_TRITON_QUANT = True
_BEST_IMPLS = {}
_A_QUANT_BUFFERS = {}
_OUT_BUFFERS = {}
_SELECT_ITERS = 5
_SELECT_WARMUP_ITERS = 1
_RTOL = 1e-2
_ATOL = 1e-2
_WRAPPER_ONLY_SHAPES = {
(64, 7168, 2048),
(256, 3072, 1536),
}
_PREFERRED_TILE_SPECS = {
(4, 2880, 512): (64, 128),
(16, 2112, 7168): (32, 128),
(32, 4096, 512): (32, 256),
(32, 2880, 512): (64, 128),
}
def _ceil_div(x: int, y: int) -> int:
return (x + y - 1) // y
def _lazy_init():
global _AITER, _CLEAR_L2_CACHE, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _ASM_KERNELS
if _AITER is not None:
return
import aiter
from aiter import QuantType, dtypes
from aiter.jit.utils.chip_info import get_gfx
from aiter.utility import fp4_utils
from utils import clear_l2_cache_large
_AITER = aiter
_CLEAR_L2_CACHE = clear_l2_cache_large
_DTYPES = dtypes
_FP4_UTILS = fp4_utils
_TRITON_QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)
_ASM_KERNELS = _load_asm_kernels(aiter, get_gfx())
def _load_asm_kernels(aiter, gfx: str):
root = Path(aiter.__file__).resolve().parents[1]
csv_path = root / "hsa" / gfx / "f4gemm" / "f4gemm_bf16_per1x32Fp4.csv"
kernels = []
with csv_path.open(newline="") as f:
for row in csv.DictReader(f):
if int(row["bpreshuffle"]) != 1:
continue
kernels.append(
(
row["knl_name"],
int(row["tile_M"]),
int(row["tile_N"]),
)
)
return tuple(kernels)
def _spec_for_tiles(tile_m: int, tile_n: int):
for kernel_name, kernel_tile_m, kernel_tile_n in _ASM_KERNELS:
if kernel_tile_m == tile_m and kernel_tile_n == tile_n:
return ("asm", kernel_name, 0)
return None
def _candidate_specs(m: int, n: int, k: int):
shape_key = (m, n, k)
wrapper_spec = ("wrapper", "wrapper", 0)
if shape_key in _WRAPPER_ONLY_SHAPES:
return (wrapper_spec,)
specs = [wrapper_spec]
seen = {wrapper_spec}
preferred_tiles = _PREFERRED_TILE_SPECS.get(shape_key)
if preferred_tiles is not None:
preferred_spec = _spec_for_tiles(preferred_tiles[0], preferred_tiles[1])
if preferred_spec is not None:
specs.append(preferred_spec)
seen.add(preferred_spec)
for tile_m in (32, 64):
for tile_n in (128, 256):
spec = _spec_for_tiles(tile_m, tile_n)
if spec is not None and spec not in seen:
specs.append(spec)
seen.add(spec)
return tuple(specs)
def _get_quant_buffers(m: int, k: int, device):
scale_n = _ceil_div(k, 32)
scale_n_pad = _ceil_div(scale_n, 8) * 8
scale_m_pad = _ceil_div(m, 256) * 256
key = (device.type, device.index, m, k, scale_m_pad, scale_n_pad)
buffers = _A_QUANT_BUFFERS.get(key)
if buffers is None:
import torch
a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
a_scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
buffers = (a_q, a_scale, scale_n, scale_m_pad, scale_n_pad)
_A_QUANT_BUFFERS[key] = buffers
return buffers
def _cached_quantize_a(a):
m = int(a.shape[0])
k = int(a.shape[1])
a_q, a_scale, scale_n, scale_m_pad, scale_n_pad = _get_quant_buffers(m, k, a.device)
block_size = 128
grid = (_ceil_div(m, block_size), scale_n_pad)
_FP4_UTILS._dynamic_mxfp4_quant_kernel_asm_layout[grid](
a,
a_q,
a_scale,
*a.stride(),
*a_q.stride(),
*a_scale.stride(),
M=m,
N=k,
scaleN=scale_n,
scaleM_pad=scale_m_pad,
scaleN_pad=scale_n_pad,
BLOCK_SIZE=block_size,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
SHUFFLE=True,
)
return a_q.view(_DTYPES.fp4x2), a_scale.view(_DTYPES.fp8_e8m0)
def _quantize_a(a):
if _USE_CACHED_TRITON_QUANT:
try:
return _cached_quantize_a(a)
except Exception:
pass
return _TRITON_QUANT_FUNC(a, shuffle=True)
def _get_out_buffer(m: int, n: int, device):
rows = _ceil_div(m, 256) * 256
key = (device.type, device.index, rows, n)
out = _OUT_BUFFERS.get(key)
if out is None:
import torch
out = torch.empty((rows, n), dtype=_DTYPES.bf16, device=device)
_OUT_BUFFERS[key] = out
return out
def _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):
kind, impl_id, log2_k_split = spec
if kind == "wrapper":
return _AITER.gemm_a4w4(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=_DTYPES.bf16,
bpreshuffle=True,
)
_AITER.gemm_a4w4_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
impl_id,
None,
bpreshuffle=True,
log2_k_split=log2_k_split,
)
return out
def _is_valid(candidate, reference, m: int):
import torch
return torch.allclose(candidate[:m], reference[:m], rtol=_RTOL, atol=_ATOL)
def _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):
import torch
for _ in range(_SELECT_WARMUP_ITERS):
_run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
torch.cuda.synchronize()
times_ms = []
for _ in range(_SELECT_ITERS):
_CLEAR_L2_CACHE()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
_run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
end.record()
torch.cuda.synchronize()
times_ms.append(start.elapsed_time(end))
times_ms.sort()
return times_ms[len(times_ms) // 2]
def _select_best_impl(a_q, b_shuffle, a_scale_sh, b_scale_sh, m: int, n: int, k: int, device):
wrapper_spec = ("wrapper", "wrapper", 0)
if (m, n, k) in _WRAPPER_ONLY_SHAPES:
return wrapper_spec
out = _get_out_buffer(m, n, device)
reference = _run_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
best_spec = wrapper_spec
best_time = _time_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
for spec in _candidate_specs(m, n, k):
if spec == wrapper_spec:
continue
try:
candidate = _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
if not _is_valid(candidate, reference, m):
continue
current_time = _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)
if current_time < best_time:
best_time = current_time
best_spec = spec
except Exception:
continue
return best_spec
def custom_kernel(data: input_t) -> output_t:
_lazy_init()
a, _b, _b_q, b_shuffle, b_scale_sh = data
a = a.contiguous()
m = int(a.shape[0])
k = int(a.shape[1])
n = int(b_shuffle.shape[0])
shape_key = (m, n, k)
a_q, a_scale_sh = _quantize_a(a)
spec = _BEST_IMPLS.get(shape_key)
if spec is None:
spec = _select_best_impl(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
m,
n,
k,
a.device,
)
_BEST_IMPLS[shape_key] = spec
if spec[0] == "wrapper":
return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, None)
out = _get_out_buffer(m, n, a.device)
return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)[:m]
scrolls · 302 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 527329.
- #!POPCORN leaderboard amd-mxfp4-mm- #!POPCORN gpu MI355X+ from __future__ import annotations- """- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).- """+ import csv+ from pathlib import Path+from task import input_t, output_t+ _AITER = None+ _CLEAR_L2_CACHE = None+ _DTYPES = None+ _FP4_UTILS = None+ _TRITON_QUANT_FUNC = None+ _ASM_KERNELS = None- def custom_kernel(data: input_t) -> output_t:- """- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.- gemm_a4w4 with bpreshuffle=True.- """+ _USE_CACHED_TRITON_QUANT = True++ _BEST_IMPLS = {}+ _A_QUANT_BUFFERS = {}+ _OUT_BUFFERS = {}++ _SELECT_ITERS = 5+ _SELECT_WARMUP_ITERS = 1+ _RTOL = 1e-2+ _ATOL = 1e-2++ _WRAPPER_ONLY_SHAPES = {+ (64, 7168, 2048),+ (256, 3072, 1536),+ }++ _PREFERRED_TILE_SPECS = {+ (4, 2880, 512): (64, 128),+ (16, 2112, 7168): (32, 128),+ (32, 4096, 512): (32, 256),+ (32, 2880, 512): (64, 128),+ }+++ def _ceil_div(x: int, y: int) -> int:+ return (x + y - 1) // y+++ def _lazy_init():+ global _AITER, _CLEAR_L2_CACHE, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _ASM_KERNELS+ if _AITER is not None:+ return+import aiterfrom aiter import QuantType, dtypes+ from aiter.jit.utils.chip_info import get_gfx+ from aiter.utility import fp4_utils+ from utils import clear_l2_cache_large- A, B, B_q, B_shuffle, B_scale_sh = data- A = A.contiguous()- B = B.contiguous()- m, k = A.shape- n, _ = B.shape+ _AITER = aiter+ _CLEAR_L2_CACHE = clear_l2_cache_large+ _DTYPES = dtypes+ _FP4_UTILS = fp4_utils+ _TRITON_QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)+ _ASM_KERNELS = _load_asm_kernels(aiter, get_gfx())- quant_func = aiter.get_triton_quant(QuantType.per_1x32)- A_q, A_scale_sh = quant_func(A, shuffle=True)- out_gemm = aiter.gemm_a4w4(- A_q,- B_shuffle,- A_scale_sh,- B_scale_sh,- dtype=dtypes.bf16,++ def _load_asm_kernels(aiter, gfx: str):+ root = Path(aiter.__file__).resolve().parents[1]+ csv_path = root / "hsa" / gfx / "f4gemm" / "f4gemm_bf16_per1x32Fp4.csv"+ kernels = []++ with csv_path.open(newline="") as f:+ for row in csv.DictReader(f):+ if int(row["bpreshuffle"]) != 1:+ continue++ kernels.append(+ (+ row["knl_name"],+ int(row["tile_M"]),+ int(row["tile_N"]),+ )+ )++ return tuple(kernels)+++ def _spec_for_tiles(tile_m: int, tile_n: int):+ for kernel_name, kernel_tile_m, kernel_tile_n in _ASM_KERNELS:+ if kernel_tile_m == tile_m and kernel_tile_n == tile_n:+ return ("asm", kernel_name, 0)+ return None+++ def _candidate_specs(m: int, n: int, k: int):+ shape_key = (m, n, k)+ wrapper_spec = ("wrapper", "wrapper", 0)+ if shape_key in _WRAPPER_ONLY_SHAPES:+ return (wrapper_spec,)++ specs = [wrapper_spec]+ seen = {wrapper_spec}++ preferred_tiles = _PREFERRED_TILE_SPECS.get(shape_key)+ if preferred_tiles is not None:+ preferred_spec = _spec_for_tiles(preferred_tiles[0], preferred_tiles[1])+ if preferred_spec is not None:+ specs.append(preferred_spec)+ seen.add(preferred_spec)++ for tile_m in (32, 64):+ for tile_n in (128, 256):+ spec = _spec_for_tiles(tile_m, tile_n)+ if spec is not None and spec not in seen:+ specs.append(spec)+ seen.add(spec)++ return tuple(specs)+++ def _get_quant_buffers(m: int, k: int, device):+ scale_n = _ceil_div(k, 32)+ scale_n_pad = _ceil_div(scale_n, 8) * 8+ scale_m_pad = _ceil_div(m, 256) * 256+ key = (device.type, device.index, m, k, scale_m_pad, scale_n_pad)+ buffers = _A_QUANT_BUFFERS.get(key)++ if buffers is None:+ import torch++ a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=device)+ a_scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)+ buffers = (a_q, a_scale, scale_n, scale_m_pad, scale_n_pad)+ _A_QUANT_BUFFERS[key] = buffers++ return buffers+++ def _cached_quantize_a(a):+ m = int(a.shape[0])+ k = int(a.shape[1])+ a_q, a_scale, scale_n, scale_m_pad, scale_n_pad = _get_quant_buffers(m, k, a.device)++ block_size = 128+ grid = (_ceil_div(m, block_size), scale_n_pad)++ _FP4_UTILS._dynamic_mxfp4_quant_kernel_asm_layout[grid](+ a,+ a_q,+ a_scale,+ *a.stride(),+ *a_q.stride(),+ *a_scale.stride(),+ M=m,+ N=k,+ scaleN=scale_n,+ scaleM_pad=scale_m_pad,+ scaleN_pad=scale_n_pad,+ BLOCK_SIZE=block_size,+ MXFP4_QUANT_BLOCK_SIZE=32,+ SCALING_MODE=0,+ SHUFFLE=True,+ )++ return a_q.view(_DTYPES.fp4x2), a_scale.view(_DTYPES.fp8_e8m0)+++ def _quantize_a(a):+ if _USE_CACHED_TRITON_QUANT:+ try:+ return _cached_quantize_a(a)+ except Exception:+ pass++ return _TRITON_QUANT_FUNC(a, shuffle=True)+++ def _get_out_buffer(m: int, n: int, device):+ rows = _ceil_div(m, 256) * 256+ key = (device.type, device.index, rows, n)+ out = _OUT_BUFFERS.get(key)+ if out is None:+ import torch++ out = torch.empty((rows, n), dtype=_DTYPES.bf16, device=device)+ _OUT_BUFFERS[key] = out+ return out+++ def _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):+ kind, impl_id, log2_k_split = spec++ if kind == "wrapper":+ return _AITER.gemm_a4w4(+ a_q,+ b_shuffle,+ a_scale_sh,+ b_scale_sh,+ dtype=_DTYPES.bf16,+ bpreshuffle=True,+ )++ _AITER.gemm_a4w4_asm(+ a_q,+ b_shuffle,+ a_scale_sh,+ b_scale_sh,+ out,+ impl_id,+ None,bpreshuffle=True,+ log2_k_split=log2_k_split,)- return out_gemm+ return out+++ def _is_valid(candidate, reference, m: int):+ import torch++ return torch.allclose(candidate[:m], reference[:m], rtol=_RTOL, atol=_ATOL)+++ def _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out):+ import torch++ for _ in range(_SELECT_WARMUP_ITERS):+ _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)+ torch.cuda.synchronize()++ times_ms = []+ for _ in range(_SELECT_ITERS):+ _CLEAR_L2_CACHE()+ start = torch.cuda.Event(enable_timing=True)+ end = torch.cuda.Event(enable_timing=True)+ start.record()+ _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)+ end.record()+ torch.cuda.synchronize()+ times_ms.append(start.elapsed_time(end))++ times_ms.sort()+ return times_ms[len(times_ms) // 2]+++ def _select_best_impl(a_q, b_shuffle, a_scale_sh, b_scale_sh, m: int, n: int, k: int, device):+ wrapper_spec = ("wrapper", "wrapper", 0)+ if (m, n, k) in _WRAPPER_ONLY_SHAPES:+ return wrapper_spec++ out = _get_out_buffer(m, n, device)+ reference = _run_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)++ best_spec = wrapper_spec+ best_time = _time_impl(wrapper_spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)++ for spec in _candidate_specs(m, n, k):+ if spec == wrapper_spec:+ continue++ try:+ candidate = _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)+ if not _is_valid(candidate, reference, m):+ continue++ current_time = _time_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)+ if current_time < best_time:+ best_time = current_time+ best_spec = spec+ except Exception:+ continue++ return best_spec+++ def custom_kernel(data: input_t) -> output_t:+ _lazy_init()++ a, _b, _b_q, b_shuffle, b_scale_sh = data+ a = a.contiguous()++ m = int(a.shape[0])+ k = int(a.shape[1])+ n = int(b_shuffle.shape[0])+ shape_key = (m, n, k)++ a_q, a_scale_sh = _quantize_a(a)++ spec = _BEST_IMPLS.get(shape_key)+ if spec is None:+ spec = _select_best_impl(+ a_q,+ b_shuffle,+ a_scale_sh,+ b_scale_sh,+ m,+ n,+ k,+ a.device,+ )+ _BEST_IMPLS[shape_key] = spec++ if spec[0] == "wrapper":+ return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, None)++ out = _get_out_buffer(m, n, a.device)+ return _run_impl(spec, a_q, b_shuffle, a_scale_sh, b_scale_sh, out)[:m]
scrolls · 326 diff lines total
Best evidence level for this revision: reported
JSON