"""Compile Kyojin's JIT kernels; --gpu also loads them through torch's HIP runtime.""" import ctypes import importlib import os from pathlib import Path import sys from exllamav3.util import hip_compiler from exllamav3.util.hip_lib import load_hip_runtime sdk = Path(os.environ["EXL3_ROCM_SDK"]) assert Path(hip_compiler.hipcc()) == sdk / "bin/hipcc" assert Path(os.environ["HIP_DEVICE_LIB_PATH"]).is_dir() assert (sdk / "amdgcn/bitcode").is_dir() if os.environ.get("KYOJIN_REQUIRE_SDK") == "1": assert hip_compiler.warning() is None, hip_compiler.warning() print(hip_compiler.describe(), flush=True) lib = None if "--gpu" in sys.argv: import torch # Also exercise the runtime bitcode path: merely loading a HIP module # does not catch incompatible device libraries in CLR/COMGR. x = torch.ones((32, 32), device="cuda", dtype=torch.float32) assert torch.equal(x @ x, torch.full_like(x, 32)) torch.cuda.synchronize() lib = load_hip_runtime() lib.hipModuleLoad.argtypes = [ctypes.POINTER(ctypes.c_void_p), ctypes.c_char_p] lib.hipModuleUnload.argtypes = [ctypes.c_void_p] for module, function, args in [ ("gr_mix_hip", "compile_hsaco", ()), ("exllamav3.modules.attention_fn.qsa_prefill_hip", "compile_hsaco", ()), ("exllamav3.modules.ple_fn.ple_hip", "compile_hsaco", ()), ("exllamav3.vendor.fla.hip.gdn_fused_h_hip", "_compile", ("", "gfx1151")), ("exllamav3.vendor.fla.hip.kda_fused_h_hip", "_compile", ("", "gfx1151")), ]: path = Path(getattr(importlib.import_module(module), function)(*args)) # HIP accepts ELF, offload bundles, and compressed offload bundles (CCOB). assert path.read_bytes().startswith((b"\x7fELF", b"__CLANG_OFFLOAD_BUNDLE__", b"CCOB")), path if lib is not None: handle = ctypes.c_void_p() assert lib.hipModuleLoad(ctypes.byref(handle), os.fsencode(path)) == 0, path assert lib.hipModuleUnload(handle) == 0, path print(f"OK: {module}", flush=True)