Files
millerson-overlay.nix/packages/kyojin/package.nix
T
alex 0813f55762
CI / check (push) Has been cancelled
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.
2026-10-08 17:17:56 +03:00

252 lines
8.6 KiB
Nix

{
lib,
stdenv,
pkgs,
python3,
fetchFromGitHub,
makeWrapper,
rocmPackages,
# Override only the serving toolchain; torch and the extension keep their ROCm.
jitRocmSdk ? null,
}:
let
py = python3.pkgs;
# Kyojin is the Yamz fork of ExLlamaV3: a Python package plus a C++ extension
# that torch's cpp_extension cross-compiles to HIP code objects. Upstream
# only supports AMD Strix Halo, so the kernels are built for gfx1151 alone.
rocmArch = "gfx1151";
version = "1.3";
src = fetchFromGitHub {
owner = "Yamz-Labs";
repo = "kyojin";
rev = "v${version}";
hash = "sha256-/DPoqW/iaLXRAH8XF5RdE+K5YS4jZT6Sic6QeCDQMlU=";
};
# setup.py passes EXL3_ROCM_DEV_INCLUDE as -I to every translation unit. The
# AMD wheels ship one devel tree holding hipsparse/ and thrust/; nixpkgs
# splits them per package, so join the include dirs. torch's public HIP
# headers also pull hipsolver/, and it no longer vendors pybind11.
rocmDevInclude = pkgs.symlinkJoin {
name = "kyojin-rocm-dev-include";
paths = map lib.getDev (
[ py.pybind11 ]
++ [
rocmPackages.clr
rocmPackages.hipblas
rocmPackages.hipblaslt
rocmPackages.hipcub
rocmPackages.hipfft
rocmPackages.hiprand
rocmPackages.hipsparse
rocmPackages.hipsolver
rocmPackages.rocblas
rocmPackages.rocprim
rocmPackages.rocrand
rocmPackages.rocsolver
rocmPackages.rocsparse
rocmPackages.rocthrust
]
);
};
# torch's HIP path in cpp_extension expects one ROCM_HOME holding bin/hipcc,
# the headers and the amdgcn bitcode; nixpkgs ships them separately.
rocmToolkit = pkgs.symlinkJoin {
name = "kyojin-rocm-toolkit";
paths = [
rocmPackages."rocm-core"
rocmPackages."rocm-runtime"
rocmPackages."rocm-device-libs"
rocmPackages."rocm-comgr"
rocmPackages.clr
(lib.getBin rocmPackages.hipcc)
rocmPackages.rocblas
rocmPackages.hipblas
rocmPackages.hipblas-common
rocmPackages.hipblaslt
rocmPackages.hipsparse
rocmPackages.hipsolver
rocmPackages.hiprand
rocmPackages.rocrand
rocmPackages.hipfft
rocmPackages.rocsparse
rocmPackages.rocsolver
rocmPackages.rocthrust
rocmPackages.rocprim
rocmPackages.hipcub
];
# The setup hooks of the joined packages would rewrite build flags; the
# compiler only needs the files.
postBuild = ''
rm -rf $out/nix-support
'';
};
# The library: exllamav3 plus the compiled exllamav3_ext module.
library = py.buildPythonPackage {
inherit src version;
pname = "kyojin";
format = "pyproject";
build-system = with py; [
setuptools
wheel
];
nativeBuildInputs = [
py.ninja
rocmPackages."rocm-runtime"
];
buildInputs = [ (lib.getBin rocmPackages.hipcc) ];
propagatedBuildInputs = with py; [
aiohttp
huggingface-hub
jinja2
llguidance
marisa-trie
numpy
pillow
pydantic
pyyaml
rich
safetensors
tokenizers
torchWithRocm
typing-extensions
];
env = {
PYTORCH_ROCM_ARCH = rocmArch;
ROCM_PATH = "${rocmToolkit}";
ROCM_HOME = "${rocmToolkit}";
HIP_PATH = "${rocmToolkit}";
HIPCC = "${rocmToolkit}/bin/hipcc";
HIP_DEVICE_LIB_PATH = "${rocmPackages."rocm-device-libs"}/amdgcn/bitcode";
# Upstream build scripts look these up to find the ROCm devel tree.
EXL3_ROCM_SDK = "${rocmToolkit}";
EXL3_ROCM_DEV_INCLUDE = "${rocmDevInclude}/include";
# Default defines of upstream build.sh (tools/strix_halo/rebuild.sh).
EXL3_HIP_DEFINES = "EXL3_HIP_STG_PAD";
};
# setup.py pulls the extension source list and the arch helper with
# `from exllamav3...import ...`, which runs the package __init__ and
# therefore exllamav3/ext.py. That module JIT-compiles the extension when no
# precompiled exllamav3_ext is importable, i.e. during every wheel build
# (and it wants a writable $HOME for it). Register an empty `exllamav3`
# package first so only the two leaf modules are imported, and load
# build_config without the package wrapper.
postPatch = ''
sed -i '1i import sys as _sys, types as _types; _exl3_pkg = _types.ModuleType("exllamav3"); _exl3_pkg.__path__ = ["exllamav3"]; _sys.modules.setdefault("exllamav3", _exl3_pkg)' setup.py
sed -i 's|^from exllamav3\.exllamav3_ext\.build_config import get_sources as _get_sources$|import sys as _exl3_sys; _exl3_sys.path.insert(0, "exllamav3/exllamav3_ext"); from build_config import get_sources as _get_sources; _exl3_sys.path.pop(0)|' setup.py
grep -q 'from build_config import' setup.py
# Upstream serves these from the checkout (pip install -e .); the wheel only
# packages what pyproject lists, and three things are read out of the
# installed package at run time: the .hip sources the JIT kernels compile,
# the qsa_proof marker gating the QSA prefill kernel, and the dense-GEMM
# tuning seed (without it every start re-tunes for ~7 minutes).
sed -i 's|^ "exllamav3_ext/\*\*/\*",$|&\n "**/*.hip",\n "**/*.ok",\n "model/dense_gemm_tune_seed.txt",|' pyproject.toml
grep -q '"\*\*/\*.hip"' pyproject.toml
'';
# tests/ needs a real gfx1151 GPU and the published model packs.
doCheck = false;
meta = with lib; {
description = "Inference engine for 100 GB-class EXL3 MoE models on AMD Strix Halo (gfx1151, ROCm 7)";
homepage = "https://github.com/Yamz-Labs/kyojin";
changelog = "https://github.com/Yamz-Labs/kyojin/releases";
license = licenses.mit;
sourceProvenance = with sourceTypes; [ fromSource ];
# HIP code objects are gfx1151-only and ROCm is Linux-only.
platforms = platforms.linux;
};
};
python = py.python.withPackages (_: [ library ]);
jitSdk = if jitRocmSdk == null then rocmToolkit else jitRocmSdk;
in
stdenv.mkDerivation {
inherit src version;
pname = "kyojin";
# Only the wrappers are installed here: the serve scripts are plain scripts,
# they import their siblings by path, and the module itself is in `library`.
#
# Upstream `tools/strix_halo/env.sh` is sourced before serving, and on a
# distro box that brings a system `cc` (the AMD Triton backend compiles a
# CPython glue module at runtime, and again per kernel launcher) plus a hipcc
# for the JIT HIP kernels (qsa prefill, gdn, kda, ple). Nix gives the wrapper
# neither, so export them here instead: Triton takes the compiler from $CC,
# exllamav3 finds hipcc under $EXL3_ROCM_SDK, and hipclang only finds the
# amdgcn bitcode through $HIP_DEVICE_LIB_PATH (hipcc's own DEVICE_LIB_PATH is
# not enough for a --genco call). The ROCm paths stay out of LD_LIBRARY_PATH
# (upstream Trap 1/5: a second libhsa segfaults torch), and PYTHONPATH gets
# `gr/`: upstream env.sh puts the repo root there so `import gr_mix_hip` (the
# hand-written gated-residual WMMA kernel the server asks for with
# EXL3_GR_HIP=1) resolves, and it is not part of the wheel either.
dontConfigure = true;
dontBuild = true;
nativeBuildInputs = [ makeWrapper ];
installPhase = ''
runHook preInstall
mkdir -p $out/bin
for model in glm mimo qwen; do
makeWrapper ${python}/bin/python $out/bin/kyojin-serve-$model \
--set PYTORCH_ROCM_ARCH ${rocmArch} \
--set CC ${stdenv.cc}/bin/cc \
--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"
done
runHook postInstall
'';
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"
"--flake"
".#kyojin"
];
};
meta = library.meta // {
mainProgram = "kyojin-serve-qwen";
};
}