submission 593457
abhicloudstalk13 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 92 lines, June 9 Researcher Reciprocity License v1.0.
improved_gemm_submission_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-593457?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:b2f657f665b39229340ccd7a05999d29d2df2f7166a02bfdeef41cc6b7556a44
license declaredunknown
license concludedunknown
authorsabhicloudstalk13
imported2026-08-26
Kernel source
improved_gemm_submission_v2.py92 lines
from task import input_t, output_t
import torch
from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
def _c(bm, bn, bk, grp, nw, ns, wpe, ksplit, cm=None):
return {
"BLOCK_SIZE_M": bm,
"BLOCK_SIZE_N": bn,
"BLOCK_SIZE_K": bk,
"GROUP_SIZE_M": grp,
"num_warps": nw,
"num_stages": ns,
"waves_per_eu": wpe,
"matrix_instr_nonkdim": 16,
"cache_modifier": cm,
"NUM_KSPLIT": ksplit,
}
_CFG_2880_4 = _c(8, 32, 256, 1, 2, 2, 2, 2, ".cg")
_CFG_2880_32 = _c(16, 32, 256, 1, 4, 2, 2, 2, ".cg")
_CFG_4096_32 = _c(32, 32, 256, 1, 4, 2, 2, 2, ".cg")
_CFG_2112_16 = _c(16, 128, 512, 1, 4, 1, 1, 14, ".cg")
_CFG_7168_64 = _c(16, 64, 512, 1, 8, 2, 4, 1, ".cg")
_CFG_3072_256 = _c(64, 32, 512, 1, 4, 2, 1, 1)
_CFG_DEFAULT = _c(32, 64, 512, 1, 8, 1, 2, 1)
_WEIGHT_VIEW_CACHE = {}
def _prepare_weight_views(weight, scale, n, k):
key = (
weight.data_ptr(),
scale.data_ptr(),
weight.shape,
scale.shape,
n,
k,
)
cached = _WEIGHT_VIEW_CACHE.get(key)
if cached is None:
packed_weight = weight.view(torch.uint8).view(n >> 4, (k >> 1) << 4)
packed_scale = scale.view(torch.uint8).view(scale.shape[0] >> 5, k)[: n >> 5]
cached = (packed_weight, packed_scale)
_WEIGHT_VIEW_CACHE[key] = cached
return cached
def _get_config(m, n, k):
if k == 512:
if n == 2880:
if m <= 8:
return _CFG_2880_4
if m <= 32:
return _CFG_2880_32
elif n == 4096 and m <= 32:
return _CFG_4096_32
elif k == 7168:
if n == 2112 and m <= 16:
return _CFG_2112_16
elif k == 2048:
if n == 7168 and m <= 64:
return _CFG_7168_64
elif k == 1536:
if n == 3072 and m <= 256:
return _CFG_3072_256
return _CFG_DEFAULT
def custom_kernel(data: input_t) -> output_t:
a = data[0]
b_shuffle = data[3]
b_scale_sh = data[4]
if a.stride(1) != 1:
a = a.contiguous()
m, k = a.shape
n = b_shuffle.shape[0]
b_w, b_scale_w = _prepare_weight_views(b_shuffle, b_scale_sh, n, k)
return gemm_a16wfp4_preshuffle(
a,
b_w,
b_scale_w,
dtype=dtypes.bf16,
config=_get_config(m, n, k),
)
scrolls · 92 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