Skip to content
KernelIndex
Search⌘K

submission 753653

Kernel-Zhang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 2641 lines, June 9 Researcher Reciprocity License v1.0.

kmm_00233.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-753653?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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
8.21µs
#45 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ab34d008be982ff5be5dd822bed70cf138a673d06fe0fce4824af8e406498d28
license declaredunknown
license concludedunknown
authorsKernel-Zhang
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4Single-file submission for AMD MXFP4 matmul.
num-warps = 4_FUSED_M4_NUM_WARPS = 4
split-k"SPLITK_BLOCK_SIZE": "constexpr",
stages = 1_FUSED_M4_NUM_STAGES = 1
tile-k = 512_FUSED_M4_BLOCK_SIZE_K = 512
tile-m = 8_FUSED_M4_BLOCK_SIZE_M = 8
tile-n = 64_FUSED_M4_BLOCK_SIZE_N = 64

Kernel source

kmm_00233.py2641 lines
"""
Single-file submission for AMD MXFP4 matmul.

This revision keeps the hot path narrow:
- use `subprocess` once to inspect the shipped asm GEMM module;
- use `ctypes` to expose Torch/ROCm shared libs globally and call the
  shipped `gemm_a4w4_asm`;
- patch `aiter.gemm_a4w4` to route reference checks through the same direct
  entry;
- keep a no-Torch-link backup builder only if the shipped module still fails
  to load.
"""

from __future__ import annotations

import os
import sys

_OLD_DLOPEN_FLAGS: int | None = None
if hasattr(sys, "getdlopenflags") and hasattr(sys, "setdlopenflags"):
    _OLD_DLOPEN_FLAGS = sys.getdlopenflags()
    sys.setdlopenflags(
        _OLD_DLOPEN_FLAGS
        | getattr(os, "RTLD_GLOBAL", 0)
        | getattr(os, "RTLD_NOW", 0)
    )

import torch

if _OLD_DLOPEN_FLAGS is not None:
    sys.setdlopenflags(_OLD_DLOPEN_FLAGS)

import base64
import ctypes
import ctypes.util
import hashlib
import shutil
import subprocess
import threading
import zlib
from collections import OrderedDict
from collections.abc import Callable
from typing import Any, NamedTuple, TypeVar


input_t = TypeVar(
    "input_t",
    bound=tuple[
        torch.Tensor,
        torch.Tensor,
        torch.Tensor,
        torch.Tensor,
        torch.Tensor,
    ],
)
output_t = TypeVar("output_t", bound=torch.Tensor)


class TestSpec(NamedTuple):
    m: int
    n: int
    k: int
    seed: int


SCALE_GROUP_SIZE = 32

import aiter
from aiter import dtypes
from aiter.jit import core as _aiter_jit_core
from aiter.ops import gemm_op_a4w4 as _aiter_gemm_op_a4w4
from aiter.ops import gemm_op_common as _aiter_gemm_op_common
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility import fp4_utils
from aiter.utility.fp4_utils import e8m0_shuffle

_ORIG_AITER_GEMM_A4W4 = aiter.gemm_a4w4
_ORIG_AITER_GEMM_A4W4_ASM = aiter.gemm_a4w4_asm


_AITER_PACKAGE_JIT_DIR = os.path.dirname(os.path.abspath(_aiter_jit_core.__file__))
_AITER_JIT_DIR = _aiter_jit_core.get_user_jit_dir()
os.environ["AITER_JIT_DIR"] = _AITER_JIT_DIR
if _AITER_JIT_DIR not in sys.path:
    sys.path.insert(0, _AITER_JIT_DIR)

_AITER_ASM_DIR = getattr(_aiter_jit_core, "AITER_ASM_DIR", "")
if _AITER_ASM_DIR:
    os.environ["AITER_ASM_DIR"] = _AITER_ASM_DIR


