From 0813f55762ba6151d0fba8be2b196b1129da9266 Mon Sep 17 00:00:00 2001 From: Alexander Miroshnichenko Date: Thu, 8 Oct 2026 17:17:56 +0300 Subject: [PATCH] fix(kyojin): Isolate JIT and runtime toolchains 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. --- README.md | 24 +++++++++++++++++++ packages/kyojin/check-jit.py | 46 ++++++++++++++++++++++++++++++++++++ packages/kyojin/default.nix | 3 ++- packages/kyojin/package.nix | 23 +++++++++++++++++- 4 files changed, 94 insertions(+), 2 deletions(-) create mode 100644 packages/kyojin/check-jit.py diff --git a/README.md b/README.md index 3bfec0f..1458c0b 100644 --- a/README.md +++ b/README.md @@ -95,6 +95,30 @@ nix run git+https://git.millerson.name/alex/millerson-overlay.nix.git#mcp-gatewa nix profile install git+https://git.millerson.name/alex/millerson-overlay.nix.git#mcp-gateway ``` +### Kyojin JIT Toolchain + +Kyojin's serving wrappers use the packaged ROCm compiler by default. Supply a +pinned SDK to compile HIP JIT kernels with another compiler without rebuilding +PyTorch or the HIP extension: + +```nix +kyojin.override { jitRocmSdk = myJitSdk; } +``` + +The SDK must provide `bin/hipcc` and matching `amdgcn/bitcode`. Its compiler +wrapper must set `HIP_DEVICE_LIB_PATH` to that bitcode **in the compiler +subprocess only**, along with any required `ROCM_PATH` and Nix C++ toolchain +flags. Keep the compiler wrapper's real path inside the SDK prefix so upstream's +compiler provenance check recognizes it. + +The serving process keeps PyTorch's ROCm device libraries. Setting TheRock's +Clang 23 bitcode globally makes the ROCm 7.2 runtime crash on a basic PyTorch +matrix multiplication. The serving wrappers also disable unsupported expandable +allocator segments. `nix build .#kyojin.tests.jit` checks compilation of five HIP +kernel modules without requiring a GPU; `packages/kyojin/check-jit.py --gpu`, run +with the package's Python environment and serving variables, additionally checks +PyTorch matrix multiplication and module loading through its existing HIP runtime. + ### nftablesbuilder Runtime Notes `nftablesbuilder` is a web interface for managing nftables. It consists of a diff --git a/packages/kyojin/check-jit.py b/packages/kyojin/check-jit.py new file mode 100644 index 0000000..b21575e --- /dev/null +++ b/packages/kyojin/check-jit.py @@ -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) diff --git a/packages/kyojin/default.nix b/packages/kyojin/default.nix index a543bca..a2bdcad 100644 --- a/packages/kyojin/default.nix +++ b/packages/kyojin/default.nix @@ -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; } diff --git a/packages/kyojin/package.nix b/packages/kyojin/package.nix index f6f9b7d..7d92642 100644 --- a/packages/kyojin/package.nix +++ b/packages/kyojin/package.nix @@ -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"