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
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.
fp4
Single-file submission for AMD MXFP4 matmul.num-warps = 4
_FUSED_M4_NUM_WARPS = 4split-k
"SPLITK_BLOCK_SIZE": "constexpr",stages = 1
_FUSED_M4_NUM_STAGES = 1tile-k = 512
_FUSED_M4_BLOCK_SIZE_K = 512tile-m = 8
_FUSED_M4_BLOCK_SIZE_M = 8tile-n = 64
_FUSED_M4_BLOCK_SIZE_N = 64Kernel 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