{ lib, stdenv, pkgs, python3, fetchFromGitHub, makeWrapper, rocmPackages, }: 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 ''; # 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 ]); 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`. 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} \ --add-flags "${src}/tools/$model/serve.py" done runHook postInstall ''; passthru = { inherit python rocmArch; pythonPackage = library; category = "AI Inference"; updateScript = [ "nix-update" "--flake" ".#kyojin" ]; }; meta = library.meta // { mainProgram = "kyojin-serve-qwen"; }; }