submission 528029
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 173 lines, June 9 Researcher Reciprocity License v1.0.
submission_rent_direct_20260311_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-528029?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:b4b06181a5d48b3099eb397d2f3a7654853db1c68e253d61b4450232717c2cf3
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
kernels.add((row["knl_name"], int(row["splitK"])))Kernel source
submission_rent_direct_20260311_v1.py173 lines
from __future__ import annotations
import csv
from pathlib import Path
from task import input_t, output_t
_AITER = None
_DTYPES = None
_FP4_UTILS = None
_TRITON_QUANT_FUNC = None
_KNOWN_ASM = None
_A_QUANT_BUFFERS = {}
_OUT_BUFFERS = {}
_PREFERRED_KERNELS = {
(4, 2880, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
0,
),
(16, 2112, 7168): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
0,
),
(32, 4096, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
0,
),
(32, 2880, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
0,
),
(64, 7168, 2048): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
0,
),
(256, 3072, 1536): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
0,
),
}
def _ceil_div(x: int, y: int) -> int:
return (x + y - 1) // y
def _lazy_init():
global _AITER, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _KNOWN_ASM
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
_AITER = aiter
_DTYPES = dtypes
_FP4_UTILS = fp4_utils
_TRITON_QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)
_KNOWN_ASM = _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 = set()
with csv_path.open(newline="") as f:
for row in csv.DictReader(f):
if int(row["bpreshuffle"]) == 1:
kernels.add((row["knl_name"], int(row["splitK"])))
return kernels
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 _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)
try:
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)
except Exception:
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 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])
a_q, a_scale_sh = _quantize_a(a)
preferred = _PREFERRED_KERNELS.get((m, n, k))
if preferred is not None and preferred in _KNOWN_ASM:
out = _get_out_buffer(m, n, a.device)
_AITER.gemm_a4w4_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
preferred[0],
None,
bpreshuffle=True,
log2_k_split=preferred[1],
)
return out[:m]
return _AITER.gemm_a4w4(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=_DTYPES.bf16,
bpreshuffle=True,
)
scrolls · 173 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 527938.
⋯ 5 unchanged linesfrom task import input_t, output_t_AITER = None- _CLEAR_L2_CACHE = None_DTYPES = None_FP4_UTILS = None_TRITON_QUANT_FUNC = None- _ASM_KERNELS = None+ _KNOWN_ASM = 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_KERNELS = {+ (4, 2880, 512): (+ "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",+ 0,+ ),+ (16, 2112, 7168): (+ "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",+ 0,+ ),+ (32, 4096, 512): (+ "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",+ 0,+ ),+ (32, 2880, 512): (+ "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",+ 0,+ ),+ (64, 7168, 2048): (+ "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",+ 0,+ ),+ (256, 3072, 1536): (+ "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",+ 0,+ ),}- _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) // ydef _lazy_init():- global _AITER, _CLEAR_L2_CACHE, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _ASM_KERNELS+ global _AITER, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _KNOWN_ASMif _AITER is not None:return⋯ 1 unchanged linesfrom aiter import QuantType, dtypesfrom aiter.jit.utils.chip_info import get_gfxfrom 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())+ _KNOWN_ASM = _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 = []-+ kernels = set()with csv_path.open(newline="") as f:for row in csv.DictReader(f):- if int(row["bpreshuffle"]) != 1:- continue+ if int(row["bpreshuffle"]) == 1:+ kernels.add((row["knl_name"], int(row["splitK"])))+ return kernels- 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⋯ 12 unchanged linesreturn buffers- def _cached_quantize_a(a):+ def _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)+ try:+ 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)+ except Exception:+ return _TRITON_QUANT_FUNC(a, shuffle=True)- _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) * 256key = (device.type, device.index, rows, n)⋯ 6 unchanged linesreturn 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()⋯ 3 unchanged linesm = 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(+ preferred = _PREFERRED_KERNELS.get((m, n, k))+ if preferred is not None and preferred in _KNOWN_ASM:+ out = _get_out_buffer(m, n, a.device)+ _AITER.gemm_a4w4_asm(a_q,b_shuffle,a_scale_sh,b_scale_sh,- m,- n,- k,- a.device,+ out,+ preferred[0],+ None,+ bpreshuffle=True,+ log2_k_split=preferred[1],)- _BEST_IMPLS[shape_key] = spec+ return out[:m]- 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]+ return _AITER.gemm_a4w4(+ a_q,+ b_shuffle,+ a_scale_sh,+ b_scale_sh,+ dtype=_DTYPES.bf16,+ bpreshuffle=True,+ )
scrolls · 353 diff lines total
Best evidence level for this revision: reported
JSON