submission 528216
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 317 lines, June 9 Researcher Reciprocity License v1.0.
submission_rent2_exhaustive_20260311_v3-1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-528216?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:fe610987513fb9f49603fba6f192342f92fb394fad69d55d7d51af2e1e363b46
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
int(row["splitK"]),Kernel source
submission_rent2_exhaustive_20260311_v3-1.py317 lines
from __future__ import annotations
import csv
from pathlib import Path
import torch
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
_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, _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"]),
int(row["splitK"]),
)
)
return tuple(kernels)
def _kernel_spec(kernel_name: str, split_k: int):
return ("asm", kernel_name, split_k)
def _candidate_specs(m: int, n: int, k: int):
shape_key = (m, n, k)
wrapper_spec = ("wrapper", "wrapper", 0)
specs = [wrapper_spec]
seen = {wrapper_spec}
preferred = _PREFERRED_KERNELS.get(shape_key)
if preferred is not None:
spec = _kernel_spec(*preferred)
specs.append(spec)
seen.add(spec)
padded_m = _ceil_div(m, 32) * 32
def _kernel_rank(entry):
kernel_name, tile_m, tile_n, split_k = entry
m_penalty = abs(tile_m - min(max(32, padded_m), 256))
n_penalty = abs(tile_n - min(max(128, n), 1024))
split_penalty = split_k
return (m_penalty, n_penalty, split_penalty, tile_m, tile_n, kernel_name)
for kernel_name, _tile_m, _tile_n, split_k in sorted(_ASM_KERNELS, key=_kernel_rank):
spec = _kernel_spec(kernel_name, split_k)
if 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, 32) * 32
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)
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
@torch.no_grad()
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 · 317 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 528029.
⋯ 2 unchanged linesimport csvfrom pathlib import Path+ import torchfrom task import input_t, output_t_AITER = None+ _CLEAR_L2_CACHE = None_DTYPES = None_FP4_UTILS = None_TRITON_QUANT_FUNC = None- _KNOWN_ASM = 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+_PREFERRED_KERNELS = {(4, 2880, 512): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",⋯ 27 unchanged linesdef _lazy_init():- global _AITER, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _KNOWN_ASM+ global _AITER, _CLEAR_L2_CACHE, _DTYPES, _FP4_UTILS, _TRITON_QUANT_FUNC, _ASM_KERNELSif _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)- _KNOWN_ASM = _load_asm_kernels(aiter, get_gfx())+ _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 = set()+ kernels = []+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+ if int(row["bpreshuffle"]) != 1:+ continue+ kernels.append(+ (+ row["knl_name"],+ int(row["tile_M"]),+ int(row["tile_N"]),+ int(row["splitK"]),+ )+ )+ return tuple(kernels)++ def _kernel_spec(kernel_name: str, split_k: int):+ return ("asm", kernel_name, split_k)+++ def _candidate_specs(m: int, n: int, k: int):+ shape_key = (m, n, k)+ wrapper_spec = ("wrapper", "wrapper", 0)+ specs = [wrapper_spec]+ seen = {wrapper_spec}++ preferred = _PREFERRED_KERNELS.get(shape_key)+ if preferred is not None:+ spec = _kernel_spec(*preferred)+ specs.append(spec)+ seen.add(spec)++ padded_m = _ceil_div(m, 32) * 32++ def _kernel_rank(entry):+ kernel_name, tile_m, tile_n, split_k = entry+ m_penalty = abs(tile_m - min(max(32, padded_m), 256))+ n_penalty = abs(tile_n - min(max(128, n), 1024))+ split_penalty = split_k+ return (m_penalty, n_penalty, split_penalty, tile_m, tile_n, kernel_name)++ for kernel_name, _tile_m, _tile_n, split_k in sorted(_ASM_KERNELS, key=_kernel_rank):+ spec = _kernel_spec(kernel_name, split_k)+ if 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 _quantize_a(a):+ 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)- 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)+ 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+ rows = _ceil_div(m, 32) * 32key = (device.type, device.index, rows, n)out = _OUT_BUFFERS.get(key)if out is None:⋯ 4 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)+ 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+++ @torch.no_grad()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)- 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(+ spec = _BEST_IMPLS.get(shape_key)+ if spec is None:+ spec = _select_best_impl(a_q,b_shuffle,a_scale_sh,b_scale_sh,- out,- preferred[0],- None,- bpreshuffle=True,- log2_k_split=preferred[1],+ m,+ n,+ k,+ a.device,)- return out[:m]+ _BEST_IMPLS[shape_key] = spec- return _AITER.gemm_a4w4(- a_q,- b_shuffle,- a_scale_sh,- b_scale_sh,- dtype=_DTYPES.bf16,- bpreshuffle=True,- )+ 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 · 323 diff lines total
Best evidence level for this revision: reported
JSON