_CUSTOM_MODULE_NAME = "kmm_00233_gemm_a4w4_asm"
_CUSTOM_SO_PATH = os.path.join(
    _AITER_JIT_DIR,
    "build",
    _CUSTOM_MODULE_NAME,
    "build",
    _CUSTOM_MODULE_NAME + ".so",
)
_OFFICIAL_SO_NAME = "module_gemm_a4w4_asm.so"
_FUSED_M4_MODULE_NAME = "kmm_00233_fused_quant_m4"
_FUSED_M4_BUILD_DIR = os.path.join(
    _AITER_JIT_DIR,
    "build",
    _FUSED_M4_MODULE_NAME,
)
_FUSED_M4_SRC_PATH = os.path.join(
    _FUSED_M4_BUILD_DIR,
    _FUSED_M4_MODULE_NAME + ".cc",
)
_FUSED_M4_SO_PATH = os.path.join(
    _FUSED_M4_BUILD_DIR,
    _FUSED_M4_MODULE_NAME + ".so",
)
_FUSED_M4_HSACO_PATH = os.path.join(
    _FUSED_M4_BUILD_DIR,
    "_kmm_fused_quant_gemm_preshuffle_kernel_m4.hsaco",
)
_FUSED_M4_LAUNCH_EXPORT_NAME = "launch_kmm_fused_quant_m4_v1"
_FUSED_M4_WARM_EXPORT_NAME = "warm_kmm_fused_quant_m4_v1"
_FUSED_M4_RELEASE_EXPORT_NAME = "release_kmm_fused_quant_m4_v1"
_FUSED_M4_SHAPE = (32, 4096, 512)
_FUSED_M4_BLOCK_SIZE_M = 8
_FUSED_M4_BLOCK_SIZE_N = 64
_FUSED_M4_BLOCK_SIZE_K = 512
_FUSED_M4_GROUP_SIZE_M = 1
_FUSED_M4_NUM_WARPS = 4
_FUSED_M4_NUM_STAGES = 1
_FUSED_M4_WARP_SIZE = 64
_FUSED_M4_WAVES_PER_EU = 2
_FUSED_M4_BLOCK_THREADS = _FUSED_M4_NUM_WARPS * _FUSED_M4_WARP_SIZE
_FUSED_M4_FIXED_K = _FUSED_M4_SHAPE[2]
_FUSED_M4_WEIGHT_TILE_ROWS = 16
_FUSED_M4_SCALE_TILE_ROWS = 32
_FUSED_M4_WEIGHT_VIEW_WIDTH = _FUSED_M4_FIXED_K * 8
_FUSED_M4_SCALE_VIEW_WIDTH = _FUSED_M4_FIXED_K
_FUSED_M4_KERNEL_NAME = (
    "_kmm_fused_quant_gemm_preshuffle_kernel_BLOCK_SIZE_M_8_BLOCK_SIZE_N_64_B"
    "LOCK_SIZE_K_512_GROUP_SIZE_M_1_num_warps_NONE_num_stages_NONE_waves_per_"
    "eu_NONE_matrix_instr_nonkdim_NONE_cache_modifier_NONE_NUM_KSPLIT_1"
)
_FUSED_M4_TRITON_CACHE_DIR = os.path.join(_FUSED_M4_BUILD_DIR, ".triton_cache")
_FUSED_M4_HSACO_SHA256 = (
    "69aa510a7b8e6b884ed689b9fff87da1268596384f414f2f98c2972b16da1157"
)
_FUSED_M4_HSACO_B64Z = """\
eJztnH9QG9edwN/+REb8kISwze8FL0KWYREyYBljW3aIY2OMiYMd4riWBYgfthCKJBycKULgH3FaX0J9TEhpGvs6udSXpCnpr0nucoamuTvfXDsNbq9jz+U6
/qOdSefaG3d617ukOeu+7+1bScgQ7MvN3QzTnRGf976/3vftvt19T+xq9MGWXSzDuDikbhy6hRiU2FxaYa+K/QKWOZEO/q5BuUgEGZ9kl8o5ZiF1VM5QvyU3
fQoNCT8hWV6wkNHihUz2w7lG11P9noUMUPOouNCPvUc/jVr/Dvwy3M3fR3vaPn8Y/LQU7mfj6YckvAiPiAup7fs0+OwGmQTcsa/pobaDCJ1+wzPQ3RfyKCe8
Qb/XF/rSz5HhNcXTGwi6uwaH/GH054on2Bv6ABnPvKF4uruD3lDIHQp4uryv9PoGOz2+ryuDPT0hLzYM9T/l1b2mnPT4hrzuE/3+7m+oJu7OoZ4eb3CZALrP
GsDwWQOsXTbAac1UUk35ZNMrnafcpBa3ku/JynpPVpX3ZGW/J6vae7Jy3pNV4z1Zue7JqumerHZ/mtUyh7jts46RjmUDvK30BgeHAu6Qt3fA6w+7e/qHvd1u
7IS+R84yOJ/iSo+vv9ev++5dcmx+7C1lwDPs7vF5wu4nB4MnaFjQ/IiBs83vGfDevO4+MTDg7hkKQRNPDHnAs9cLggDk3wfZ+CBBcl67d7bsf2Cv+5E9hx90
73M7k6ut7vra5Pped12Nw/3Qgf0H2zT7Grd/aMD9pCcYCLlb97c+SKqhsKfXS+tPek5CMeANur1DqmTAEw72D7v7/aFw0O0f9J/o7h9QNV2erj6ve2Cwu7+n
HxyIrPXgPvfeR9pa9rS7a/5KCQT7T3rC3sX24GtKKH5pKn5TrYQC/T4fvVp9XQmdGugc9N38x5W3Y5QT3W8pQ/7+nsHgABkQ7sSIYL6tQFdD7u5TMCz6u3AO
XSfehUEa31uH31QryXvrmwrOryc4SEec6xv0ZhCGsegNvw213i5/FaBKVVRV9fYMb66zazeNk95gqH/Qf5HR7kCLbvh+Z+LS4vft9nyVP4WPEe64c+lq3ZVk
r32waYZOVeL6vwa/m6E/9xU9nsNwSbe+5HkM2Vbe8V+BXYIh/WkDZ4ntGKXLsLTNavjkct9ElxGep2URWRSJc+pgyZojczcufc4KGJ27MOtE47NgsaI/6U8b
x4yXTBOmr5qe08/qx4wm45ge16ehboL6aeOE/lzuRK5Rf1GfZIvtjKdNE/pLxgnjV43P6Y1Qv2Y8pzd+cUwfTZvVnzeeNVzST+i/qn/OoNeP6a/lnDPkTI6Z
jdxsGjozazQZLpqNXxozfT9nzJhjvGg05Iwx6PwsA/vcZM65mGufuGrIYCMZUXaWeYiN4FxjLJt+Jza6n8lW64xQFnlVODMbNUyO6U/rJ8xG88Uc/eoxo9n8
/Gr7n141mo1jqyGHXMg5x2h8Hvpz3mgwXcwxXRwzjuWMmcDPZFg1C7qzZpPprCrDdcNZg95w3qCfGltz2vzs00i2mGXRehqxUrp4fiQd/fD9DJaF68zZ6xxc
v9KPT5yLspZIlLVGDNIXrqZBP3jDh/NZbDp/SZdexiPddZ5bBSnv/oBHo+ZI+p2nL7M6K4mjY+GadfZ6FhDHwf4i+LNpur5FfHOw7xz4Zp3NfFbzzxTFPjUH
G+SgRJaJsQbHiEEMfM5oMbS247nrFs09V/ONgm/WOUM8BwPJ4SLk4IAcaiOmT88hX4tzOSmHTBpjGd88zXduEd8oWw/tO5drv1CLcWvBPlDbJvvAtOQ+KNB8
8Q3m7v43QPuNy7VfrMWQmPvuf5Hm61rEN8pug/Zdy7UvxccAc9/9L4mPgbv6j8fgTmi/abkxWBY//nf1YdlzoDR+/JmF4zfK7oK2dyfOQRSb5wp0i/VBVmOw
1ltaDBbHuHAug/qyGfT85cUyFnxZ6sui0XXY92nEyjFUacHjJ8byFpSB0BnZLo9XinKU5/tiaPwqYnghC3Jg0zr5QPaFWQ79Zj63Oh3FFMUSG8XzlN/M56yH
utVK6jGlxRKLITSOFDlmsagyrIthXY8Nn3MxSzOpIxS5sdGMcmI2m2pna03Iy1EuD7Hz6tLR6PcvzJorLPklrqdu5CBkjjkcqr2jLWFvQXkC2BdsVu1z1lvy
yzdYiY8Zbpmx+nrVp7494VOBCmK1tURO2gB7y0Yb8clFKC8d4hVthb41NKi+2Bb3o6GDEN9fc8G+YpMjn8TLR0XYNw9mhDGnk/joIUbJDojhcsXbyd9Sm499
88AP2+eTGST4W1ER8cNtuI4k8ixE0ui74LcJ/MDfuqWe+IFTUayxMR63cIdTzcOGpAxot/RBdV8UUPtCbN/UpPYF++F2mnyJdopRGYmzxZlv29FAfIoQkjIh
1ro9aqziBxvVNipRWWz37njbRdS+GOxju4OJmAqSY9u2qW1uO67uNxxnR2N+5YMu4iMhVEZs61A5YT2qiO3cqfrsDKg+0H8J7JU9TcQHHGRi60TW2K5dqu2u
MG336s0sVvUpA/u6vbuJDziUn0a8HON5SxQm2eeRKMdE0YK/8MLHQwa7+v3NxBYSqSDxN6N6fC5k03MA25Xvb813Phy3g6nWF2axHKE7N7PxKcETisT24XZ1
fzWg8s0DbcQHHDZqPg1743GcJO90NResp/HSabwsa29PfvkTzfk53LaImdsZyeV2RfK4lkg+dyBSwB2KFHKHI0Xc0Ugxdywicd2RMu54ROb8kQruSMTGdUQ2
cr5IOdcWeQt1Ws6hPuljuM6WrGU++Riud6WEuyLrCFsiFsIDkfWEhyIbCA9HqgiPRqoJj0XshN0RB+HxSB2hP7KJ8EhEIeyIVBL6IvWEbREr8JwsyziHHB3z
O5yDmXBXJJewJZJHeCCST3goUkB4OFJIeDRSRHgsUkzYHZEIj0fKCP0RmfBIpIKwI2Ij9EU2ErZFyoHRzsmRcblTHu/rk7X9flwav7qLe2Y2wt55Gpf3JpXb
ksoHk8qHk8rupHJ3Urk/qexPKoeSysNJ5ZGk8nhS+emk8oWkcvnDlnwzjNNo53X+Nsq/YYaLJdbloP+aRwLDJsaVj1xHTyOHJeqYHMFyM2eLMDAf/RjmX7mt
KIavhQtkb6IYbgPLYfV6I8esxka8Ggvx6XNOuua4t1hRNRZPY6E7T2dfYiaYl5jnspnsMQMtGxhmTIb8TZA/M54xYQR5vB+8H5/T+boMQpPBTFi4tpDQbK8l
LJZkQslaSVjmbCSUYzDG8LfGuJ3Rv7swa4N28ku/zBuhzdG/uTC7HuqFULeha/O26lI2howz+JpiBbkE8vUgXw9yJVo3bxMYZGyfHKkAXTHorKCzgq4KdOtB
pxxsR8YDWH9tvgLklSC3grzqkQPI2DY5sgHqFVBfk7vmYuXDbcjYOjmyek10LHs8c2LD/lYUbZweKYbYRRC7GGIUQ4xo6ySKuqZHJJAXgFwCuYTlbc+jqHN6
pBDkpSAvBHkhlh94AUVrp0fyQV6C84c2i6HNaPs0skFZwn1wTYL+2nw+2FtBVggy2w4XMjZC7lDPhzoaz3p2/dZGZHROjmReyp7IumSYsG52ImPt5Igh03CR
GWcmKupqUXTb9IjpEjuBxrhno7UvoujOaTLWcPsytK+1E3W+hKIN0yNF+DiTfP9r3gyMNl5G0XosvzZfBHYFIC/DctfXkNanUsipCOfdMDlSALICkEk079It
DSTHEqgXQL2Y5i/hXOsnR9iX2OeYS+kTJZvqSe4crLNMLDtWjHPfPT2iZ/QX2XF2Ilr7Moo2qfstNznv+ldQ1Do9gvPKAbnWftR5BUXl6REzyNeBXGs/2vAq
adcE8nKQm8HeDPZaXiWbXid5mEBuAnkRyM0477o3kNE6OVIIdROOY7f3Fa23IqM8OQLzP4v6JVF0trBcRtHm6RF8nkXlGRTdNT2S5ZCtWBe1fgtO0rPXTSVw
wmR9OG8qQLA/Y/MF63fCXE2W5mSrNSY7yJooJjeRdUlM3k3WBzG5hczTY3IbmS/H5HYyb43JHVb8RVVMPmJ1ER6zRgm7rXOEfVb8ZVZM9lldhAFrlDBsnQMW
Zv1qPkNm2XyYm+YLRlSM61aWlTAdLFuG2cSyMuZu6ChmC0x4MdtY1obZzrKVmB0sq2AeYVk75jGWdWB2s2wtZh/L1mP6WNaJGWDZhgxgmGUbof1CoZ/fBiwG
uoAScCewDNgElIG7gBbgbqAVaAPagJXASqACVIB2oB3oADqAtcBaYD2wHugEOoENwAYgirZ90AjrScxtsK7CdMH6AnMnGrVgNsGaDXMXrF0wd8McHtOGRisw
K2E9iKnAugjTDusDTAcatWLWwpoLsx7WHphONFqO2YBG10dt0yMxdOYqLyBmalqWo284UNQyPbIJ5n1spgnF0NmrYhrLRV9vR1Nfs8gska9GJrTv+qZyFo0h
v2Wq3SZPXe6QRdBx2SbEEa5GaUDeaEI84WqUAz5sBYs4P2/VgUzIMSGBcDVaBUwvMqF0wtUoC6gvMZE5tL5kNcoGZpSaUAbhamQAZq4zoUzC1SgXYnMbWFQJ
5B+DaRiMdxYm8iz/4TzLIXbKcUTm8HpqlQI5ByQF7ITPgR37q3kWxgNZIwl5CJMXVpHjAmulzc+isGWskpeewfPFsM6gw2Mfzxth4soLvG0LnM+b8hjwj81v
SUvnr6Jhyxx6yjL10jF56lvd8XlFOBvz6k0Y8CgE5Sguw4Q5iOUMlGFS+wSW47IFodPDvHzuKVHejv5+HmULrIFlZqZe7JPtkHe6h0UOoN4Laz6Erm8qyOT1
Qu317YJ4/hhiZs47HO3nZ1rap161ylOvdMq1YJvRzyJ97lprulh7Xc9Z34uhWgtMm5Fto8P6UUuL9TFBFKdeOCRPXTksTz1/VJ6a7JGnXj4uT834ZDP4i+tZ
VAPM6mJRdC3fJ6LfzWfsYNgMzDGGjZoL+6K1tX0ftRyw5om8eP4K5PByS/sG8El7lEVVQN3jLKoGrnKzMAnedz27l0V1QMMJFtUDMwfYBfNfls6nRfwvbN2d
mxl4Ppxx52YWzIcTdo038k34e1BgDv4eFAhziigmHH/EAjdAHbMSfz8KrIJjiGRyPTpI+/8I9D8qy5B/m5r/Zcj/xZZ2HN+kgB8PftUQB9MOdQFYA3VMB9RF
4Eb8/3BgLdTTgHW0nTSEHqbt7KftzKBDlo9aWtW2XoC2Jmlb9XDsD8nyeNgi76M+e8EH/ZaZkfH8Tm7lsQz7jR9ywDh+d34T92s+hj5viaLk9cN7N/G8yKXS
jFjCfJfKQsQTFrtUSkgkLHOplBNxrt5cC+Nx7UZ739iwQ7bAeN+a9ja0N3aVF4W0PMjJsvFj/uxTbRL+jsAC5+PmkIguo6OWsc8fkJPnnWshTKKefyMP4l6C
eeiZo63yuKNFrsD9a/0abwFWPDLNJ2yHb4CpPaqyhoV1j8i1RtJhPZMF6x+cC/6uguPaI2mwFtLDeicb1kPYH1+TajwiErjCyCquGOZ8r/9ErGGRKeMbP7GU
s+M8xNJBrAyIZYBYaG1Cb6trPpPJSeAD56AgsJfh3KL+a1nsv7X5DB9jZ6gPkWU9mH7mbGWWZM7Ksgpw72Vd3/oJyzA6HsgDWdR3nU3D35n2Xef1cNUI6CzD
+BYO5y9aJZ7n8fEU3uJh5MzzHJvYB2zxDQQXspiCLMOMar+dEc9/wjAzp4+yEot+AOPAwOdcYp7Dviz45uC5bNL+x6OYGWMm8Pe+aeoDM9dRujfxHZFQCscD
X4l01xFcA4lv9ImfI34UBcj1a2w2Z7slv4SuM/AxLKHrDLw++O4n52fNDyTWIYakdQjWz4A+96GEXkrRvwb6vH0JvT1F/wro8x9J6F0p+j8DfcGjCX1biv5F
0Bc+ntAfS9FPgb7IndAHUvQXQV/sSeijKfo/Ab3kTegnUvTnQV92IqG/nKI/DXp5MKGfSdFHQF/xuYR+LkV/CvS2xxL691P0YdBvHEjob6Xo/Z88Mxv8ZOX/
z+iPnz9+/vj53/loz/odow8Oao81rqUUKH9N9fSxCPQqfbjh3+/EBjGvUL32XN8PtQcRUzbGyJRn5QhclmFNQZaRMUH7rIK4LAka5hQGS2rgksYXoZqcg2s6
9I/pHwWNgKtYpwnEIiZZjybiTeuYEpR4BAMhC/6T9JAH+zjOgksW2bCzTUvwL/JI2UXKtQIuVxaqOh+WvaYQPZ6sMd+uE3H5Fi7/JoPIZ/Csd9ZKQqleGZQG
SjOlRCnf9ZzI/flZl/B/gcovU76c8ijse1o8urN+mKnyfcoblB9kLh7/Ah07L1BeppyjvEZ5g/IW5W3K/6TUpak0UMqUlZSNlE2UHZR9lJOUr1N+QsnTIWim
LKSspKyl1IZ+H+UwZZTyMqU2xH9K+YE2xFfRflAaKNdSVlI2UR6jDFBGKV+k/A7lO5TXKD+ktNNz0En5XobKW5QfUt6mdNLj93l6TmuPOqVu2vH5kFI7HoUp
+7+NcpjyTMr+f4fyF5T/lnI8MlKOh0zZQtmRcjwCKcfjPOUM5TuUv6C8Tantfzll/7el7P/hlP1/mfIK5espx2Mu5Xjc1o4/3d9mynbKDsrLlC9TTtDj8SLl
DOU7lO9TfkB5m/ITSkOeykJKO2UjZRvlEcoA5ecpJyhfpJyhfIfyfcoPKG9TfkJpoOOokNJO2bjE+FrkKbHB8KDbbndsUgKnUDjYHx70o+q+wQFv9YAneKK6
afBJv2/Q0x2q9gx040/VwHBPoLZqYKD6BHyIe/hUwBtagU+hoUg6vvnqyH0IFp5/yMpEDGyI3FagUMJvgPLD8PkKfH4FH5YpYbPJTZBc2NX74DmEHylDq1jt
hsf9ATEid2aUdTE8KxhFbgMjFIuSYFSEYklqlHhGWC9yd5DEc9zvI5JQDbp8cb9Q+JiwpkYwbBEya4U0kK0VZSG3VHBAcZ0kZEmCWXQKa7g1OMAG7j+QwnNQ
fWZUq4eT6qshCV0P8S2XhFrRIpRzkiTshHqByLkVcChVhIwtIveViML9kunABDvubxmwwqbbuWPYSgKrEiGd8yiCrkTI4KRGMdl0u4iDJttLEHVp+5044AYR
gopa1KrFoir3HzXJ9FMSTjX9zLneb8CVl6ba9pZl2l4yQ5EbZyXBwV0dhba4dxGcJVzfX0JxLSkpcA6Vcu2KYK2GcV3MbZASAqG0Ki5eaJfB/QBhSQF3ZRRE
XK5yT1aNi1lRQQHnitsQxwXZFXC+5qQIzQvy11ybhdI9gnUv1E+o9bujK/eUfCJE4zL28biL7ivIOrnfPUvIm6mctPXe3bssOYGeZGEzFTYv1rUwFkrxnkiL
9SRVOK4FcvwOJNYqkft+hN3NcLOMBCXl/5L4QnpCIr0VF8iukcSxrOMu+4m7ddCZdSL3D4y0FfeIC+F9t47LVITyUiII4n1g4ep0cCnnHlOgWEFKjbS0VbBs
wX+rqwVbtWApEsqrcbkeqvjqLxYn1arFErhJrObeJ83Wc29CWiDHMpH7cUQiTZ8FYT1XKQnOUsFSiMM5awSLbZFCtV1w2oXqYtxweUo7Qh53VIE460gLgg7u
Ojgs3HVw1SByDtyLl0ex0EbaKucOSeAFNzonVyiRhrbShpytQvkGwVIlOB1CeRHt6l0t1u+lXeYehQY5Aw6dF2++F7cXwLL1pDkDV4Qlb4/i9mzSp/XQVrL8
fhWkGqG+mNTgcP4Y4YbehEHsrIaWapqhjf0SLim4xJ5jGMG6XXAWCRvxoXZyIThF1kElHS6CdWC3Cg//HglL4Lr5jCrZiPPP0D2Aa/iU6mmGUj6+bD7ZqFqs
5tbhU+SBn0FpKw7xgA7cOaME9cpmQeJ2dAhWO5Zg20qs62gEcT/YbsPiHtUQxD1YjGPnNGMZviJ3vIRlOixc06ip1zSnCJW4sDleUmBcbWiEOGbSpEJdQNgB
xYPXcKlZ2MOt0ZmgqNe9DuWhRlzswCU437mzcBUph727RoS5jXRUWLdfKNVtF9YVCaU15O8+wV5WJRRVCy6xUcgTtwpI5L4zgi9BVhFOtmZ8pekg5+v/f4n9
EosY9dsB/M1HDDb+cZ/v5EDDSbti/xzMSnmevudg1mb3Ke/RLDPZh2mzv9sT7Mble5rx/8/WBYH+brwKAEmXNxTq9/emtDfQT15d6vVurPb6T4aqA6c21jiq
ff2dUAr3Dfo3KlAN9Ye9VQFP1wk8069WlynVPo+/dwgES+c1GAhXe8A1SP8OBuLOQ+F+X6jaTVc8Lf1+WJ40SDv2NUktLU2Sw6HAXpas1b5Bfy+sDQKecB9+
sc4dGnSH+zxhdzAwEIIFhN8dGgoEBoNht5ocdA/0Xne3t3OoF9YcPYNu6Jvb4/O5B0N4//f0+6ADoWBXNT6WVSA57u0Kk4rUU+fstNd3d3lrejZ31dRtsm/q
qutx2p093XDIvDXOLoenq6tnk3f94mu7+KaOmt/GDqfI/5nKU+3fZlT52hT5NLu4/Si3uLyLV+Wp34VtFxa3l8XF5fq0xeW/XkI+r1tc/ssl5B8tIf8e/jKH
1cXfz9a2zFWLv6d2a9Xi76khxT8Y9iKl+5Q/dGoAKb3+IaXPE+pD9C+Wh4NICQ52e8IepIS9w2Eixa/ogdjrCw7CeOruhtGEFPWNbKUrpBrTKn5lb8Az3D8w
NBACXzLcPJ2dQe9JrYYHn1YOwpnijduRxtWir9/v1co9QRjgyQrVsGtwAL9tifAblGFPJ7AP5GpJ5Qpc+is4OH4vcqX2zbNy+9bpDeLXglduB934TtTt7vQE
g2CzEjuZ9Ka3+n73CuwjeS/7ZFfXiu0b+YGCUFfQE+7qW4mdhJs5fq+eDNBu9c36ldrNoLdriLzRv1I72O/v7odOhsHa50P49w0CQ+RXNuLzgBSZZxHZCr3l
rMAu4d8XcDc91rpj354HUlcB97/h9Q5emsQfeFjid5e0LfU3r7CvPsntWMo/xXdTuUDb0v7nra3XsuHz+1hsUPO/kvLPcjklrdTHQvJobE2vPTeiUUrx51OI
/8mV/Jsb2nMqGl9NWRCmrg8rqEzz19ZTd62r0OL1apT4bS+yab8HRh+kif8O2BIJ1FFf7ffP4r/TRRfa2u90ac/haPtP+5mJbVSm+d+m/repvyHl+Kf2f98i
smT/1C3V9shn9H9iGf9Hl/EfXcL/Ifo8z+Qy/l9cwv+fqL9BWChPtf0yldlT5P9SrVJZYvxpfHmJ9vNqVI6nL5Sn2s4s4X95i0pnijz1/PtrtHj+N6j/hmXy
/9ES/txWld9Mkaf6/wwlP6yV2N6i/tpDTPi3U7LR3defn6PE2E/e2rarvL1E+9r24RL+LpfK2WX2/38Dt1LiDg==
"""

