submission 690209
Tecahens · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 277 lines, June 9 Researcher Reciprocity License v1.0.
submissionv1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-690209?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:65198b9188f95bb5b53d220032f41ce08de0bb1ce6736e70c24b9f724c68fb05
license declaredunknown
license concludedunknown
authorsTecahens
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submissionv1.py277 lines
"""
MXFP4 GEMM: A16 preshuffle with reused config strings and workspace buffers.
"""
import torch
import triton
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle_
from aiter.ops.triton.utils.common_utils import serialize_dict
_DIRECT_CFG = {
(4, 2880, 512): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(32, 4096, 512): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(32, 2880, 512): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(16, 2112, 7168): {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 14,
},
(64, 7168, 2048): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(256, 3072, 1536): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
}
_A16_CONFIGS = {
(7168, 2048): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(3072, 1536): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
}
_DIRECT_STATE = {}
_A16_WORKSPACE = {}
_DIRECT_REDUCE_CFG = {
(16, 2112, 7168): (16, 16),
}
def _workspace_key(tensor, n, k, config_override):
cfg_key = tuple(sorted(config_override.items())) if config_override else None
return (str(tensor.device), tensor.shape[0], n, k, cfg_key)
def _get_a16_workspace(a, n, k, config_override):
key = _workspace_key(a, n, k, config_override)
workspace = _A16_WORKSPACE.get(key)
if workspace is not None:
return workspace
workspace = {
"config": serialize_dict(config_override) if config_override is not None else None,
"out": torch.empty((a.shape[0], n), dtype=torch.bfloat16, device=a.device),
}
_A16_WORKSPACE[key] = workspace
return workspace
def _prepare_direct_entry(m, n, k, config):
logical_k = k >> 1
direct_config = dict(config)
if direct_config["NUM_KSPLIT"] > 1:
splitk_block_size, block_size_k, num_ksplit = get_splitk(
logical_k, direct_config["BLOCK_SIZE_K"], direct_config["NUM_KSPLIT"]
)
direct_config["SPLITK_BLOCK_SIZE"] = splitk_block_size
direct_config["BLOCK_SIZE_K"] = block_size_k
direct_config["NUM_KSPLIT"] = num_ksplit
if direct_config["BLOCK_SIZE_K"] >= 2 * logical_k:
direct_config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * logical_k)
direct_config["NUM_KSPLIT"] = 1
direct_config["BLOCK_SIZE_N"] = max(direct_config["BLOCK_SIZE_N"], 32)
if direct_config["BLOCK_SIZE_N"] % 32 != 0 or direct_config["BLOCK_SIZE_K"] % 256 != 0:
return None
direct_config["SPLITK_BLOCK_SIZE"] = (
2 * logical_k
if direct_config["NUM_KSPLIT"] == 1
else direct_config["SPLITK_BLOCK_SIZE"]
)
grid = (
direct_config["NUM_KSPLIT"]
* triton.cdiv(m, direct_config["BLOCK_SIZE_M"])
* triton.cdiv(n, direct_config["BLOCK_SIZE_N"]),
)
return {
"config": direct_config,
"grid": (grid,),
}
def _get_direct_state(a, n, k):
key = (str(a.device), a.shape[0], n, k)
state = _DIRECT_STATE.get(key)
if state is not None:
return state
config = _DIRECT_CFG.get((a.shape[0], n, k))
if config is None:
return None
state = _prepare_direct_entry(a.shape[0], n, k, config)
if state is None:
return None
state["out"] = torch.empty((a.shape[0], n), dtype=torch.bfloat16, device=a.device)
if state["config"]["NUM_KSPLIT"] > 1:
state["y_pp"] = torch.empty(
(state["config"]["NUM_KSPLIT"], a.shape[0], n),
dtype=torch.float32,
device=a.device,
)
_DIRECT_STATE[key] = state
return state
def _run_direct_shape(a, w, w_scales, n, k):
state = _get_direct_state(a, n, k)
if state is None:
return None
logical_k = k >> 1
out = state.get("y_pp", state["out"])
_gemm_a16wfp4_preshuffle_kernel[state["grid"][0]](
a,
w,
out,
w_scales,
a.shape[0],
n,
logical_k,
a.stride(0),
a.stride(1),
w.stride(0),
w.stride(1),
0 if "y_pp" not in state else state["y_pp"].stride(0),
state["out"].stride(0) if "y_pp" not in state else state["y_pp"].stride(1),
state["out"].stride(1) if "y_pp" not in state else state["y_pp"].stride(2),
w_scales.stride(0),
w_scales.stride(1),
PREQUANT=True,
**state["config"],
)
if "y_pp" in state:
reduce_m, reduce_n = _DIRECT_REDUCE_CFG.get((a.shape[0], n, k), (16, 64))
actual_ksplit = triton.cdiv(
logical_k, state["config"]["SPLITK_BLOCK_SIZE"] >> 1
)
_gemm_afp4wfp4_reduce_kernel[
(triton.cdiv(a.shape[0], reduce_m), triton.cdiv(n, reduce_n))
](
state["y_pp"],
state["out"],
a.shape[0],
n,
state["y_pp"].stride(0),
state["y_pp"].stride(1),
state["y_pp"].stride(2),
state["out"].stride(0),
state["out"].stride(1),
reduce_m,
reduce_n,
actual_ksplit,
triton.next_power_of_2(state["config"]["NUM_KSPLIT"]),
)
return state["out"]
def custom_kernel(data):
a, _, _, b_shuffle, b_scale_sh = data
k = a.shape[1]
n = b_shuffle.shape[0]
w = b_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
w_scales = b_scale_sh.view(torch.uint8).reshape(b_scale_sh.shape[0] // 32, k)
out = _run_direct_shape(a, w, w_scales, n, k)
if out is not None:
return out
workspace = _get_a16_workspace(a, n, k, _A16_CONFIGS.get((n, k)))
return gemm_a16wfp4_preshuffle_(
a,
w,
w_scales,
True,
torch.bfloat16,
workspace["out"],
workspace["config"],
False,
)
scrolls · 277 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