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:
@@ -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)
|
||||
@@ -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; }
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user