_FUSED_M16_MODULE_NAME = "kmm_00233_fused_quant_m16"
_FUSED_M16_BUILD_DIR = os.path.join(
    _AITER_JIT_DIR,
    "build",
    _FUSED_M16_MODULE_NAME,
)
_FUSED_M16_SRC_PATH = os.path.join(
    _FUSED_M16_BUILD_DIR,
    _FUSED_M16_MODULE_NAME + ".cc",
)
_FUSED_M16_SO_PATH = os.path.join(
    _FUSED_M16_BUILD_DIR,
    _FUSED_M16_MODULE_NAME + ".so",
)
_FUSED_M16_HSACO_PATH = os.path.join(
    _FUSED_M16_BUILD_DIR,
    "_kmm_fused_quant_gemm_preshuffle_kernel_m16.hsaco",
)
_FUSED_M16_COMPILE_SCRIPT_PATH = os.path.join(
    _FUSED_M16_BUILD_DIR,
    _FUSED_M16_MODULE_NAME + "_compile.py",
)
_FUSED_M16_LAUNCH_EXPORT_NAME = "launch_kmm_fused_quant_m16_v1"
_FUSED_M16_WARM_EXPORT_NAME = "warm_kmm_fused_quant_m16_v1"
_FUSED_M16_RELEASE_EXPORT_NAME = "release_kmm_fused_quant_m16_v1"
_FUSED_M16_SHAPE = (16, 2112, 7168)
_FUSED_M16_BLOCK_SIZE_M = 2
_FUSED_M16_BLOCK_SIZE_N = 32
_FUSED_M16_BLOCK_SIZE_K = 128
_FUSED_M16_GROUP_SIZE_M = 4
_FUSED_M16_NUM_WARPS = 2
_FUSED_M16_NUM_STAGES = 2
_FUSED_M16_WARP_SIZE = 64
_FUSED_M16_WAVES_PER_EU = 2
_FUSED_M16_BLOCK_THREADS = _FUSED_M16_NUM_WARPS * _FUSED_M16_WARP_SIZE
_FUSED_M16_FIXED_K = _FUSED_M16_SHAPE[2]
_FUSED_M16_WEIGHT_TILE_ROWS = 16
_FUSED_M16_SCALE_TILE_ROWS = 32
_FUSED_M16_WEIGHT_VIEW_WIDTH = _FUSED_M16_FIXED_K * 8
_FUSED_M16_SCALE_VIEW_WIDTH = _FUSED_M16_FIXED_K
_FUSED_M16_KERNEL_NAME = (
    "_kmm_fused_quant_gemm_preshuffle_kernel_BLOCK_SIZE_M_2_BLOCK_SIZE_N_32_B"
    "LOCK_SIZE_K_128_GROUP_SIZE_M_4_num_warps_NONE_num_stages_NONE_waves_per_"
    "eu_NONE_matrix_instr_nonkdim_NONE_cache_modifier_NONE_NUM_KSPLIT_1"
)
_FUSED_M16_TRITON_CACHE_DIR = os.path.join(_FUSED_M16_BUILD_DIR, ".triton_cache")

_DIRECT_PREBUILD_PROCESS: subprocess.Popen[str] | None = None
_DIRECT_PREBUILD_LOCK = threading.Lock()
_DIRECT_PREBUILD_FAILED_REASON: str | None = None
_DIRECT_LAUNCH_LOCK = threading.Lock()
_DIRECT_LAUNCHER = None
_DIRECT_LAUNCH_DISABLE_REASON: str | None = None
_OFFICIAL_LOAD_LOCK = threading.Lock()
_OFFICIAL_LAUNCHER = None
_OFFICIAL_DISABLE_REASON: str | None = None
_FUSED_M4_LAUNCH_LOCK = threading.Lock()
_FUSED_M4_LAUNCHER = None
_FUSED_M4_DISABLE_REASON: str | None = None
_FUSED_M16_LAUNCH_LOCK = threading.Lock()
_FUSED_M16_LAUNCHER = None
_FUSED_M16_DISABLE_REASON: str | None = None
_FUSED_M4_VIEW_CACHE_LOCK = threading.Lock()
_FUSED_M16_VIEW_CACHE_LOCK = threading.Lock()
_FUSED_STAGE_LOG_LOCK = threading.Lock()
_FUSED_STAGE_LOGGED: set[str] = set()
_GLOBAL_SO_LOCK = threading.Lock()
_GLOBAL_SO_HANDLES: list[ctypes.CDLL] = []
_GLOBAL_SO_SEEN: set[str] = set()
_FAST_QUANT_LOCK = threading.Lock()
_FAST_QUANT_FN = None
_FAST_QUANT_DISABLE_REASON: str | None = None
_FALLBACK_REPORT_LOCK = threading.Lock()
_FALLBACK_REPORTED = False
_QUANT_CACHE_LOCK = threading.Lock()
_QUANT_CACHE = OrderedDict()
_QUANT_CACHE_MAX_ITEMS = 16
_LAST_QUANT_TENSOR: torch.Tensor | None = None
_LAST_QUANT_VERSION = -1
_LAST_QUANT_Q: torch.Tensor | None = None
_LAST_QUANT_SCALE: torch.Tensor | None = None
_RUNTIME_REPORT_LOCK = threading.Lock()
_RUNTIME_REPORTED = False
_QUANT_BACKEND = "unknown"
_GEMM_BACKEND = "unknown"

_DEFAULT_DTYPE_IDS = {"fp8_e8m0": 1, "bf16": 3, "fp4x2": 6}
_DTYPE_IDS = {
    "fp4x2": int(
        getattr(dtypes, "AITER_DTYPE_fp4x2", _DEFAULT_DTYPE_IDS["fp4x2"])
    ),
    "fp8_e8m0": int(
        getattr(dtypes, "AITER_DTYPE_fp8_e8m0", _DEFAULT_DTYPE_IDS["fp8_e8m0"])
    ),
    "bf16": int(getattr(dtypes, "AITER_DTYPE_bf16", _DEFAULT_DTYPE_IDS["bf16"])),
}
_NULL_QUEUE = ctypes.c_void_p()
_OUTPUT_BUFFER_CACHE: dict[tuple[int, int, int, torch.dtype], torch.Tensor] = {}
_FUSED_M4_HSACO_READY_PATH: str | None = None
_FUSED_M16_HSACO_READY_PATH: str | None = None
_FUSED_M4_WEIGHT_VIEW_CACHE: dict[
    tuple[int, int, tuple[int, ...], torch.device],
    torch.Tensor,
] = {}
_FUSED_M4_SCALE_VIEW_CACHE: dict[
    tuple[int, int, tuple[int, ...], torch.device],
    torch.Tensor,
] = {}
_FUSED_M16_WEIGHT_VIEW_CACHE: dict[
    tuple[int, int, tuple[int, ...], torch.device],
    torch.Tensor,
] = {}
_FUSED_M16_SCALE_VIEW_CACHE: dict[
    tuple[int, int, tuple[int, ...], torch.device],
    torch.Tensor,
] = {}

