fix(kyojin): Isolate JIT and runtime toolchains
CI / check (push) Has been cancelled

Allow a pinned JIT SDK without rebuilding PyTorch or the extension.
Keep runtime device bitcode unchanged: exporting TheRock bitcode into
the serving process crashes CLR 7.2 on a basic tensor operation.

Disable unsupported expandable segments and check all five HIP kernel
modules, with an optional GPU matmul and module-loading regression test.
This commit is contained in:
2026-10-08 17:17:56 +03:00
parent 4c24610aae
commit 0813f55762
4 changed files with 94 additions and 2 deletions
+46
View File
@@ -0,0 +1,46 @@
"""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)
+2 -1
View File
@@ -1,6 +1,7 @@
{
pkgs,
inputs,
jitRocmSdk ? null,
...
}:
# The HIP extension has to be compiled against a ROCm build of PyTorch, and the
@@ -9,4 +10,4 @@
# ships torch 2.13, whose ROCm build fails to configure in nixpkgs and is not
# on cache.nixos.org.
inputs.nixpkgs-torch211.legacyPackages.${pkgs.stdenv.hostPlatform.system}.callPackage ./package.nix
{ }
{ inherit jitRocmSdk; }
+22 -1
View File
@@ -6,6 +6,8 @@
fetchFromGitHub,
makeWrapper,
rocmPackages,
# Override only the serving toolchain; torch and the extension keep their ROCm.
jitRocmSdk ? null,
}:
let
@@ -170,6 +172,7 @@ let
};
python = py.python.withPackages (_: [ library ]);
jitSdk = if jitRocmSdk == null then rocmToolkit else jitRocmSdk;
in
stdenv.mkDerivation {
inherit src version;
@@ -203,8 +206,9 @@ stdenv.mkDerivation {
makeWrapper ${python}/bin/python $out/bin/kyojin-serve-$model \
--set PYTORCH_ROCM_ARCH ${rocmArch} \
--set CC ${stdenv.cc}/bin/cc \
--set EXL3_ROCM_SDK ${rocmToolkit} \
--set EXL3_ROCM_SDK ${jitSdk} \
--set HIP_DEVICE_LIB_PATH ${rocmPackages."rocm-device-libs"}/amdgcn/bitcode \
--set EXL3_EXPANDABLE_SEGMENTS 0 \
--prefix PATH : ${lib.makeBinPath [ stdenv.cc ]} \
--prefix PYTHONPATH : ${src}/gr \
--add-flags "${src}/tools/$model/serve.py"
@@ -216,6 +220,23 @@ stdenv.mkDerivation {
passthru = {
inherit python rocmArch;
pythonPackage = library;
jitRocmSdk = jitSdk;
tests.jit =
pkgs.runCommand "kyojin-jit-check"
{
EXL3_ROCM_SDK = jitSdk;
# CLR/COMGR also reads this variable. Newer bitcode crashes torch
# matmul; a custom hipcc must select its own bitcode in its subprocess.
HIP_DEVICE_LIB_PATH = "${rocmPackages."rocm-device-libs"}/amdgcn/bitcode";
EXL3_EXPANDABLE_SEGMENTS = "0";
PYTHONPATH = "${src}/gr";
KYOJIN_REQUIRE_SDK = lib.boolToString (jitRocmSdk != null);
}
''
export HOME="$TMPDIR"
${python}/bin/python ${./check-jit.py}
touch "$out"
'';
category = "AI Inference";
updateScript = [
"nix-update"