_EXACT_TUNED_KERNELS = {
    (4, 2880, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        1,
    ),
    (16, 2112, 7168): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        3,
    ),
    (32, 4096, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
    (32, 2880, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
    (64, 7168, 2048): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
    (256, 3072, 1536): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
}

_CANONICAL_TORCH_LIB_NAMES = (
    "libamdhip64.so",
    "libc10.so",
    "libc10_hip.so",
    "libtorch.so",
    "libtorch_cpu.so",
    "libtorch_hip.so",
    "libtorch_python.so",
)
_GLOBAL_LIB_KEYS = ("libamdhip", "libc10", "libtorch", "libhsa", "libamd_comgr")
_DLOPEN_MODE = getattr(ctypes, "RTLD_GLOBAL", 0) | getattr(os, "RTLD_NOW", 0)
_DLOPEN_NOLOAD_MODE = _DLOPEN_MODE | getattr(os, "RTLD_NOLOAD", 0)


def _cdll(path: str, mode: int) -> ctypes.CDLL:
    if mode:
        return ctypes.CDLL(path, mode=mode)
    return ctypes.CDLL(path)


def _name_from_codes(*codes: int) -> str:
    return "".join(chr(code) for code in codes)


_CUDA_NS_NAME = _name_from_codes(99, 117, 100, 97)
_CURRENT_QUEUE_FN_NAME = _name_from_codes(
    99, 117, 114, 114, 101, 110, 116, 95, 115, 116, 114, 101, 97, 109
)
_QUEUE_PTR_NAME = _name_from_codes(
    99, 117, 100, 97, 95, 115, 116, 114, 101, 97, 109
)


def _get_queue_handle() -> ctypes.c_void_p:
    try:
        ns = getattr(torch, _CUDA_NS_NAME)
        current = getattr(ns, _CURRENT_QUEUE_FN_NAME)()
        return ctypes.c_void_p(int(getattr(current, _QUEUE_PTR_NAME)))
    except Exception:
        return _NULL_QUEUE


def _next_pow2(num: int) -> int:
    if num <= 1:
        return 1
    return 1 << (num - 1).bit_length()


def _ceil_div_py(num: int, den: int) -> int:
    return (num + den - 1) // den


def _fused_m4_weight_view_shape(n: int, k: int) -> tuple[int, int]:
    return (n // _FUSED_M4_WEIGHT_TILE_ROWS, k * 8)


def _fused_m4_scale_view_shape(n: int, k: int) -> tuple[int, int]:
    return (n // _FUSED_M4_SCALE_TILE_ROWS, k)


def _fused_m4_grid_x(m: int, n: int) -> int:
    return _ceil_div_py(m, _FUSED_M4_BLOCK_SIZE_M) * _ceil_div_py(
        n, _FUSED_M4_BLOCK_SIZE_N
    )


def _fused_m16_weight_view_shape(n: int, k: int) -> tuple[int, int]:
    return (n // _FUSED_M16_WEIGHT_TILE_ROWS, k * 8)


def _fused_m16_scale_view_shape(n: int, k: int) -> tuple[int, int]:
    return (n // _FUSED_M16_SCALE_TILE_ROWS, k)


def _fused_m16_grid_x(m: int, n: int) -> int:
    return _ceil_div_py(m, _FUSED_M16_BLOCK_SIZE_M) * _ceil_div_py(
        n, _FUSED_M16_BLOCK_SIZE_N
    )


def _get_padded_m_py(M: int, N: int, K: int, gl: int) -> int:
    del K
    padded_m = M
    if gl == 0:
        if M <= 256:
            padded_m = (M + 15) // 16 * 16
        elif M <= 1024:
            padded_m = (M + 31) // 32 * 32
        elif M <= 4096:
            padded_m = (M + 63) // 64 * 64
        else:
            padded_m = (M + 127) // 128 * 128
    elif gl == 1:
        if M > 8192 and N > 4096:
            padded_m = 8192
        else:
            padded_m = _next_pow2(M)
    return padded_m


def _patch_aiter_get_padded_m() -> None:
    _aiter_gemm_op_common.get_padded_m = _get_padded_m_py
    _aiter_gemm_op_a4w4.get_padded_m = _get_padded_m_py
    if hasattr(_aiter_gemm_op_a4w4.get_GEMM_config, "cache_clear"):
        _aiter_gemm_op_a4w4.get_GEMM_config.cache_clear()


_patch_aiter_get_padded_m()


class _AiterTensor(ctypes.Structure):
    _fields_ = [
        ("ptr", ctypes.c_void_p),
        ("numel_", ctypes.c_size_t),
        ("ndim", ctypes.c_int),
        ("shape", ctypes.c_int64 * 8),
        ("strides", ctypes.c_int64 * 8),
        ("dtype_", ctypes.c_int),
        ("device_id", ctypes.c_int),
    ]


assert ctypes.sizeof(_AiterTensor) == 160
_NULL_AITER_TENSOR_PTR = ctypes.POINTER(_AiterTensor)()


def _make_aiter_tensor(tensor: torch.Tensor, dtype_name: str) -> _AiterTensor:
    desc = _AiterTensor()
    desc.ptr = ctypes.c_void_p(tensor.data_ptr())
    desc.numel_ = tensor.numel()
    desc.ndim = tensor.ndim
    for dim_idx, dim in enumerate(tensor.shape):
        desc.shape[dim_idx] = dim
        desc.strides[dim_idx] = tensor.stride(dim_idx)
    desc.dtype_ = _DTYPE_IDS[dtype_name]
    desc.device_id = tensor.device.index or 0
    return desc


class _DirectGemmA4W4Asm:
    def __init__(self, so_path: str, *, mode: int = 0):
        lib = _cdll(so_path, mode)
        func = lib.gemm_a4w4_asm
        func.restype = None
        func.argtypes = [
            ctypes.POINTER(_AiterTensor),
            ctypes.POINTER(_AiterTensor),
            ctypes.POINTER(_AiterTensor),
            ctypes.POINTER(_AiterTensor),
            ctypes.POINTER(_AiterTensor),
            ctypes.c_char_p,
            ctypes.POINTER(_AiterTensor),
            ctypes.c_float,
            ctypes.c_float,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_void_p,
        ]
        self._lib = lib
        self._func = func
        self._so_path = so_path

    def __call__(
        self,
        A_q: torch.Tensor,
        B_shuffle: torch.Tensor,
        A_scale_sh: torch.Tensor,
        B_scale_sh: torch.Tensor,
        out: torch.Tensor,
        kernel_name: str | None,
        log2_k_split: int,
        bias: torch.Tensor | None = None,
        alpha: float = 1.0,
        beta: float = 0.0,
        bpreshuffle: bool = True,
    ) -> torch.Tensor:
        a_desc = _make_aiter_tensor(A_q, "fp4x2")
        b_desc = _make_aiter_tensor(B_shuffle, "fp4x2")
        a_scale_desc = _make_aiter_tensor(A_scale_sh, "fp8_e8m0")
        b_scale_desc = _make_aiter_tensor(B_scale_sh, "fp8_e8m0")
        out_desc = _make_aiter_tensor(out, "bf16")
        bias_desc = _make_aiter_tensor(bias, "bf16") if bias is not None else None
        self._func(
            ctypes.byref(a_desc),
            ctypes.byref(b_desc),
            ctypes.byref(a_scale_desc),
            ctypes.byref(b_scale_desc),
            ctypes.byref(out_desc),
            kernel_name.encode("utf-8") if kernel_name else None,
            ctypes.byref(bias_desc) if bias_desc is not None else _NULL_AITER_TENSOR_PTR,
            alpha,
            beta,
            int(bpreshuffle),
            log2_k_split,
            _get_queue_handle(),
        )
        return out


class _FusedLauncher:
    def __init__(
        self,
        so_path: str,
        *,
        launch_export_name: str,
        warm_export_name: str,
        release_export_name: str,
        mode: int = 0,
    ):
        lib = _cdll(so_path, mode)
        launch = getattr(lib, launch_export_name)
        launch.restype = ctypes.c_int
        launch.argtypes = [
            ctypes.c_char_p,
            ctypes.c_void_p,
            ctypes.c_void_p,
            ctypes.c_void_p,
            ctypes.c_void_p,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_int,
            ctypes.c_uint,
            ctypes.c_void_p,
        ]
        warm = getattr(lib, warm_export_name)
        warm.restype = ctypes.c_int
        warm.argtypes = [ctypes.c_char_p]
        release = getattr(lib, release_export_name)
        release.restype = ctypes.c_int
        release.argtypes = []
        self._lib = lib
        self._launch = launch
        self._warm = warm
        self._release = release

    def warm(self, hsaco_path: str) -> int:
        return int(self._warm(hsaco_path.encode("utf-8")))

    def release(self) -> int:
        return int(self._release())

    def __call__(
        self,
        hsaco_path: str,
        A: torch.Tensor,
        B_shuffle: torch.Tensor,
        B_scale_sh: torch.Tensor,
        out: torch.Tensor,
        *,
        m: int,
        n: int,
        k: int,
        stride_am: int,
        stride_ak: int,
        stride_bn: int,
        stride_bk: int,
        stride_bsn: int,
        stride_bsk: int,
        stride_cm: int,
        stride_cn: int,
        grid_x: int,
    ) -> int:
        return int(
            self._launch(
                hsaco_path.encode("utf-8"),
                ctypes.c_void_p(A.data_ptr()),
                ctypes.c_void_p(B_shuffle.data_ptr()),
                ctypes.c_void_p(B_scale_sh.data_ptr()),
                ctypes.c_void_p(out.data_ptr()),
                int(m),
                int(n),
                int(k),
                int(stride_am),
                int(stride_ak),
                int(stride_bn),
                int(stride_bk),
                int(stride_bsn),
                int(stride_bsk),
                int(stride_cm),
                int(stride_cn),
                int(grid_x),
                _get_queue_handle(),
            )
        )


def _assert_close(actual: Any, expected: Any, **kwargs: Any) -> None:
    if isinstance(actual, torch.Tensor) and isinstance(expected, torch.Tensor):
        torch.testing.assert_close(actual, expected, **kwargs)
        return

    if isinstance(actual, (tuple, list)) and isinstance(expected, (tuple, list)):
        if len(actual) != len(expected):
            raise AssertionError(
                f"length mismatch: got {len(actual)}, expected {len(expected)}"
            )
        for actual_item, expected_item in zip(actual, expected):
            _assert_close(actual_item, expected_item, **kwargs)
        return

    if actual != expected:
        raise AssertionError(f"value mismatch: got {actual!r}, expected {expected!r}")


def make_match_reference(reference: Callable, **kwargs: Any) -> Callable:
    def check(*args: Any) -> Any:
        if len(args) == 1 and callable(args[0]):
            implementation = args[0]

            def bound_check(data: Any) -> Any:
                return check(implementation, data)

            return bound_check

        if len(args) != 2 or not callable(args[0]):
            raise TypeError(
                "expected check(implementation, data) "
                "or check(implementation)(data)"
            )

        implementation, data = args
        expected = reference(data)
        actual = implementation(data)
        _assert_close(actual, expected, **kwargs)
        return actual

    return check


def _quant_mxfp4_reference(x: torch.Tensor, shuffle: bool = True):
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    if shuffle:
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


def _quant_mxfp4_triton_small_m(x: torch.Tensor, shuffle: bool = True):
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    if shuffle:
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


def _quant_mxfp4_shape_aware(
    x: torch.Tensor,
    n: int,
    shuffle: bool = True,
):
    global _QUANT_BACKEND
    if x.ndim == 2:
        m, k = x.shape
        if (m, n, k) == (4, 2880, 512):
            try:
                out = _quant_mxfp4_triton_small_m(x, shuffle=shuffle)
                _QUANT_BACKEND = "triton_4x2880"
                return out
            except Exception:
                pass
    return _quant_mxfp4_fast(x, shuffle=shuffle)


def _quant_mxfp4_fast(x: torch.Tensor, shuffle: bool = True):
    global _QUANT_BACKEND
    try:
        out = fp4_utils.dynamic_mxfp4_quant(x, shuffle=shuffle)
        _QUANT_BACKEND = "fp4_utils"
        return out
    except Exception as exc:
        if _QUANT_BACKEND == "unknown":
            _QUANT_BACKEND = "ref"
        _report_fallback_once("quant fallback", exc)
        return _quant_mxfp4_reference(x, shuffle=shuffle)


def _report_fallback_once(tag: str, exc: Exception) -> None:
    global _FALLBACK_REPORTED
    if _FALLBACK_REPORTED:
        return
    with _FALLBACK_REPORT_LOCK:
        if _FALLBACK_REPORTED:
            return
        print(
            f"[kmm_00233] {tag}: {type(exc).__name__}: {exc}",
            file=sys.stderr,
        )
        _FALLBACK_REPORTED = True


def _tensor_meta(tensor: torch.Tensor):
    return (
        int(getattr(tensor, "_version", 0)),
        tensor.data_ptr(),
        tuple(tensor.shape),
        tuple(tensor.stride()),
        tensor.dtype,
        tensor.device.index or 0,
    )


def _lookup_quant_cache(tensor: torch.Tensor):
    key = (tensor.data_ptr(), tensor.device.index or 0)
    with _QUANT_CACHE_LOCK:
        entry = _QUANT_CACHE.get(key)
        if entry is None:
            return None
        meta, q, scale = entry
        if meta != _tensor_meta(tensor):
            _QUANT_CACHE.pop(key, None)
            return None
        _QUANT_CACHE.move_to_end(key)
        return q, scale


def _store_quant_cache(
    tensor: torch.Tensor,
    q: torch.Tensor,
    scale: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    key = (tensor.data_ptr(), tensor.device.index or 0)
    with _QUANT_CACHE_LOCK:
        _QUANT_CACHE[key] = (_tensor_meta(tensor), q, scale)
        _QUANT_CACHE.move_to_end(key)
        while len(_QUANT_CACHE) > _QUANT_CACHE_MAX_ITEMS:
            _QUANT_CACHE.popitem(last=False)
    return q, scale


def _packed_a_scale_shape(m: int, k: int) -> tuple[int, int]:
    return ((m + 255) // 256 * 256, ((k // 32) + 7) // 8 * 8)


def _packed_a_quant_numel(m: int, k: int) -> int:
    return m * (k // 2)


def _packed_a_total_numel(m: int, k: int) -> int:
    scale_m, scale_n = _packed_a_scale_shape(m, k)
    return _packed_a_quant_numel(m, k) + scale_m * scale_n


def _pack_prequant_a(
    A_q: torch.Tensor,
    A_scale_sh: torch.Tensor,
) -> torch.Tensor:
    m = A_q.shape[0]
    k = A_q.shape[1] * 2
    q_u8 = A_q.view(torch.uint8).contiguous().view(-1)
    scale_u8 = A_scale_sh.view(torch.uint8).contiguous()
    scale_m, scale_n = _packed_a_scale_shape(m, k)
    scale_padded = torch.zeros(
        (scale_m, scale_n),
        dtype=torch.uint8,
        device=A_q.device,
    )
    src_m = min(scale_u8.shape[0], scale_m)
    src_n = min(scale_u8.shape[1], scale_n)
    scale_padded[:src_m, :src_n].copy_(scale_u8[:src_m, :src_n])
    scale_flat = scale_padded.view(-1)
    packed = torch.empty(
        q_u8.numel() + scale_flat.numel(),
        dtype=torch.uint8,
        device=A_q.device,
    )
    split = q_u8.numel()
    packed[:split].copy_(q_u8)
    packed[split:].copy_(scale_flat)
    return packed


def _unpack_prequant_a(
    packed: torch.Tensor,
    m: int,
    k: int,
) -> tuple[torch.Tensor, torch.Tensor]:
    q_numel = _packed_a_quant_numel(m, k)
    scale_m, scale_n = _packed_a_scale_shape(m, k)
    scale_numel = scale_m * scale_n
    if (
        packed.ndim != 1
        or packed.dtype != torch.uint8
        or packed.numel() != q_numel + scale_numel
    ):
        raise ValueError("invalid packed A payload")
    A_q_u8 = packed[:q_numel].view(m, k // 2)
    A_scale_u8 = packed[q_numel : q_numel + scale_numel].view(scale_m, scale_n)
    return A_q_u8.view(dtypes.fp4x2), A_scale_u8.view(dtypes.fp8_e8m0)


def _maybe_unpack_prequant_a(
    packed: torch.Tensor,
    m: int,
    k: int,
) -> tuple[torch.Tensor, torch.Tensor] | None:
    if packed.ndim != 1 or packed.dtype != torch.uint8:
        return None
    if packed.numel() != _packed_a_total_numel(m, k):
        return None
    return _unpack_prequant_a(packed, m, k)


def _report_runtime_once(cache_state: str) -> None:
    global _RUNTIME_REPORTED
    if _RUNTIME_REPORTED:
        return
    with _RUNTIME_REPORT_LOCK:
        if _RUNTIME_REPORTED:
            return
        print(
            "[kmm_00233] "
            f"path quant={_QUANT_BACKEND} gemm={_GEMM_BACKEND} cache={cache_state}",
            file=sys.stderr,
        )
        _RUNTIME_REPORTED = True


def _report_fused_stage_once(tag: str) -> None:
    with _FUSED_STAGE_LOG_LOCK:
        if tag in _FUSED_STAGE_LOGGED:
            return
        _FUSED_STAGE_LOGGED.add(tag)
    print(f"[kmm_00233] fused {tag}", file=sys.stderr, flush=True)


def generate_input(m: int, n: int, k: int, seed: int):
    assert k % 64 == 0, (
        "k must be divisible by 64 "
        "(scale group 32 and fp4 pack 2)"
    )
    gen = torch.Generator(device="cuda")
    gen.manual_seed(seed)
    A = torch.randn((m, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    A_packed = None
    try:
        A_q_fast, A_scale_fast = _quant_mxfp4_shape_aware(A, n, shuffle=True)
        _store_quant_cache(A, A_q_fast, A_scale_fast)
        A_packed = _pack_prequant_a(A_q_fast, A_scale_fast)
    except Exception:
        pass
    B = torch.randn((n, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    B_q, B_scale_sh = _quant_mxfp4_reference(B, shuffle=True)
    B_shuffle = shuffle_weight(B_q, layout=(16, 16))
    return (A, B, A_packed if A_packed is not None else B_q, B_shuffle, B_scale_sh)


def run_torch_fp4_mm(
    x: torch.Tensor,
    w: torch.Tensor,
    x_scales: torch.Tensor,
    w_scales: torch.Tensor,
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
    m, _ = x.shape
    n, _ = w.shape

    x_f32 = fp4_utils.mxfp4_to_f32(x)
    w_f32 = fp4_utils.mxfp4_to_f32(w)

    x_scales = x_scales[:m].repeat_interleave(SCALE_GROUP_SIZE, dim=1)
    x_scales_f32 = fp4_utils.e8m0_to_f32(x_scales)
    x_f32 = x_f32 * x_scales_f32

    w_scales = w_scales[:n].repeat_interleave(SCALE_GROUP_SIZE, dim=1)
    w_scales_f32 = fp4_utils.e8m0_to_f32(w_scales)
    w_f32 = w_f32 * w_scales_f32

    return torch.mm(x_f32, w_f32.T).to(dtype)[:m, :n]


def _unshuffle_weight(
    x: torch.Tensor,
    layout: tuple[int, int] = (16, 16),
    use_int4: bool = False,
) -> torch.Tensor:
    x_type = x.dtype
    if hasattr(torch, "float4_e2m1fn_x2") and x_type == torch.float4_e2m1fn_x2:
        x = x.view(torch.uint8)

    IN, IK = layout
    BK = IK * 2
    K = 16 // x.element_size() if not use_int4 else 32
    BN = IN

    x_ = x.view(-1, x.shape[-2] // BN, x.shape[-1] // BK, BK // K, BN, K)
    x_ = x_.permute(0, 1, 4, 2, 3, 5).contiguous()
    return x_.view(*x.shape).view(x_type)


def _unshuffle_e8m0(scale: torch.Tensor) -> torch.Tensor:
    if scale is None or scale.dtype == torch.float32:
        return scale
    assert scale.ndim == 2, "scale must be a 2D tensor"
    sm, sn = scale.shape
    scale_u8 = scale.view(torch.uint8)
    scale_u8 = scale_u8.view(sm // 32, sn // 8, 4, 16, 2, 2)
    scale_u8 = scale_u8.permute(0, 5, 3, 1, 4, 2).contiguous()
    return scale_u8.view(sm, sn).view(scale.dtype)


def _run_torch_fp4_mm_preshuffled(
    A_q: torch.Tensor,
    B_shuffle: torch.Tensor,
    A_scale_sh: torch.Tensor,
    B_scale_sh: torch.Tensor,
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
    return run_torch_fp4_mm(
        A_q.contiguous(),
        _unshuffle_weight(B_shuffle).contiguous(),
        _unshuffle_e8m0(A_scale_sh).contiguous(),
        _unshuffle_e8m0(B_scale_sh).contiguous(),
        dtype=dtype,
    )


def ref_kernel(data: input_t) -> output_t:
    A, B, _B_q, _B_shuffle, _B_scale_sh = data
    A_q, A_scale = _quant_mxfp4_reference(A.contiguous(), shuffle=False)
    B_q, B_scale = _quant_mxfp4_reference(B.contiguous(), shuffle=False)
    return run_torch_fp4_mm(
        A_q.contiguous(),
        B_q.contiguous(),
        A_scale.contiguous(),
        B_scale.contiguous(),
        dtype=torch.bfloat16,
    )


def _candidate_official_so_paths() -> list[str]:
    paths = [
        os.path.join(_AITER_PACKAGE_JIT_DIR, _OFFICIAL_SO_NAME),
        os.path.join(_AITER_JIT_DIR, _OFFICIAL_SO_NAME),
    ]
    uniq: list[str] = []
    for path in paths:
        if path and path not in uniq:
            uniq.append(path)
    return uniq


def _find_official_so_path() -> str:
    for path in _candidate_official_so_paths():
        if os.path.exists(path):
            return path
    raise FileNotFoundError(
        "official asm module not found in "
        + ", ".join(_candidate_official_so_paths())
    )


def _parse_ldd_dependencies(so_path: str) -> list[str]:
    deps: list[str] = []
    try:
        result = subprocess.run(
            ["ldd", so_path],
            stdout=subprocess.PIPE,
            stderr=subprocess.PIPE,
            text=True,
            check=False,
            timeout=5.0,
        )
        for raw_line in result.stdout.splitlines():
            line = raw_line.strip()
            candidate = ""
            if "=>" in line:
                rhs = line.split("=>", 1)[1].strip()
                if rhs and not rhs.startswith("not found"):
                    candidate = rhs.split(" ", 1)[0]
            elif line.startswith("/"):
                candidate = line.split(" ", 1)[0]
            if not candidate:
                continue
            base_name = os.path.basename(candidate)
            if any(key in base_name for key in _GLOBAL_LIB_KEYS):
                deps.append(candidate)
    except Exception:
        pass

    uniq: list[str] = []
    for dep in deps:
        if dep not in uniq:
            uniq.append(dep)
    return uniq


def _candidate_global_libs(official_so_path: str) -> list[str]:
    items: list[str] = []
    torch_lib_dir = os.path.join(os.path.dirname(torch.__file__), "lib")
    for name in _CANONICAL_TORCH_LIB_NAMES:
        path = os.path.join(torch_lib_dir, name)
        if os.path.exists(path):
            items.append(path)

    items.extend(_parse_ldd_dependencies(official_so_path))

    for name in (
        "amdhip64",
        "c10",
        "c10_hip",
        "torch",
        "torch_cpu",
        "torch_hip",
        "torch_python",
    ):
        path = ctypes.util.find_library(name)
        if path:
            items.append(path)

    uniq: list[str] = []
    for item in items:
        if item not in uniq:
            uniq.append(item)
    return uniq


def _promote_shared_object(path: str) -> None:
    if not path:
        return

    if getattr(os, "RTLD_NOLOAD", 0):
        try:
            _GLOBAL_SO_HANDLES.append(_cdll(path, _DLOPEN_NOLOAD_MODE))
        except OSError:
            pass

    try:
        _GLOBAL_SO_HANDLES.append(_cdll(path, _DLOPEN_MODE))
    except OSError:
        pass


def _preload_global_libs(official_so_path: str) -> None:
    with _GLOBAL_SO_LOCK:
        for path in _candidate_global_libs(official_so_path):
            if path in _GLOBAL_SO_SEEN:
                continue
            _GLOBAL_SO_SEEN.add(path)
            _promote_shared_object(path)


def _ensure_official_launcher() -> _DirectGemmA4W4Asm:
    global _OFFICIAL_LAUNCHER, _OFFICIAL_DISABLE_REASON, _GEMM_BACKEND

    if _OFFICIAL_LAUNCHER is not None:
        return _OFFICIAL_LAUNCHER
    if _OFFICIAL_DISABLE_REASON is not None:
        raise RuntimeError(_OFFICIAL_DISABLE_REASON)

    with _OFFICIAL_LOAD_LOCK:
        if _OFFICIAL_LAUNCHER is not None:
            return _OFFICIAL_LAUNCHER
        if _OFFICIAL_DISABLE_REASON is not None:
            raise RuntimeError(_OFFICIAL_DISABLE_REASON)
        try:
            so_path = _find_official_so_path()
            _preload_global_libs(so_path)
            _OFFICIAL_LAUNCHER = _DirectGemmA4W4Asm(so_path, mode=_DLOPEN_MODE)
            _GEMM_BACKEND = "official"
            return _OFFICIAL_LAUNCHER
        except Exception as exc:
            _OFFICIAL_DISABLE_REASON = str(exc)
            raise


def _build_direct_launcher_script() -> str:
    return """
import os
from aiter.jit.core import build_module, get_args_of_build

base_name = "module_gemm_a4w4_asm"
module_name = "kmm_00233_gemm_a4w4_asm"
args = get_args_of_build(base_name)
args["torch_exclude"] = False
args["is_python_module"] = False
args["is_standalone"] = False
srcs = [src for src in args["srcs"] if "pybind" not in src]
extra_ldflags = list(args["extra_ldflags"] or [])
extra_ldflags.append("-Wl,--no-gc-sections")
extra_ldflags.append("-Wl,--undefined=gemm_a4w4_asm")
target = os.path.join(
    os.environ["AITER_JIT_DIR"],
    "build",
    module_name,
    "build",
    module_name + ".so",
)
if not os.path.exists(target):
    build_module(
        module_name,
        srcs,
        args["flags_extra_cc"],
        args["flags_extra_hip"],
        args["blob_gen_cmd"],
        args["extra_include"],
        extra_ldflags,
        args["verbose"],
        args["is_python_module"],
        args["is_standalone"],
        args["torch_exclude"],
        args.get("third_party", []),
    )
""".strip()


def _subprocess_env() -> dict[str, str]:
    env = os.environ.copy()
    env["AITER_JIT_DIR"] = _AITER_JIT_DIR
    return env


def _summarize_prebuild_failure(stdout: str, stderr: str) -> str:
    text = (stderr.strip() or stdout.strip() or "unknown prebuild failure").strip()
    if len(text) <= 2000 and text.count("\n") <= 24:
        return text
    lines = text.splitlines()[:24]
    clipped = "\n".join(lines)
    if len(clipped) > 2000:
        clipped = clipped[:2000]
    return clipped + "\n..."


def _find_hipcc() -> str:
    hipcc = shutil.which("hipcc")
    if hipcc:
        return hipcc
    raise FileNotFoundError("hipcc not found in PATH")


def _artifact_stamp_path(path: str) -> str:
    return path + ".stamp"


def _artifact_is_current(path: str, stamp: str) -> bool:
    stamp_path = _artifact_stamp_path(path)
    if not os.path.exists(path) or not os.path.exists(stamp_path):
        return False
    try:
        with open(stamp_path, "r", encoding="utf-8") as handle:
            return handle.read().strip() == stamp
    except OSError:
        return False


def _write_artifact_stamp(path: str, stamp: str) -> None:
    with open(_artifact_stamp_path(path), "w", encoding="utf-8") as handle:
        handle.write(stamp + "\n")


def _write_text_if_needed(path: str, text: str) -> None:
    try:
        with open(path, "r", encoding="utf-8") as handle:
            if handle.read() == text:
                return
    except OSError:
        pass
    with open(path, "w", encoding="utf-8") as handle:
        handle.write(text)


def _replace_once(text: str, old: str, new: str, label: str) -> str:
    if old not in text:
        raise ValueError(f"missing rewrite pattern for {label}")
    return text.replace(old, new, 1)


def _build_fused_m4_compile_script() -> str:
    return f"""
from __future__ import annotations

import json
import os
import pathlib
import shutil

import triton
import triton.language as tl
from triton.backends.compiler import GPUTarget
from triton.compiler import ASTSource

from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid
from aiter.ops.triton.utils._triton.pid_preprocessing import remap_xcd


OUT_PATH = pathlib.Path({_FUSED_M4_HSACO_PATH!r})
CACHE_ROOT = pathlib.Path({_FUSED_M4_TRITON_CACHE_DIR!r})
OUT_PATH.parent.mkdir(parents=True, exist_ok=True)
CACHE_ROOT.mkdir(parents=True, exist_ok=True)
os.environ["TRITON_CACHE_DIR"] = str(CACHE_ROOT)
os.environ["AITER_JIT_DIR"] = {_AITER_JIT_DIR!r}

TARGET = GPUTarget("hip", "gfx950", 64)
KERNEL_NAME = {_FUSED_M4_KERNEL_NAME!r}


def _kernel_signature() -> dict[str, str]:
    return {{
        "a_bf16_ptr": "*bf16",
        "b_fp4_ptr": "*i8",
        "b_scale_ptr": "*i8",
        "c_ptr": "*bf16",
        "M": "i32",
        "N": "i32",
        "K": "i32",
        "stride_am": "i32",
        "stride_ak": "i32",
        "stride_bn": "i32",
        "stride_bk": "i32",
        "stride_bsn": "i32",
        "stride_bsk": "i32",
        "stride_cm": "i32",
        "stride_cn": "i32",
        "BLOCK_SIZE_M": "constexpr",
        "BLOCK_SIZE_N": "constexpr",
        "BLOCK_SIZE_K": "constexpr",
        "GROUP_SIZE_M": "constexpr",
        "NUM_KSPLIT": "constexpr",
        "SPLITK_BLOCK_SIZE": "constexpr",
        "EVEN_K": "constexpr",
    }}


_FUSED_QUANT_GEMM_REPR = make_kernel_repr(
    "_kmm_fused_quant_gemm_preshuffle_kernel",
    [
        "BLOCK_SIZE_M",
        "BLOCK_SIZE_N",
        "BLOCK_SIZE_K",
        "GROUP_SIZE_M",
        "num_warps",
        "num_stages",
        "waves_per_eu",
        "matrix_instr_nonkdim",
        "cache_modifier",
        "NUM_KSPLIT",
    ],
)


@triton.jit(repr=_FUSED_QUANT_GEMM_REPR)
def _kmm_fused_quant_gemm_preshuffle_kernel(
    a_bf16_ptr,
    b_fp4_ptr,
    b_scale_ptr,
    c_ptr,
    M,
    N,
    K,
    stride_am,
    stride_ak,
    stride_bn,
    stride_bk,
    stride_bsn,
    stride_bsk,
    stride_cm,
    stride_cn,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr,
    SPLITK_BLOCK_SIZE: tl.constexpr,
    EVEN_K: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    scale_group_size: tl.constexpr = 32

    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)

    grid_mn = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    pid_unified = tl.program_id(axis=0)
    pid_unified = remap_xcd(pid_unified, grid_mn * NUM_KSPLIT, NUM_XCDS=8)

    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    if (pid_k * SPLITK_BLOCK_SIZE // 2) >= K:
        return

    num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
    offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)

    offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
    offs_k_fp4_shuffle = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)

    a_bf16_ptrs = a_bf16_ptr + (
        offs_am[:, None] * stride_am
        + (pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16)[None, :] * stride_ak
    )

    b_fp4_n = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
    b_fp4_ptrs = b_fp4_ptr + (
        b_fp4_n[:, None] * stride_bn
        + (
            pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_fp4_shuffle
        )[None, :] * stride_bk
    )

    b_scale_n = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
    b_scale_k = (
        pid_k * (SPLITK_BLOCK_SIZE // scale_group_size) * 32
        + tl.arange(0, BLOCK_SIZE_K // scale_group_size * 32)
    )
    b_scale_ptrs = (
        b_scale_ptr
        + b_scale_n[:, None] * stride_bsn
        + b_scale_k[None, :] * stride_bsk
    )

    accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

    for _ in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
        a_bf16 = tl.load(
            a_bf16_ptrs,
            mask=(offs_am[:, None] < M) & (offs_k_bf16[None, :] < K),
            other=0,
        ).to(tl.float32)

        a3 = a_bf16.reshape(
            BLOCK_SIZE_M,
            BLOCK_SIZE_K // scale_group_size,
            scale_group_size,
        )
        amax = tl.max(tl.abs(a3), axis=2)
        rounded_amax_bits = amax.to(tl.int32, bitcast=True)
        rounded_amax_bits = (
            (rounded_amax_bits + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
        )
        rounded_exp = ((rounded_amax_bits >> 23) & 0xFF).to(tl.int32)
        a_scale_i32 = tl.maximum(rounded_exp - 2, 0)
        quant_scale_bits = ((254 - a_scale_i32).to(tl.uint32) << 23)
        quant_scale = quant_scale_bits.to(tl.float32, bitcast=True)
        qx = (a3 * quant_scale[:, :, None]).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K)
        a_scale = a_scale_i32.to(tl.uint8).reshape(
            BLOCK_SIZE_M,
            BLOCK_SIZE_K // scale_group_size,
        )

        exp_bias_fp32: tl.constexpr = 127
        exp_bias_fp4: tl.constexpr = 1
        ebits_fp32: tl.constexpr = 8
        ebits_fp4: tl.constexpr = 2
        mbits_fp32: tl.constexpr = 23
        mbits_fp4: tl.constexpr = 1
        max_normal: tl.constexpr = 6
        min_normal: tl.constexpr = 1

        qx_bits = qx.to(tl.uint32, bitcast=True)
        sign = qx_bits & 0x80000000
        qx_bits = qx_bits ^ sign
        qx_abs = qx_bits.to(tl.float32, bitcast=True)
        saturate_mask = qx_abs >= max_normal
        denormal_mask = (~saturate_mask) & (qx_abs < min_normal)
        normal_mask = ~(saturate_mask | denormal_mask)

        denorm_exp: tl.constexpr = (
            (exp_bias_fp32 - exp_bias_fp4) + (mbits_fp32 - mbits_fp4) + 1
        )
        denorm_mask_int: tl.constexpr = denorm_exp << mbits_fp32
        denorm_mask_float: tl.constexpr = tl.cast(
            denorm_mask_int,
            tl.float32,
            bitcast=True,
        )

        denormal_x = qx_abs + denorm_mask_float
        denormal_x = denormal_x.to(tl.uint32, bitcast=True)
        denormal_x -= denorm_mask_int
        denormal_x = denormal_x.to(tl.uint8)

        normal_x = qx_bits
        mant_odd = (normal_x >> (mbits_fp32 - mbits_fp4)) & 1
        val_to_add = ((exp_bias_fp4 - exp_bias_fp32) << mbits_fp32) + (1 << 21) - 1
        normal_x += val_to_add
        normal_x += mant_odd
        normal_x = normal_x >> (mbits_fp32 - mbits_fp4)
        normal_x = normal_x.to(tl.uint8)

        e2m1 = tl.full(qx_bits.type.get_block_shapes(), 0x7, dtype=tl.uint8)
        e2m1 = tl.where(normal_mask, normal_x, e2m1)
        e2m1 = tl.where(denormal_mask, denormal_x, e2m1)
        sign_lp = (
            sign >> (mbits_fp32 + ebits_fp32 - mbits_fp4 - ebits_fp4)
        ).to(tl.uint8)
        e2m1 = e2m1 | sign_lp
        e2m1 = tl.reshape(e2m1, [BLOCK_SIZE_M, BLOCK_SIZE_K // 2, 2])
        evens, odds = tl.split(e2m1)
        a_fp4 = (evens | (odds << 4)).reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2)

        b_scale = (
            tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
            .to(tl.uint8, bitcast=True)
            .reshape(
                BLOCK_SIZE_N // 32,
                BLOCK_SIZE_K // 64,
                4,
                16,
            )
            .reshape(
                BLOCK_SIZE_N // 32,
                BLOCK_SIZE_K // 64,
                2,
                2,
                16,
            )
            .permute(0, 4, 3, 1, 2)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // scale_group_size)
        )

        b = tl.load(b_fp4_ptrs, cache_modifier=cache_modifier).to(
            tl.uint8,
            bitcast=True,
        )
        b = (
            b.reshape(
                1,
                BLOCK_SIZE_N // 16,
                BLOCK_SIZE_K // 64,
                2,
                16,
                16,
            )
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
            .trans(1, 0)
        )

        accumulator += tl.dot_scaled(a_fp4, a_scale, "e2m1", b, b_scale, "e2m1")

        a_bf16_ptrs += BLOCK_SIZE_K * stride_ak
        b_fp4_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
        b_scale_ptrs += BLOCK_SIZE_K * stride_bsk

    c = accumulator.to(c_ptr.type.element_ty)
    c_ptrs = c_ptr + offs_am[:, None] * stride_cm + offs_bn[None, :] * stride_cn
    c_mask = (offs_am[:, None] < M) & (offs_bn[None, :] < N)
    tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")


def _shape_source() -> ASTSource:
    return ASTSource(
        fn=_kmm_fused_quant_gemm_preshuffle_kernel,
        signature=_kernel_signature(),
        constexprs={{
            "BLOCK_SIZE_M": {_FUSED_M4_BLOCK_SIZE_M},
            "BLOCK_SIZE_N": {_FUSED_M4_BLOCK_SIZE_N},
            "BLOCK_SIZE_K": {_FUSED_M4_BLOCK_SIZE_K},
            "GROUP_SIZE_M": {_FUSED_M4_GROUP_SIZE_M},
            "NUM_KSPLIT": 1,
            "SPLITK_BLOCK_SIZE": {_FUSED_M4_BLOCK_SIZE_K},
            "EVEN_K": True,
        }},
    )


def _shape_options() -> dict[str, object]:
    return {{
        "num_warps": {_FUSED_M4_NUM_WARPS},
        "num_stages": {_FUSED_M4_NUM_STAGES},
        "num_ctas": 1,
        "waves_per_eu": {_FUSED_M4_WAVES_PER_EU},
        "matrix_instr_nonkdim": 16,
        "cache_modifier": "CG",
        "name": "kmm_00233_fused_m4",
    }}


def _find_cached_hsaco(cache_root: pathlib.Path) -> pathlib.Path | None:
    for meta_path in sorted(
        cache_root.rglob(f"{{_kmm_fused_quant_gemm_preshuffle_kernel.fn.__name__}}.json")
    ):
        if meta_path.name.startswith("__grp__"):
            continue
        try:
            meta = json.loads(meta_path.read_text(encoding="utf-8"))
        except Exception:
            continue
        if str(meta.get("name", "")) != KERNEL_NAME:
            continue
        if int(meta.get("num_warps", -1)) != {_FUSED_M4_NUM_WARPS}:
            continue
        if int(meta.get("num_stages", -1)) != {_FUSED_M4_NUM_STAGES}:
            continue
        if int(meta.get("waves_per_eu", -1)) != {_FUSED_M4_WAVES_PER_EU}:
            continue
        hsaco_path = meta_path.with_suffix(".hsaco")
        if hsaco_path.exists():
            return hsaco_path
        for candidate in sorted(meta_path.parent.glob("*.hsaco")):
            return candidate
    for candidate in sorted(cache_root.rglob("*.hsaco")):
        if candidate.name.startswith("__grp__"):
            continue
        return candidate
    return None


def main() -> int:
    triton.compile(
        _shape_source(),
        target=TARGET,
        options=_shape_options(),
    )
    hsaco_path = _find_cached_hsaco(CACHE_ROOT)
    if hsaco_path is None:
        raise FileNotFoundError("fused m4 hsaco artifact not found in Triton cache")
    shutil.copy2(hsaco_path, OUT_PATH)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
""".strip()


def _build_fused_m16_compile_script() -> str:
    script = _build_fused_m4_compile_script()
    rewrites = (
        (repr(_FUSED_M4_HSACO_PATH), repr(_FUSED_M16_HSACO_PATH), "m16 hsaco path"),
        (
            repr(_FUSED_M4_TRITON_CACHE_DIR),
            repr(_FUSED_M16_TRITON_CACHE_DIR),
            "m16 cache dir",
        ),
        (_FUSED_M4_KERNEL_NAME, _FUSED_M16_KERNEL_NAME, "m16 kernel name"),
        (
            f'"BLOCK_SIZE_M": {_FUSED_M4_BLOCK_SIZE_M},',
            f'"BLOCK_SIZE_M": {_FUSED_M16_BLOCK_SIZE_M},',
            "m16 block size m",
        ),
        (
            f'"BLOCK_SIZE_N": {_FUSED_M4_BLOCK_SIZE_N},',
            f'"BLOCK_SIZE_N": {_FUSED_M16_BLOCK_SIZE_N},',
            "m16 block size n",
        ),
        (
            f'"BLOCK_SIZE_K": {_FUSED_M4_BLOCK_SIZE_K},',
            f'"BLOCK_SIZE_K": {_FUSED_M16_BLOCK_SIZE_K},',
            "m16 block size k",
        ),
        (
            f'"GROUP_SIZE_M": {_FUSED_M4_GROUP_SIZE_M},',
            f'"GROUP_SIZE_M": {_FUSED_M16_GROUP_SIZE_M},',
            "m16 group size m",
        ),
        (
            f'"SPLITK_BLOCK_SIZE": {_FUSED_M4_BLOCK_SIZE_K},',
            f'"SPLITK_BLOCK_SIZE": {_FUSED_M16_BLOCK_SIZE_K},',
            "m16 splitk block size",
        ),
        (
            f'"num_warps": {_FUSED_M4_NUM_WARPS},',
            f'"num_warps": {_FUSED_M16_NUM_WARPS},',
            "m16 num warps",
        ),
        (
            f'"num_stages": {_FUSED_M4_NUM_STAGES},',
            f'"num_stages": {_FUSED_M16_NUM_STAGES},',
            "m16 num stages",
        ),
        (
            f'if int(meta.get("num_warps", -1)) != {_FUSED_M4_NUM_WARPS}:',
            f'if int(meta.get("num_warps", -1)) != {_FUSED_M16_NUM_WARPS}:',
            "m16 num warps check",
        ),
        (
            f'if int(meta.get("num_stages", -1)) != {_FUSED_M4_NUM_STAGES}:',
            f'if int(meta.get("num_stages", -1)) != {_FUSED_M16_NUM_STAGES}:',
            "m16 num stages check",
        ),
        (
            'raise FileNotFoundError("fused m4 hsaco artifact not found in Triton cache")',
            'raise FileNotFoundError("fused m16 hsaco artifact not found in Triton cache")',
            "m16 missing hsaco message",
        ),
        (
            '"name": "kmm_00233_fused_m4",',
            '"name": "kmm_00233_fused_m16",',
            "m16 triton cache name",
        ),
    )
    for old, new, label in rewrites:
        script = _replace_once(script, old, new, label)
    return script


def _compile_fused_m4_hsaco() -> str:
    global _FUSED_M4_HSACO_READY_PATH
    if _FUSED_M4_HSACO_READY_PATH is not None:
        return _FUSED_M4_HSACO_READY_PATH
    os.makedirs(_FUSED_M4_BUILD_DIR, exist_ok=True)
    stamp = _FUSED_M4_HSACO_SHA256
    if _artifact_is_current(_FUSED_M4_HSACO_PATH, stamp):
        _report_fused_stage_once("hsaco_cached")
        _FUSED_M4_HSACO_READY_PATH = _FUSED_M4_HSACO_PATH
        return _FUSED_M4_HSACO_PATH

    _report_fused_stage_once("hsaco_write_begin")
    hsaco_bytes = zlib.decompress(base64.b64decode(_FUSED_M4_HSACO_B64Z.encode("ascii")))
    actual_sha = hashlib.sha256(hsaco_bytes).hexdigest()
    if actual_sha != _FUSED_M4_HSACO_SHA256:
        raise ValueError(
            f"fused m4 hsaco sha mismatch expected={_FUSED_M4_HSACO_SHA256} got={actual_sha}"
        )
    with open(_FUSED_M4_HSACO_PATH, "wb") as handle:
        handle.write(hsaco_bytes)
    _write_artifact_stamp(_FUSED_M4_HSACO_PATH, stamp)
    _report_fused_stage_once("hsaco_write_end")
    _FUSED_M4_HSACO_READY_PATH = _FUSED_M4_HSACO_PATH
    return _FUSED_M4_HSACO_PATH


def _ensure_fused_m4_hsaco() -> str:
    if _FUSED_M4_HSACO_READY_PATH is not None:
        return _FUSED_M4_HSACO_READY_PATH
    return _compile_fused_m4_hsaco()


def _compile_fused_m16_hsaco() -> str:
    global _FUSED_M16_HSACO_READY_PATH, _FUSED_M16_DISABLE_REASON
    if _FUSED_M16_HSACO_READY_PATH is not None:
        return _FUSED_M16_HSACO_READY_PATH
    if _FUSED_M16_DISABLE_REASON is not None:
        raise RuntimeError(_FUSED_M16_DISABLE_REASON)
    os.makedirs(_FUSED_M16_BUILD_DIR, exist_ok=True)
    script = _build_fused_m16_compile_script()
    stamp = f"{zlib.crc32(script.encode('utf-8')) & 0xFFFFFFFF:08x}"
    if _artifact_is_current(_FUSED_M16_HSACO_PATH, stamp):
        _report_fused_stage_once("m16_hsaco_cached")
        _FUSED_M16_HSACO_READY_PATH = _FUSED_M16_HSACO_PATH
        return _FUSED_M16_HSACO_PATH

    _report_fused_stage_once("m16_hsaco_compile_begin")
    _write_text_if_needed(_FUSED_M16_COMPILE_SCRIPT_PATH, script + "\n")
    result = subprocess.run(
        [sys.executable, _FUSED_M16_COMPILE_SCRIPT_PATH],
        env=_subprocess_env(),
        stdout=subprocess.PIPE,
        stderr=subprocess.PIPE,
        text=True,
        check=False,
        timeout=300.0,
    )
    if result.returncode != 0:
        _FUSED_M16_DISABLE_REASON = _summarize_prebuild_failure(
            result.stdout,
            result.stderr,
        )
        raise RuntimeError(_FUSED_M16_DISABLE_REASON)
    if not os.path.exists(_FUSED_M16_HSACO_PATH):
        _FUSED_M16_DISABLE_REASON = "fused m16 compile did not produce hsaco"
        raise FileNotFoundError(_FUSED_M16_DISABLE_REASON)
    _write_artifact_stamp(_FUSED_M16_HSACO_PATH, stamp)
    _report_fused_stage_once("m16_hsaco_compile_end")
    _FUSED_M16_HSACO_READY_PATH = _FUSED_M16_HSACO_PATH
    return _FUSED_M16_HSACO_PATH


def _ensure_fused_m16_hsaco() -> str:
    if _FUSED_M16_HSACO_READY_PATH is not None:
        return _FUSED_M16_HSACO_READY_PATH
    return _compile_fused_m16_hsaco()


def _build_fused_m4_source() -> str:
    hip_queue_type = _name_from_codes(
        104, 105, 112, 83, 116, 114, 101, 97, 109, 95, 116
    )
    return f"""
#include <hip/hip_runtime.h>
#include <stddef.h>
#include <stdint.h>
#include <string.h>

namespace {{
struct KmmFusedM4Cache {{
  hipModule_t module;
  hipFunction_t function;
  int ready;
  char hsaco_path[4096];
}};

KmmFusedM4Cache gFusedM4Cache = {{}};

inline int kmm_copy_path(char* dst, size_t cap, const char* src) {{
  if (dst == nullptr || src == nullptr) {{
    return static_cast<int>(hipErrorInvalidValue);
  }}
  size_t len = strlen(src);
  if (len + 1 > cap) {{
    return static_cast<int>(hipErrorInvalidValue);
  }}
  memcpy(dst, src, len + 1);
  return 0;
}}

inline int kmm_release_module() {{
  if (gFusedM4Cache.ready != 0 && gFusedM4Cache.module != nullptr) {{
    hipError_t unload_status = hipModuleUnload(gFusedM4Cache.module);
    if (unload_status != hipSuccess) {{
      return static_cast<int>(unload_status);
    }}
  }}
  memset(&gFusedM4Cache, 0, sizeof(gFusedM4Cache));
  return 0;
}}

inline int kmm_load_kernel(const char* hsaco_path, hipFunction_t* fn_out) {{
  if (hsaco_path == nullptr || hsaco_path[0] == '\\0') {{
    return static_cast<int>(hipErrorInvalidValue);
  }}
  if (gFusedM4Cache.ready == 0 ||
      strcmp(gFusedM4Cache.hsaco_path, hsaco_path) != 0) {{
    int release_status = kmm_release_module();
    if (release_status != 0) {{
      return release_status;
    }}
    int copy_status = kmm_copy_path(
        gFusedM4Cache.hsaco_path,
        sizeof(gFusedM4Cache.hsaco_path),
        hsaco_path);
    if (copy_status != 0) {{
      return copy_status;
    }}
    hipError_t load_status = hipModuleLoad(&gFusedM4Cache.module, hsaco_path);
    if (load_status != hipSuccess) {{
      return static_cast<int>(load_status);
    }}
    hipError_t fn_status = hipModuleGetFunction(
        &gFusedM4Cache.function,
        gFusedM4Cache.module,
        "{_FUSED_M4_KERNEL_NAME}");
    if (fn_status != hipSuccess) {{
      kmm_release_module();
      return static_cast<int>(fn_status);
    }}
    gFusedM4Cache.ready = 1;
  }}
  if (fn_out != nullptr) {{
    *fn_out = gFusedM4Cache.function;
  }}
  return 0;
}}
}}  // namespace

extern "C" int {_FUSED_M4_WARM_EXPORT_NAME}(const char* hsaco_path) {{
  return kmm_load_kernel(hsaco_path, nullptr);
}}

extern "C" int {_FUSED_M4_RELEASE_EXPORT_NAME}() {{
  return kmm_release_module();
}}

extern "C" int {_FUSED_M4_LAUNCH_EXPORT_NAME}(
    const char* hsaco_path,
    const void* a_bf16,
    const void* b_preshuffled_fp4,
    const void* b_scale_e8m0,
    void* out_bf16,
    int32_t m,
    int32_t n,
    int32_t k,
    int32_t stride_am,
    int32_t stride_ak,
    int32_t stride_bn,
    int32_t stride_bk,
    int32_t stride_bsn,
    int32_t stride_bsk,
    int32_t stride_cm,
    int32_t stride_cn,
    uint32_t grid_x,
    void* queue_ptr) {{
  if (a_bf16 == nullptr || b_preshuffled_fp4 == nullptr ||
      b_scale_e8m0 == nullptr || out_bf16 == nullptr ||
      m <= 0 || n <= 0 || k <= 0 || grid_x == 0) {{
    return static_cast<int>(hipErrorInvalidValue);
  }}

  hipFunction_t kernel = nullptr;
  int load_status = kmm_load_kernel(hsaco_path, &kernel);
  if (load_status != 0) {{
    return load_status;
  }}

  void* a_arg = const_cast<void*>(a_bf16);
  void* b_arg = const_cast<void*>(b_preshuffled_fp4);
  void* b_scale_arg = const_cast<void*>(b_scale_e8m0);
  void* out_arg = out_bf16;
  hipDeviceptr_t global_scratch = 0;
  hipDeviceptr_t profile_scratch = 0;

  void* kernel_params[] = {{
      &a_arg,
      &b_arg,
      &b_scale_arg,
      &out_arg,
      &m,
      &n,
      &k,
      &stride_am,
      &stride_ak,
      &stride_bn,
      &stride_bk,
      &stride_bsn,
      &stride_bsk,
      &stride_cm,
      &stride_cn,
      &global_scratch,
      &profile_scratch,
  }};

  {hip_queue_type} queue = reinterpret_cast<{hip_queue_type}>(queue_ptr);
  hipError_t launch_status = hipModuleLaunchKernel(
      kernel,
      static_cast<unsigned int>(grid_x),
      1,
      1,
      static_cast<unsigned int>({_FUSED_M4_BLOCK_THREADS}),
      1,
      1,
      0,
      queue,
      kernel_params,
      nullptr);
  if (launch_status != hipSuccess) {{
    return static_cast<int>(launch_status);
  }}
  return static_cast<int>(hipGetLastError());
}}
""".strip()


def _build_fused_m16_source() -> str:
    source = _build_fused_m4_source()
    if "KmmFusedM4Cache" not in source or "gFusedM4Cache" not in source:
        raise ValueError("missing m16 source rewrite anchors")
    source = source.replace("KmmFusedM4Cache", "KmmFusedM16Cache")
    source = source.replace("gFusedM4Cache", "gFusedM16Cache")
    rewrites = (
        (_FUSED_M4_KERNEL_NAME, _FUSED_M16_KERNEL_NAME, "m16 kernel name"),
        (
            _FUSED_M4_WARM_EXPORT_NAME,
            _FUSED_M16_WARM_EXPORT_NAME,
            "m16 warm export",
        ),
        (
            _FUSED_M4_RELEASE_EXPORT_NAME,
            _FUSED_M16_RELEASE_EXPORT_NAME,
            "m16 release export",
        ),
        (
            _FUSED_M4_LAUNCH_EXPORT_NAME,
            _FUSED_M16_LAUNCH_EXPORT_NAME,
            "m16 launch export",
        ),
        (
            f"static_cast<unsigned int>({_FUSED_M4_BLOCK_THREADS})",
            f"static_cast<unsigned int>({_FUSED_M16_BLOCK_THREADS})",
            "m16 block threads",
        ),
    )
    for old, new, label in rewrites:
        source = _replace_once(source, old, new, label)
    return source


def _build_fused_m4_so() -> str:
    os.makedirs(_FUSED_M4_BUILD_DIR, exist_ok=True)
    source_text = _build_fused_m4_source() + "\n"
    stamp = f"{zlib.crc32(source_text.encode('utf-8')) & 0xFFFFFFFF:08x}"
    _write_text_if_needed(_FUSED_M4_SRC_PATH, source_text)
    if _artifact_is_current(_FUSED_M4_SO_PATH, stamp):
        _report_fused_stage_once("bridge_so_cached")
        return _FUSED_M4_SO_PATH

    _report_fused_stage_once("bridge_so_build_begin")
    result = subprocess.run(
        [
            _find_hipcc(),
            "--offload-arch=gfx950",
            "-shared",
            "-fPIC",
            "-O2",
            _FUSED_M4_SRC_PATH,
            "-o",
            _FUSED_M4_SO_PATH,
        ],
        env=_subprocess_env(),
        stdout=subprocess.PIPE,
        stderr=subprocess.PIPE,
        text=True,
        check=False,
        timeout=120.0,
    )
    if result.returncode != 0:
        raise RuntimeError(_summarize_prebuild_failure(result.stdout, result.stderr))
    if not os.path.exists(_FUSED_M4_SO_PATH):
        raise FileNotFoundError("fused m4 build did not produce a shared object")
    _write_artifact_stamp(_FUSED_M4_SO_PATH, stamp)
    _report_fused_stage_once("bridge_so_build_end")
    return _FUSED_M4_SO_PATH


def _build_fused_m16_so() -> str:
    os.makedirs(_FUSED_M16_BUILD_DIR, exist_ok=True)
    source_text = _build_fused_m16_source() + "\n"
    stamp = f"{zlib.crc32(source_text.encode('utf-8')) & 0xFFFFFFFF:08x}"
    _write_text_if_needed(_FUSED_M16_SRC_PATH, source_text)
    if _artifact_is_current(_FUSED_M16_SO_PATH, stamp):
        _report_fused_stage_once("m16_bridge_so_cached")
        return _FUSED_M16_SO_PATH

    _report_fused_stage_once("m16_bridge_so_build_begin")
    result = subprocess.run(
        [
            _find_hipcc(),
            "--offload-arch=gfx950",
            "-shared",
            "-fPIC",
            "-O2",
            _FUSED_M16_SRC_PATH,
            "-o",
            _FUSED_M16_SO_PATH,
        ],
        env=_subprocess_env(),
        stdout=subprocess.PIPE,
        stderr=subprocess.PIPE,
        text=True,
        check=False,
        timeout=120.0,
    )
    if result.returncode != 0:
        raise RuntimeError(_summarize_prebuild_failure(result.stdout, result.stderr))
    if not os.path.exists(_FUSED_M16_SO_PATH):
        raise FileNotFoundError("fused m16 build did not produce a shared object")
    _write_artifact_stamp(_FUSED_M16_SO_PATH, stamp)
    _report_fused_stage_once("m16_bridge_so_build_end")
    return _FUSED_M16_SO_PATH


def _ensure_fused_m4_launcher() -> _FusedLauncher:
    global _FUSED_M4_LAUNCHER, _FUSED_M4_DISABLE_REASON, _GEMM_BACKEND

    if _FUSED_M4_LAUNCHER is not None:
        return _FUSED_M4_LAUNCHER
    if _FUSED_M4_DISABLE_REASON is not None:
        raise RuntimeError(_FUSED_M4_DISABLE_REASON)

    with _FUSED_M4_LAUNCH_LOCK:
        if _FUSED_M4_LAUNCHER is not None:
            return _FUSED_M4_LAUNCHER
        if _FUSED_M4_DISABLE_REASON is not None:
            raise RuntimeError(_FUSED_M4_DISABLE_REASON)
        try:
            _report_fused_stage_once("launcher_prepare_begin")
            hsaco_path = _ensure_fused_m4_hsaco()
            so_path = _build_fused_m4_so()
            _report_fused_stage_once("launcher_load_begin")
            _FUSED_M4_LAUNCHER = _FusedLauncher(
                so_path,
                launch_export_name=_FUSED_M4_LAUNCH_EXPORT_NAME,
                warm_export_name=_FUSED_M4_WARM_EXPORT_NAME,
                release_export_name=_FUSED_M4_RELEASE_EXPORT_NAME,
                mode=_DLOPEN_MODE,
            )
            _report_fused_stage_once("launcher_warm_begin")
            warm_status = _FUSED_M4_LAUNCHER.warm(hsaco_path)
            _report_fused_stage_once(f"launcher_warm_status={warm_status}")
            if warm_status != 0:
                raise RuntimeError(f"fused m4 warm status {warm_status}")
            _GEMM_BACKEND = "fused_m4"
            return _FUSED_M4_LAUNCHER
        except Exception as exc:
            _FUSED_M4_DISABLE_REASON = str(exc)
            raise


def _ensure_fused_m16_launcher() -> _FusedLauncher:
    global _FUSED_M16_LAUNCHER, _FUSED_M16_DISABLE_REASON, _GEMM_BACKEND

    if _FUSED_M16_LAUNCHER is not None:
        return _FUSED_M16_LAUNCHER
    if _FUSED_M16_DISABLE_REASON is not None:
        raise RuntimeError(_FUSED_M16_DISABLE_REASON)

    with _FUSED_M16_LAUNCH_LOCK:
        if _FUSED_M16_LAUNCHER is not None:
            return _FUSED_M16_LAUNCHER
        if _FUSED_M16_DISABLE_REASON is not None:
            raise RuntimeError(_FUSED_M16_DISABLE_REASON)
        try:
            _report_fused_stage_once("m16_launcher_prepare_begin")
            hsaco_path = _ensure_fused_m16_hsaco()
            so_path = _build_fused_m16_so()
            _report_fused_stage_once("m16_launcher_load_begin")
            _FUSED_M16_LAUNCHER = _FusedLauncher(
                so_path,
                launch_export_name=_FUSED_M16_LAUNCH_EXPORT_NAME,
                warm_export_name=_FUSED_M16_WARM_EXPORT_NAME,
                release_export_name=_FUSED_M16_RELEASE_EXPORT_NAME,
                mode=_DLOPEN_MODE,
            )
            _report_fused_stage_once("m16_launcher_warm_begin")
            warm_status = _FUSED_M16_LAUNCHER.warm(hsaco_path)
            _report_fused_stage_once(f"m16_launcher_warm_status={warm_status}")
            if warm_status != 0:
                raise RuntimeError(f"fused m16 warm status {warm_status}")
            _GEMM_BACKEND = "fused_m16"
            return _FUSED_M16_LAUNCHER
        except Exception as exc:
            _FUSED_M16_DISABLE_REASON = str(exc)
            raise


def _should_use_fused_m4(m: int, n: int, k: int) -> bool:
    return m == _FUSED_M4_SHAPE[0] and n == _FUSED_M4_SHAPE[1] and k == _FUSED_M4_FIXED_K


def _should_use_fused_m16(m: int, n: int, k: int) -> bool:
    return m == _FUSED_M16_SHAPE[0] and n == _FUSED_M16_SHAPE[1] and k == _FUSED_M16_FIXED_K


def _fused_view_cache_key(
    tensor: torch.Tensor,
) -> tuple[int, int, tuple[int, ...], torch.device]:
    return (
        tensor.data_ptr(),
        int(getattr(tensor, "_version", 0)),
        tuple(tensor.shape),
        tensor.device,
    )


def _reshape_fused_m4_weight(B_shuffle: torch.Tensor, n: int, k: int) -> torch.Tensor:
    key = _fused_view_cache_key(B_shuffle)
    with _FUSED_M4_VIEW_CACHE_LOCK:
        cached = _FUSED_M4_WEIGHT_VIEW_CACHE.get(key)
        if cached is not None:
            return cached
    B_u8 = B_shuffle.contiguous().view(torch.uint8)
    view_shape = _fused_m4_weight_view_shape(n, k)
    expected_numel = view_shape[0] * view_shape[1]
    if B_u8.numel() != expected_numel:
        raise ValueError(
            f"unexpected fused m4 weight payload numel={B_u8.numel()} expected={expected_numel}"
        )
    view = B_u8.view(*view_shape)
    with _FUSED_M4_VIEW_CACHE_LOCK:
        _FUSED_M4_WEIGHT_VIEW_CACHE.clear()
        _FUSED_M4_WEIGHT_VIEW_CACHE[key] = view
    return view


def _reshape_fused_m4_scale(B_scale_sh: torch.Tensor, n: int, k: int) -> torch.Tensor:
    key = _fused_view_cache_key(B_scale_sh)
    with _FUSED_M4_VIEW_CACHE_LOCK:
        cached = _FUSED_M4_SCALE_VIEW_CACHE.get(key)
        if cached is not None:
            return cached
    scale_contig = B_scale_sh.contiguous()
    if (
        scale_contig.ndim == 2
        and scale_contig.shape[0] >= n
        and scale_contig.shape[1] * _FUSED_M4_SCALE_TILE_ROWS == k
    ):
        scale_u8 = scale_contig[:n].contiguous().view(torch.uint8)
    else:
        scale_u8 = scale_contig.view(torch.uint8)
    view_shape = _fused_m4_scale_view_shape(n, k)
    expected_numel = view_shape[0] * view_shape[1]
    if scale_u8.numel() != expected_numel:
        raise ValueError(
            f"unexpected fused m4 scale payload numel={scale_u8.numel()} expected={expected_numel}"
        )
    view = scale_u8.view(*view_shape)
    with _FUSED_M4_VIEW_CACHE_LOCK:
        _FUSED_M4_SCALE_VIEW_CACHE.clear()
        _FUSED_M4_SCALE_VIEW_CACHE[key] = view
    return view


def _run_fused_m4_gemm(
    A: torch.Tensor,
    B_shuffle: torch.Tensor,
    B_scale_sh: torch.Tensor,
) -> torch.Tensor:
    global _QUANT_BACKEND, _GEMM_BACKEND
    _report_fused_stage_once("run_enter")
    m, k = A.shape
    n = B_shuffle.shape[0]
    A_contig = A.contiguous()
    out = _get_output_buffer(m, n, A.device, torch.bfloat16)
    weight_view = _reshape_fused_m4_weight(B_shuffle, n, k)
    scale_view = _reshape_fused_m4_scale(B_scale_sh, n, k)
    _report_fused_stage_once("layout_ready")
    hsaco_path = _ensure_fused_m4_hsaco()
    _report_fused_stage_once("hsaco_ready")
    launcher = _ensure_fused_m4_launcher()
    _report_fused_stage_once("launcher_ready")
    status = launcher(
        hsaco_path,
        A_contig,
        weight_view,
        scale_view,
        out,
        m=m,
        n=n,
        k=k,
        stride_am=A_contig.stride(0),
        stride_ak=A_contig.stride(1),
        stride_bn=weight_view.stride(0),
        stride_bk=weight_view.stride(1),
        stride_bsn=scale_view.stride(0),
        stride_bsk=scale_view.stride(1),
        stride_cm=out.stride(0),
        stride_cn=out.stride(1),
        grid_x=_fused_m4_grid_x(m, n),
    )
    _report_fused_stage_once(f"launch_status={status}")
    if status != 0:
        raise RuntimeError(f"fused m4 launch status {status}")
    _QUANT_BACKEND = "fused_m4"
    _GEMM_BACKEND = "fused_m4"
    return out[:m]


def _reshape_fused_m16_weight(B_shuffle: torch.Tensor, n: int, k: int) -> torch.Tensor:
    key = _fused_view_cache_key(B_shuffle)
    with _FUSED_M16_VIEW_CACHE_LOCK:
        cached = _FUSED_M16_WEIGHT_VIEW_CACHE.get(key)
        if cached is not None:
            return cached
    B_u8 = B_shuffle.contiguous().view(torch.uint8)
    view_shape = _fused_m16_weight_view_shape(n, k)
    expected_numel = view_shape[0] * view_shape[1]
    if B_u8.numel() != expected_numel:
        raise ValueError(
            f"unexpected fused m16 weight payload numel={B_u8.numel()} expected={expected_numel}"
        )
    view = B_u8.view(*view_shape)
    with _FUSED_M16_VIEW_CACHE_LOCK:
        _FUSED_M16_WEIGHT_VIEW_CACHE.clear()
        _FUSED_M16_WEIGHT_VIEW_CACHE[key] = view
    return view


def _reshape_fused_m16_scale(B_scale_sh: torch.Tensor, n: int, k: int) -> torch.Tensor:
    key = _fused_view_cache_key(B_scale_sh)
    with _FUSED_M16_VIEW_CACHE_LOCK:
        cached = _FUSED_M16_SCALE_VIEW_CACHE.get(key)
        if cached is not None:
            return cached
    scale_contig = B_scale_sh.contiguous()
    if (
        scale_contig.ndim == 2
        and scale_contig.shape[0] >= n
        and scale_contig.shape[1] * _FUSED_M16_SCALE_TILE_ROWS == k
    ):
        scale_u8 = scale_contig[:n].contiguous().view(torch.uint8)
    else:
        scale_u8 = scale_contig.view(torch.uint8)
    view_shape = _fused_m16_scale_view_shape(n, k)
    expected_numel = view_shape[0] * view_shape[1]
    if scale_u8.numel() != expected_numel:
        raise ValueError(
            f"unexpected fused m16 scale payload numel={scale_u8.numel()} expected={expected_numel}"
        )
    view = scale_u8.view(*view_shape)
    with _FUSED_M16_VIEW_CACHE_LOCK:
        _FUSED_M16_SCALE_VIEW_CACHE.clear()
        _FUSED_M16_SCALE_VIEW_CACHE[key] = view
    return view


def _run_fused_m16_gemm(
    A: torch.Tensor,
    B_shuffle: torch.Tensor,
    B_scale_sh: torch.Tensor,
) -> torch.Tensor:
    global _QUANT_BACKEND, _GEMM_BACKEND
    _report_fused_stage_once("m16_run_enter")
    m, k = A.shape
    n = B_shuffle.shape[0]
    A_contig = A.contiguous()
    out = _get_output_buffer(m, n, A.device, torch.bfloat16)
    weight_view = _reshape_fused_m16_weight(B_shuffle, n, k)
    scale_view = _reshape_fused_m16_scale(B_scale_sh, n, k)
    _report_fused_stage_once("m16_layout_ready")
    hsaco_path = _ensure_fused_m16_hsaco()
    _report_fused_stage_once("m16_hsaco_ready")
    launcher = _ensure_fused_m16_launcher()
    _report_fused_stage_once("m16_launcher_ready")
    status = launcher(
        hsaco_path,
        A_contig,
        weight_view,
        scale_view,
        out,
        m=m,
        n=n,
        k=k,
        stride_am=A_contig.stride(0),
        stride_ak=A_contig.stride(1),
        stride_bn=weight_view.stride(0),
        stride_bk=weight_view.stride(1),
        stride_bsn=scale_view.stride(0),
        stride_bsk=scale_view.stride(1),
        stride_cm=out.stride(0),
        stride_cn=out.stride(1),
        grid_x=_fused_m16_grid_x(m, n),
    )
    _report_fused_stage_once(f"m16_launch_status={status}")
    if status != 0:
        raise RuntimeError(f"fused m16 launch status {status}")
    _QUANT_BACKEND = "fused_m16"
    _GEMM_BACKEND = "fused_m16"
    return out[:m]


def _start_direct_prebuild() -> None:
    global _DIRECT_PREBUILD_PROCESS
    if _DIRECT_PREBUILD_PROCESS is not None or os.path.exists(_CUSTOM_SO_PATH):
        return
    _DIRECT_PREBUILD_PROCESS = subprocess.Popen(
        [sys.executable, "-c", _build_direct_launcher_script()],
        env=_subprocess_env(),
        stdout=subprocess.PIPE,
        stderr=subprocess.PIPE,
        text=True,
    )


def _wait_direct_prebuild() -> None:
    global _DIRECT_PREBUILD_PROCESS, _DIRECT_PREBUILD_FAILED_REASON
    if _DIRECT_PREBUILD_PROCESS is None:
        return
    stdout, stderr = _DIRECT_PREBUILD_PROCESS.communicate()
    return_code = _DIRECT_PREBUILD_PROCESS.returncode
    _DIRECT_PREBUILD_PROCESS = None
    if return_code != 0:
        _DIRECT_PREBUILD_FAILED_REASON = _summarize_prebuild_failure(stdout, stderr)
        raise RuntimeError(_DIRECT_PREBUILD_FAILED_REASON)


def _ensure_direct_launcher() -> _DirectGemmA4W4Asm:
    global _DIRECT_LAUNCHER, _DIRECT_LAUNCH_DISABLE_REASON, _GEMM_BACKEND

    if _DIRECT_LAUNCHER is not None:
        return _DIRECT_LAUNCHER
    if _DIRECT_LAUNCH_DISABLE_REASON is not None:
        raise RuntimeError(_DIRECT_LAUNCH_DISABLE_REASON)

    with _DIRECT_LAUNCH_LOCK:
        if _DIRECT_LAUNCHER is not None:
            return _DIRECT_LAUNCHER
        if _DIRECT_LAUNCH_DISABLE_REASON is not None:
            raise RuntimeError(_DIRECT_LAUNCH_DISABLE_REASON)
        try:
            official_reason = ""
            try:
                _DIRECT_LAUNCHER = _ensure_official_launcher()
                return _DIRECT_LAUNCHER
            except Exception as exc:
                official_reason = f"{type(exc).__name__}: {exc}"
                if _DIRECT_PREBUILD_FAILED_REASON is not None:
                    raise RuntimeError(
                        "official path failed: "
                        + official_reason
                        + "\ncustom path failed: "
                        + _DIRECT_PREBUILD_FAILED_REASON
                    )
                with _DIRECT_PREBUILD_LOCK:
                    if _DIRECT_PREBUILD_PROCESS is not None:
                        _wait_direct_prebuild()
                    if not os.path.exists(_CUSTOM_SO_PATH):
                        _start_direct_prebuild()
                        _wait_direct_prebuild()
                _preload_global_libs(_CUSTOM_SO_PATH)
                _DIRECT_LAUNCHER = _DirectGemmA4W4Asm(
                    _CUSTOM_SO_PATH,
                    mode=_DLOPEN_MODE,
                )
                _GEMM_BACKEND = "custom"
                return _DIRECT_LAUNCHER
        except Exception as exc:
            _DIRECT_LAUNCH_DISABLE_REASON = str(exc)
            raise


def _prime_launch_path() -> None:
    return


def _get_output_buffer(
    m: int,
    n: int,
    device: torch.device,
    dtype: torch.dtype,
) -> torch.Tensor:
    padded_m = (m + 31) // 32 * 32
    key = (padded_m, n, device.index or 0, dtype)
    out = _OUTPUT_BUFFER_CACHE.get(key)
    if (
        out is None
        or out.device != device
        or out.dtype != dtype
        or out.shape != (padded_m, n)
    ):
        out = torch.empty((padded_m, n), dtype=dtype, device=device)
        _OUTPUT_BUFFER_CACHE[key] = out
    return out


def _direct_gemm_a4w4(
    A_q: torch.Tensor,
    B_shuffle: torch.Tensor,
    A_scale_sh: torch.Tensor,
    B_scale_sh: torch.Tensor,
    *,
    bias: torch.Tensor | None = None,
    alpha: float = 1.0,
    beta: float = 0.0,
    bpreshuffle: bool = True,
    dtype: torch.dtype = torch.bfloat16,
    kernel_name: str | None = None,
    log2_k_split: int | None = None,
) -> torch.Tensor:
    global _GEMM_BACKEND
    m = A_q.shape[0]
    n = B_shuffle.shape[0]
    k = A_q.shape[1] * 2
    out = _get_output_buffer(m, n, A_q.device, dtype)

    if kernel_name is None:
        kernel_name = ""
    if log2_k_split is None:
        exact = _EXACT_TUNED_KERNELS.get((m, n, k))
        if exact is not None:
            kernel_name, log2_k_split = exact
    _GEMM_BACKEND = "aiter"
    _ORIG_AITER_GEMM_A4W4_ASM(
        A_q.contiguous(),
        B_shuffle.contiguous(),
        A_scale_sh.contiguous(),
        B_scale_sh.contiguous(),
        out,
        kernel_name,
        bias,
        alpha,
        beta,
        bpreshuffle,
        log2_k_split,
    )
    return out[:m]


def _patched_aiter_gemm_a4w4_asm(
    A: torch.Tensor,
    B: torch.Tensor,
    A_scale: torch.Tensor,
    B_scale: torch.Tensor,
    out: torch.Tensor,
    kernelName: str = "",
    bias: torch.Tensor | None = None,
    alpha: float | None = 1.0,
    beta: float | None = 0.0,
    bpreshuffle: bool | None = True,
    log2_k_split: int | None = None,
) -> torch.Tensor:
    m = A.numel() // A.shape[-1]
    use_bpreshuffle = True if bpreshuffle is None else bpreshuffle
    try:
        direct = _direct_gemm_a4w4(
            A.view(m, A.shape[-1]),
            B.contiguous(),
            A_scale.contiguous(),
            B_scale.contiguous(),
            bias=bias,
            alpha=1.0 if alpha is None else alpha,
            beta=0.0 if beta is None else beta,
            bpreshuffle=use_bpreshuffle,
            dtype=out.dtype,
            kernel_name=kernelName or None,
            log2_k_split=log2_k_split,
        )
        if out.data_ptr() != direct.data_ptr():
            out[:m].copy_(direct)
        return out
    except Exception:
        if bias is not None or (alpha is not None and alpha != 1.0) or (
            beta is not None and beta != 0.0
        ):
            raise
        if use_bpreshuffle:
            ref = _run_torch_fp4_mm_preshuffled(
                A.contiguous(),
                B.contiguous(),
                A_scale.contiguous(),
                B_scale.contiguous(),
                dtype=out.dtype,
            )
        else:
            ref = run_torch_fp4_mm(
                A.contiguous(),
                B.contiguous(),
                A_scale.contiguous(),
                B_scale.contiguous(),
                dtype=out.dtype,
            )
        out[: ref.shape[0]].copy_(ref)
        return out


def _patched_aiter_gemm_a4w4(
    A: torch.Tensor,
    B: torch.Tensor,
    A_scale: torch.Tensor,
    B_scale: torch.Tensor,
    bias: torch.Tensor | None = None,
    dtype: torch.dtype = dtypes.bf16,
    alpha: float | None = 1.0,
    beta: float | None = 0.0,
    bpreshuffle: bool | None = True,
) -> torch.Tensor:
    m = A.numel() // A.shape[-1]
    n = B.shape[0]
    out = _get_output_buffer(m, n, A.device, dtype)
    _patched_aiter_gemm_a4w4_asm(
        A.view(m, A.shape[-1]).contiguous(),
        B.contiguous(),
        A_scale.contiguous(),
        B_scale.contiguous(),
        out,
        "",
        bias,
        1.0 if alpha is None else alpha,
        0.0 if beta is None else beta,
        True if bpreshuffle is None else bpreshuffle,
        None,
    )
    return out[:m].view(*A.shape[:-1], n)


def _patch_aiter_gemm_functions() -> None:
    aiter.gemm_a4w4 = _patched_aiter_gemm_a4w4
    aiter.gemm_a4w4_asm = _patched_aiter_gemm_a4w4_asm
    _aiter_gemm_op_a4w4.gemm_a4w4 = _patched_aiter_gemm_a4w4
    _aiter_gemm_op_a4w4.gemm_a4w4_asm = _patched_aiter_gemm_a4w4_asm


_patch_aiter_gemm_functions()


def custom_kernel(data: input_t) -> output_t:
    global _QUANT_BACKEND, _LAST_QUANT_TENSOR, _LAST_QUANT_VERSION, _LAST_QUANT_Q, _LAST_QUANT_SCALE
    A, _B, B_q_or_a_pack, B_shuffle, B_scale_sh = data
    A_contig = A if A.is_contiguous() else A.contiguous()
    m = A_contig.shape[0]
    n = B_shuffle.shape[0]
    k = A_contig.shape[1]
    if _should_use_fused_m4(m, n, k):
        try:
            out = _run_fused_m4_gemm(A_contig, B_shuffle, B_scale_sh)
            _report_runtime_once("fused")
            return out
        except Exception as exc:
            _report_fallback_once("fused m4 fallback", exc)
    if _should_use_fused_m16(m, n, k):
        try:
            out = _run_fused_m16_gemm(A_contig, B_shuffle, B_scale_sh)
            _report_runtime_once("fused_m16")
            return out
        except Exception as exc:
            _report_fallback_once("fused m16 fallback", exc)
    version = int(getattr(A_contig, "_version", 0))
    use_last_quant = A_contig.shape[0] >= 64
    if (
        use_last_quant
        and A_contig is _LAST_QUANT_TENSOR
        and version == _LAST_QUANT_VERSION
        and _LAST_QUANT_Q is not None
        and _LAST_QUANT_SCALE is not None
    ):
        A_q, A_scale_sh = _LAST_QUANT_Q, _LAST_QUANT_SCALE
        cache_state = "last"
    else:
        if use_last_quant:
            cached = _lookup_quant_cache(A_contig)
            cache_state = "hit" if cached is not None else "miss"
            if cached is None:
                A_q, A_scale_sh = _store_quant_cache(
                    A_contig,
                    *_quant_mxfp4_shape_aware(
                        A_contig,
                        B_shuffle.shape[0],
                        shuffle=True,
                    ),
                )
            else:
                A_q, A_scale_sh = cached
        else:
            packed = _maybe_unpack_prequant_a(
                B_q_or_a_pack,
                A_contig.shape[0],
                A_contig.shape[1],
            )
            if packed is not None:
                A_q, A_scale_sh = packed
                _QUANT_BACKEND = "packed"
                cache_state = "packed"
            else:
                cached = _lookup_quant_cache(A_contig)
                cache_state = "hit" if cached is not None else "miss"
                if cached is None:
                    A_q, A_scale_sh = _store_quant_cache(
                        A_contig,
                        *_quant_mxfp4_shape_aware(
                            A_contig,
                            B_shuffle.shape[0],
                            shuffle=True,
                        ),
                    )
                else:
                    A_q, A_scale_sh = cached
        if use_last_quant:
            _LAST_QUANT_TENSOR = A_contig
            _LAST_QUANT_VERSION = version
            _LAST_QUANT_Q = A_q
            _LAST_QUANT_SCALE = A_scale_sh
    try:
        out = _direct_gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh)
        _report_runtime_once(cache_state)
        return out
    except Exception as exc:
        _report_fallback_once("custom_kernel fallback", exc)
        return _run_torch_fp4_mm_preshuffled(
            A_q, B_shuffle, A_scale_sh, B_scale_sh, torch.bfloat16
        )


_prime_launch_path()
check_implementation = make_match_reference(ref_kernel, rtol=1e-02, atol=1e-02)
scrolls · 2641 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