diff --git a/.gitattributes b/.gitattributes index b9d3042d888990496cf6308cade39963747919e5..d802057e8664dd6f2607aaa1f38acfd135fb0094 100644 --- a/.gitattributes +++ b/.gitattributes @@ -34,3 +34,13 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text bundle/models/target/tokenizer.json filter=lfs diff=lfs merge=lfs -text +assets/ciru-halo-agent.png filter=lfs diff=lfs merge=lfs -text +bundle/native/libornith_attention_iu4.so filter=lfs diff=lfs merge=lfs -text +bundle/native/libornith_head_i8_tile.so filter=lfs diff=lfs merge=lfs -text +bundle/native/libornith_persistent_iu4.so filter=lfs diff=lfs merge=lfs -text +bundle/native/libornith_routed_direct.so filter=lfs diff=lfs merge=lfs -text +bundle/native/libornith_routed_n32.so filter=lfs diff=lfs merge=lfs -text +bundle/native/libornith_routed_storage_n32.so filter=lfs diff=lfs merge=lfs -text +runtime/aiter-jit-gfx1151/module_aiter_core.so filter=lfs diff=lfs merge=lfs -text +runtime/wheels/amd_aiter-0.1.0rc1-cp314-cp314-linux_x86_64.whl filter=lfs diff=lfs merge=lfs -text +runtime/wheels/vllm-0.1.0rc2.dev9+g9255fd9fb9.rocm100-cp314-cp314-linux_x86_64.whl filter=lfs diff=lfs merge=lfs -text diff --git a/INSTALL.md b/INSTALL.md new file mode 100644 index 0000000000000000000000000000000000000000..6c136b5b203e82445d68fb074bf88a03d82d3558 --- /dev/null +++ b/INSTALL.md @@ -0,0 +1,75 @@ +# Install and run Ornith1.5 Ciru Halo Agent + +This release includes the target model, trained DFlash2 drafter, native kernels, custom vLLM plugin, and exact vLLM/AITER runtime wheels and source archives. **Use this runtime; stock `pip install vllm` does not provide the custom quantization or serving path.** + +## Hardware and platform + +- AMD Ryzen AI Max+ Strix Halo, gfx1151, with 128 GB unified memory. +- Linux x86-64 with a working AMD GPU driver, readable/writable `/dev/kfd` and render nodes. +- Budget roughly 100 GB available system memory for the loaded production profile. The recorded whole-host peak was 95.35 GB; other applications also use that memory. +- Allow at least 60 GB free disk for the 24.3 GB model assets, runtime installation and caches; source rebuilds need additional space. +- Validated platform: NixOS, Linux 7.2.2, glibc 2.42. A clean isolated runtime installation was tested on Strix Halo. The Ubuntu recipe below is provided for deployment and has not been independently validated on Ubuntu. Do not replace your system glibc to run this model. + +## Download + +Install [uv](https://docs.astral.sh/uv/getting-started/installation/) and Git, then: + +```bash +uvx --from huggingface_hub hf download \ + jcbtc/Ornith1.5-Ciru-Halo-Agent-vllm-strix-halo \ + --local-dir ./ciru-halo-agent +cd ciru-halo-agent +``` + +## Ubuntu 26.04 LTS prerequisites + +Ubuntu 26.04 supplies [glibc 2.43](https://packages.ubuntu.com/resolute/libc6). This is the mainstream distro recipe; the validated host remains NixOS. + +```bash +sudo apt-get update +sudo apt-get install -y build-essential git cmake ninja-build pkg-config xxd \ + curl ca-certificates tar libnuma-dev libdrm-dev libelf-dev libssl-dev \ + zlib1g-dev libvulkan-dev +``` + +Ensure your user can access `/dev/kfd` and `/dev/dri/renderD*`; GPU access must work before launching. The installer obtains the pinned ROCm SDK and Python packages; a stock distro vLLM package is not needed. + +```bash +bash runtime/INSTALL-ORNITH-RUNTIME.sh "$PWD/installed-runtime" +bash bundle/serve.sh --dry-run +bash bundle/serve.sh --host 127.0.0.1 --port 8000 +``` + +The runtime installer refuses to overwrite an existing installation. `uv` must be on PATH. Installation downloads pinned dependencies from the AMD wheel index and Python package index. + +## NixOS + +Enable the standard dynamic loader with `programs.nix-ld.enable = true;` and working AMD GPU device access. The provided shell supplies the compiler and host libraries: + +```bash +nix-shell runtime/shell.nix --run \ + 'bash runtime/INSTALL-ORNITH-RUNTIME.sh "$PWD/installed-runtime"' +nix-shell runtime/shell.nix --run \ + 'bash bundle/serve.sh --host 127.0.0.1 --port 8000' +``` + +## API and agent clients + +The server exposes an OpenAI-compatible API at `http://127.0.0.1:8000/v1`, model ID **`ciru-halo-agent`**. Point Hermes or another compatible agent client to that URL. Use `--host 0.0.0.0` only when you intend to expose it on your network; the launcher provides no authentication by default. + +```bash +curl http://127.0.0.1:8000/health +curl http://127.0.0.1:8000/v1/chat/completions \ + -H 'Content-Type: application/json' \ + -d '{"model":"ciru-halo-agent","messages":[{"role":"user","content":"Write a Python function that merges overlapping intervals."}],"temperature":0.6,"top_p":0.95,"max_tokens":4096,"chat_template_kwargs":{"enable_thinking":false}}' +``` + +The default profile supplies **262,144 tokens of per-request context capacity**, **eight active sequences**, **44 GiB shared KV/state pool**, prefix caching and adaptive speculation. Input and output share the context window; your agent client must reserve output space and compact history before filling it. Eight independent, fully populated 256K histories are not promised. The server is text-only in this release; older optional vision experiments are not presented as current-profile validation. + +First startup compiles/loads GPU kernels and creates caches. Wait for `/health` before sending work. Keep `bundle/cache` writable. Change runtime/model locations with `ORNITH_RUNTIME_ROOT`, `ORNITH_MODEL`, and `ORNITH_DRAFT`. Ordinary users do not need to change quantization or draft-policy settings. + +## Build from source + +The [Ciru source repository](https://github.com/ciru-ai/ornith-ciru-halo-agent) contains the model plugin, all eight native kernel sources, and the corresponding build script. See its `BUILD.md` for the native rebuild command. Runtime source archives are provided in this Hugging Face repository under `runtime/`; they include the matching vLLM/AITER source, licenses and release overlay notes. The binary installation above is the tested way to assemble the pinned engine; native source rebuilding is separate from retraining or requantizing the model. + +Pinned runtime: vLLM `0.1.0rc2.dev9+g9255fd9fb9.rocm100` (base `9255fd9fb9fedf4b29d574a8d8bb21d93892cc98` plus supplied cache overlay), AITER `0.1.0rc1`, Python 3.14.3, PyTorch `2.13.0+rocm10.0.0`, ROCm SDK 10.0.0 and Transformers 5.16.1. Preserve included third-party licenses when redistributing. diff --git a/assets/ciru-halo-agent.png b/assets/ciru-halo-agent.png new file mode 100644 index 0000000000000000000000000000000000000000..13953cfbda946ca3e39d04b98379b74c5c7a4795 --- /dev/null +++ b/assets/ciru-halo-agent.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2dda3ce144f139947260b33cd0216a2d231bbbb6ba1801752633b8f4082700f4 +size 2424752 diff --git a/bundle/native/libornith_attention_iu4.so b/bundle/native/libornith_attention_iu4.so new file mode 100644 index 0000000000000000000000000000000000000000..265bc1c8388f9ac46393d7f9c77db65ee4e97281 --- /dev/null +++ b/bundle/native/libornith_attention_iu4.so @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2905806824bece62ce9d9140608859df869f3e3d3663c67bf79ad4ea495cd61f +size 186640 diff --git a/bundle/native/libornith_dense_g256.so b/bundle/native/libornith_dense_g256.so new file mode 100644 index 0000000000000000000000000000000000000000..1e0d34b43624229bf9fc644a9574ccc4c0dee0d7 Binary files /dev/null and b/bundle/native/libornith_dense_g256.so differ diff --git a/bundle/native/libornith_dense_g256_n32.so b/bundle/native/libornith_dense_g256_n32.so new file mode 100644 index 0000000000000000000000000000000000000000..ea31a858d06d2a998720aacd4c08bd0db0910b77 Binary files /dev/null and b/bundle/native/libornith_dense_g256_n32.so differ diff --git a/bundle/native/libornith_head_i8_tile.so b/bundle/native/libornith_head_i8_tile.so new file mode 100644 index 0000000000000000000000000000000000000000..70b866b65b2f7484377402fed6971a9c1edc13d1 --- /dev/null +++ b/bundle/native/libornith_head_i8_tile.so @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2c03de5ed015a052b407436c7e42b50fa33ef76cb46dccab9b226546ca1d1e09 +size 128952 diff --git a/bundle/native/libornith_persistent_iu4.so b/bundle/native/libornith_persistent_iu4.so new file mode 100644 index 0000000000000000000000000000000000000000..c629c5e6074cad528d230e1cc10419a369739a83 --- /dev/null +++ b/bundle/native/libornith_persistent_iu4.so @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:40e5cf35157c3e716de8067f7eef1312a3b9e787dd52a6f2e232ccd2d07fa68c +size 103792 diff --git a/bundle/native/libornith_routed_direct.so b/bundle/native/libornith_routed_direct.so new file mode 100644 index 0000000000000000000000000000000000000000..01704e81ba6b2a9c2e32dd1fce14ae8d73eece2d --- /dev/null +++ b/bundle/native/libornith_routed_direct.so @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5b3ca6826547e0ff1c417e733229e104721cec403ccd8f88f7fcdb1726d6f4ed +size 156392 diff --git a/bundle/native/libornith_routed_n32.so b/bundle/native/libornith_routed_n32.so new file mode 100644 index 0000000000000000000000000000000000000000..96651e089ce58e73ac0b4c817b47ce29c67663e0 --- /dev/null +++ b/bundle/native/libornith_routed_n32.so @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:48415e71f6c85e3aefcecf99d72607bb05900244b4aaa9a7ba2c58378d92f5fb +size 151928 diff --git a/bundle/native/libornith_routed_storage_n32.so b/bundle/native/libornith_routed_storage_n32.so new file mode 100644 index 0000000000000000000000000000000000000000..cb0a5061fe40c856af0ba18531be421a58a77932 --- /dev/null +++ b/bundle/native/libornith_routed_storage_n32.so @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3ea288cd0e37a4c373038744c3cc4cd4eac09675ac137e953d5ae7a6cd43dd3d +size 158536 diff --git a/bundle/packaging/serve.sh b/bundle/packaging/serve.sh new file mode 100644 index 0000000000000000000000000000000000000000..0c850da37b160ee55dc337adb60bfe1e028fbf8b --- /dev/null +++ b/bundle/packaging/serve.sh @@ -0,0 +1,31 @@ +#!/usr/bin/env bash +# Copyright 2026 Ciru. Source only the explicitly selected installed runtime. +set -euo pipefail +repo_root=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd) +runtime_root=${ORNITH_RUNTIME_ROOT:-} +plugin_site=${ORNITH_PLUGIN_SITE:-$repo_root/.runtime/plugin-site} +cache_directory=${XDG_CACHE_HOME:-$HOME/.cache}/ornith-g256 +launch_args=() +while (($#)); do + case "$1" in + --runtime-root) runtime_root=${2:?--runtime-root requires a directory}; shift 2 ;; + --plugin-site) plugin_site=${2:?--plugin-site requires a directory}; shift 2 ;; + --cache-directory) cache_directory=${2:?--cache-directory requires a directory}; launch_args+=("$1" "$2"); shift 2 ;; + *) launch_args+=("$1"); shift ;; + esac +done +if [[ -z "$runtime_root" ]]; then + echo 'Set --runtime-root DIR (installed Ciru vLLM runtime) or ORNITH_RUNTIME_ROOT.' >&2 + exit 2 +fi +test -f "$runtime_root/runtime-env.sh" +test -x "$runtime_root/venv/bin/python" +test -d "$plugin_site/ornith_g256" +export VLLM_SOURCE="$runtime_root/vllm" VLLM_VENV="$runtime_root/venv" +export AITER_SOURCE="$runtime_root/aiter" +export XDG_CACHE_HOME="$cache_directory" AITER_JIT_DIR="$cache_directory/aiter" +# shellcheck source=/dev/null +source "$runtime_root/runtime-env.sh" +unset VLLM_SOURCE VLLM_VENV +export PYTHONPATH="$plugin_site${PYTHONPATH:+:$PYTHONPATH}" +exec "$runtime_root/venv/bin/python" -m ornith_g256.launch "${launch_args[@]}" --cache-directory "$cache_directory" diff --git a/bundle/paths.env b/bundle/paths.env new file mode 100644 index 0000000000000000000000000000000000000000..5ae323088a9cb5bb4b833f514b77aec323788b03 --- /dev/null +++ b/bundle/paths.env @@ -0,0 +1,4 @@ +# Environment overrides for the packaged release. +export ORNITH_MODEL="${ORNITH_MODEL:-$bundle_root/models/target}" +export ORNITH_DRAFT="${ORNITH_DRAFT:-$bundle_root/models/draft}" +export ORNITH_RUNTIME_ROOT="${ORNITH_RUNTIME_ROOT:-$bundle_root/../installed-runtime}" diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/INSTALLER b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/INSTALLER new file mode 100644 index 0000000000000000000000000000000000000000..5c69047b2eb8235994febeeae1da4a82365a240a --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/INSTALLER @@ -0,0 +1 @@ +uv \ No newline at end of file diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/METADATA b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/METADATA new file mode 100644 index 0000000000000000000000000000000000000000..be7fa3c31fed04dea0b0f51edc5bf80601674002 --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/METADATA @@ -0,0 +1,8 @@ +Metadata-Version: 2.4 +Name: ciru-ornith-g256 +Version: 0.0.2a0 +Summary: Self-contained project adapter for Ornith G256 and DFlash2 on Ciru vLLM +Author-email: Ciru +Requires-Python: >=3.10 +License-File: LICENSE-APACHE-2.0 +Dynamic: license-file diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/RECORD b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/RECORD new file mode 100644 index 0000000000000000000000000000000000000000..64ef437886cf9bce8ba745418b67c58dabb49f36 --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/RECORD @@ -0,0 +1,32 @@ +bin/ornith-g256-serve,sha256=gw9CG161mmR0ebQFSvD64ZhMtGPuBb9gDv9lhy6oHHw,319 +ciru_ornith_g256-0.0.2a0.dist-info/INSTALLER,sha256=5hhM4Q4mYTT9z6QB6PGpUAW81PGNFrYrdXMj4oM_6ak,2 +ciru_ornith_g256-0.0.2a0.dist-info/METADATA,sha256=rTkRa0xu1BdLKZsDFr6laFc4K679lMWJ8BlBNhUBdnc,256 +ciru_ornith_g256-0.0.2a0.dist-info/RECORD,, +ciru_ornith_g256-0.0.2a0.dist-info/REQUESTED,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0 +ciru_ornith_g256-0.0.2a0.dist-info/WHEEL,sha256=SmOxYU7pzNKBqASvQJ7DjX3XGUF92lrGhMb3R6_iiqI,91 +ciru_ornith_g256-0.0.2a0.dist-info/direct_url.json,sha256=MwFMTtCGcu6fn-TruM1lo2XsI7q0gT0JZbfW7sb3OEA,105 +ciru_ornith_g256-0.0.2a0.dist-info/entry_points.txt,sha256=ESe1wopdkiGfr5DNWZcSemlgPr3uvRre1OyHkl88QqU,121 +ciru_ornith_g256-0.0.2a0.dist-info/licenses/LICENSE-APACHE-2.0,sha256=z8d0m5b2O9McPEK1xHG_dWgUBT6EfBDz6wA0F7xSPTA,11358 +ciru_ornith_g256-0.0.2a0.dist-info/top_level.txt,sha256=EdCMLXnn8tDNBUR3JhkMvOi8hEw6IuhYFYEcfdhekxU,12 +ciru_ornith_g256-0.0.2a0.dist-info/uv_build.json,sha256=RBNvo1WzZ4oRRq0W9-hknpT7T8If536DEMBg9hyq_4o,2 +ciru_ornith_g256-0.0.2a0.dist-info/uv_cache.json,sha256=L73WKGonFia8yJ8vK90qJmbQxtUhlvCNsq7D4umUJX4,137 +ornith_g256/__init__.py,sha256=vk_ZSh7Mt3Drbg5zyLpEGPvSGuG2MMiKV7_hwsp4-TM,1004 +ornith_g256/attention.py,sha256=L908DthNywUqno-d9AaN34CpaZMb6gopXgItNrdTzaE,772 +ornith_g256/attention_fast.py,sha256=J7GhM8rbh_86Y8AnBJ4dqB5BEv1dUmCkRA_ehUkfg4o,12274 +ornith_g256/attention_tile.py,sha256=_OaqYdO7YhAobtgDAdQmoFy-g5s7DF8PGO25VuY9E3g,2157 +ornith_g256/attention_verify.py,sha256=OjKVmWlCZfS63swvm3CbKA2ZYLln1xkvOVmtng2MYHc,3755 +ornith_g256/column_backend.py,sha256=oDW4OTaeczsF9ozoPfDQIstiRVAW_M8wwpeBJ6s8VAM,5752 +ornith_g256/column_kernel.py,sha256=J_sFPZcBhHMqRiEgJq_TFS7Aeg38PcTFzFzA1TnxfIE,9422 +ornith_g256/config.py,sha256=IyENi2Je83Vjji1gQ6kqJhgaW18cU32dSzTJG3EJ828,2869 +ornith_g256/dflash_spec.py,sha256=yUi5nqsi0q9cX8DKn3ZKCFWj2GrKBgk2K_VHSpvv8bM,1393 +ornith_g256/gdn_spec.py,sha256=CTGxs2Adk1u2rN_LZdMH2uQQVkXNOHvZslbdelJTdqM,5696 +ornith_g256/launch.py,sha256=o_aMA9-DiJBB6FH2aOhEgynFUrl5GXU9YY9gTRQm_0w,9265 +ornith_g256/lifecycle.py,sha256=MoVVIXfLjOigFJoB-BaNK8RKc5qTlcqE4VrMy7IKC3M,8269 +ornith_g256/loader.py,sha256=qFPG0BQhSnrfKYVI-RdbbhjdBTOJF67Bbry0_Dw-yxM,2933 +ornith_g256/method.py,sha256=0cvBQvPLd1vgTOh8zE6AOWmNaLJYhmkFZ0ZWvdHAcaM,7340 +ornith_g256/moe_base.py,sha256=JvtRdBXySXhZ8uP_Gu2PMefFakqszPs0xmOOZdffB8Q,3683 +ornith_g256/native.py,sha256=U14h8Uy3BlRBBfRNWXK05mVtwRmUgbvWLQ3QUjJvPYU,6901 +ornith_g256/prefix_cache.py,sha256=tdbO44VRim0oVvy-ELoJOYqSgcC_MGFKkUzncIrovZs,4817 +ornith_g256/runtime.py,sha256=xCdlAa-_jiOLw0l8JK4A0XK--_HQdjAvskFlXKHWG2E,4997 +ornith_g256/worker.py,sha256=UpYQcXpzv4awhAKYiPsBSuk-z9-5rV2DUb_q-qUPw0A,9492 +ornith_g256/worker_base.py,sha256=oA7xKVFy5lXIs8fXkM_QMWEUqUkuSD3t_eoDbTt7-Gg,4601 diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/REQUESTED b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/REQUESTED new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/WHEEL b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/WHEEL new file mode 100644 index 0000000000000000000000000000000000000000..8acb95590701b87bf84eec079cf4e3989f63b098 --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/WHEEL @@ -0,0 +1,5 @@ +Wheel-Version: 1.0 +Generator: setuptools (79.0.1) +Root-Is-Purelib: true +Tag: py3-none-any + diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/direct_url.json b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/direct_url.json new file mode 100644 index 0000000000000000000000000000000000000000..ee840bf369bfa97be00e26bc5b514ee94be4bf1c --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/direct_url.json @@ -0,0 +1 @@ +{"url":"file:///srv/ssd/sn850x/scratch/crown/ornith-prefix-64k-v1/src/runtime/ornith_g256","dir_info":{}} \ No newline at end of file diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/entry_points.txt b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/entry_points.txt new file mode 100644 index 0000000000000000000000000000000000000000..706feaaf40c9feb2edaddda423c8839277393959 --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/entry_points.txt @@ -0,0 +1,5 @@ +[console_scripts] +ornith-g256-serve = ornith_g256.launch:main + +[vllm.general_plugins] +ornith_g256 = ornith_g256:register diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/licenses/LICENSE-APACHE-2.0 b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/licenses/LICENSE-APACHE-2.0 new file mode 100644 index 0000000000000000000000000000000000000000..d645695673349e3947e8e5ae42332d0ac3164cd7 --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/licenses/LICENSE-APACHE-2.0 @@ -0,0 +1,202 @@ + + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/top_level.txt b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/top_level.txt new file mode 100644 index 0000000000000000000000000000000000000000..3387e57824f90b70180371d49b8b223a7ef7433c --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/top_level.txt @@ -0,0 +1 @@ +ornith_g256 diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/uv_build.json b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/uv_build.json new file mode 100644 index 0000000000000000000000000000000000000000..9e26dfeeb6e641a33dae4961196235bdb965b21b --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/uv_build.json @@ -0,0 +1 @@ +{} \ No newline at end of file diff --git a/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/uv_cache.json b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/uv_cache.json new file mode 100644 index 0000000000000000000000000000000000000000..b221de5455c0c45daa6ea222178c9e5c3d88af02 --- /dev/null +++ b/bundle/plugin-site/ciru_ornith_g256-0.0.2a0.dist-info/uv_cache.json @@ -0,0 +1 @@ +{"timestamp":{"secs_since_epoch":1788727205,"nanos_since_epoch":456565118},"commit":null,"tags":null,"env":{},"directories":{"src":null}} \ No newline at end of file diff --git a/bundle/plugin-site/ornith_g256/__init__.py b/bundle/plugin-site/ornith_g256/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..141f428a1bf57dbe906a0f5fa74f0e64e25a0c90 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/__init__.py @@ -0,0 +1,26 @@ +"""Ciru G256 prototype; no installed vLLM files are modified.""" + + +def register(): + from .cache_full1120 import install as install_full1120 + install_full1120() + from vllm.model_executor.layers.quantization import ( + _CUSTOMIZED_METHOD_TO_QUANT_CONFIG, register_quantization_config, + ) + from vllm.v1.attention.backends.registry import AttentionBackendEnum, register_backend + from .config import OrnithG256Config + + existing = _CUSTOMIZED_METHOD_TO_QUANT_CONFIG.get("ornith_g256") + if existing is None: + register_quantization_config("ornith_g256")(OrnithG256Config) + elif existing is not OrnithG256Config: + raise RuntimeError("Another plugin owns ornith_g256") + register_backend(AttentionBackendEnum.CUSTOM, "ornith_g256.attention_fast.OrnithG256SelectableAttentionBackend") + # Install in every vLLM process: cache groups are also built in the engine + # core, which does not instantiate our worker. The wrapper is a no-op for + # prefix-off runs and models outside this project. + from .prefix_cache import install + install() + + from .adaptive_c1 import install as install_adaptive_c1 + install_adaptive_c1() diff --git a/bundle/plugin-site/ornith_g256/adaptive_c1.py b/bundle/plugin-site/ornith_g256/adaptive_c1.py new file mode 100644 index 0000000000000000000000000000000000000000..f63f9ca40537b250051ed487d015a9b72964a3fb --- /dev/null +++ b/bundle/plugin-site/ornith_g256/adaptive_c1.py @@ -0,0 +1,340 @@ +"""Private reversible C1 policy pilot; fixed controls and C2–8 DF7 retained.""" +from collections import deque +from functools import wraps +from pathlib import Path +import json +import os +import time + +CONTROL = Path(__file__).resolve().parents[3] / 'policy-mode' +_INSTALLED = False + + +class RequestCost: + """Request-local, reversible C1 probe policy. Thresholds are pilot settings.""" + MAX_CONTEXT = 32768 + INITIAL_FLOOR_MS = 19.0 + WINDOW = 8 + GAIN = .90 + MAX_PROBE_CYCLES = 4 + + def __init__(self): + self.mode = os.environ.get('ORNITH_C1_POLICY') or CONTROL.read_text().strip() + if self.mode not in ('k0', 'k7', 'k15', 'auto'): + raise ValueError(f'Unknown C1 mode {self.mode!r}') + self.depth = 15 if self.mode == 'auto' else int(self.mode[1:]) + self.cycles = 0 + self.samples = deque(maxlen=self.WINDOW) + self.floor_ms = self.INITIAL_FLOOR_MS + self.cycle_ms = {} + self.tokens = 0 + self.floor_tokens = 0 + self.warm = 2 + self.bad_windows = 0 + self.probe = None + self.retry_after = {0: 0, 7: 0, 15: 0} + self.last_15_probe = 0 + self.last_7_probe = 0 + self.paused = False + self.long_context = False + self.rid = None + + def event(self, kind, **values): + print('ORNITH_C1_' + kind + ' ' + json.dumps(dict( + request_id=self.rid, mode=self.mode, depth=self.depth, + output_progress=self.tokens, **values)), flush=True) + + def reset_to_15(self, reason): + old = self.depth + self.depth = 15 + self.probe = None + self.samples.clear() + self.warm = 2 + self.bad_windows = 0 + self.floor_tokens = 0 + if old != 15: + self.event('RECOVERY', previous=old, reason=reason) + + def budget(self, context): + if self.mode != 'auto': + return self.depth + if context > self.MAX_CONTEXT: + if not self.long_context: + self.reset_to_15('context_above_32768') + self.long_context = True + return 15 + if self.paused: + self.paused = False + self.reset_to_15('return_from_concurrency') + return self.depth + + def start_probe(self, candidate, base_cost, reason): + incumbent = self.depth + self.probe = dict(incumbent=incumbent, candidate=candidate, + baseline_ms=base_cost, samples=[], excess_ms=0.0) + self.depth = candidate + self.samples.clear() + self.warm = 1 + self.bad_windows = 0 + if candidate == 15: + self.last_15_probe = self.tokens + if candidate == 7: + self.last_7_probe = self.tokens + self.event('PROBE', incumbent=incumbent, candidate=candidate, + baseline_ms_per_token=base_cost, reason=reason) + + def finish_probe(self, accept, cost, reason): + probe = self.probe + assert probe is not None + candidate, incumbent = probe['candidate'], probe['incumbent'] + self.depth = candidate if accept else incumbent + self.probe = None + self.samples.clear() + self.warm = 1 + self.floor_tokens = 0 + if not accept: + self.retry_after[candidate] = self.tokens + 64 + if accept and incumbent == 15: + self.last_15_probe = self.tokens + self.retry_after[15] = self.tokens + 32 + self.event('PROBE_RESULT', incumbent=incumbent, candidate=candidate, + accepted=accept, cost_ms_per_token=cost, + reference_ms_per_token=probe['baseline_ms'], + observed_cycles=len(probe['samples']), excess_ms=probe['excess_ms'], + reason=reason) + # Failure of DF7 must not block an independently promising floor. + # The first paired run exposed repeated 15->7 rejections with K0 never + # reachable, despite measured DF15 costs far above ordinary decode. + if (not accept and candidate == 7 and incumbent == 15 + and probe['baseline_ms'] >= self.floor_ms / self.GAIN + and self.tokens >= self.retry_after[0]): + self.start_probe(0, probe['baseline_ms'], 'rejected_df7_try_floor_directly') + return + # Full seven-token proposals are direct evidence of a saturated DF7 + # block. Probe DF15 promptly; this still requires a measured gain. + if accept and candidate == 7 and incumbent == 0: + filled = sum(progress == 8 for _, progress in probe['samples']) + if filled >= 2 and self.tokens >= self.retry_after[15]: + self.start_probe(15, cost, 'recovered_full_df7_proposals') + + def observe(self, elapsed_ms, progressed, rid, actual, future, context): + self.rid = rid + self.tokens += progressed + if actual == 0: + self.floor_tokens += progressed + self.cycles += 1 + if self.mode == 'auto': + self.budget(context) + steady = actual == future == self.depth + if self.probe is not None: + # Include transition/warm work as probe overhead, but never use it + # as a stationary mode-cost estimate. + self.probe['excess_ms'] += max( + 0.0, elapsed_ms - progressed * self.probe['baseline_ms']) + if not steady: + self.event('TRANSITION', actual_depth=actual, future_depth=future, + elapsed_ms=elapsed_ms, progressed=progressed) + return + if self.mode == 'auto' and self.long_context: + return + if self.warm: + self.warm -= 1 + return + self.samples.append((elapsed_ms, progressed)) + if self.mode != 'auto': + if self.cycles in (2, 8, 16) or self.cycles % 32 == 0: + ms = sum(x[0] for x in self.samples) + progress = sum(x[1] for x in self.samples) + self.event('WINDOW', actual_depth=actual, cycle=self.cycles, + elapsed_ms=ms, progressed=progress, + tg=1000*progress/ms if ms else None) + return + if self.probe is not None: + probe = self.probe + probe['samples'].append((elapsed_ms, progressed)) + n = len(probe['samples']) + ms = sum(x[0] for x in probe['samples']) + progress = sum(x[1] for x in probe['samples']) + cost = ms / progress + if n < 2: + return + if n == 2 and cost <= self.GAIN*probe['baseline_ms'] and actual == 7 and all(p == 8 for _, p in probe['samples']): + self.cycle_ms[actual] = ms/n + self.finish_probe(True, cost, 'two_full_df7_blocks') + elif cost >= 1.25*probe['baseline_ms'] or (probe['excess_ms'] >= 100.0 and cost >= probe['baseline_ms']): + self.finish_probe(False, cost, 'quick_loss_or_probe_cost_cap') + elif n >= self.MAX_PROBE_CYCLES: + self.cycle_ms[actual] = ms/n + if actual == 0: + self.floor_ms = cost + self.finish_probe(cost <= self.GAIN*probe['baseline_ms'], cost, 'four_steady_cycles') + return + # Bound recovery latency by output progress, including warm K0 work. + if actual == 0 and self.floor_tokens >= 63: + ms = sum(x[0] for x in self.samples) + progress = sum(x[1] for x in self.samples) + self.floor_ms = ms/progress + self.floor_tokens = 0 + self.start_probe(7, self.floor_ms, 'floor_recovery_after_63_tokens') + return + if len(self.samples) < self.WINDOW: + return + samples = list(self.samples) + self.samples.clear() + ms = sum(x[0] for x in samples) + progress = sum(x[1] for x in samples) + cost = ms/progress + self.cycle_ms[actual] = ms/len(samples) + self.event('WINDOW', actual_depth=actual, cycle=self.cycles, + elapsed_ms=ms, progressed=progress, tg=1000/cost, + floor_ms_per_token=self.floor_ms) + if actual == 0: + self.floor_ms = cost + return + if actual == 15: + # Prefix truncation predicts opportunity only, never a measured + # speed claim: Q8 and Q16 noncausal proposals need not be identical. + cycle7 = self.cycle_ms.get(7, self.cycle_ms[15]*.76) + predicted7 = cycle7*len(samples)/sum(min(p, 8) for _, p in samples) + promising = predicted7 <= self.GAIN*cost and cost >= .85*self.floor_ms + self.bad_windows = self.bad_windows + 1 if promising else 0 + if self.bad_windows >= 2 and self.tokens >= self.retry_after[7]: + self.start_probe(7, cost, 'two_windows_predict_at_least_10pct_gain') + return + if actual == 7: + filled = sum(p == 8 for _, p in samples) + if (self.tokens >= self.retry_after[15] and + (filled >= 2 or self.tokens-self.last_15_probe >= 64)): + self.start_probe(15, cost, 'df7_full_blocks_or_periodic_df15_recovery') + elif cost >= 1.05*self.floor_ms and self.tokens >= self.retry_after[0]: + self.start_probe(0, cost, 'df7_cost_above_measured_floor') + + +class Policy: + def __init__(self): + self.requests = {} + self.pending = None + self.observing = None + self.previous_mode = None + + def budget(self, ids, scheduler=None): + if len(ids) != 1: + if len(ids) > 1: + for rid in ids: + state = self.requests.get(rid) + if state is not None and state.mode == 'auto': + state.paused = True + return 7 + rid = ids[0] + if rid not in self.requests: + self.requests[rid] = RequestCost() + state = self.requests[rid] + state.rid = rid + context = 0 if scheduler is None else scheduler.requests[rid].num_computed_tokens + return state.budget(context) + + +def _policy(scheduler): + if not scheduler.vllm_config.additional_config.get('ornith_g256', {}).get('adaptive_c1_fallback'): + return None + if not hasattr(scheduler, '_ornith_c1_cost_policy'): + scheduler._ornith_c1_cost_policy = Policy() + return scheduler._ornith_c1_cost_policy + + +def install(): + global _INSTALLED + if _INSTALLED: + return + from vllm.v1.core.sched.scheduler import Scheduler + schedule_original = Scheduler.schedule + update_original = Scheduler.update_from_output + stats_original = Scheduler.make_spec_decoding_stats + draft_original = Scheduler.update_draft_token_ids + + @wraps(schedule_original) + def schedule(self, *args, **kwargs): + policy = _policy(self) + if policy is None: + return schedule_original(self, *args, **kwargs) + policy.requests = {rid: state for rid, state in policy.requests.items() + if rid in self.requests and not self.requests[rid].is_finished()} + started = time.monotonic() + output = schedule_original(self, *args, **kwargs) + ids = list(output.num_scheduled_tokens) + future = policy.budget(ids, self) + output.num_spec_tokens_to_schedule = future + policy.pending = dict(started=started, ids=ids, future=future) + mode = (len(ids), future) + if ids and mode != policy.previous_mode: + print('ORNITH_C1_MODE ' + json.dumps(dict(requests=len(ids), draft=future)), flush=True) + policy.previous_mode = mode + return output + + @wraps(update_original) + def update(self, scheduler_output, model_runner_output): + policy = _policy(self) + if policy is None: + return update_original(self, scheduler_output, model_runner_output) + ids = list(scheduler_output.num_scheduled_tokens) + rid = ids[0] if len(ids) == 1 else None + cached = scheduler_output.scheduled_cached_reqs + depth = len(scheduler_output.scheduled_spec_decode_tokens.get(rid, ())) + eligible = (rid is not None and not scheduler_output.scheduled_new_reqs + and rid in cached.req_ids and not cached.is_context_phase(rid) + and depth in (0, 7, 15) + and scheduler_output.num_scheduled_tokens[rid] == depth + 1) + policy.observing = dict(rid=rid, accepted=0 if depth == 0 else None) if eligible else None + try: + result = update_original(self, scheduler_output, model_runner_output) + seen = policy.observing + if seen is not None and seen['accepted'] is not None and policy.pending is not None: + policy.pending['observation'] = (rid, seen['accepted'] + 1, depth) + return result + finally: + policy.observing = None + + @wraps(stats_original) + def stats(self, spec_decoding_stats, num_draft_tokens, num_accepted_tokens, + num_invalid_spec_tokens, request_id): + result = stats_original(self, spec_decoding_stats, num_draft_tokens, + num_accepted_tokens, num_invalid_spec_tokens, request_id) + policy = getattr(self, '_ornith_c1_cost_policy', None) + seen = None if policy is None else policy.observing + if (seen is not None and request_id == seen['rid'] + and 0 <= num_accepted_tokens <= num_draft_tokens + and not (num_invalid_spec_tokens or {}).get(request_id, 0)): + seen['accepted'] = num_accepted_tokens + return result + + @wraps(draft_original) + def draft(self, draft_token_ids): + result = draft_original(self, draft_token_ids) + policy = _policy(self) + if policy is None: + return result + pending = policy.pending + policy.pending = None + if pending is None: + return result + observation = pending.get('observation') + if observation is not None: + rid, progressed, depth = observation + state = policy.requests.get(rid) + request = self.requests.get(rid) + if state is not None and request is not None and not request.is_finished(): + state.observe((time.monotonic()-pending['started'])*1000, progressed, rid, depth, pending['future'], request.num_computed_tokens) + # Truncate only before reservation, retaining the already tested K0 + # context updates and homogeneous Q8 batch path on later arrivals. + live = [rid for rid in pending['ids'] if rid in self.requests + and not self.requests[rid].is_finished()] + budget = policy.budget(live, self) + if len(live) == 1: + del self.requests[live[0]].spec_token_ids[budget:] + return result + + Scheduler.schedule = schedule + Scheduler.update_from_output = update + Scheduler.make_spec_decoding_stats = stats + Scheduler.update_draft_token_ids = draft + _INSTALLED = True diff --git a/bundle/plugin-site/ornith_g256/attention.py b/bundle/plugin-site/ornith_g256/attention.py new file mode 100644 index 0000000000000000000000000000000000000000..8711bc6aa6c70d20c6d45b7d9cc8b960a1866e01 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention.py @@ -0,0 +1,19 @@ +"""Reuse the retained exact-parent column kernel without its old admission campaign.""" +# Copyright 2026 Ciru. +from collections import Counter +from vllm.v1.attention.backends.rocm_attn import RocmAttentionBackend, RocmAttentionImpl +from .column_backend import OrnithColumnAttentionImpl + + +class OrnithG256AttentionBackend(RocmAttentionBackend): + @staticmethod + def get_name(): return 'CUSTOM' + @staticmethod + def get_impl_cls(): return OrnithG256AttentionImpl + + +class OrnithG256AttentionImpl(OrnithColumnAttentionImpl): + def __init__(self, *args, **kwargs): + RocmAttentionImpl.__init__(self, *args, **kwargs) + self.dispatch_counts, self.capture_by_C, self.eager_by_C = Counter(), Counter(), Counter() + self.context_bounds = [None, None] diff --git a/bundle/plugin-site/ornith_g256/attention_compact.py b/bundle/plugin-site/ornith_g256/attention_compact.py new file mode 100644 index 0000000000000000000000000000000000000000..b413684e5c552abb262e1ad6bf218063fafb6159 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_compact.py @@ -0,0 +1,204 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright 2026 Ciru. +"""Shared dense KV scratch for target prefill with installed AMD Triton flash. + +Configure once before profiling/capture. Calls must be serialized across target +layers, as for the existing native arena. The backend retains ownership of KV +writes. This module neither installs a hook nor changes decode dispatch. +""" +import ast +import inspect +import os + +os.environ['FLASH_ATTENTION_TRITON_AMD_AUTOTUNE'] = '0' + +import torch +from vllm.logger import init_logger +from vllm.triton_utils import tl, triton + +logger = init_logger('vllm.ornith_g256.attention_compact') +_arena = None +_prefill = None + + +@triton.jit +def _kv_boundaries(Starts, Seq, CuK, REQUESTS: tl.constexpr, BLOCK: tl.constexpr): + req = tl.arange(0, BLOCK) + start = tl.load(Starts + req, req < REQUESTS, other=0) + end = tl.load(Starts + req + 1, req < REQUESTS, other=0) + length = tl.load(Seq + req, req < REQUESTS, other=0) + length = tl.where(end > start, length, 0) + boundaries = tl.cumsum(length, 0) + tl.store(CuK, 0) + tl.store(CuK + req + 1, boundaries, req < REQUESTS) + + +@triton.jit +def _gather(K, V, Table, CuK, DenseK, DenseV, + KBS: tl.constexpr, KH: tl.constexpr, KD: tl.constexpr, + KT: tl.constexpr, KX: tl.constexpr, + VBS: tl.constexpr, VH: tl.constexpr, VD: tl.constexpr, VT: tl.constexpr, + TABLE_STRIDE: tl.constexpr, TOKENS: tl.constexpr, PAGE_SIZE: tl.constexpr): + token = tl.program_id(0) * TOKENS + tl.arange(0, TOKENS) + head = tl.program_id(1) + req = tl.program_id(2) + first = tl.load(CuK + req) + length = tl.load(CuK + req + 1) - first + if tl.program_id(0) * TOKENS >= length: + return + dim = tl.arange(0, 256) + valid = token < length + block = tl.load(Table + req * TABLE_STRIDE + token // PAGE_SIZE, + valid, other=0).to(tl.int64) + within = token % PAGE_SIZE + k_offset = (block[:, None] * KBS + head * KH + (dim[None, :] // 8) * KD + + within[:, None] * KT + (dim[None, :] % 8) * KX) + v_offset = (block[:, None] * VBS + head * VH + dim[None, :] * VD + + within[:, None] * VT) + key = tl.load(K + k_offset, valid[:, None], other=0) + value = tl.load(V + v_offset, valid[:, None], other=0) + offset = ((first + token[:, None]) * 2 + head) * 256 + dim[None, :] + tl.store(DenseK + offset, key, valid[:, None]) + tl.store(DenseV + offset, value, valid[:, None]) + + +def _graph_safe_prefill(): + from aiter.ops.triton._triton_kernels.flash_attn_triton_amd import fwd_prefill as module + if module.AUTOTUNE != 'off': + raise RuntimeError('Set FLASH_ATTENTION_TRITON_AMD_AUTOTUNE=0 before importing AMD flash attention') + tree = ast.parse(inspect.getsource(module.attention_forward_prefill_triton_impl)) + + class RemoveDeviceBoundaryAssertions(ast.NodeTransformer): + removed = 0 + + def visit_Assert(self, node): + # vLLM owns/validates the GPU metadata. The four wrapper checks + # read cu[0]/cu[-1] on the host, preventing graph capture. Capacity + # tails are also legal here: cu[-1] need not equal tensor capacity. + if any(isinstance(child, ast.Subscript) + and isinstance(child.value, ast.Name) + and child.value.id in ('cu_seqlens_q', 'cu_seqlens_k') + for child in ast.walk(node.test)): + self.removed += 1 + return None + return node + + transform = RemoveDeviceBoundaryAssertions() + tree = transform.visit(tree) + if transform.removed != 4: + raise RuntimeError('Unsupported AMD flash varlen boundary assertions') + namespace = dict(vars(module)) + exec(compile(ast.fix_missing_locations(tree), __file__ + ':varlen', 'exec'), namespace) + return namespace['attention_forward_prefill_triton_impl'] + + +def configure(*, max_num_seqs, max_model_len, max_num_batched_tokens, device, + iu4_prefill_library=None): + """Allocate one reusable arena, returning its accounted tensor bytes.""" + global _arena, _prefill + device = torch.device(device) + if device.index is None: + device = torch.device('cuda', torch.cuda.current_device()) + library = (os.path.realpath(os.path.expanduser(os.fspath(iu4_prefill_library))) + if iu4_prefill_library is not None else None) + identity = (max_num_seqs, max_model_len, max_num_batched_tokens, device) + if _arena is not None: + if _arena['identity'] != identity or _arena['iu4_library'] != library: + raise RuntimeError('Compact attention arena was configured differently') + return _arena['bytes'] + if (not 1 <= max_num_seqs <= 8 or not 1 <= max_model_len <= 262144 + or not 1 <= max_num_batched_tokens <= 2048 + or torch.cuda.is_current_stream_capturing()): + raise ValueError('Configure target compact prefill before capture, within C8/256K/2048 tokens') + _prefill = _graph_safe_prefill() + key = torch.empty((max_num_seqs * max_model_len, 2, 256), + dtype=torch.bfloat16, device=device) + value = torch.empty_like(key) + cu_k = torch.empty(max_num_seqs + 1, dtype=torch.int32, device=device) + lse = torch.empty((16, max_num_batched_tokens), dtype=torch.float32, device=device) + tensors = (key, value, cu_k, lse) + iu4 = None + if library is not None: + from .attention_iu4 import NativeAttention + iu4 = NativeAttention(device, library, max_model_len=max_model_len) + size = sum(t.numel() * t.element_size() for t in tensors) + if iu4 is not None: + size += iu4.bytes + _arena = dict(identity=identity, key=key, value=value, cu_k=cu_k, lse=lse, + bytes=size, iu4=iu4, iu4_library=library) + return size + + +def try_forward(query, key_cache, value_cache, output, block_table, + query_start_loc, seq_lens, *, max_query_len, max_seq_len, sm_scale): + """Return False for an unsupported call; otherwise write output and return True. + + Supports variable-length C1..C8 prefill, mixed small queries, empty request + slots and trailing graph token padding. CUDA metadata stays on the GPU. + Target eligibility (causal, no window/sinks/bias/output scaling) is the + caller's responsibility. The native cache writer must have run first. + """ + page_size = key_cache.shape[3] if key_cache.ndim == 5 else 0 + if (max_query_len <= 8 or query.ndim != 3 or query.shape[1:] != (16, 256) + or page_size not in (1120, 2240) or key_cache.shape[1:] != (2, 32, page_size, 8) + or value_cache.ndim != 4 or value_cache.shape[1:] != (2, 256, page_size) + or output.shape != query.shape or query.stride(2) != 1 + or output.stride(2) != 1 or query_start_loc.ndim != 1 + or seq_lens.ndim != 1 or block_table.ndim != 2 + or block_table.stride(1) != 1 + or not query_start_loc.is_contiguous() or not seq_lens.is_contiguous() + or query_start_loc.numel() != seq_lens.numel() + 1 + or block_table.shape[0] < seq_lens.numel() + or block_table.shape[1] * page_size < max_seq_len + or any(t.dtype != torch.bfloat16 for t in (query, key_cache, value_cache, output)) + or any(t.dtype != torch.int32 for t in (block_table, query_start_loc, seq_lens)) + or any(t.device != query.device for t in (key_cache, value_cache, output, + block_table, query_start_loc, seq_lens))): + return False + if _arena is None: + raise RuntimeError('Configure compact attention before memory profiling') + max_reqs, max_length, max_tokens, device = _arena['identity'] + requests, rows = seq_lens.numel(), query.shape[0] + if (not 1 <= requests <= max_reqs or not 1 <= max_seq_len <= max_length + or not 1 <= rows <= max_tokens or query.device != device): + return False + # Capacity-sized views avoid reading the final cumulative GPU length. + # The original flash kernel bounds every request by cu_q/cu_k instead. + key = _arena['key'][:requests * max_seq_len] + value = _arena['value'][:requests * max_seq_len] + cu_k = _arena['cu_k'][:requests + 1] + lse = _arena['lse'][:, :rows] + _kv_boundaries[(1,)](query_start_loc, seq_lens, cu_k, REQUESTS=requests, + BLOCK=triton.next_power_of_2(requests), num_warps=1) + _gather[(triton.cdiv(max_seq_len, 32), 2, requests)]( + key_cache, value_cache, block_table, cu_k, key, value, + *key_cache.stride(), *value_cache.stride(), + TABLE_STRIDE=block_table.stride(0), TOKENS=32, PAGE_SIZE=page_size, + num_warps=8, num_stages=1) + # Graph-padded rows are not real requests and receive defined zero output. + output.zero_() + if (_arena['iu4'] is not None and requests == 1 and max_query_len == 1120 + and max_seq_len >= 4096 and max_seq_len % 32 == 0 + and max_query_len <= rows and sm_scale == 256**-.5): + _arena['iu4'].forward(query, key, value, output, query_start_loc, + seq_lens, max_seq_len) + logger.info_once('Ornith C1 IU4 prefill active: BF16 cache gather + ' + 'normalized Q/K H256 and P/V H32, signed IU4 QK/PV; ' + 'Q1120 and aligned K>=4096, existing BF16 fallback elsewhere') + return True + _prefill(q=query, k=key, v=value, o=output, softmax_lse=lse, + sd_mask=None, sm_scale=sm_scale, alibi_slopes=None, causal=True, + window_size_left=-1, window_size_right=-1, bias=None, layout='thd', + cu_seqlens_q=query_start_loc, cu_seqlens_k=cu_k, + # Varlen reads actual Q/K lengths from cu_q/cu_k. These are only + # launch bounds (Q also sets the rectangular grid), yet upstream + # specializes both as constexpr. Stable arena bounds avoid compiling + # a new flash kernel at every growing-context prefill chunk. Empty Q + # tiles return before the attention loop in the installed kernel. + max_seqlens_q=max_tokens, max_seqlens_k=max_length, + dropout_p=0.0, philox_seed=0, philox_offset=0, + return_scores=False, use_exp2=True, + q_descale=None, k_descale=None, v_descale=None) + logger.info_once('Ornith compact target prefill active: shared KV gather + AMD ' + 'Triton varlen flash; <=8-query decode unchanged') + return True diff --git a/bundle/plugin-site/ornith_g256/attention_fast.py b/bundle/plugin-site/ornith_g256/attention_fast.py new file mode 100644 index 0000000000000000000000000000000000000000..e6116e5f8420d8b3850acc63587289fcf00660ae --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_fast.py @@ -0,0 +1,216 @@ +"""Optional retained FP32 segmented decode; stock ROCm handles multiple queries.""" +# Copyright 2026 Ciru. +import ctypes +from collections import Counter +from types import SimpleNamespace + +import torch +from vllm.config import get_current_vllm_config +from vllm.v1.attention.backend import AttentionType +from vllm.v1.attention.backends.rocm_attn import RocmAttentionBackend, RocmAttentionImpl +from vllm.v1.attention.ops.paged_attn import PagedAttention + + +class PagedLayout(ctypes.Structure): + _fields_ = [('bytes', ctypes.c_size_t), ('valid_lengths', ctypes.c_size_t), + ('partial', ctypes.c_size_t), ('segments', ctypes.c_int)] + + +class PagedStrides(ctypes.Structure): + _fields_ = [(name, ctypes.c_int64) for name in + ('q0', 'q1', 'k0', 'k1', 'k2', 'k3', 'k4', 'v0', 'v1', 'v2', 'v3', 'table0')] + + +class NativeFastFP32: + """ABI1 binding only: same library, strides and scratch layout as paged_torch.""" + def __init__(self, library, *, Ccap, Lcap, device): + if not 1 <= Ccap <= 8 or not 1 <= Lcap <= 8192: + raise ValueError('FP32 paged decode supports <=8 sequences and <=8192 context') + self.Ccap, self.Lcap, self.device = Ccap, Lcap, torch.device(device) + self.library = str(library) + self.lib = ctypes.CDLL(self.library) + self.lib.ornith_paged_abi_version.argtypes = [] + self.lib.ornith_paged_abi_version.restype = ctypes.c_uint32 + if self.lib.ornith_paged_abi_version() != 1: + raise ValueError('Expected retained paged attention ABI1') + self.lib.ornith_paged_get_layout.argtypes = [ctypes.c_int, ctypes.c_int, ctypes.POINTER(PagedLayout)] + self.lib.ornith_paged_get_layout.restype = ctypes.c_int + self.lib.ornith_paged_launch.argtypes = ([ctypes.c_void_p] * 6 + [ctypes.c_size_t] + + [ctypes.c_void_p] * 3 + [ctypes.c_int] * 4 + [PagedStrides, ctypes.c_void_p]) + self.lib.ornith_paged_launch.restype = ctypes.c_int + layout = PagedLayout() + if self.lib.ornith_paged_get_layout(Ccap, Lcap, ctypes.byref(layout)): + raise ValueError('Paged attention workspace query failed') + self.workspace_bytes = layout.bytes + self.output_f32_offset = (layout.bytes + 255) // 256 * 256 + self.arena_bytes = self.output_f32_offset + Ccap * 16 * 256 * 4 + + def bind_arena(self, arena): + if (arena.dtype != torch.uint8 or arena.ndim != 1 or not arena.is_contiguous() + or arena.numel() < self.arena_bytes or arena.data_ptr() % 256 + or arena.device != self.device): + raise ValueError('Invalid shared attention arena') + output = arena[self.output_f32_offset:self.arena_bytes].view(torch.float32).view(self.Ccap, 16, 256) + return SimpleNamespace(workspace=arena[:self.workspace_bytes], + outputs=tuple(output[:c] for c in range(self.Ccap + 1))) + + def launch_out(self, query, key, value, table, lengths, buffers, flag, output): + strides = PagedStrides(*query.stride()[:2], *key.stride(), *value.stride(), table.stride(0)) + stream = torch.cuda.current_stream(query.device) + status = self.lib.ornith_paged_launch( + *[tensor.data_ptr() for tensor in (query, key, value, table, lengths, buffers.workspace)], + buffers.workspace.numel(), buffers.outputs[query.shape[0]].data_ptr(), output.data_ptr(), flag.data_ptr(), + query.shape[0], self.Ccap, self.Lcap, key.shape[0], strides, stream.cuda_stream) + if status: + raise RuntimeError(f'FP32 paged attention launch failed: {status}') + + +class OrnithG256SelectableAttentionBackend(RocmAttentionBackend): + @staticmethod + def get_name(): return 'CUSTOM' + + @staticmethod + def get_impl_cls(): + settings = get_current_vllm_config().additional_config.get('ornith_g256', {}) + mode = settings.get('attention_mode', 'column') + if mode == 'stock_rocm': + return RocmAttentionImpl + if mode == 'column': + from .attention import OrnithG256AttentionImpl + return OrnithG256AttentionImpl + if mode == 'fast_fp32': + if not settings.get('attention_library'): + raise ValueError('fast_fp32 requires additional_config.ornith_g256.attention_library') + return OrnithG256FastAttentionImpl + raise ValueError(f'Unknown ornith_g256 attention_mode: {mode}') + + +class OrnithG256FastAttentionImpl(RocmAttentionImpl): + implementation = 'ciru.ornith.g256.fast_fp32.abi1' + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self._ornith_fast = None + self._ornith_verify = None + self.dispatch_counts = Counter() + self.capture_by_C, self.eager_by_C = Counter(), Counter() + self.context_bounds = [None, None] + + def static_native_support(self): + return (self.attn_type == AttentionType.DECODER + and (self.num_heads, self.num_kv_heads, self.head_size) == (16, 2, 256) + and self.scale == 0.0625 and self.kv_cache_dtype in ('auto', 'bfloat16') + and self.alibi_slopes is None and self.sliding_window == (-1, -1) + and self.logits_soft_cap == 0 and self.sinks is None + and self.kv_sharing_target_layer_name is None) + + def bind_native(self, backend, buffers, flags, slot, prefix): + if self._ornith_fast is not None: + raise RuntimeError('Attention storage already bound') + if not self.static_native_support(): + raise ValueError('Unsupported Ornith FP32 attention geometry') + if (flags.dtype != torch.int32 or flags.ndim != 1 or flags.device != backend.device + or not 0 <= slot < flags.numel()): + raise ValueError('Invalid attention flag slot') + self._ornith_fast = (backend, buffers, flags[slot:slot + 1], slot, str(prefix)) + + def bind_verify(self, backend, buffers): + if self._ornith_fast is None or self._ornith_verify is not None: + raise RuntimeError('Bind verification once after normal attention storage') + self._ornith_verify = (backend, buffers) + + def _try_verify(self, query, kv_cache, m, output, output_scale, output_block_scale): + if (self._ornith_verify is None or m is None or not self.static_native_support() + or m.use_cascade or m.causal is not True or output_scale is not None + or output_block_scale is not None or not 1 <= m.max_query_len <= 8): + return False + backend, buffers = self._ornith_verify + count, sequences = query.shape[0], m.seq_lens.shape[0] + if (not 1 <= count <= backend.Ccap or not 1 <= sequences <= 8 + or m.max_seq_len > backend.Lcap or query.shape != (count, 16, 256) + or output.shape != query.shape or not 0 <= m.num_actual_tokens <= count + or m.block_table.shape[0] != sequences + or m.query_start_loc.shape != (sequences + 1,) + or m.block_table.shape[1] < backend.table_cols + or query.dtype != torch.bfloat16 or output.dtype != torch.bfloat16 + or kv_cache.dtype != torch.bfloat16 or m.block_table.dtype != torch.int32 + or m.seq_lens.dtype != torch.int32 or m.query_start_loc.dtype != torch.int32 + or query.stride(2) != 1 or not output.is_contiguous() + or m.block_table.stride(1) != 1 or not m.seq_lens.is_contiguous() + or not m.query_start_loc.is_contiguous()): + return False + kcache, vcache = PagedAttention.split_kv_cache(kv_cache.transpose(0, 1), 2, 256) + if kcache.shape[1:] != (2, 32, 1104, 8) or vcache.shape[1:] != (2, 256, 1104): + return False + # The inherited stock do_kv_cache_update executes before this attention + # operation. GPU query lengths mask the newly written future positions. + backend.launch_out(query, kcache, vcache, m.block_table, m.seq_lens, + m.query_start_loc, m.num_actual_tokens, buffers, + self._ornith_fast[2], output) + capturing = torch.cuda.is_current_stream_capturing() + self.dispatch_counts['capture_verify_calls' if capturing else 'eager_verify_calls'] += 1 + self.dispatch_counts['verify_calls'] += 1 + (self.capture_by_C if capturing else self.eager_by_C)[str(count)] += 1 + return True + + def forward(self, layer, query, key, value, kv_cache, attn_metadata, output, + output_scale=None, output_block_scale=None): + m = attn_metadata + if self._try_verify(query, kv_cache, m, output, output_scale, output_block_scale): + return output + reason = None + if m is None: + reason = 'profile' + elif not self.static_native_support(): + reason = 'static_feature' + elif (m.use_cascade or m.causal is not True or output_scale is not None + or output_block_scale is not None): + reason = 'metadata_feature' + elif m.max_query_len != 1: + reason = 'prefill_or_multiquery' + elif self._ornith_fast is None: + raise RuntimeError('FP32 decode reached an unbound target attention layer') + else: + backend, buffers, flag, _, _ = self._ornith_fast + count = m.seq_lens.shape[0] + if (not 1 <= count <= backend.Ccap or m.max_seq_len > backend.Lcap + or query.shape != (count, 16, 256) or output.shape != query.shape + or not 0 < m.num_actual_tokens <= count + or m.block_table.shape[0] != count or m.query_start_loc.shape != (count + 1,) + or m.block_table.shape[1] < (backend.Lcap + 1055) // 1056): + reason = 'capacity_or_query_mapping' + elif (query.dtype != torch.bfloat16 or output.dtype != torch.bfloat16 + or kv_cache.dtype != torch.bfloat16 or m.block_table.dtype != torch.int32 + or m.seq_lens.dtype != torch.int32 or m.query_start_loc.dtype != torch.int32): + reason = 'dtype' + elif (query.stride(2) != 1 or not output.is_contiguous() + or m.block_table.stride(1) != 1 or not m.seq_lens.is_contiguous() + or not m.query_start_loc.is_contiguous()): + reason = 'stride' + else: + kcache, vcache = PagedAttention.split_kv_cache(kv_cache.transpose(0, 1), 2, 256) + if kcache.shape[1:] != (2, 32, 1056, 8) or vcache.shape[1:] != (2, 256, 1056): + reason = 'cache_page_shape' + else: + # The existing opaque Attention op owns this current-stream + # launch and the preceding stock cache update. No host KV read. + backend.launch_out(query, kcache, vcache, m.block_table, m.seq_lens, + buffers, flag, output) + capturing = torch.cuda.is_current_stream_capturing() + self.dispatch_counts['capture_native_calls' if capturing else 'eager_native_calls'] += 1 + (self.capture_by_C if capturing else self.eager_by_C)[str(count)] += 1 + self.dispatch_counts['native_calls'] += 1 + lo, hi = self.context_bounds + self.context_bounds = [m.max_seq_len if lo is None else min(lo, m.max_seq_len), + m.max_seq_len if hi is None else max(hi, m.max_seq_len)] + return output + self.dispatch_counts['fallback_' + reason] += 1 + return super().forward(layer, query, key, value, kv_cache, m, output, + output_scale, output_block_scale) + + def inspect_dispatch(self, prefix=None): + bound = self._ornith_fast + return dict(implementation=self.implementation, bound=bound is not None, + prefix=prefix or (bound[4] if bound else None), flag_slot=bound[3] if bound else None, + counts=dict(self.dispatch_counts), capture_by_C=dict(self.capture_by_C), + eager_by_C=dict(self.eager_by_C), native_max_seq_len_host_bounds=list(self.context_bounds)) diff --git a/bundle/plugin-site/ornith_g256/attention_folded.py b/bundle/plugin-site/ornith_g256/attention_folded.py new file mode 100644 index 0000000000000000000000000000000000000000..e83f107b58c3a1c5b7d8de05f9e05376b6c48439 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_folded.py @@ -0,0 +1,170 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright 2026 Ciru. Optional native query/GQA folding. +# Generated arithmetic retains vLLM/SGLang Apache-2.0 code: Copyright +# contributors to the vLLM project; Copyright2025 vLLM Team; +# Copyright2023-2024 SGLang Team. See installed triton_decode_attention.py. +"""Optional native K/V sharing in eight-query tiles times eight GQA heads. + +No installation hook: callers explicitly select this operation. +""" +from functools import lru_cache +import inspect +import linecache + +import torch +from vllm.triton_utils import tl, triton +from . import attention_partition as current + + +@triton.jit +def _reduction_lengths(Starts, Seq, RowSeq, + REQUESTS: tl.constexpr, REQUEST_BLOCK: tl.constexpr): + row = tl.program_id(0) + req = tl.arange(0, REQUEST_BLOCK) + first = tl.load(Starts + req, req < REQUESTS, other=0) + end = tl.load(Starts + req + 1, req < REQUESTS, other=0) + owns = (req < REQUESTS) & (row >= first) & (row < end) + owner = tl.max(tl.where(owns, req + 1, 0), 0) - 1 + seq = tl.load(Seq + owner, owner >= 0, other=0) + # All queries share the request's split boundaries. Causality is applied + # per row in stage1, rather than changing which partials stage2 visits. + tl.store(RowSeq + row, seq) + + +@lru_cache(maxsize=2) +def _kernel(query_tiles=1): + if query_tiles not in (1, 2): + raise ValueError('Folded attention supports one or two eight-query tiles') + native, reduce = current._kernels() + source = inspect.getsource(native.fn) + replacements = { + 'def _fwd_grouped_kernel_stage1(': 'def _folded_query_stage1(', + ' B_Seqlen,\n': ' B_Seqlen,\n Query_Start,\n', + ''' cur_head_id = tl.program_id(1) + cur_kv_head = cur_head_id // tl.cdiv(kv_group_num, BLOCK_H) + split_kv_id = tl.program_id(2) + + VALID_BLOCK_H: tl.constexpr = BLOCK_H if kv_group_num > BLOCK_H else kv_group_num + cur_head = cur_head_id * VALID_BLOCK_H + tl.arange(0, BLOCK_H) + mask_h = cur_head < (cur_head_id + 1) * VALID_BLOCK_H + mask_h = mask_h & (cur_head < q_head_num)''': + ''' cur_kv_head = tl.program_id(1) + split_kv_id = tl.program_id(2) + query_first = tl.load(Query_Start + cur_batch) + query_end = tl.load(Query_Start + cur_batch + 1) + query_count = query_end - query_first + if query_count <= 0: + return + offs_m = tl.arange(0, BLOCK_H) + cur_query = query_first + offs_m // 8 + cur_head = cur_kv_head * 8 + offs_m % 8 + mask_h = offs_m // 8 < query_count''', + ' cur_batch_seq_len = tl.load(B_Seqlen + cur_batch)\n': + ' cur_batch_seq_len = tl.load(B_Seqlen + cur_batch)\n' + ' causal_len = cur_batch_seq_len - query_count + offs_m // 8 + 1\n', + 'cur_batch * stride_qbs + cur_head[:, None] * stride_qh + offs_d[None, :]': + 'cur_query[:, None] * stride_qbs + cur_head[:, None] * stride_qh + offs_d[None, :]', + 'mask_h[:, None] & (offs_n[None, :] < split_kv_end), qk, float("-inf")': + 'mask_h[:, None] & (offs_n[None, :] < split_kv_end) & (offs_n[None, :] < causal_len[:, None]), qk, float("-inf")', + 'cur_batch * stride_mid_ob\n + cur_head[:, None] * stride_mid_oh': + 'cur_query[:, None] * stride_mid_ob\n + cur_head[:, None] * stride_mid_oh', + 'cur_batch * stride_mid_ob\n + cur_head * stride_mid_oh': + 'cur_query * stride_mid_ob\n + cur_head * stride_mid_oh', + ' acc / e_sum[:, None],': + ' tl.where(causal_len[:, None] > split_kv_start, acc / e_sum[:, None], 0.),', + ' e_max + tl.log(e_sum),': + ' tl.where(causal_len > split_kv_start, e_max + tl.log(e_sum), -float("inf")),', + } + for old, new in replacements.items(): + if source.count(old) != 1: + raise RuntimeError('Unsupported folded-query source anchor: ' + old[:80]) + source = source.replace(old, new) + kernel_name = '_folded_query_stage1' + if query_tiles == 2: + # Keep the original <=8-query specialization byte-for-byte. For up to + # sixteen queries, each request has two otherwise identical 64-row + # tiles. query_count remains the TOTAL request query count: shrinking + # it to the tile size would shift the causal positions of tile zero. + tiled_replacements = { + ' cur_batch = tl.program_id(0)\n': + ' cur_batch = tl.program_id(0) // 2\n' + ' query_tile = tl.program_id(0) % 2\n', + ' if query_count <= 0:\n': + ' if query_count <= query_tile * 8:\n', + ' cur_query = query_first + offs_m // 8\n': + ' cur_query = query_first + query_tile * 8 + offs_m // 8\n', + ' mask_h = offs_m // 8 < query_count\n': + ' mask_h = query_tile * 8 + offs_m // 8 < query_count\n', + ' causal_len = cur_batch_seq_len - query_count + offs_m // 8 + 1\n': + ' causal_len = cur_batch_seq_len - query_count + query_tile * 8 + offs_m // 8 + 1\n', + 'def _folded_query_stage1(': 'def _folded_query_tiles_stage1(', + } + for old, new in tiled_replacements.items(): + if source.count(old) != 1: + raise RuntimeError('Unsupported folded query-tile source anchor: ' + old[:80]) + source = source.replace(old, new) + kernel_name = '_folded_query_tiles_stage1' + filename = __file__ + ('.generated' if query_tiles == 1 else '.query_tiles.generated') + linecache.cache[filename] = (len(source), None, source.splitlines(True), filename) + namespace = dict(native.fn.__globals__) + exec(compile(source, filename, 'exec'), namespace) + return namespace[kernel_name], reduce + + +def forward(query, key_cache, value_cache, output, block_table, query_start_loc, + seq_lens, sm_scale, k_scale, v_scale, *, max_query_len=8): + """Match attention_partition.forward, with native current KV prewritten. + + Caller eligibility: causal target attention, no sinks/alibi/window/output + scaling, at most8 requests and128 total rows, max_query_len<=16. Query + starts include empty slots; trailing padded rows receive zero outputs. + Tensor metadata is trusted as in the ROCm backend; no host reads. + Scratch remains [query,16,32,257] FP32 plus LSE and per-query reducer extent. + This is R*526404 bytes and does not copy the block table or KV pool. + """ + page_size = key_cache.shape[3] if key_cache.ndim == 5 else 0 + if (query.ndim != 3 or query.shape[1:] != (16, 256) + or page_size not in (1120, 2240) or key_cache.shape[1:] != (2, 32, page_size, 8) + or value_cache.ndim != 4 or value_cache.shape[1:] != (2, 256, page_size) + or key_cache.shape[0] != value_cache.shape[0] + or key_cache.shape[0] == 0 or output.shape != query.shape): + raise ValueError("Folded attention requires target H16/KV2/D256/page1120 or2240") + if (not 1 <= max_query_len <= 16 or query.shape[0] > 128 + or not 1 <= seq_lens.numel() <= 8 + or block_table.ndim != 2 or block_table.shape[0] < seq_lens.numel() + or block_table.shape[1] < 1 or query_start_loc.numel() != seq_lens.numel() + 1 + or not seq_lens.is_contiguous() or not query_start_loc.is_contiguous() + or block_table.stride(1) != 1 or query.stride(2) != 1 or output.stride(2) != 1): + raise ValueError("Unsupported query metadata or noncontiguous inner dimensions") + if (not query.is_cuda or any(t.device != query.device for t in + (key_cache, value_cache, output, block_table, query_start_loc, seq_lens, k_scale, v_scale)) + or any(t.dtype != torch.bfloat16 for t in (query, key_cache, value_cache, output)) + or any(t.dtype not in (torch.int32, torch.int64) for t in + (block_table, query_start_loc, seq_lens))): + raise ValueError("Partitioned attention requires BF16 and integer metadata on one GPU") + rows = query.shape[0] + if not rows: + return output + query_tiles = 1 if max_query_len <= 8 else 2 + stage1, reduce = _kernel(query_tiles) + row_seq = torch.empty((rows,), dtype=torch.int32, device=query.device) + _reduction_lengths[(rows,)](query_start_loc, seq_lens, row_seq, + REQUESTS=seq_lens.numel(), REQUEST_BLOCK=triton.next_power_of_2(seq_lens.numel()), + num_warps=4) + splits = 32 + logits = torch.empty((rows, 16, splits, 257), dtype=torch.float32, device=query.device) + lse = torch.empty((rows, 16), dtype=torch.float32, device=query.device) + stage1[(seq_lens.numel() * query_tiles, 2, splits)]( + query, key_cache, value_cache, sm_scale, block_table, seq_lens, query_start_loc, logits, + block_table.stride(0), query.stride(0), query.stride(1), + key_cache.stride(0), key_cache.stride(3), key_cache.stride(1), + value_cache.stride(0), value_cache.stride(3), value_cache.stride(1), + logits.stride(0), logits.stride(1), logits.stride(2), k_scale, v_scale, + kv_group_num=8, q_head_num=16, BLOCK_DMODEL=256, BLOCK_DPE=0, + BLOCK_DV=256, BLOCK_N=16, BLOCK_H=64, NUM_KV_SPLITS=splits, + PAGE_SIZE=page_size, logit_cap=0., Lk=256, Lv=256, IS_MLA=False, + stride_buf_kds=key_cache.stride(2), stride_buf_kxs=key_cache.stride(4), + stride_buf_vds=value_cache.stride(2), num_warps=4, num_stages=1, + waves_per_eu=1, matrix_instr_nonkdim=16, kpack=2) + reduce(logits, query, output, lse, value_cache.transpose(2, 3), row_seq, splits) + return output diff --git a/bundle/plugin-site/ornith_g256/attention_iu4.py b/bundle/plugin-site/ornith_g256/attention_iu4.py new file mode 100644 index 0000000000000000000000000000000000000000..fc66610298413fe103cae497efed3291e303212d --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_iu4.py @@ -0,0 +1,41 @@ +# Copyright 2026 Ciru. Isolated C1 attention binding; no installed runtime edits. +import ctypes +import torch + + +class NativeAttention: + """Preallocated signed IU4 scratch and stream-only native enqueue calls.""" + + def __init__(self, device, library_path, max_model_len=65536): + if not 1120 <= max_model_len <= 262144: + raise ValueError("IU4 prefill capacity must be1120..262144") + self.max_keys = 65536 if max_model_len <= 65536 else 262144 + groups = self.max_keys // 32 + self.library = ctypes.CDLL(str(library_path)) + self.library.iu4_prepare.argtypes = [ctypes.c_void_p] * 9 + [ctypes.c_int] * 3 + [ctypes.c_void_p] + self.library.iu4_prepare.restype = ctypes.c_int + self.library.iu4_attention.argtypes = [ctypes.c_void_p] * 7 + [ctypes.c_int] + [ctypes.c_void_p] + self.library.iu4_attention.restype = ctypes.c_int + specs = [((16, 1120, 32), torch.int32), ((16, 1120), torch.float16), + ((2, groups, 32, 32), torch.int32), ((2, self.max_keys), torch.float16), + ((2, groups, 4, 256), torch.int32), ((2, groups, 256), torch.float16)] + self.packed = [torch.empty(shape, dtype=dtype, device=device) for shape, dtype in specs] + self.pointers = [tensor.data_ptr() for tensor in self.packed] + self.output = torch.empty((1120, 16, 256), dtype=torch.bfloat16, device=device) + self.bytes = self.output.numel() * self.output.element_size() + sum(tensor.numel() * tensor.element_size() for tensor in self.packed) + + def forward(self, query, key, value, output, starts, lengths, max_keys): + # C1 max_query_len1120 and exact CPU metadata max_seq_len are required. + # The unchanged consumer uses its original contiguous fixed1120 output. + # No allocation, tensor copy to CPU, synchronization or cache mutation here. + if not 1120 <= max_keys <= self.max_keys or max_keys % 32: + raise ValueError("IU4 prefill requires aligned keys within reserved capacity") + stream = torch.cuda.current_stream().cuda_stream + rc = self.library.iu4_prepare(query.data_ptr(), key.data_ptr(), value.data_ptr(), + *self.pointers, max_keys, query.stride(0), query.stride(1), stream) + if rc: + raise RuntimeError(f'IU4 attention preparation launch failed: HIP {rc}') + rc = self.library.iu4_attention(*self.pointers, self.output.data_ptr(), max_keys, stream) + if rc: + raise RuntimeError(f'IU4 attention launch failed: HIP {rc}') + output[:1120].copy_(self.output) diff --git a/bundle/plugin-site/ornith_g256/attention_iu4_persistent.py b/bundle/plugin-site/ornith_g256/attention_iu4_persistent.py new file mode 100644 index 0000000000000000000000000000000000000000..ec2470dc3017167c7aab06c9e982f546f907e15f --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_iu4_persistent.py @@ -0,0 +1,254 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright 2026 Ciru. Isolated optional persistent long-context attention. +"""Page1120 only; canonical BF16 cache owns page identity and prefix lifecycle.""" +import ctypes +from functools import lru_cache +import inspect +import linecache +import os +import types + +import numpy as np +import torch + +_state = None +_installed = False +_library = None +_threshold = 32768 + + +def _compile(source, namespace, suffix): + name = __file__ + suffix + linecache.cache[name] = (len(source), None, source.splitlines(True), name) + exec(compile(source, name, 'exec'), namespace) + + +@lru_cache(maxsize=2) +def _bf16_kernels(query_tiles): + """Original BF16 arithmetic; short requests only, GPU-gated at replay.""" + from . import attention_folded as folded + native, reduce = folded._kernel(query_tiles) + source = inspect.getsource(native.fn) + anchor = ' cur_batch_seq_len = tl.load(B_Seqlen + cur_batch)\n' + assert source.count(anchor) == 1 + source = source.replace(anchor, anchor + + f' if cur_batch_seq_len >= {_threshold}:\n return\n') + namespace = dict(native.fn.__globals__) + _compile(source, namespace, f'.bf16_{query_tiles}.generated') + stage1 = namespace[native.fn.__name__] + namespace = dict(reduce.__globals__) + stage2 = namespace['_fwd_kernel_stage2'] + source = inspect.getsource(stage2.fn) + assert source.count(anchor) == 1 + # Padded outputs belong exclusively to the native guarded reducer. + source = source.replace(anchor, anchor + + f' if cur_batch_seq_len >= {_threshold} or cur_batch_seq_len <= 0:\n return\n') + _compile(source, namespace, '.bf16_reduce.generated') + reduce = types.FunctionType(reduce.__code__, namespace, reduce.__name__, + reduce.__defaults__, reduce.__closure__) + return stage1, reduce + + +class Bank: + def __init__(self, cache): + self.cache = cache + n = cache.shape[0] + groups = n * 35 + self.packed = [torch.empty(shape, dtype=dtype, device=cache.device) + for shape, dtype in [ + ((n, 35, 2, 32, 32), torch.int32), + ((n, 35, 2, 32), torch.float16), + ((n, 35, 2, 4, 256), torch.int32), + ((n, 35, 2, 256), torch.int32) # FP16 nonDC scale + FP16 normalized V_DC, + ]] + self.valid = torch.zeros(groups, dtype=torch.int32, device=cache.device) + self.busy = torch.zeros_like(self.valid) + + +class State: + def __init__(self, runner): + self.lib = ctypes.CDLL(_library) + P, I = ctypes.c_void_p, ctypes.c_int + for name, args in [ + ('update_cache', [P] * 9 + [I, P]), + ('prepare_decode', [P] * 9 + [I] * 5 + [P]), + ('cached_attention', [P] * 12 + [I] * 4 + [P]), + ('tail_reduce', [P] * 9 + [I] * 7 + [P]), + ('cache_reset', [P] * 3 + [I, P]), + ('cache_copy', [P] * 7 + [I, P]), + ]: + fn = getattr(self.lib, name) + fn.argtypes, fn.restype = args, I + self.banks = {} + targets = 0 + for layer in runner.get_model().modules(): + impl = getattr(layer, 'impl', None) + if (getattr(impl, 'num_heads', None), getattr(impl, 'num_kv_heads', None), + getattr(impl, 'head_size', None)) != (16, 2, 256): + continue + cache = layer.kv_cache + if (not isinstance(cache, torch.Tensor) or cache.ndim != 4 + or cache.shape[1:] != (2, 1120, 512) + or cache.dtype != torch.bfloat16 + or cache.stride() != (1196032, 573440, 512, 1)): + raise RuntimeError('Persistent IU4 requires bound layer-major target page1120 BF16') + targets += 1 + if cache.data_ptr() not in self.banks: + self.banks[cache.data_ptr()] = Bank(cache) + if targets != 10 or len(self.banks) != 5: + raise RuntimeError(f'Expected ten target layers aliasing five banks, got {targets}/{len(self.banks)}') + self.device = next(iter(self.banks.values())).cache.device + device = self.device + self.groups = torch.empty(2048, dtype=torch.int32, device=device) + self.contexts = torch.empty(8, dtype=torch.int32, device=device) + self.counts = torch.empty(8, dtype=torch.int32, device=device) + self.owners = torch.empty(64, dtype=torch.int32, device=device) + self.stats = torch.zeros(3, dtype=torch.int64, device=device) + self.qp = torch.empty((16, 64, 32), dtype=torch.int32, device=device) + self.qs = torch.empty((16, 64), dtype=torch.float16, device=device) + # Both arms use disjoint rows of the same preallocated partial storage. + self.parts = torch.empty((64, 16, 33, 257), dtype=torch.float32, device=device) + self.row_seq = torch.empty(64, dtype=torch.int32, device=device) + self.lse = torch.empty((64, 16), dtype=torch.float32, device=device) + self.reset_events = self.copy_events = 0 + + def call(self, name, tensors, *ints): + rc = getattr(self.lib, name)(*[t.data_ptr() for t in tensors], *ints, + torch.cuda.current_stream(self.device).cuda_stream) + if rc: + raise RuntimeError(f'Persistent IU4 {name} launch failed: {rc}') + + def update(self, cache, slots): + bank = self.banks.get(cache.data_ptr()) + if bank is None: + return + if slots.numel() > 2048 or slots.dtype != torch.int64: + raise RuntimeError('Unexpected persistent IU4 slot capacity/type') + self.call('update_cache', [cache, slots, self.groups, bank.busy, + bank.valid, *bank.packed], slots.numel()) + + def events(self, scheduler_output): + from vllm.utils.torch_utils import async_tensor_h2d + zeros = scheduler_output.new_block_ids_to_zero + copies = scheduler_output.kv_cache_block_copies + if zeros: + ids = async_tensor_h2d(np.asarray(zeros, dtype=np.int64), device=self.device) + for bank in self.banks.values(): + self.call('cache_reset', [bank.valid, bank.busy, ids], len(zeros)) + self.reset_events += len(zeros) + if copies: + pairs_cpu = np.asarray(copies, dtype=np.int64).reshape(-1, 2) + # Usual CoW destinations are fresh. For dependency chains, retain + # upstream snapshot-copy semantics using its existing helper. + if set(pairs_cpu[:, 0]) & set(pairs_cpu[:, 1]): + from vllm.v1.worker.utils import copy_kv_cache_blocks_inplace + for bank in self.banks.values(): + copy_kv_cache_blocks_inplace( + [*bank.packed, bank.valid.view(-1, 35), bank.busy.view(-1, 35)], + bank.cache.shape[0], copies) + else: + pairs = async_tensor_h2d(pairs_cpu, device=self.device) + for bank in self.banks.values(): + self.call('cache_copy', [*bank.packed, bank.valid, bank.busy, pairs], len(copies)) + self.copy_events += len(copies) + + def forward(self, query, key_cache, value_cache, output, block_table, + starts, seq, sm_scale, k_scale, v_scale, max_query_len): + from vllm.triton_utils import triton + from .attention_folded import _reduction_lengths + rows, requests = query.shape[0], seq.numel() + bank = self.banks[key_cache.data_ptr()] + tiles = 1 if max_query_len <= 8 else 2 + self.call('prepare_decode', [query, starts, seq, self.contexts, + self.counts, self.owners, self.stats, self.qp, self.qs], + requests, rows, query.stride(0), query.stride(1), _threshold) + _reduction_lengths[(rows,)](starts, seq, self.row_seq, + REQUESTS=requests, REQUEST_BLOCK=triton.next_power_of_2(requests), num_warps=4) + logits = self.parts[:rows, :, :32, :] + stage1, reduce = _bf16_kernels(tiles) + stage1[(requests * tiles, 2, 32)]( + query, key_cache, value_cache, sm_scale, block_table, seq, starts, logits, + block_table.stride(0), query.stride(0), query.stride(1), + key_cache.stride(0), key_cache.stride(3), key_cache.stride(1), + value_cache.stride(0), value_cache.stride(3), value_cache.stride(1), + logits.stride(0), logits.stride(1), logits.stride(2), k_scale, v_scale, + kv_group_num=8, q_head_num=16, BLOCK_DMODEL=256, BLOCK_DPE=0, + BLOCK_DV=256, BLOCK_N=16, BLOCK_H=64, NUM_KV_SPLITS=32, + PAGE_SIZE=1120, logit_cap=0., Lk=256, Lv=256, IS_MLA=False, + stride_buf_kds=key_cache.stride(2), stride_buf_kxs=key_cache.stride(4), + stride_buf_vds=value_cache.stride(2), num_warps=4, num_stages=1, + waves_per_eu=1, matrix_instr_nonkdim=16, kpack=2) + reduce(logits, query, output, self.lse[:rows], value_cache.transpose(2, 3), self.row_seq[:rows], 32) + self.call('cached_attention', [self.qp, self.qs, *bank.packed, self.parts, + block_table, self.contexts, self.counts, starts, seq], requests, + block_table.stride(0), tiles, _threshold) + self.call('tail_reduce', [query, bank.cache, block_table, self.contexts, + starts, self.owners, seq, self.parts, output], rows, + block_table.stride(0), query.stride(0), query.stride(1), + output.stride(0), output.stride(1), _threshold) + return output + + +def bind(runner): + global _state + if _state is not None: + return 0 + before = torch.cuda.memory_allocated(runner.device) + _state = State(runner) + added = torch.cuda.memory_allocated(runner.device) - before + print('ORNITH_PERSISTENT_IU4_BOUND', len(_state.banks), + next(iter(_state.banks.values())).cache.shape[0], added, _threshold, flush=True) + return added + + +def snapshot(): + if _state is None: + return None + return {'bank_count': len(_state.banks), 'gpu_long_short_requests_long_queries': _state.stats.cpu().tolist(), + 'reset_page_events': _state.reset_events, 'copy_page_events': _state.copy_events, + 'threshold': _threshold} + + +def install(): + global _installed, _library, _threshold + if _installed: + return + _library = os.environ['ORNITH_PERSISTENT_IU4_LIBRARY'] + _threshold = int(os.environ.get('ORNITH_PERSISTENT_IU4_MIN_SEQ', '32768')) + if not 32768 <= _threshold <= 63000: + raise ValueError('Persistent IU4 experiment requires a long-only threshold32768..63000') + from vllm.v1.attention.backends.rocm_attn import RocmAttentionImpl + from vllm.v1.worker.gpu_model_runner import GPUModelRunner + from . import attention_folded + original_update = RocmAttentionImpl.do_kv_cache_update + original_states = GPUModelRunner._update_states + original_forward = attention_folded.forward + + def update(self, layer, key, value, kv_cache, slot_mapping): + result = original_update(self, layer, key, value, kv_cache, slot_mapping) + if _state is not None: + _state.update(kv_cache, slot_mapping) + return result + + def states(self, scheduler_output): + if _state is not None: + _state.events(scheduler_output) + return original_states(self, scheduler_output) + + def forward(query, key_cache, value_cache, output, block_table, query_start_loc, + seq_lens, sm_scale, k_scale, v_scale, *, max_query_len=8): + if (_state is None or key_cache.data_ptr() not in _state.banks + or query.shape[0] > 64 or seq_lens.numel() > 8 + or sm_scale != .0625 or query.dtype != torch.bfloat16 + or output.dtype != torch.bfloat16 or query.stride(2) != 1 + or output.stride(2) != 1 or not 1 <= max_query_len <= 16): + return original_forward(query, key_cache, value_cache, output, + block_table, query_start_loc, seq_lens, sm_scale, k_scale, v_scale, + max_query_len=max_query_len) + return _state.forward(query, key_cache, value_cache, output, block_table, + query_start_loc, seq_lens, sm_scale, k_scale, v_scale, max_query_len) + + RocmAttentionImpl.do_kv_cache_update = update + GPUModelRunner._update_states = states + attention_folded.forward = forward + _installed = True diff --git a/bundle/plugin-site/ornith_g256/attention_mixed.py b/bundle/plugin-site/ornith_g256/attention_mixed.py new file mode 100644 index 0000000000000000000000000000000000000000..5cc6018155434be318c9ee6a101be58fcd928c3f --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_mixed.py @@ -0,0 +1,75 @@ +"""Isolated mixed-target dispatch using existing kernels and CPU phase bounds. + +The backend writes current K/V once before this helper. This helper only creates +views and rebases query starts on the device; it neither changes cache ownership +nor reads GPU metadata on the host. Not installed in any measured bundle. +""" + + +_diagnostic_signatures = set() + +def try_forward(dispatch, query, key, value, output, kv_cache_dtype, + key_cache, value_cache, block_table, query_start_loc, seq_lens, + max_seq_len, max_query_len, k_scale, v_scale, + alibi_slopes, sliding_window, sm_scale, output_scale, sinks, + is_block_table_ptr, causal): + from . import native + requests = getattr(native, '_ATTENTION_PHASE_REQUESTS', ()) + # Subcalls have fewer requests, preventing recursion without global state. + # Dummy capture, draft attention, and unmatched metadata use the old route. + n = len(requests) + if (n < 2 or n != seq_lens.numel() or query_start_loc.numel() != n + 1 + or block_table.shape[0] < n or output.shape != query.shape + or requests[0][0] != 0 or requests[-1][1] > query.shape[0] + or max(end-start for start, end, _, _ in requests) != max_query_len + or not all(end > start for start, end, _, _ in requests) + or not any(r[2] for r in requests) or all(r[2] for r in requests)): + return False + groups = [] + for req, (_, _, prefill, _) in enumerate(requests): + if not prefill and groups and not groups[-1][2]: + first, _, phase = groups[-1] + groups[-1] = (first, req + 1, phase) + else: + # Keep each prefill C1, enabling its existing Q1120 IU4 path. + groups.append((req, req + 1, prefill)) + # Folded decode keeps its current <=128-row contract. Do not partially + # dispatch before deciding whether the entire mixed call is supported. + if any(not phase and requests[last-1][1]-requests[first][0] > 128 + for first, last, phase in groups): + return False + logical_end = requests[-1][1] + if logical_end < output.shape[0]: + output[logical_end:].zero_() + diagnostic = [] + for first, last, prefill in groups: + start, end = requests[first][0], requests[last-1][1] + # GPU subtraction rebases cu_q to the sliced query, without .item(), + # .cpu(), or tensor reconstruction from GPU values. + starts = query_start_loc[first:last+1] - start + max_q = max(r[1]-r[0] for r in requests[first:last]) + max_k = max(r[3] for r in requests[first:last]) + dispatch(query[start:end], None if key is None else key[start:end], + None if value is None else value[start:end], output[start:end], + kv_cache_dtype, key_cache, value_cache, + block_table[first:last], starts, seq_lens[first:last], + max_k, max_q, k_scale, v_scale, alibi_slopes, sliding_window, + sm_scale, output_scale, sinks, is_block_table_ptr, causal) + if len(_diagnostic_signatures) < 4: + from . import attention_compact, attention_iu4_persistent + iu4_prefill = (prefill and last-first == 1 and max_q == 1120 + and max_k >= 4096 and max_k % 32 == 0 + and attention_compact._arena['iu4'] is not None) + diagnostic.append({'rows': (start, end), 'prefill': prefill, + 'max_q': max_q, 'max_k': max_k, + 'iu4_prefill_guard': iu4_prefill, + 'persistent_decode_installed': attention_iu4_persistent._installed}) + # Successful subdispatch receipts only; CPU metadata, no GPU read or sync. + # At most four distinct query/path signatures, so repeated layers stay quiet. + if diagnostic: + signature = tuple((d['max_q'], d['prefill'], d['iu4_prefill_guard']) + for d in diagnostic) + if signature not in _diagnostic_signatures: + _diagnostic_signatures.add(signature) + print('ORNITH_MIXED_ATTENTION_SUCCESS ' + str(diagnostic), flush=True) + return True diff --git a/bundle/plugin-site/ornith_g256/attention_partition.py b/bundle/plugin-site/ornith_g256/attention_partition.py new file mode 100644 index 0000000000000000000000000000000000000000..fbf25954d5b8bc44328d80ef6f81082fc56a2368 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_partition.py @@ -0,0 +1,162 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright 2026 Ciru. +# The generated kernels retain the installed vLLM/SGLang Apache-2.0 code: +# Copyright contributors to the vLLM project; Copyright 2025 vLLM Team; +# Copyright 2023-2024 SGLang Team. See vLLM triton_decode_attention.py. +"""Native-page partitioned attention for target H16/KV2/D256/page1120 or2240. + +The backend must write current K/V before calling ``forward``. Each query row +becomes a causal decode request; existing grouped split/reduction arithmetic is +retained. This module has no installation hook and never converts the KV pool. +""" + +from functools import lru_cache +import inspect +import linecache + +import torch +from vllm.triton_utils import tl, triton + + +@triton.jit +def _query_metadata( + Starts, Seq, Table, RowSeq, RowTable, + TABLE_STRIDE: tl.constexpr, COLS: tl.constexpr, + REQUESTS: tl.constexpr, REQUEST_BLOCK: tl.constexpr, COL_BLOCK: tl.constexpr, + PAGE_SIZE: tl.constexpr, +): + row = tl.program_id(0) + req = tl.arange(0, REQUEST_BLOCK) + start = tl.load(Starts + req, req < REQUESTS, other=0) + end = tl.load(Starts + req + 1, req < REQUESTS, other=0) + owns = (req < REQUESTS) & (start <= row) & (row < end) + owner = tl.max(tl.where(owns, req + 1, 0), 0) - 1 + active = owner >= 0 + first = tl.sum(tl.where(owns, start, 0), 0) + count = tl.sum(tl.where(owns, end - start, 0), 0) + seq = tl.load(Seq + owner, active, other=0) + causal_len = tl.where(active, seq - count + row - first + 1, 0) + tl.store(RowSeq + row, causal_len) + cols = tl.arange(0, COL_BLOCK) + page = tl.load(Table + owner * TABLE_STRIDE + cols, + active & (cols < COLS) & (cols * PAGE_SIZE < causal_len), other=0) + tl.store(RowTable + row * COLS + cols, page, cols < COLS) + + +def _compile(source, namespace, suffix): + filename = __file__ + suffix + linecache.cache[filename] = (len(source), None, source.splitlines(True), filename) + exec(compile(source, filename, "exec"), namespace) + + +@lru_cache(maxsize=1) +def _kernels(): + from vllm.v1.attention.ops import triton_decode_attention as module + + namespace = dict(vars(module)) + source = inspect.getsource(module._fwd_grouped_kernel_stage1.fn) + replacements = { + " IS_MLA: tl.constexpr = False,": + " stride_buf_kds: tl.constexpr,\n stride_buf_kxs: tl.constexpr,\n" + " stride_buf_vds: tl.constexpr,\n IS_MLA: tl.constexpr = False,", + "base_offs_k = cur_kv_head * stride_buf_kh + offs_d[:, None]": + "base_offs_k = cur_kv_head * stride_buf_kh + (offs_d[:, None] // 8) * stride_buf_kds + (offs_d[:, None] % 8) * stride_buf_kxs", + "base_offs_v = cur_kv_head * stride_buf_vh + offs_dv[None, :]": + "base_offs_v = cur_kv_head * stride_buf_vh + offs_dv[None, :] * stride_buf_vds", + } + for old, new in replacements.items(): + if source.count(old) != 1: + raise RuntimeError("Unsupported installed partitioned decoder addressing") + source = source.replace(old, new) + _compile(source, namespace, ".stage1.generated") + + # A graph-padded row owns no request. Stage1 then reads no KV or scratch; + # bypass the reducer's undefined 0/0 result without changing active math. + source = inspect.getsource(module._fwd_kernel_stage2.fn) + old = " cur_batch_seq_len = tl.load(B_Seqlen + cur_batch)\n" + if source.count(old) != 1: + raise RuntimeError("Unsupported installed partitioned decoder reduction") + source = source.replace(old, old + + " if cur_batch_seq_len <= 0:\n" + " pad_d = tl.arange(0, BLOCK_DV)\n" + " tl.store(o + cur_batch * stride_obs + cur_head * stride_oh + pad_d, 0., pad_d < Lv)\n" + " tl.store(lse + cur_batch * stride_lse_bs + cur_head, -float('inf'))\n" + " return\n") + _compile(source, namespace, ".stage2.generated") + _compile(inspect.getsource(module._decode_softmax_reducev_fwd), namespace, + ".reduce.generated") + return namespace["_fwd_grouped_kernel_stage1"], namespace["_decode_softmax_reducev_fwd"] + + +def _metadata(block_table, query_start_loc, seq_lens, rows, page_size=1120): + """Allocate and fill per-query metadata without a device-to-host read.""" + row_table = torch.empty((rows, block_table.shape[1]), dtype=torch.int32, + device=block_table.device) + row_seq = torch.empty((rows,), dtype=torch.int32, device=block_table.device) + if rows: + _query_metadata[(rows,)]( + query_start_loc, seq_lens, block_table, row_seq, row_table, + TABLE_STRIDE=block_table.stride(0), COLS=block_table.shape[1], + REQUESTS=seq_lens.numel(), REQUEST_BLOCK=triton.next_power_of_2(seq_lens.numel()), + COL_BLOCK=triton.next_power_of_2(block_table.shape[1]), PAGE_SIZE=page_size, + num_warps=4) + return row_table, row_seq + + +def forward(query, key_cache, value_cache, output, block_table, query_start_loc, + seq_lens, sm_scale, k_scale, v_scale, *, max_query_len=8): + """Write and return ``output``; allocate scratch only for this call. + + Caller eligibility: BF16 causal target attention, no sinks/alibi/window/output + scaling, at most 8 requests and max_query_len <= 8, at most 64 query rows + including padding. ``query_start_loc`` has + len(seq_lens)+1 entries and includes zero-query slots; trailing query padding + beyond its final entry is allowed and gets zero output. Metadata and cached + page IDs must be valid, as in the ROCm backend. No host synchronization occurs. + + Scratch bytes for R query rows and B table columns: R*(526404 + 4*B), + including float32 split outputs/LSE and int32 causal lengths/block tables. + """ + page_size = key_cache.shape[3] if key_cache.ndim == 5 else 0 + if (query.ndim != 3 or query.shape[1:] != (16, 256) + or page_size not in (1120, 2240) or key_cache.shape[1:] != (2, 32, page_size, 8) + or value_cache.ndim != 4 or value_cache.shape[1:] != (2, 256, page_size) + or key_cache.shape[0] != value_cache.shape[0] + or key_cache.shape[0] == 0 or output.shape != query.shape): + raise ValueError("Partitioned attention requires target H16/KV2/D256/page1120 or2240") + if (not 1 <= max_query_len <= 8 or query.shape[0] > 64 + or not 1 <= seq_lens.numel() <= 8 + or block_table.ndim != 2 or block_table.shape[0] < seq_lens.numel() + or block_table.shape[1] < 1 or query_start_loc.numel() != seq_lens.numel() + 1 + or not seq_lens.is_contiguous() or not query_start_loc.is_contiguous() + or block_table.stride(1) != 1 or query.stride(2) != 1 or output.stride(2) != 1): + raise ValueError("Unsupported query metadata or noncontiguous inner dimensions") + if (not query.is_cuda or any(t.device != query.device for t in + (key_cache, value_cache, output, block_table, query_start_loc, seq_lens, k_scale, v_scale)) + or any(t.dtype != torch.bfloat16 for t in (query, key_cache, value_cache, output)) + or any(t.dtype not in (torch.int32, torch.int64) for t in + (block_table, query_start_loc, seq_lens))): + raise ValueError("Partitioned attention requires BF16 and integer metadata on one GPU") + rows = query.shape[0] + if not rows: + return output + stage1, reduce = _kernels() + table, seq = _metadata(block_table, query_start_loc, seq_lens, rows, page_size) + splits = 32 + logits = torch.empty((rows, 16, splits, 257), dtype=torch.float32, device=query.device) + lse = torch.empty((rows, 16), dtype=torch.float32, device=query.device) + stage1[(rows, 2, splits)]( + query, key_cache, value_cache, sm_scale, table, seq, logits, + table.stride(0), query.stride(0), query.stride(1), + key_cache.stride(0), key_cache.stride(3), key_cache.stride(1), + value_cache.stride(0), value_cache.stride(3), value_cache.stride(1), + logits.stride(0), logits.stride(1), logits.stride(2), k_scale, v_scale, + kv_group_num=8, q_head_num=16, BLOCK_DMODEL=256, BLOCK_DPE=0, + BLOCK_DV=256, BLOCK_N=16, BLOCK_H=16, NUM_KV_SPLITS=splits, + PAGE_SIZE=page_size, logit_cap=0., Lk=256, Lv=256, IS_MLA=False, + stride_buf_kds=key_cache.stride(2), stride_buf_kxs=key_cache.stride(4), + stride_buf_vds=value_cache.stride(2), num_warps=4, num_stages=1, + waves_per_eu=1, matrix_instr_nonkdim=16, kpack=2) + # The reduction helper only reads v_buffer.shape[-1], never cache data. + reduce(logits, query, output, lse, value_cache.transpose(2, 3), seq, splits) + return output diff --git a/bundle/plugin-site/ornith_g256/attention_storage.py b/bundle/plugin-site/ornith_g256/attention_storage.py new file mode 100644 index 0000000000000000000000000000000000000000..6c0e33f8c4fb5f73927286768357734c98c029c5 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_storage.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright 2026 Ciru. +"""Store target KV token-major inside the existing vLLM cache pages. + +Allocation, block IDs, prefix ownership and cache size remain backend-owned. +The writer receives true BSHD views; stride-aware attention receives the +equivalent historical 5D/4D views. No second cache or history copy is created. +""" +import types + +import torch +from vllm.logger import init_logger + +logger = init_logger('vllm.ornith_g256.attention_storage') +_installed = False + + +def _flash_views(cache): + # Bound backend cache is [blocks, K/V, tokens, heads * dim]. + page_size = cache.shape[2] if cache.ndim == 4 else 0 + if (page_size not in (1120, 2240) or cache.shape[1:] != (2, page_size, 512) + or cache.dtype != torch.bfloat16 + or cache.stride(2) != 512 or cache.stride(3) != 1): + raise ValueError('Token-major target KV requires BF16 [B,2,page,512], page1120/2240') + blocks = cache.shape[0] + return (cache[:, 0].view(blocks, page_size, 2, 256), + cache[:, 1].view(blocks, page_size, 2, 256)) + + +def _split_for_attention(cache, num_heads, head_size): + # RocmAttentionImpl passes the K/V-first transpose to this seam. + if (num_heads, head_size) != (2, 256): + raise ValueError('Token-major cache views are target-specific') + key, value = _flash_views(cache.transpose(0, 1)) + blocks, page_size = key.shape[:2] + key_view = key.as_strided((blocks, 2, 32, page_size, 8), + (key.stride(0), 256, 8, 512, 1)) + value_view = value.as_strided((blocks, 2, 256, page_size), + (value.stride(0), 256, 1, 512)) + return key_view, value_view + + +def install(): + """Install only in an isolated worker before loading/capturing the model.""" + global _installed + if _installed: + return + from vllm._aiter_ops import rocm_aiter_ops + from vllm.v1.attention.backend import AttentionType + from vllm.v1.attention.backends import rocm_attn + from vllm.v1.attention.ops.triton_reshape_and_cache_flash import ( + triton_reshape_and_cache_flash, + ) + + if rocm_aiter_ops.is_enabled(): + raise ValueError('Token-major KV currently uses the unfused ROCm cache writer') + impl = rocm_attn.RocmAttentionImpl + original_forward = impl.forward + original_update = impl.do_kv_cache_update + namespace = dict(original_forward.__globals__) + namespace['PagedAttention'] = types.SimpleNamespace( + split_kv_cache=_split_for_attention) + forward_with_views = types.FunctionType( + original_forward.__code__, namespace, 'forward_token_major', + original_forward.__defaults__, original_forward.__closure__) + forward_with_views.__kwdefaults__ = original_forward.__kwdefaults__ + + def target(self): + return (self.attn_type == AttentionType.DECODER + and (self.num_heads, self.num_kv_heads, self.head_size) == (16, 2, 256) + and self.kv_cache_dtype == 'auto' + and self.alibi_slopes is None and self.sinks is None + and self.sliding_window == (-1, -1) + and self.logits_soft_cap == 0) + + def forward(self, layer, query, key, value, kv_cache, attn_metadata, + output, output_scale=None, output_block_scale=None): + selected = forward_with_views if target(self) else original_forward + return selected(self, layer, query, key, value, kv_cache, attn_metadata, + output, output_scale, output_block_scale) + + def update(self, layer, key, value, kv_cache, slot_mapping): + if not target(self): + return original_update(self, layer, key, value, kv_cache, slot_mapping) + key_cache, value_cache = _flash_views(kv_cache) + triton_reshape_and_cache_flash( + key, value, key_cache, value_cache, slot_mapping, + self.kv_cache_dtype, layer._k_scale, layer._v_scale) + logger.info_once('Ornith target KV uses token-major pages in the existing cache allocation') + + impl.forward = forward + impl.do_kv_cache_update = update + _installed = True diff --git a/bundle/plugin-site/ornith_g256/attention_tile.py b/bundle/plugin-site/ornith_g256/attention_tile.py new file mode 100644 index 0000000000000000000000000000000000000000..8b570a05af9009edfd90f82098c306223f434822 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_tile.py @@ -0,0 +1,160 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright 2026 Ciru. Wrappers around the installed vLLM attention implementation. +"""Small query tiles plus bounded sliding-window context loads for the drafter.""" +import hashlib +import inspect +import types +import torch + +from vllm.v1.attention.ops import prefix_prefill +from vllm.v1.attention.ops.chunked_prefill_paged_decode import chunked_prefill_paged_decode +from .attention_window import build_query_kernel, build_window_kernel +from .attention_partition import forward as partitioned_attention + +_compact_prefill = False +_folded_decode = False +_folded_decode_max_queries = 8 + +original_source = inspect.getsource(prefix_prefill.context_attention_fwd) +replacement = ' BLOCK_M = 16 if 1 <= max_input_len <= 8 else 32' +assert original_source.count(' BLOCK_M = 32') == 1 +modified_source = original_source.replace(' BLOCK_M = 32', replacement) +assert modified_source.count('_fwd_kernel[') == 1 +modified_source = modified_source.replace( + '_fwd_kernel[', + '(_window_kernel if sliding_window > 0 else _query_kernel)[') +helper_namespace = dict(vars(prefix_prefill)) +helper_namespace['_window_kernel'] = build_window_kernel(prefix_prefill) +helper_namespace['_query_kernel'] = build_query_kernel(prefix_prefill) +exec(compile(modified_source, __file__ + ':context_attention_fwd', 'exec'), helper_namespace) +context_attention_fwd_tile16 = helper_namespace['context_attention_fwd'] +caller_namespace = dict(chunked_prefill_paged_decode.__globals__) +caller_namespace['context_attention_fwd'] = context_attention_fwd_tile16 +_tiled_attention = types.FunctionType( + chunked_prefill_paged_decode.__code__, caller_namespace, + 'chunked_prefill_paged_decode_tile16', chunked_prefill_paged_decode.__defaults__, + chunked_prefill_paged_decode.__closure__) +_tiled_attention.__kwdefaults__ = chunked_prefill_paged_decode.__kwdefaults__ + + +def chunked_prefill_paged_decode_tile16( + query, key, value, output, kv_cache_dtype, key_cache, value_cache, + block_table, query_start_loc, seq_lens, max_seq_len, max_query_len, + k_scale, v_scale, alibi_slopes=None, sliding_window=None, sm_scale=None, + output_scale=None, sinks=None, is_block_table_ptr=False, causal=True, +): + page_size = key_cache.shape[3] if key_cache.ndim == 5 else 0 + target_cache = (page_size in (1120, 2240) + and key_cache.shape[1:] == (2, 32, page_size, 8) + and value_cache.shape[1:] == (2, 256, page_size)) + # Current K/V has already been written by the ROCm backend. Split only + # mixed target calls, preserving the existing per-shape attention choices. + if (_compact_prefill and _folded_decode and target_cache + and query.shape[1:] == (16, 256) + and query.dtype == key_cache.dtype == value_cache.dtype == output.dtype == torch.bfloat16 + and kv_cache_dtype == 'auto' and causal and not is_block_table_ptr + and alibi_slopes is None and sinks is None and output_scale is None + and (sliding_window is None or sliding_window <= 0)): + from .attention_mixed import try_forward as mixed_forward + if mixed_forward(chunked_prefill_paged_decode_tile16, + query, key, value, output, kv_cache_dtype, key_cache, value_cache, + block_table, query_start_loc, seq_lens, max_seq_len, max_query_len, + k_scale, v_scale, alibi_slopes, sliding_window, sm_scale, + output_scale, sinks, is_block_table_ptr, causal): + return output + # Q1/C1 uses the existing partitioned decoder; shape-only eligibility + # is identical during graph capture and replay as sequence length grows. + if (_folded_decode and max_query_len == 1 + and query.shape == (1,16,256) and seq_lens.numel() == 1 + and target_cache + and query.dtype == key_cache.dtype == value_cache.dtype == output.dtype == torch.bfloat16 + and kv_cache_dtype == 'auto' and causal and not is_block_table_ptr + and alibi_slopes is None and sinks is None and output_scale is None + and (sliding_window is None or sliding_window <= 0)): + return partitioned_attention(query,key_cache,value_cache,output, + block_table,query_start_loc,seq_lens, + sm_scale if sm_scale is not None else 256**-.5, + k_scale,v_scale,max_query_len=1) + # A 16-token verification block is two efficient eight-query/GQA tiles. + # Route it before generic prefill; the latter wastes the intended decode + # reuse and should not determine this block size's performance potential. + if (_folded_decode and 8 < max_query_len <= _folded_decode_max_queries + and query.shape[0] <= 128 and seq_lens.numel() <= 8 + and query.shape[1:] == (16, 256) + and target_cache + and query.dtype == key_cache.dtype == value_cache.dtype == output.dtype == torch.bfloat16 + and kv_cache_dtype == 'auto' and causal and not is_block_table_ptr + and alibi_slopes is None and sinks is None and output_scale is None + and (sliding_window is None or sliding_window <= 0)): + from .attention_folded import forward + return forward(query, key_cache, value_cache, output, block_table, + query_start_loc, seq_lens, + sm_scale if sm_scale is not None else 256**-.5, + k_scale, v_scale, max_query_len=max_query_len) + if (_compact_prefill and max_query_len > 8 + and query.shape[1:] == (16, 256) + and target_cache + and query.dtype == key_cache.dtype == value_cache.dtype == output.dtype == torch.bfloat16 + and kv_cache_dtype == 'auto' and causal and not is_block_table_ptr + and alibi_slopes is None and sinks is None and output_scale is None + and (sliding_window is None or sliding_window <= 0)): + from .attention_compact import try_forward + if try_forward(query, key_cache, value_cache, output, block_table, + query_start_loc, seq_lens, max_query_len=max_query_len, + max_seq_len=max_seq_len, + sm_scale=sm_scale if sm_scale is not None else 256**-.5): + return output + raise RuntimeError('Enabled compact prefill rejected target metadata: ' + str({ + name: (tuple(t.shape), tuple(t.stride()), str(t.dtype), str(t.device)) + for name, t in [('query', query), ('output', output), ('table', block_table), + ('starts', query_start_loc), ('seq_lens', seq_lens)] + }) + f'; max_query_len={max_query_len}, max_seq_len={max_seq_len}') + # Parallel KV partitions hide the long-latency accesses of dispersed hybrid + # cache pages. Larger query batches reuse K/V better in the tiled prefill + # kernel; Q106 was slower with virtual decode, so keep the measured cutoff. + if (1 < max_query_len <= 8 and max_seq_len >= 8192 + and query.shape[0] <= 64 and seq_lens.numel() <= 8 + and query.shape[1:] == (16, 256) + and target_cache + and query.dtype == key_cache.dtype == value_cache.dtype == output.dtype == torch.bfloat16 + and kv_cache_dtype == 'auto' and causal and not is_block_table_ptr + and alibi_slopes is None and sinks is None and output_scale is None + and (sliding_window is None or sliding_window <= 0)): + if _folded_decode: + from .attention_folded import forward as decode + else: + decode = partitioned_attention + return decode( + query, key_cache, value_cache, output, block_table, query_start_loc, + seq_lens, sm_scale if sm_scale is not None else 256**-.5, + k_scale, v_scale, max_query_len=max_query_len) + return _tiled_attention( + query, key, value, output, kv_cache_dtype, key_cache, value_cache, + block_table, query_start_loc, seq_lens, max_seq_len, max_query_len, + k_scale, v_scale, alibi_slopes, sliding_window, sm_scale, output_scale, + sinks, is_block_table_ptr, causal) + + +provenance = dict(source_file=prefix_prefill.__file__, + original_helper_sha256=hashlib.sha256(original_source.encode()).hexdigest(), + modified_helper_sha256=hashlib.sha256(modified_source.encode()).hexdigest(), + exact_change=replacement.strip(), + window_change='Skip and mask context K/V outside the earliest query-row window', + query_change='Return before context work for query tiles with no output rows', + decode_change='Target page1120/2240, context>=8192, maxQ2..8: native-layout32-way KV partitioning; optional maxQ16 folding', + unchanged='Full-attention arithmetic, BLOCK_N32, cache tile32, four warps, one stage') + + +def install(*, compact_prefill=False, folded_decode=False, folded_decode_max_queries=8): + """Select this helper for target and draft ROCm attention in this worker.""" + global _compact_prefill, _folded_decode, _folded_decode_max_queries + if folded_decode_max_queries not in (8, 16): + raise ValueError('Folded verification supports eight or sixteen query positions') + _compact_prefill = compact_prefill + _folded_decode = folded_decode + _folded_decode_max_queries = folded_decode_max_queries + from vllm.v1.attention.backends import rocm_attn + current = rocm_attn.chunked_prefill_paged_decode + if current not in (chunked_prefill_paged_decode, chunked_prefill_paged_decode_tile16): + raise RuntimeError('Another extension replaced the ROCm prefill helper') + rocm_attn.chunked_prefill_paged_decode = chunked_prefill_paged_decode_tile16 diff --git a/bundle/plugin-site/ornith_g256/attention_verify.py b/bundle/plugin-site/ornith_g256/attention_verify.py new file mode 100644 index 0000000000000000000000000000000000000000..883a8f7ad1a22c0b5f45dc18390aacd983fcf5d5 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_verify.py @@ -0,0 +1,64 @@ +"""Optional page1104 multi-query verification, with prebound GPU metadata.""" +# Copyright 2026 Ciru. +import ctypes +from types import SimpleNamespace + +import torch +from .attention_fast import PagedLayout, PagedStrides + + +class NativeVerifyFP32: + page_size = 1104 + + def __init__(self, library, *, Ccap, Lcap, device): + if not 1 <= Ccap <= 64 or not 1 <= Lcap <= 8192: + raise ValueError('Verify attention supports <=64 queries, context <=8192') + self.Ccap, self.Lcap, self.device = Ccap, Lcap, torch.device(device) + self.library = str(library) + self.lib = ctypes.CDLL(self.library) + self.lib.ornith_verify_paged_abi_version.restype = ctypes.c_uint32 + if self.lib.ornith_verify_paged_abi_version() != 1: + raise ValueError('Expected verification ABI1') + self.lib.ornith_verify_paged_get_layout.argtypes = [ctypes.c_int, ctypes.c_int, ctypes.POINTER(PagedLayout)] + self.lib.ornith_verify_paged_get_layout.restype = ctypes.c_int + self.lib.ornith_verify_paged_launch.argtypes = ([ctypes.c_void_p] * 6 + [ctypes.c_size_t] + + [ctypes.c_void_p] * 3 + [ctypes.c_int] * 4 + [PagedStrides, ctypes.c_void_p]) + self.lib.ornith_verify_paged_launch.restype = ctypes.c_int + self.lib.ornith_verify_metadata_launch.argtypes = ([ctypes.c_void_p] * 6 + + [ctypes.c_int] * 4 + [ctypes.c_int64, ctypes.c_void_p]) + self.lib.ornith_verify_metadata_launch.restype = ctypes.c_int + layout = PagedLayout() + if self.lib.ornith_verify_paged_get_layout(Ccap, Lcap, ctypes.byref(layout)): + raise ValueError('Verify workspace query failed') + self.workspace_bytes = layout.bytes + self.output_offset = (layout.bytes + 255) // 256 * 256 + self.table_offset = self.output_offset + Ccap * 16 * 256 * 4 + self.table_cols = (Lcap + self.page_size - 1) // self.page_size + self.lengths_offset = self.table_offset + Ccap * self.table_cols * 4 + self.arena_bytes = self.lengths_offset + Ccap * 4 + + def bind_arena(self, arena): + if (arena.dtype != torch.uint8 or arena.ndim != 1 or not arena.is_contiguous() + or arena.numel() < self.arena_bytes or arena.data_ptr() % 256 + or arena.device != self.device): + raise ValueError('Invalid verification attention arena') + return SimpleNamespace( + workspace=arena[:self.workspace_bytes], + output=arena[self.output_offset:self.table_offset].view(torch.float32), + table=arena[self.table_offset:self.lengths_offset].view(torch.int32), + lengths=arena[self.lengths_offset:self.arena_bytes].view(torch.int32)) + + def launch_out(self, query, key, value, table, lengths, starts, actual, buffers, flag, output): + stream = torch.cuda.current_stream(query.device).cuda_stream + status = self.lib.ornith_verify_metadata_launch( + *[t.data_ptr() for t in (table, lengths, starts, buffers.table, buffers.lengths, flag)], + query.shape[0], lengths.shape[0], actual, self.table_cols, table.stride(0), stream) + if status: + raise RuntimeError(f'Verify metadata launch failed: {status}') + strides = PagedStrides(*query.stride()[:2], *key.stride(), *value.stride(), self.table_cols) + status = self.lib.ornith_verify_paged_launch( + *[t.data_ptr() for t in (query, key, value, buffers.table, buffers.lengths, buffers.workspace)], + buffers.workspace.numel(), buffers.output.data_ptr(), output.data_ptr(), flag.data_ptr(), + query.shape[0], self.Ccap, self.Lcap, key.shape[0], strides, stream) + if status: + raise RuntimeError(f'Verify attention launch failed: {status}') diff --git a/bundle/plugin-site/ornith_g256/attention_window.py b/bundle/plugin-site/ornith_g256/attention_window.py new file mode 100644 index 0000000000000000000000000000000000000000..75b5a7480fcf0e584d03cd077687e33e8d35412d --- /dev/null +++ b/bundle/plugin-site/ornith_g256/attention_window.py @@ -0,0 +1,74 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright 2026 Ciru. Derived at runtime from the installed vLLM kernel. +"""Skip empty query tiles and context outside the drafter's sliding window. + +Masking probabilities alone is insufficient when an evicted cache entry points +to the null page: zero times a nonfinite V still contaminates the accumulator. +The remaining per-row attention mask preserves the installed window semantics. +""" +import inspect +import linecache + + +def _skip_empty_queries(source): + marker = ' block_start_loc = BLOCK_M * start_m\n' + if source.count(marker) != 1: + raise RuntimeError('Unsupported installed prefix kernel: query tile structure') + return source.replace(marker, marker + + ' if block_start_loc >= cur_batch_query_len:\n' + ' return\n') + + +def _compile_kernel(prefix_prefill, source, variant): + # Triton obtains source through inspect; keep generated sources available + # under distinct project-owned filenames without modifying installed vLLM. + filename = __file__ + '.' + variant + '.generated' + linecache.cache[filename] = (len(source), None, source.splitlines(True), filename) + namespace = dict(vars(prefix_prefill)) + exec(compile(source, filename, 'exec'), namespace) + return namespace['_fwd_kernel'] + + +def build_query_kernel(prefix_prefill): + """Preserve full attention arithmetic, omitting tiles with no output rows.""" + source = _skip_empty_queries(inspect.getsource(prefix_prefill._fwd_kernel.fn)) + return _compile_kernel(prefix_prefill, source, 'query') + + +def build_window_kernel(prefix_prefill): + source = _skip_empty_queries(inspect.getsource(prefix_prefill._fwd_kernel.fn)) + old_loop = ''' # compute query against context (no causal mask here) + for start_n in tl.range( + 0, cur_batch_ctx_len, BLOCK_SIZE, loop_unroll_factor=num_unroll_cache + ):''' + new_loop = ''' # No row in this query tile can attend earlier context. Align the loop + # to the cache tile and mask its boundary loads as well as probabilities. + first_context_token = 0 + if SLIDING_WINDOW > 0: + first_context_token = tl.maximum( + 0, cur_batch_ctx_len + block_start_loc - SLIDING_WINDOW + 1 + ) + first_context_tile = (first_context_token // BLOCK_SIZE) * BLOCK_SIZE + for start_n in tl.range( + first_context_tile, cur_batch_ctx_len, BLOCK_SIZE, + loop_unroll_factor=num_unroll_cache + ):''' + replacements = [(old_loop, new_loop)] + # These clauses occur once each for K and V in the context loop only. + old_condition = ''' start_n + BLOCK_SIZE > cur_batch_ctx_len + or BLOCK_DMODEL != BLOCK_DMODEL_PADDED''' + if source.count(old_condition) != 2: + raise RuntimeError('Unsupported installed prefix kernel: context load conditions') + source = source.replace(old_condition, ''' start_n < first_context_token + or start_n + BLOCK_SIZE > cur_batch_ctx_len + or BLOCK_DMODEL != BLOCK_DMODEL_PADDED''') + for indices in ('offs_bs_n[None, :]', 'offs_bs_n[:, None]'): + old = f'& ((start_n + {indices}) < cur_batch_ctx_len),' + new = (f'& ((start_n + {indices}) < cur_batch_ctx_len)\n' + f' & ((start_n + {indices}) >= first_context_token),') + replacements.append((old, new)) + for old, new in replacements: + if source.count(old) != 1: + raise RuntimeError('Unsupported installed prefix kernel: context window structure') + source = source.replace(old, new) + return _compile_kernel(prefix_prefill, source, 'window') diff --git a/bundle/plugin-site/ornith_g256/cache_full1120.py b/bundle/plugin-site/ornith_g256/cache_full1120.py new file mode 100644 index 0000000000000000000000000000000000000000..0d12865da10744d56be5448dc7ba24f2ed9774c6 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/cache_full1120.py @@ -0,0 +1,72 @@ +"""Isolated full-state1120 cache geometry; original tensors and arithmetic.""" +import logging +from dataclasses import replace + +PAGE_BYTES = 2_392_064 +TARGET_BLOCK = 1120 +DRAFT_BLOCK = 560 +_upstream_align = None +logger = logging.getLogger(__name__) + + +def enabled(config): + return (getattr(config.model_config, 'quantization', None) == 'ornith_g256' + and config.cache_config.enable_prefix_caching + and config.additional_config.get('ornith_g256', {}).get('dynamic_spec_profile', False)) + + +def _align(cls, vllm_config, backend_cls): + _upstream_align(cls, vllm_config, backend_cls) + if not enabled(vllm_config): + return + cache = vllm_config.cache_config + spec = vllm_config.speculative_config + assert cache.mamba_cache_mode == 'align' and cache.prefix_match_unit is None + assert spec.method == 'dflash' and spec.num_speculative_tokens == 15 + cache.block_size = TARGET_BLOCK + cache.mamba_block_size = TARGET_BLOCK + cache.mamba_page_size_padded = PAGE_BYTES + logger.info('Ornith full1120 platform: target/Mamba=%s draft=%s physical_page=%s', + TARGET_BLOCK, DRAFT_BLOCK, PAGE_BYTES) + + +def padded_specs(config, specs): + if not enabled(config): + return specs + from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec, SlidingWindowSpec + from .prefix_cache import EXPECTED_DRAFT_NAMES + drafts = {name for name in specs if name.startswith('model.layers.')} + targets = set(specs) - drafts + assert drafts == EXPECTED_DRAFT_NAMES + assert len(targets) == 40 and all(name.startswith('language_model.model.layers.') for name in targets) + assert sum(isinstance(specs[name], MambaSpec) for name in targets) == 30 + assert sum(isinstance(specs[name], FullAttentionSpec) for name in targets) == 10 + result = {} + for name, spec in specs.items(): + if name in drafts: + assert isinstance(spec, SlidingWindowSpec) and spec.sliding_window == 4096 + block = DRAFT_BLOCK + else: + block = TARGET_BLOCK + if isinstance(spec, MambaSpec): + assert spec.real_page_size_bytes == PAGE_BYTES + updated = replace(spec, block_size=block, page_size_padded=PAGE_BYTES) + assert updated.page_size_bytes == PAGE_BYTES + result[name] = updated + logger.info('Ornith full1120 specs: target_attention=10x1120 Mamba=30x1120 ' + 'draft=6x560 all_page_bytes=%s recurrent_dtype=%s', PAGE_BYTES, + next(spec.dtypes for spec in result.values() if isinstance(spec, MambaSpec))) + return result + + +def install(): + from vllm.platforms.interface import Platform + global _upstream_align + current = Platform._align_hybrid_block_size.__func__ + if current is _align: + return + if _upstream_align is None: + _upstream_align = current + elif current is not _upstream_align: + raise RuntimeError('Another extension replaced hybrid cache alignment') + Platform._align_hybrid_block_size = classmethod(_align) diff --git a/bundle/plugin-site/ornith_g256/column_backend.py b/bundle/plugin-site/ornith_g256/column_backend.py new file mode 100644 index 0000000000000000000000000000000000000000..c9c2ba4f5cf7ca2cec8a3c929d82b2b963c34e86 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/column_backend.py @@ -0,0 +1,93 @@ +"""Parent32tile semantics with four output-column workgroups per KV head.""" +# Copyright2026 Ciru. +from collections import Counter +import torch +from vllm.config import get_current_vllm_config +from vllm.v1.attention.backend import AttentionType +from vllm.v1.attention.ops.paged_attn import PagedAttention +from vllm.v1.attention.backends.rocm_attn import RocmAttentionBackend,RocmAttentionImpl +KERNEL_SHA256='27fb053d970184732a46212026afd3152ec07a0dfc3dc4c5cc5cc0d539f17c81' +from .column_kernel import ornith_column_paged_attention + + +class OrnithColumnAttentionBackend(RocmAttentionBackend): + @staticmethod + def get_name(): return 'CUSTOM' + @staticmethod + def get_impl_cls(): return OrnithColumnAttentionImpl + + +class OrnithColumnAttentionImpl(RocmAttentionImpl): + implementation='ciru.ornith.column.backend.v1' + def __init__(self,*args,**kwargs): + super().__init__(*args,**kwargs) + self.dispatch_counts=Counter() + self.capture_by_C=Counter() + self.eager_by_C=Counter() + self.context_bounds=[None,None] + + def static_native_support(self): + return (self.attn_type==AttentionType.DECODER + and (self.num_heads,self.num_kv_heads,self.head_size)==(16,2,256) + and self.scale==.0625 and self.kv_cache_dtype in ('auto','bfloat16') + and self.alibi_slopes is None and self.sliding_window==(-1,-1) + and self.logits_soft_cap==0 and self.sinks is None + and self.kv_sharing_target_layer_name is None) + + def forward(self,layer,query,key,value,kv_cache,attn_metadata,output, + output_scale=None,output_block_scale=None): + m=attn_metadata;reason=None + if m is None: reason='profile' + elif not self.static_native_support(): reason='static_feature' + elif (m.use_cascade or m.causal is not True or output_scale is not None + or output_block_scale is not None): reason='metadata_feature' + elif m.max_query_len!=1: reason='prefill_or_mixed' + else: + C=m.seq_lens.shape[0] + if (C<1 or query.shape!=(C,16,256) or output.shape!=query.shape + or not 01:self.dispatch_counts['eager_cached_native_calls']+=1 + lo,hi=self.context_bounds + self.context_bounds=[m.max_seq_len if lo is None else min(lo,m.max_seq_len), + m.max_seq_len if hi is None else max(hi,m.max_seq_len)] + return output + self.dispatch_counts['fallback_'+reason]+=1 + return super().forward(layer,query,key,value,kv_cache,m,output,output_scale,output_block_scale) + + def inspect_dispatch(self,prefix=None): + return {'implementation':self.implementation,'kernel_sha256':KERNEL_SHA256,'prefix':prefix, + 'counts':dict(self.dispatch_counts),'capture_by_C':dict(self.capture_by_C), + 'eager_by_C':dict(self.eager_by_C),'native_max_seq_len_host_bounds':list(self.context_bounds), + 'counter_semantics':'Host dispatch counts; capture keyed by actual C/prefix, graph replay counted separately', + 'capacity_policy':'inherits parent context/batch capacity; tested coverage recorded externally'} diff --git a/bundle/plugin-site/ornith_g256/column_kernel.py b/bundle/plugin-site/ornith_g256/column_kernel.py new file mode 100644 index 0000000000000000000000000000000000000000..ba98c327fcdc90c9198c4f4bcf4bc3de7f7cc754 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/column_kernel.py @@ -0,0 +1,263 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +# Authors: +# - Burkhard Ringlein +# - Jan van Lunteren +# - Chih-Chieh Yang +# - Thomas Parnell + +# Ciru2026 modification: four contiguous64-column output partitions. +# Full D256 QK/all8 query rows (padded16) and32tile arithmetic are retained. +import torch +from vllm.platforms import current_platform +from vllm.triton_utils import tl, triton +float8_info = torch.finfo(current_platform.fp8_dtype()) + +@triton.jit +def cdiv_fn(x, y): + return (x + y - 1) // y + + +@triton.jit +def ornith_column_paged_attention( + output_ptr, # [num_tokens, num_query_heads, head_size] + query_ptr, # [num_tokens, num_query_heads, head_size] + key_cache_ptr, # [num_blks, num_kv_heads, head_size // x, blk_size, x] + value_cache_ptr, # [num_blks, num_kv_heads, head_size, blk_size] + sink_ptr, # [num_query_heads] + block_tables_ptr, # [num_seqs, max_num_blocks_per_seq] + seq_lens_ptr, # [num_seqs] + alibi_slopes_ptr, # [num_query_heads] + scale, # float32 + k_scale, # float32 + v_scale, # float32 + out_scale_inv, + num_query_heads: tl.constexpr, # int + num_queries_per_kv: tl.constexpr, # int + num_queries_per_kv_padded: tl.constexpr, # int + block_table_stride: tl.int64, # int + query_stride_0: tl.int64, # int + query_stride_1: tl.int64, # int, should be equal to head_size + output_stride_0: tl.int64, # int + output_stride_1: tl.int64, # int, should be equal to head_size + BLOCK_SIZE: tl.constexpr, # int + PHYSICAL_BLOCK_SIZE: tl.constexpr, # int + HEAD_SIZE: tl.constexpr, # int + HEAD_SIZE_PADDED: tl.constexpr, # int, must be power of 2 + USE_ALIBI_SLOPES: tl.constexpr, # bool + SLIDING_WINDOW: tl.constexpr, # int + x: tl.constexpr, # int + stride_k_cache_0: tl.int64, # int + stride_k_cache_1: tl.int64, # int + stride_k_cache_2: tl.int64, # int + stride_k_cache_3: tl.int64, # int + stride_k_cache_4: tl.int64, # int + stride_v_cache_0: tl.int64, # int + stride_v_cache_1: tl.int64, # int + stride_v_cache_2: tl.int64, # int + stride_v_cache_3: tl.int64, # int + filter_by_query_len: tl.constexpr, # bool + query_start_len_ptr, # [num_seqs+1] + USE_SINKS: tl.constexpr, # bool + USE_FP8: tl.constexpr, + FP8_MIN: tl.constexpr = float8_info.min, + FP8_MAX: tl.constexpr = float8_info.max, +): + seq_idx = tl.program_id(0) + kv_head_idx = tl.program_id(1) + output_partition = tl.program_id(2) + offs_v = output_partition * 64 + tl.arange(0, 64) + v_dim_mask = offs_v < HEAD_SIZE + + if filter_by_query_len: + cur_batch_in_all_start_index = tl.load(query_start_len_ptr + seq_idx) + cur_batch_in_all_stop_index = tl.load(query_start_len_ptr + seq_idx + 1) + cur_batch_query_len = cur_batch_in_all_stop_index - cur_batch_in_all_start_index + if cur_batch_query_len > 1: + return + else: + cur_batch_in_all_start_index = seq_idx + + query_head_idx = kv_head_idx * num_queries_per_kv + tl.arange( + 0, num_queries_per_kv_padded + ) + + query_offset = ( + cur_batch_in_all_start_index * query_stride_0 + + query_head_idx[:, None] * query_stride_1 + ) + + head_mask = query_head_idx < (kv_head_idx + 1) * num_queries_per_kv + head_mask = head_mask & (query_head_idx < num_query_heads) + + dim_mask = tl.where(tl.arange(0, HEAD_SIZE_PADDED) < HEAD_SIZE, 1, 0).to(tl.int1) + + # Q : (num_queries_per_kv, HEAD_SIZE,) + Q = tl.load( + query_ptr + query_offset + tl.arange(0, HEAD_SIZE_PADDED)[None, :], + mask=dim_mask[None, :] & head_mask[:, None], + other=0.0, + ) + + block_table_offset = seq_idx * block_table_stride + + if not USE_SINKS: + M = tl.full([num_queries_per_kv_padded], float("-inf"), dtype=tl.float32) + L = tl.zeros([num_queries_per_kv_padded], dtype=tl.float32) + else: + M = tl.load( + sink_ptr + query_head_idx, + mask=head_mask, + other=float("-inf"), + ).to(dtype=tl.float32) + L = tl.where(float("-inf") < M, 1.0, 0.0) + + acc = tl.zeros([num_queries_per_kv_padded, 64], dtype=tl.float32) + + # sequence len for this particular sequence + seq_len = tl.load(seq_lens_ptr + seq_idx) + + # alibi slope for this head + if USE_ALIBI_SLOPES: + alibi_slope = tl.load( + alibi_slopes_ptr + query_head_idx, mask=head_mask, other=0.0 + ) + + num_blocks = cdiv_fn(seq_len, BLOCK_SIZE) + + offs_n = tl.arange(0, BLOCK_SIZE) + offs_d = tl.arange(0, HEAD_SIZE_PADDED) + # iterate through tiles + for j in range(0, num_blocks): + start_n = j * BLOCK_SIZE + # Calculate the logical location within a non-standard physical block, + # such as 544 in Qwen/Qwen3-Next-80B-A3B-Thinking. + # Supports non-contiguous mapping + # from logical blocks to physical blocks + abs_token_idx = start_n + offs_n + l_block_idx = abs_token_idx // PHYSICAL_BLOCK_SIZE + # Vectorized loading of physical block IDs + p_block_idx = tl.load(block_tables_ptr + block_table_offset + l_block_idx) + internal_offsets = abs_token_idx % PHYSICAL_BLOCK_SIZE + + # 5D addressing logic of K + k_offset = ( + p_block_idx[None, :] * stride_k_cache_0 + + kv_head_idx * stride_k_cache_1 + + (offs_d[:, None] // x) * stride_k_cache_2 + + internal_offsets[None, :] * stride_k_cache_3 + + (offs_d[:, None] % x) * stride_k_cache_4 + ) + + # 4D addressing logic of V (Slot is innermost) + v_offset = ( + p_block_idx[:, None] * stride_v_cache_0 + + kv_head_idx * stride_v_cache_1 + + offs_v[None, :] * stride_v_cache_2 + + internal_offsets[:, None] * stride_v_cache_3 + ) + + # Only the final tile can straddle seq_len. Slots >= seq_len are + # unwritten KV cache that may hold NaN/garbage; they are score-masked + # below, but 0 * NaN = NaN would still poison the output, so mask them + # out of the K/V loads too. Earlier tiles are fully written, so they + # use the cheaper token-uniform dim_mask (matching the pre-0.25.0 fast + # path) and skip the per-token predicate entirely. + # K : (HEAD_SIZE, BLOCK_SIZE), V : (BLOCK_SIZE, HEAD_SIZE) + if j == num_blocks - 1: + kv_load_mask = abs_token_idx < seq_len + K_load = tl.load( + key_cache_ptr + k_offset, + mask=dim_mask[:, None] & kv_load_mask[None, :], + other=0.0, + eviction_policy="evict_last", + ) + V_load = tl.load( + value_cache_ptr + v_offset, + mask=v_dim_mask[None, :] & kv_load_mask[:, None], + other=0.0, + eviction_policy="evict_last", + ) + else: + K_load = tl.load( + key_cache_ptr + k_offset, + mask=dim_mask[:, None], + other=0.0, + eviction_policy="evict_last", + ) + V_load = tl.load( + value_cache_ptr + v_offset, + mask=v_dim_mask[None, :], + other=0.0, + eviction_policy="evict_last", + ) + + if K_load.dtype.is_fp8(): + K = (K_load.to(tl.float32) * tl.load(k_scale)).to(Q.dtype) + else: + K = K_load + + if V_load.dtype.is_fp8(): + V = (V_load.to(tl.float32) * tl.load(v_scale)).to(Q.dtype) + else: + V = V_load + + seq_offset = j * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + boundary = tl.full([BLOCK_SIZE], seq_len, dtype=tl.int32) + seq_mask = seq_offset[None, :] < boundary + + # First calculate the dot, then apply the mask. + qk = scale * tl.dot(Q, K) + S = tl.where(head_mask[:, None] & seq_mask, qk, float("-inf")) + + context_len = seq_len - 1 + + if SLIDING_WINDOW > 0: + S = tl.where((context_len - seq_offset) < SLIDING_WINDOW, S, -10000) + + if USE_ALIBI_SLOPES: + S += alibi_slope[:, None] * (seq_offset - context_len) + + # compute running maximum + # m_j : (num_queries_per_kv,) + m_j = tl.maximum(M, tl.max(S, axis=1)) + + # P : (num_queries_per_kv, BLOCK_SIZE,) + p = tl.exp(S - m_j[:, None]) + p = tl.where(m_j[:, None] == float("-inf"), 0.0, p) + + # l_j : (num_queries_per_kv,) + l_j = tl.sum(p, axis=1) + + # alpha : (num_queries_per_kv, ) + alpha = tl.exp(M - m_j) + alpha = tl.where(float("-inf") == M, 0.0, alpha) + + # acc : (num_queries_per_kv, BLOCK_SIZE,) + acc = acc * alpha[:, None] + + # update constants + L = L * alpha + l_j + M = m_j + + # acc : (num_queries_per_kv, BLOCK_SIZE,) + acc += tl.dot(p.to(V.dtype), V) + + # epilogue + acc = acc / (L[:, None] + 1e-10) + if USE_FP8: + acc = acc * tl.load(out_scale_inv) + acc = tl.clamp(acc, FP8_MIN, FP8_MAX) + + output_offset = ( + cur_batch_in_all_start_index * output_stride_0 + + query_head_idx * output_stride_1 + ) + + tl.store( + output_ptr + output_offset[:, None] + offs_v[None, :], + acc, + mask=v_dim_mask[None, :] & head_mask[:, None], + ) + diff --git a/bundle/plugin-site/ornith_g256/config.py b/bundle/plugin-site/ornith_g256/config.py new file mode 100644 index 0000000000000000000000000000000000000000..9f9614e29f04bad2c3ed3f30eee8d1302488c91e --- /dev/null +++ b/bundle/plugin-site/ornith_g256/config.py @@ -0,0 +1,58 @@ +"""Explicit large-projection policy; routing and small GDN controls stay BF16.""" +# Copyright 2026 Ciru. +import re +import torch +from vllm.model_executor.layers.quantization.base_config import QuantizationConfig + +QUANT_CONFIG = dict(quant_method="ornith_g256", group_size=256, + transform_block=128, activation_bits=8) +DENSE_SUFFIXES = (".linear_attn.in_proj_qkvz", ".linear_attn.out_proj", + ".self_attn.qkv_proj", ".self_attn.o_proj", + ".mlp.shared_expert.gate_up_proj", ".mlp.shared_expert.down_proj") + + +class OrnithG256Config(QuantizationConfig): + def __init__(self, config): + super().__init__() + if any(config.get(k) != v for k, v in QUANT_CONFIG.items()): + raise ValueError(f"Expected {QUANT_CONFIG}") + self.checkpoint_config = dict(config) + self.activation_bits, self.transform_block = 8, 128 + + @classmethod + def get_name(cls): return "ornith_g256" + def get_supported_act_dtypes(self): return [torch.bfloat16] + @classmethod + def get_min_capability(cls): return 0 + @staticmethod + def get_config_filenames(): return ["quantize_config.json"] + @classmethod + def from_config(cls, config): return cls(config) + + def get_quant_method(self, layer, prefix): + from vllm.model_executor.layers.fused_moe.routed_experts import RoutedExperts + from vllm.model_executor.layers.linear import LinearBase, UnquantizedLinearMethod + from vllm.model_executor.layers.vocab_parallel_embedding import ( + ParallelLMHead, VocabParallelEmbedding, UnquantizedEmbeddingMethod, + ) + from .method import G256MoEMethod, G256LinearMethod, W8HeadMethod + # MTP weights retain BF16; its temporary head is replaced with the + # target W8 head by the installed proposer after loading. + if prefix.startswith('mtp.'): + if isinstance(layer, RoutedExperts): + from vllm.model_executor.layers.fused_moe.unquantized_fused_moe_method import UnquantizedFusedMoEMethod + return UnquantizedFusedMoEMethod(layer.moe_config) + if isinstance(layer, LinearBase): + return UnquantizedLinearMethod() + if isinstance(layer, RoutedExperts): + if not re.search(r"layers\.[0-9]+\.mlp\.experts$", prefix): + raise ValueError(f"Unsupported routed module {prefix}") + G256MoEMethod.validate_layer_scope(layer) + return G256MoEMethod(layer.moe_config, self, prefix) + if isinstance(layer, ParallelLMHead): + return W8HeadMethod(prefix) + if isinstance(layer, LinearBase): + return G256LinearMethod(prefix) if prefix.endswith(DENSE_SUFFIXES) else UnquantizedLinearMethod() + if isinstance(layer, VocabParallelEmbedding): + return UnquantizedEmbeddingMethod() + return None diff --git a/bundle/plugin-site/ornith_g256/dense_n32.py b/bundle/plugin-site/ornith_g256/dense_n32.py new file mode 100644 index 0000000000000000000000000000000000000000..20e9c4b11b73a22a9b5faf2065beaade8080240c --- /dev/null +++ b/bundle/plugin-site/ornith_g256/dense_n32.py @@ -0,0 +1,53 @@ +"""Isolated M1 dense N32 shadows; retain N16 verification and prefill.""" +# Copyright 2026 Ciru. +import ctypes as C +import json +from pathlib import Path +import torch +_INSTALLED=False +_BANKS={} +_TOTAL=0 +_LIB=None +_LAUNCH=None + +def install(): + global _INSTALLED,_LIB,_LAUNCH + if _INSTALLED:return + from . import native + from .method import G256LinearMethod + library=Path(__file__).resolve().parents[2]/'native/libornith_dense_g256_n32.so' + _LIB=C.CDLL(str(library.resolve(strict=True))) + _LAUNCH=_LIB.ornith_dense_g256_launch_n32 + _LAUNCH.argtypes=[C.c_void_p]*4+[C.c_size_t]+[C.c_void_p]*2+[C.c_int]*7+[C.c_void_p] + _LAUNCH.restype=C.c_int + original_load=G256LinearMethod.process_weights_after_loading + def load(self,layer): + global _TOTAL + original_load(self,layer) + if self.fields != ('g256_codes','g256_metadata'):return + n,k=self.n,self.k + if n%32 or not 0=1 and geometry==2: + bank=_BANKS.get(codes.data_ptr()) + if bank is None:raise RuntimeError('Missing dense N32 shadow for M1') + native.validate(x,out,capacity,k) + native.check(_LAUNCH(native.ptr(x),native.ptr(bank[0]),native.ptr(bank[1]),native.ptr(workspace),workspace.numel(),native.ptr(out),native.ptr(flags),1,capacity,n,k,8,128,2,native.stream(x)),'G256 dense N32 M1') + return + return original_dense(x,codes,metadata,workspace,out,flags,capacity,n,k,geometry,a8_max_rows) + native._dense_impl=dense + _INSTALLED=True diff --git a/bundle/plugin-site/ornith_g256/dense_source.py b/bundle/plugin-site/ornith_g256/dense_source.py new file mode 100644 index 0000000000000000000000000000000000000000..66535b1c0c4f26d1b6de0a0ce3eb9ab401222e73 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/dense_source.py @@ -0,0 +1,37 @@ +"""Original BF16 selected dense weights for prefill; verification stays packed.""" +from pathlib import Path +import torch +from safetensors import safe_open +_INSTALLED=False +_WEIGHTS={} +def install(): + global _INSTALLED + if _INSTALLED:return + from . import native + from .method import G256LinearMethod + path=Path(__file__).resolve().parents[2]/'models/source-prefill.safetensors' + original_load=G256LinearMethod.process_weights_after_loading + def load(self,layer): + original_load(self,layer) + if not self.prefix.endswith(('.linear_attn.in_proj_qkvz','.mlp.shared_expert.gate_up_proj','.mlp.shared_expert.down_proj')):return + name=self.prefix.replace('language_model.model.','model.language_model.',1) + with safe_open(str(path),framework='pt',device='cpu') as f:w=f.get_tensor(name) + assert w.dtype==torch.bfloat16 and tuple(w.shape)==(self.n,self.k) + w=w.to(device=layer.g256_codes.device) + layer.register_buffer('_source_prefill_weight',w,persistent=False) + _WEIGHTS[layer.g256_codes.data_ptr()]=w + if len(_WEIGHTS) in (1,110):print('ORNITH_SOURCE_PREFILL',len(_WEIGHTS),sum(t.numel()*t.element_size() for t in _WEIGHTS.values()),flush=True) + G256LinearMethod.process_weights_after_loading=load + original_dense=native._dense_impl + def dense(x,codes,metadata,workspace,out,flags,capacity,n,k,geometry,a8_max_rows): + # M64 complete-operation measurements favor these two source shapes. + # The resident weights already exist for prefill; no additional allocation. + if x.shape[0]>a8_max_rows or (x.shape[0]==64 and + (n,k) in ((12288,2048),(1024,2048))): + weight=_WEIGHTS.get(codes.data_ptr()) + if weight is not None: + torch.mm(x,weight.t(),out=out) + return + return original_dense(x,codes,metadata,workspace,out,flags,capacity,n,k,geometry,a8_max_rows) + native._dense_impl=dense + _INSTALLED=True diff --git a/bundle/plugin-site/ornith_g256/dflash_conv_boundary.py b/bundle/plugin-site/ornith_g256/dflash_conv_boundary.py new file mode 100644 index 0000000000000000000000000000000000000000..295ac6b96c296aff1ff57dfadbf43e2dd9704bc8 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/dflash_conv_boundary.py @@ -0,0 +1,123 @@ +"""Keep DFlash convolution request boundaries dynamic inside token-count graphs.""" +# Copyright 2026 Ciru. Original convolution arithmetic from vLLM, Apache-2.0. +from contextvars import ContextVar +import inspect +import linecache +import logging +import torch +import torch.nn.functional as F + +_ACTIVE_MASK = ContextVar('ornith_dflash_query_mask', default=None) +_INSTALLED = False +_ORIGINAL_CONVOLVE = None +_INPUT_KERNEL = None +logger = logging.getLogger(__name__) + + +def grouped_conv_dynamic(hidden_states, delta, base, query_mask, num_groups, group_size, taps): + # Float expressions and evaluation order are unchanged from upstream. + blocks = hidden_states.unflatten(-1, (num_groups, group_size)) + coefficients = base.view(1, taps, num_groups, group_size) + delta.unsqueeze(-1) + output = coefficients[:, 0] * blocks + position = torch.arange(hidden_states.shape[0], device=hidden_states.device) + position = position & query_mask + for tap in range(1, taps): + shifted = F.pad(blocks[:-tap], (0, 0, 0, 0, tap, 0)) + output += coefficients[:, tap] * shifted * (position >= tap).view(-1, 1, 1) + return output.flatten(-2) + + +def convolve(self, hidden_states, delta, side): + mask = getattr(self, '_ornith_query_mask', None) + if mask is None: + return _ORIGINAL_CONVOLVE(self, hidden_states, delta, side) + return grouped_conv_dynamic(hidden_states, delta, self.base_kernel[side], mask, + self.num_groups, self.group_size, self.taps) + + +def build_input_kernel(): + """Add one scalar store to the existing input-preparation launch.""" + global _INPUT_KERNEL + if _INPUT_KERNEL is not None: + return _INPUT_KERNEL + from vllm.v1.spec_decode import utils + original = utils.copy_and_expand_dflash_inputs_kernel + source = inspect.getsource(original.fn) + arg = ' out_token_indices_ptr, # [num_reqs * num_speculative_tokens] (output)\n' + body = ' block_idx = tl.program_id(axis=1)\n' + if source.count(arg) != 1 or source.count(body) != 1: + raise RuntimeError('Unsupported DFlash input-kernel source') + source = source.replace('def copy_and_expand_dflash_inputs_kernel(', + 'def _copy_and_expand_with_query_mask(', 1) + source = source.replace(arg, arg + ' out_query_mask_ptr, # persistent scalar for captured convolutions\n', 1) + source = source.replace(body, body + ' tl.store(out_query_mask_ptr, num_query_per_req - 1,\n' + ' mask=(req_idx == 0) & (block_idx == 0))\n', 1) + filename = __file__ + '.input.generated' + linecache.cache[filename] = (len(source), None, source.splitlines(True), filename) + namespace = dict(vars(utils)) + exec(compile(source, filename, 'exec'), namespace) + _INPUT_KERNEL = namespace['_copy_and_expand_with_query_mask'] + return _INPUT_KERNEL + + +class _InputKernelProxy: + def __init__(self, original): + self.original = original + + def __getitem__(self, grid): + def launch(*args, **kwargs): + mask = _ACTIVE_MASK.get() + if mask is None: + return self.original[grid](*args, **kwargs) + return build_input_kernel()[grid](*args, out_query_mask_ptr=mask, **kwargs) + return launch + + +def install(): + global _INSTALLED, _ORIGINAL_CONVOLVE + if _INSTALLED: + return + from vllm.v1.spec_decode import dflash + from vllm.model_executor.models.qwen3_dflash2 import DFlashGroupedConv + original_load = dflash.DFlashProposer.load_model + original_inputs = dflash.DFlashProposer.set_inputs_first_pass + _ORIGINAL_CONVOLVE = DFlashGroupedConv._convolve + + def load_model(self, target_model): + result = original_load(self, target_model) + if not self.is_dflash2: + return result + width = 1 + self.speculative_config.num_speculative_tokens + if width not in (8, 16): + raise ValueError('Current boundary fix supports original Q8/Q16 only') + self._ornith_query_mask = torch.full((1,), width - 1, dtype=torch.int32, device=self.device) + model = self.model.unwrap() if hasattr(self.model, 'unwrap') else self.model + count = 0 + for module in model.modules(): + if isinstance(module, DFlashGroupedConv): + module.register_buffer('_ornith_query_mask', self._ornith_query_mask, persistent=False) + count += 1 + if count != 12: + raise RuntimeError(f'Expected current six-layer DFlash2 with12 convolutions, got {count}') + logger.info('DFlash boundary correction:12 convs share one4-byte GPU mask; dynamicQ8/Q16') + return result + + def set_inputs_first_pass(self, *args, **kwargs): + mask = getattr(self, '_ornith_query_mask', None) + if mask is None: + return original_inputs(self, *args, **kwargs) + width = 1 + self.num_speculative_tokens + if width not in (8, 16): + raise ValueError(f'Unexpected current DFlash query width{width}') + token = _ACTIVE_MASK.set(mask) + try: + return original_inputs(self, *args, **kwargs) + finally: + _ACTIVE_MASK.reset(token) + + build_input_kernel() + dflash.copy_and_expand_dflash_inputs_kernel = _InputKernelProxy(dflash.copy_and_expand_dflash_inputs_kernel) + dflash.DFlashProposer.load_model = load_model + dflash.DFlashProposer.set_inputs_first_pass = set_inputs_first_pass + DFlashGroupedConv._convolve = convolve + _INSTALLED = True diff --git a/bundle/plugin-site/ornith_g256/dflash_spec.py b/bundle/plugin-site/ornith_g256/dflash_spec.py new file mode 100644 index 0000000000000000000000000000000000000000..2d082dd6fdbcfa30a1f0f322da8dc3a8c6b6445d --- /dev/null +++ b/bundle/plugin-site/ornith_g256/dflash_spec.py @@ -0,0 +1,29 @@ +"""Keep DFlash query slots after the runner supplies target context metadata.""" +# Copyright 2026 Ciru. +from vllm.v1.spec_decode.dflash import DFlashProposer + + +_upstream_set_inputs_first_pass = DFlashProposer.set_inputs_first_pass + + +def set_inputs_first_pass(self, *args, **kwargs): + # The runner seeds per-group metadata with target/context slots. DFlash + # then generates separate context and query slots from the draft block + # table. The base proposer otherwise prefers the stale context slots to + # the new query view, overwriting the query mapping (and overflowing its + # 64-token buffer on a 128-token prefill). + if any(group.kv_cache_group_id != self.kv_cache_gid + for group in self.draft_attn_groups): + raise ValueError('Ornith DFlash requires one draft KV cache group') + result = _upstream_set_inputs_first_pass(self, *args, **kwargs) + _, _, query_metadata = result + self._per_group_slot_mappings[self.kv_cache_gid] = query_metadata.slot_mapping + return result + + +def install(): + """Patch only DFlash's input hook; leave installed vLLM files untouched.""" + current = DFlashProposer.set_inputs_first_pass + if current not in (_upstream_set_inputs_first_pass, set_inputs_first_pass): + raise RuntimeError('Another extension replaced DFlash input preparation') + DFlashProposer.set_inputs_first_pass = set_inputs_first_pass diff --git a/bundle/plugin-site/ornith_g256/dynamic_graphs.py b/bundle/plugin-site/ornith_g256/dynamic_graphs.py new file mode 100644 index 0000000000000000000000000000000000000000..93945eb8a9df0833bcd581bf0d2ccadde0cbb647 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/dynamic_graphs.py @@ -0,0 +1,266 @@ +"""C1 adaptive fallback: exact Q1/C1 target graph; retain Q16/Q8 keys. + +Copy into the experimental ornith_g256 package and call install() before +GPUModelRunner construction. Installed vLLM files and drafter hooks stay intact. +""" +# Copyright 2026 Ciru. +from contextlib import contextmanager +from functools import wraps + + +_INSTALLED = False +_CONFIG_HOOK_INSTALLED = False + + +def _requested(config): + return config.additional_config.get("ornith_g256", {}).get( + "dynamic_spec_profile") is True + + +def _profile_errors(config): + """Report actual values instead of silently opting an explicit profile out.""" + spec = config.speculative_config + parallel = config.parallel_config + cache = config.cache_config + schedule = getattr(spec, "num_speculative_tokens_per_batch_size", None) + values = { + "speculative.method": (getattr(spec, "method", None), "dflash"), + "speculative.num_speculative_tokens": ( + getattr(spec, "num_speculative_tokens", None), 15), + "speculative.schedule": ( + tuple(tuple(row) for row in schedule) if schedule else None, + ((1, 1, 15), (2, 8, 7))), + "max_num_seqs": (config.scheduler_config.max_num_seqs, 8), + "async_scheduling": (config.scheduler_config.async_scheduling, False), + "quantization": (config.model_config.quantization, "ornith_g256"), + "enable_prefix_caching": (cache.enable_prefix_caching, True), + "lora_config": (config.lora_config, None), + "use_v2_model_runner": (config.use_v2_model_runner, False), + "use_ubatching": (parallel.use_ubatching, False), + } + values.update((field, (getattr(parallel, field), 1)) for field in ( + "tensor_parallel_size", "pipeline_parallel_size", "data_parallel_size")) + errors = [f"{name}={actual!r} (expected {expected!r})" + for name, (actual, expected) in values.items() if actual != expected] + # Retention changes CPU checkpoint eligibility, not page geometry or graph buffers. + # This private repair profile permits frontier-only retention with the same DF15/7 graphs. + if cache.prefix_cache_retention_interval not in (0, 1120): + errors.append(f"prefix_cache_retention_interval={cache.prefix_cache_retention_interval!r} (expected 0 or 1120)") + if not 65536 <= config.model_config.max_model_len <= 262144: + errors.append("max_model_len outside 65536..262144") + # EngineCore normalizes the shared config to the smallest participating + # group before worker KV initialization: draft560, target/Mamba1120. + # Physical block geometry remains unchanged for both supported retention policies. + if cache.block_size not in (560, 1120): + errors.append(f"cache.block_size={cache.block_size!r} (expected 560 or 1120)") + return errors + + +def _reject(errors): + from vllm.logger import init_logger + + message = "Dynamic target graph profile rejected: " + "; ".join(errors) + init_logger("vllm.ornith_dynamic_graphs").error(message) + raise ValueError(message) + + +def install_config_hook(): + """Call before AsyncEngineArgs builds its VllmConfig, and in the worker.""" + global _CONFIG_HOOK_INSTALLED + if _CONFIG_HOOK_INSTALLED: + return + from vllm.config import CUDAGraphMode, VllmConfig + + original = VllmConfig._maybe_override_dynamic_sd_cudagraph_mode + + @wraps(original) + def keep_supported_full(config): + if not _requested(config): + return original(config) + # Only suppress the installed automatic FULL->PIECEWISE rewrite for + # our exact two-shape target profile. Other modes retain their rules. + if config.compilation_config.cudagraph_mode != CUDAGraphMode.FULL_DECODE_ONLY: + return original(config) + errors = _profile_errors(config) + sizes = config.compilation_config.cudagraph_capture_sizes + if sizes != [16, 32, 48, 64]: + errors.append(f"cudagraph_capture_sizes={sizes!r} (expected [16, 32, 48, 64])") + if errors: + _reject(errors) + from vllm.logger import init_logger + + init_logger("vllm.ornith_dynamic_graphs").info( + "Dynamic target config retains FULL_DECODE_ONLY for C1/Q16 and Q8/C1..8") + + VllmConfig._maybe_override_dynamic_sd_cudagraph_mode = keep_supported_full + _CONFIG_HOOK_INSTALLED = True + + +def _enabled(runner): + config = runner.vllm_config + if not _requested(config): + return False + cached = getattr(runner, "_ornith_dynamic_graphs_enabled", None) + if cached is not None: + return cached + errors = _profile_errors(config) + if errors: + _reject(errors) + import torch + + props = torch.cuda.get_device_properties(runner.device) + arch = getattr(props, "gcnArchName", "") + if not arch.startswith("gfx1151"): + _reject([f"gcnArchName={arch!r} (expected gfx1151)"]) + runner._ornith_dynamic_graphs_enabled = True + return True + + +def _supported_query_len(num_tokens, num_reqs, max_query_len): + if num_tokens != num_reqs * max_query_len: + return None + if max_query_len == 1 and num_reqs == 1: + return 1 + if max_query_len == 16 and num_reqs == 1: + return 16 + if max_query_len == 8 and 1 <= num_reqs <= 8: + return 8 + return None + + +@contextmanager +def _query_len(runner, query_len): + # This profile is synchronous TP1/PP1/DP1 without microbatching. Nested + # capture -> dummy -> dispatch scopes restore their caller's values. + dispatcher = runner.cudagraph_dispatcher + runner_previous = runner.uniform_decode_query_len + dispatcher_previous = dispatcher.uniform_decode_query_len + runner.uniform_decode_query_len = query_len + dispatcher.uniform_decode_query_len = query_len + try: + yield + finally: + runner.uniform_decode_query_len = runner_previous + dispatcher.uniform_decode_query_len = dispatcher_previous + + +def _replace_full_keys(runner): + from vllm.config import CUDAGraphMode + from vllm.forward_context import BatchDescriptor + + dispatcher = runner.cudagraph_dispatcher + if dispatcher.cudagraph_mode != CUDAGraphMode.FULL_DECODE_ONLY: + _reject([f"resolved cudagraph_mode={dispatcher.cudagraph_mode.name}; " + "expected FULL_DECODE_ONLY (install_config_hook must run before config creation)"]) + # Reuse vLLM's resolved capture-size lookup, including its Q16 rounding. + # Thus [16,32,48,64] yields Q8 captures at C2/C4/C6/C8; odd C pad up. + lookup = dispatcher._bs_to_padded_graph_size + if len(lookup) <= 64 or lookup[16] != 16 or lookup[64] != 64: + raise ValueError("Dynamic target graphs require capture sizes 16 and 64") + # Do not change global capture sizes: those also configure DFlash. + # Add only an exact target Q1/C1 key after upstream's Q16 initialization. + lookup[1] = 1 + lookup[8] = 8 + keys = {BatchDescriptor(num_tokens=1, num_reqs=1, uniform=True), + BatchDescriptor(num_tokens=16, num_reqs=1, uniform=True)} + for num_reqs in range(1, 9): + padded = lookup[num_reqs * 8] + if padded % 8 or not 8 <= padded <= 64: + raise ValueError("Dynamic Q8 graph padding must stay within C1..8") + keys.add(BatchDescriptor( + num_tokens=padded, num_reqs=padded // 8, uniform=True)) + dispatcher.cudagraph_keys[CUDAGraphMode.FULL].clear() + dispatcher.cudagraph_keys[CUDAGraphMode.FULL].update(keys) + from vllm.logger import init_logger + + init_logger("vllm.ornith_dynamic_graphs").info( + "Dynamic target FULL graph shapes (tokens, requests): %s; " + "cache_config.block_size=%s, KV group block sizes=%s", + sorted((key.num_tokens, key.num_reqs) for key in keys), + runner.vllm_config.cache_config.block_size, + [group.kv_cache_spec.block_size for group in runner.kv_cache_config.kv_cache_groups], + ) + + +def install(): + """Install process-local, config- and hardware-scoped target wrappers.""" + global _INSTALLED + if _INSTALLED: + return + # The worker also creates replaced VllmConfigs while loading DFlash; + # dataclasses.replace reruns __post_init__ and would downgrade them again. + install_config_hook() + from vllm.v1.worker.gpu_model_runner import GPUModelRunner + + original_resolve = GPUModelRunner._check_and_update_cudagraph_mode + original_dispatch = GPUModelRunner._determine_batch_execution_and_padding + original_capture = GPUModelRunner._warmup_and_capture + + @wraps(original_resolve) + def resolve(runner, *args, **kwargs): + result = original_resolve(runner, *args, **kwargs) + if _enabled(runner): + _replace_full_keys(runner) + return result + + @wraps(original_dispatch) + def dispatch(runner, num_tokens, num_reqs, num_scheduled_tokens_np, + max_num_scheduled_tokens, *args, **kwargs): + query_len = _supported_query_len( + num_tokens, num_reqs, max_num_scheduled_tokens) + if query_len == 1 and kwargs.get('force_uniform_decode') is None: + # A one-token prompt/tail is not proof of decode. CPU metadata is + # already available; no GPU read or synchronization is added. + batch = runner.input_batch + if not (batch.num_reqs == 1 and + batch.num_computed_tokens_cpu[0] >= batch.num_prompt_tokens[0]): + query_len = None + if query_len is None or not _enabled(runner): + return original_dispatch( + runner, num_tokens, num_reqs, num_scheduled_tokens_np, + max_num_scheduled_tokens, *args, **kwargs) + with _query_len(runner, query_len): + result = original_dispatch( + runner, num_tokens, num_reqs, num_scheduled_tokens_np, + max_num_scheduled_tokens, *args, **kwargs) + if kwargs.get('force_uniform_decode') is None: + seen = getattr(runner, '_ornith_dynamic_dispatch_seen', set()) + key = (query_len, num_reqs, result[0].name) + if key not in seen: + from vllm.logger import init_logger + init_logger("vllm.ornith_dynamic_graphs").info('Dynamic target dispatch Q%d/C%d: %s', *key) + seen.add(key) + runner._ornith_dynamic_dispatch_seen = seen + return result + + @wraps(original_capture) + def capture(runner, desc, *args, **kwargs): + query_len = None + if desc.uniform and desc.num_reqs: + query_len = _supported_query_len( + desc.num_tokens, desc.num_reqs, desc.num_tokens // desc.num_reqs) + if query_len is None or not _enabled(runner): + return original_capture(runner, desc, *args, **kwargs) + if query_len == 1: + # Only target Q1 is captured. The K0 diagnostic never uses the + # drafter, whose trained convolution cannot warm up a one-row input. + # Retain all its ordinary Q8/Q16 warmup and construction unchanged. + drafter = runner.drafter + had_override = 'dummy_run' in drafter.__dict__ + previous_dummy = drafter.__dict__.get('dummy_run') + drafter.dummy_run = lambda *a, **k: None + try: + with _query_len(runner, query_len): + return original_capture(runner, desc, *args, **kwargs) + finally: + if had_override: + drafter.dummy_run = previous_dummy + else: + del drafter.dummy_run + with _query_len(runner, query_len): + return original_capture(runner, desc, *args, **kwargs) + + GPUModelRunner._check_and_update_cudagraph_mode = resolve + GPUModelRunner._determine_batch_execution_and_padding = dispatch + GPUModelRunner._warmup_and_capture = capture + _INSTALLED = True diff --git a/bundle/plugin-site/ornith_g256/gdn_compact.py b/bundle/plugin-site/ornith_g256/gdn_compact.py new file mode 100644 index 0000000000000000000000000000000000000000..ed089dcb13c9fec53600b5057f80ee06ee95a4a5 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/gdn_compact.py @@ -0,0 +1,99 @@ +"""C8 recurrent update logs; reconstruct before upstream cache alignment.""" +# Copyright 2026 Ciru. +from functools import wraps +import inspect +import torch +from . import gdn_compact_kernel as kernel +_ACTIVE=False +_CAPTURE=False +_INSTALLED=False +_BANKS={} +_LAYER=None + +def install(): + global _INSTALLED + if _INSTALLED:return + from vllm.v1.worker.gpu_model_runner import GPUModelRunner + from vllm.v1.worker import mamba_utils + from vllm.model_executor.layers.mamba.gdn import qwen_gdn_linear_attn + cls=qwen_gdn_linear_attn.QwenGatedDeltaNetAttention + core_original=cls._forward_core + @wraps(core_original) + def core(self,*args,**kwargs): + global _LAYER + previous=_LAYER;_LAYER=self.prefix + try:return core_original(self,*args,**kwargs) + finally:_LAYER=previous + cls._forward_core=core + original=qwen_gdn_linear_attn.fused_sigmoid_gating_delta_rule_update + signature=inspect.signature(original) + def update(*args,**kwargs): + if not _ACTIVE:return original(*args,**kwargs) + bound=signature.bind(*args,**kwargs);bound.apply_defaults();d=dict(bound.arguments) + state=d['initial_state'];q=d['q'];v=d['v'];cu=d['cu_seqlens'];ids=d['ssm_state_indices'] + if cu is None or ids is None or ids.ndim!=2 or d['is_kda']: + raise RuntimeError('Compact GDN requires varlen indexed scalar decay') + N=len(cu)-1;HV=v.shape[2];V=v.shape[-1];K=q.shape[-1] + assert N<=8 and q.shape[1]<=128 and ids.shape[1]<=16 and (HV,V,K)==(32,128,128) + key=_LAYER + assert key is not None + if key not in _BANKS: + print('ORNITH_GDN_BANK',key,state.data_ptr(),tuple(state.shape),tuple(state.stride()),flush=True) + make=lambda *shape:torch.empty(*shape,device=q.device,dtype=torch.float32) + _BANKS[key]=dict(base=make(8,HV,V,K),keys=make(128,HV,K),values=make(128,HV,V),decays=make(128,HV), + cu=torch.empty(9,device=q.device,dtype=cu.dtype), + ids=torch.empty(8,16,device=q.device,dtype=ids.dtype),state=state) + bank=_BANKS[key] + bank['cu'][:N+1].copy_(cu) + bank['ids'][:N,:ids.shape[1]].copy_(ids) + return kernel.fused_sigmoid_gating_delta_rule_update(**d,compact_base=bank['base'], + compact_k=bank['keys'],compact_v=bank['values'],compact_g=bank['decays']) + qwen_gdn_linear_attn.fused_sigmoid_gating_delta_rule_update=update + execute_original=GPUModelRunner.execute_model + capture_original=GPUModelRunner._warmup_and_capture + dispatch_original=GPUModelRunner._determine_batch_execution_and_padding + post_original=mamba_utils.postprocess_mamba_align_gpu + @wraps(execute_original) + def execute(self,scheduler_output,*args,**kwargs): + global _ACTIVE + counts=scheduler_output.num_scheduled_tokens;cached=scheduler_output.scheduled_cached_reqs + _ACTIVE=(bool(counts) and not scheduler_output.scheduled_new_reqs and + all(c==8 for c in counts.values()) and + all(r in cached.req_ids and not cached.is_context_phase(r) for r in counts)) + return execute_original(self,scheduler_output,*args,**kwargs) + @wraps(capture_original) + def capture(self,desc,*args,**kwargs): + global _ACTIVE,_CAPTURE + previous=(_ACTIVE,_CAPTURE) + _CAPTURE=True + _ACTIVE=bool(desc.uniform and desc.num_reqs and desc.num_tokens==desc.num_reqs*8) + try:return capture_original(self,desc,*args,**kwargs) + finally:_ACTIVE,_CAPTURE=previous + @wraps(dispatch_original) + def dispatch(self,num_tokens,num_reqs,num_scheduled_tokens_np,max_num_scheduled_tokens,*args,**kwargs): + # A prompt whose shape coincides with Q8 must not replay a compact decode graph. + if not _CAPTURE and not _ACTIVE and max_num_scheduled_tokens==8: + kwargs['force_eager']=True + return dispatch_original(self,num_tokens,num_reqs,num_scheduled_tokens_np,max_num_scheduled_tokens,*args,**kwargs) + @wraps(post_original) + def post(**kwargs): + global _ACTIVE + if _ACTIVE: + assert len(_BANKS)==30, f'Expected30targetGDN layers, got{len(_BANKS)}' + ctx=kwargs['bufs'].postprocess_align + for bank in _BANKS.values(): + state=bank['state'] + kernel.replay_physical[(4,kwargs['num_reqs']*32)]( + bank['base'],bank['keys'],bank['values'],bank['decays'], + kwargs['num_accepted_tokens_gpu'],bank['cu'],bank['ids'],state, + ctx.num_computed_tokens_buf.gpu,ctx.num_scheduled_tokens_buf.gpu,ctx.num_draft_tokens_buf.gpu, + 32,128,128,state.stride(0),1120,32,num_warps=4,num_stages=3) + if not getattr(post,'logged',False): + print('ORNITH_COMPACT_GDN replayed30layers before cache alignment',flush=True);post.logged=True + try:return post_original(**kwargs) + finally:_ACTIVE=False + GPUModelRunner.execute_model=execute + GPUModelRunner._warmup_and_capture=capture + GPUModelRunner._determine_batch_execution_and_padding=dispatch + mamba_utils.postprocess_mamba_align_gpu=post + _INSTALLED=True diff --git a/bundle/plugin-site/ornith_g256/gdn_compact_kernel.py b/bundle/plugin-site/ornith_g256/gdn_compact_kernel.py new file mode 100644 index 0000000000000000000000000000000000000000..3d69f1ad1d902a306065def32ef0ff231dfe2814 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/gdn_compact_kernel.py @@ -0,0 +1,302 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Songlin Yang, Yu Zhang +# +# This file contains code copied from the flash-linear-attention project. +# The original source code was licensed under the MIT license and included +# the following copyright notice: +# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang + +import torch + +from vllm.triton_utils import tl, triton + + +@triton.heuristics( + { + "USE_INITIAL_STATE": lambda args: args["h0"] is not None, + "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, + "IS_CONTINUOUS_BATCHING": lambda args: args["ssm_state_indices"] is not None, + "IS_SPEC_DECODING": lambda args: args["num_accepted_tokens"] is not None, + } +) +@triton.jit(do_not_specialize=["N", "T"]) +def fused_sigmoid_gating_delta_rule_update_kernel( + A_log, + a, + b, + dt_bias, + beta, + threshold, + q, + k, + v, + o, + h0, + ht, + compact_base, compact_k, compact_v, compact_g, + cu_seqlens, + ssm_state_indices, + num_accepted_tokens, + scale, + N: tl.int64, # num of sequences + T: tl.int64, # num of tokens + B: tl.constexpr, + H: tl.constexpr, + HV: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + stride_init_state_token: tl.constexpr, + stride_final_state_token: tl.constexpr, + stride_indices_seq: tl.constexpr, + stride_indices_tok: tl.constexpr, + USE_INITIAL_STATE: tl.constexpr, # whether to use initial state + INPLACE_FINAL_STATE: tl.constexpr, # whether to store final state inplace + USE_QK_L2NORM_IN_KERNEL: tl.constexpr, + IS_VARLEN: tl.constexpr, + IS_CONTINUOUS_BATCHING: tl.constexpr, + IS_SPEC_DECODING: tl.constexpr, + IS_KDA: tl.constexpr, +): + i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_n, i_hv = i_nh // HV, i_nh % HV + i_h = i_hv // (HV // H) + if IS_VARLEN: + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int64), + tl.load(cu_seqlens + i_n + 1).to(tl.int64), + ) + all = T + T = eos - bos + else: + bos, eos = i_n * T, i_n * T + T + all = B * T + + if T == 0: + # no tokens to process for this sequence + return + + o_k = i_k * BK + tl.arange(0, BK) + o_v = i_v * BV + tl.arange(0, BV) + + p_q = q + (bos * H + i_h) * K + o_k + p_k = k + (bos * H + i_h) * K + o_k + p_v = v + (bos * HV + i_hv) * V + o_v + + p_A_log = A_log + i_hv + if not IS_KDA: + p_a = a + bos * HV + i_hv + p_dt_bias = dt_bias + i_hv + else: + p_a = a + (bos * HV + i_hv) * K + o_k + p_dt_bias = dt_bias + i_hv * K + o_k + + p_b = b + bos * HV + i_hv + p_o = o + ((i_k * all + bos) * HV + i_hv) * V + o_v + + mask_k = o_k < K + mask_v = o_v < V + mask_h = mask_v[:, None] & mask_k[None, :] + + b_h = tl.zeros([BV, BK], dtype=tl.float32) + if USE_INITIAL_STATE: + if IS_CONTINUOUS_BATCHING: + if IS_SPEC_DECODING: + i_t = tl.load(num_accepted_tokens + i_n).to(tl.int64) - 1 + else: + i_t = 0 + # Load state index and check for invalid entries + state_idx = tl.load(ssm_state_indices + i_n * stride_indices_seq + i_t).to( + tl.int64 + ) + # Skip if state index is invalid (NULL_BLOCK_ID=0) + if state_idx <= 0: + return + p_h0 = h0 + state_idx * stride_init_state_token + else: + p_h0 = h0 + bos * HV * V * K + p_h0 = p_h0 + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + b_h += tl.load(p_h0, mask=mask_h, other=0).to(tl.float32) + p_base=compact_base + i_n*HV*V*K + i_hv*V*K + o_v[:,None]*K+o_k[None,:] + tl.store(p_base,b_h,mask=mask_h) + + + for i_t in range(0, T): + b_q = tl.load(p_q, mask=mask_k, other=0).to(tl.float32) + b_k = tl.load(p_k, mask=mask_k, other=0).to(tl.float32) + b_v = tl.load(p_v, mask=mask_v, other=0).to(tl.float32) + b_b = tl.load(p_b).to(tl.float32) + + # If the model is loaded in fp16, without the .float() here, A might be -inf + x = tl.load(p_a).to(tl.float32) + tl.load(p_dt_bias).to(tl.float32) + softplus_x = tl.where( + beta * x <= threshold, (1 / beta) * tl.log(1 + tl.exp(beta * x)), x + ) + b_g = -tl.exp(tl.load(p_A_log).to(tl.float32)) * softplus_x + + # compute beta_output = sigmoid(b) + b_beta = tl.sigmoid(b_b.to(tl.float32)) + + if USE_QK_L2NORM_IN_KERNEL: + b_q = b_q * (tl.rsqrt(tl.sum(b_q * b_q) + 1e-6)) + b_k = b_k * (tl.rsqrt(tl.sum(b_k * b_k) + 1e-6)) + b_q = b_q * scale + # [BV, BK] + if not IS_KDA: + b_decay=tl.exp(b_g) + b_h *= b_decay + else: + b_h *= tl.exp(b_g[None, :]) + # [BV] + b_v -= tl.sum(b_h * b_k[None, :], 1) + b_v *= b_beta + log_token=bos+i_t + if i_v == 0: + tl.store(compact_k+(log_token*HV+i_hv)*K+o_k,b_k,mask=mask_k) + tl.store(compact_g+log_token*HV+i_hv,b_decay) + tl.store(compact_v+(log_token*HV+i_hv)*V+o_v,b_v,mask=mask_v) + + # [BV, BK] + b_h += b_v[:, None] * b_k[None, :] + # [BV] + b_o = tl.sum(b_h * b_q[None, :], 1) + tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v) + + # Update pointers for next timestep + p_q += H * K + p_k += H * K + p_o += HV * V + p_v += HV * V + p_b += HV + p_a += HV + + +def fused_sigmoid_gating_delta_rule_update( + A_log: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + dt_bias: torch.Tensor, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + beta: float = 1.0, + threshold: float = 20.0, + scale: float = None, + initial_state: torch.Tensor = None, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, + is_kda: bool = False, + compact_base=None, compact_k=None, compact_v=None, compact_g=None, +): + """ + Fused triton implementation of sigmoid gating delta rule update. + This function uses a single fused kernel that combines both sigmoid gating + computation and the recurrent delta rule update for better performance. + """ + B, T, H, K, V = *k.shape, v.shape[-1] + HV = v.shape[2] + N = B if cu_seqlens is None else len(cu_seqlens) - 1 + BK, BV = triton.next_power_of_2(K), min(triton.next_power_of_2(V), 32) + NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV) + assert NK == 1, "NK > 1 is not supported yet" + num_stages = 3 + num_warps = 4 + + if cu_seqlens is not None and q.shape[0] != 1: + raise ValueError( + f"The batch size is expected to be 1 rather than {q.shape[0]}" + f" when using `cu_seqlens`. Please flatten variable-length" + f" inputs before processing." + ) + if scale is None: + scale = k.shape[-1] ** -0.5 + else: + assert scale > 0, "scale must be positive" + + o = q.new_zeros(NK, *v.shape) + if inplace_final_state: + final_state = initial_state + else: + final_state = q.new_empty(T, HV, V, K, dtype=initial_state.dtype) + + stride_init_state_token = initial_state.stride(0) + stride_final_state_token = final_state.stride(0) + + if ssm_state_indices is None: + stride_indices_seq, stride_indices_tok = 1, 1 + elif ssm_state_indices.ndim == 1: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride(0), 1 + else: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride() + + grid = (NK, NV, N * HV) + fused_sigmoid_gating_delta_rule_update_kernel[grid]( + A_log=A_log, + a=a.contiguous(), + b=b.contiguous(), + dt_bias=dt_bias, + beta=beta, + threshold=threshold, + q=q.contiguous(), + k=k.contiguous(), + v=v.contiguous(), + o=o, + h0=initial_state, + ht=final_state, + compact_base=compact_base, compact_k=compact_k, compact_v=compact_v, compact_g=compact_g, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + scale=scale, + N=N, + T=T, + B=B, + H=H, + HV=HV, + K=K, + V=V, + BK=BK, + BV=BV, + stride_init_state_token=stride_init_state_token, + stride_final_state_token=stride_final_state_token, + stride_indices_seq=stride_indices_seq, + stride_indices_tok=stride_indices_tok, + INPLACE_FINAL_STATE=inplace_final_state, + USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, + IS_KDA=is_kda, + num_warps=num_warps, + num_stages=num_stages, + ) + o = o.squeeze(0) + return o, final_state + + +@triton.jit +def replay_physical(base,keys,values,decays,accepted,cu,indices,states, + computed,scheduled,drafted, + HV:tl.constexpr,K:tl.constexpr,V:tl.constexpr, + STATE_STRIDE:tl.constexpr,BLOCK:tl.constexpr,BV:tl.constexpr): + iv,nh=tl.program_id(0),tl.program_id(1) + n,h=nh//HV,nh%HV + kk=tl.arange(0,K);vv=iv*BV+tl.arange(0,BV) + offset=h*V*K+vv[:,None]*K+kk[None,:] + count=tl.load(accepted+n) + bos=tl.load(cu+n);eos=tl.load(cu+n+1) + count=tl.minimum(count,eos-bos) + if count<=0:return + state=tl.load(base+n*HV*V*K+offset) + running=tl.load(computed+n)+tl.load(scheduled+n)-tl.load(drafted+n) + for t in range(count): + token=bos+t + key=tl.load(keys+(token*HV+h)*K+kk) + value=tl.load(values+(token*HV+h)*V+vv) + decay=tl.load(decays+token*HV+h) + state=tl.fma(value[:,None],key[None,:],state*decay) + if t==count-1 or (running+t)%BLOCK==0: + dest=tl.load(indices+n*16+t).to(tl.int64) + if dest>0:tl.store(states+dest*STATE_STRIDE+offset,state) diff --git a/bundle/plugin-site/ornith_g256/gdn_spec.py b/bundle/plugin-site/ornith_g256/gdn_spec.py new file mode 100644 index 0000000000000000000000000000000000000000..22435b54f0a7699a4c242add210ddc07eb6f1f1a --- /dev/null +++ b/bundle/plugin-site/ornith_g256/gdn_spec.py @@ -0,0 +1,151 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Songlin Yang, Yu Zhang +# +# This file contains code copied from the flash-linear-attention project. +# The original source code was licensed under the MIT license and included +# the following copyright notice: +# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang +# +# Ciru modification (2026): zero-initialize output for skipped NULL state rows. +# Wrapper copied from vLLM third_party/flash_linear_attention/ops/ +# fused_sigmoid_gating.py; the installed Triton kernel remains unchanged. +"""Defined GDN speculative padding outputs, using the original Triton kernel.""" +import torch +from vllm.forward_context import get_forward_context +from vllm.logger import init_logger +from vllm.triton_utils import triton +from vllm.third_party.flash_linear_attention.ops.fused_sigmoid_gating import ( + fused_sigmoid_gating_delta_rule_update as upstream_update, + fused_sigmoid_gating_delta_rule_update_kernel, +) + +logger = init_logger(__name__) +_upstream_rocm_core = None + + +def initialized_rocm_warmup(self, qkvz, ba, z_out, core_attn_out): + """Define buffers skipped by the metadata-free upstream ROCm warmup.""" + metadata = get_forward_context().attn_metadata + no_metadata = not isinstance(metadata, dict) or metadata.get(self.prefix) is None + result = _upstream_rocm_core(self, qkvz, ba, z_out, core_attn_out) + if no_metadata: + z_out.zero_() + core_attn_out.zero_() + logger.info_once('Ornith GDN metadata-free ROCm warmup initialized z/core outputs') + return result + + +def fused_sigmoid_gating_delta_rule_update( + A_log: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + dt_bias: torch.Tensor, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + beta: float = 1.0, + threshold: float = 20.0, + scale: float = None, + initial_state: torch.Tensor = None, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, + is_kda: bool = False, +): + B, T, H, K, V = *k.shape, v.shape[-1] + HV = v.shape[2] + N = B if cu_seqlens is None else len(cu_seqlens) - 1 + BK, BV = triton.next_power_of_2(K), min(triton.next_power_of_2(V), 32) + NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV) + assert NK == 1, "NK > 1 is not supported yet" + num_stages = 3 + num_warps = 4 + + if cu_seqlens is not None and q.shape[0] != 1: + raise ValueError( + f"The batch size is expected to be 1 rather than {q.shape[0]}" + f" when using `cu_seqlens`. Please flatten variable-length" + f" inputs before processing." + ) + if scale is None: + scale = k.shape[-1] ** -0.5 + else: + assert scale > 0, "scale must be positive" + + # NULL_BLOCK_ID causes the original kernel to return without an output + # store. Valid rows are overwritten exactly as in the upstream wrapper. + o = q.new_zeros((NK, *v.shape)) + if inplace_final_state: + final_state = initial_state + else: + final_state = q.new_empty(T, HV, V, K, dtype=initial_state.dtype) + + stride_init_state_token = initial_state.stride(0) + stride_final_state_token = final_state.stride(0) + + if ssm_state_indices is None: + stride_indices_seq, stride_indices_tok = 1, 1 + elif ssm_state_indices.ndim == 1: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride(0), 1 + else: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride() + + grid = (NK, NV, N * HV) + fused_sigmoid_gating_delta_rule_update_kernel[grid]( + A_log=A_log, + a=a.contiguous(), + b=b.contiguous(), + dt_bias=dt_bias, + beta=beta, + threshold=threshold, + q=q.contiguous(), + k=k.contiguous(), + v=v.contiguous(), + o=o, + h0=initial_state, + ht=final_state, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + scale=scale, + N=N, + T=T, + B=B, + H=H, + HV=HV, + K=K, + V=V, + BK=BK, + BV=BV, + stride_init_state_token=stride_init_state_token, + stride_final_state_token=stride_final_state_token, + stride_indices_seq=stride_indices_seq, + stride_indices_tok=stride_indices_tok, + INPLACE_FINAL_STATE=inplace_final_state, + USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, + IS_KDA=is_kda, + num_warps=num_warps, + num_stages=num_stages, + ) + o = o.squeeze(0) + return o, final_state + + +def install(): + """Rebind only Qwen GDN's imported helper before MTP model construction.""" + from vllm.model_executor.layers.mamba.gdn import qwen_gdn_linear_attn + current = qwen_gdn_linear_attn.fused_sigmoid_gating_delta_rule_update + if current not in (upstream_update, fused_sigmoid_gating_delta_rule_update): + raise RuntimeError('Another extension replaced the Qwen GDN update helper') + qwen_gdn_linear_attn.fused_sigmoid_gating_delta_rule_update = fused_sigmoid_gating_delta_rule_update + global _upstream_rocm_core + cls = qwen_gdn_linear_attn.QwenGatedDeltaNetAttention + current_core = cls._forward_core_rocm + if _upstream_rocm_core is None: + _upstream_rocm_core = current_core + elif current_core not in (_upstream_rocm_core, initialized_rocm_warmup): + raise RuntimeError('Another extension replaced the Qwen GDN ROCm core') + cls._forward_core_rocm = initialized_rocm_warmup diff --git a/bundle/plugin-site/ornith_g256/launch.py b/bundle/plugin-site/ornith_g256/launch.py new file mode 100644 index 0000000000000000000000000000000000000000..020c37b465874d4dac7f0559773f38e47612c5d3 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/launch.py @@ -0,0 +1,299 @@ +"""Launch the retained Ornith G256 runtime from explicit local artifact paths.""" +# Copyright 2026 Ciru. +import argparse +import json +import os +from pathlib import Path +from .cache_full1120 import install as install_full1120 +install_full1120() + + +LIBRARIES = { + 'dense_library': 'libornith_dense_g256.so', + 'routed_library': 'libornith_routed_direct.so', + 'head_library': 'libornith_head_i8_tile.so', + 'attention_library': 'libornith_paged-v2.so', +} + + +def parser(): + p = argparse.ArgumentParser(description=__doc__) + p.add_argument('--model', type=Path, required=True) + p.add_argument('--draft', type=Path, help='DFlash2 checkpoint (required for dflash)') + p.add_argument('--native-library-directory', type=Path, required=True) + p.add_argument('--mode', choices=('dflash', 'ar', 'mtp1'), default='dflash') + p.add_argument('--draft-tokens', type=int, choices=(3, 7, 15), default=7, + help='DFlash draft tokens; 15 uses the tiled 16-token verification path') + p.add_argument('--port', type=int, default=8000) + p.add_argument('--host', default='127.0.0.1') + p.add_argument('--served-name', default='ornith-g256-dflash2') + p.add_argument('--enable-images', action='store_true', + help='Enable one image per request, preserving max8 concurrent sequences') + p.add_argument('--image-max-pixels', type=int, default=1048576, + help='Image preprocessing pixel budget when --enable-images is set') + p.add_argument('--enable-tools', action='store_true', + help='Enable automatic tool calling with the original Qwen3 XML format') + p.add_argument('--context', type=int, default=4096) + p.add_argument('--max-seqs', type=int, default=8) + p.add_argument('--cache-gib', type=float, default=32) + p.add_argument('--block-size', type=int, help='Optional vLLM KV block size; default keeps runtime selection') + p.add_argument('--prefix-cache', action='store_true', + help='Experimental DFlash7 aligned prefix reuse; selects KV block size1120') + p.add_argument('--fine-prefix-cache', action='store_true', + help='With --prefix-cache, retain draft history for8-token prompt-tail matching') + p.add_argument('--compact-prefill', action='store_true', help='Experimental compact KV FlashAttention prefill') + p.add_argument('--iu4-prefill-library', type=Path, + help='Optional fixed1120 signed IU4 prefill library; requires the current agents64k profile') + p.add_argument('--folded-decode', action='store_true', help='Experimental query/head shared KV decode') + p.add_argument('--token-major-kv', action='store_true', help='Store target KV token-major in existing cache pages') + p.add_argument('--routed-prefill-a4', action='store_true', help='Experimental one-pass A4 experts for rows above64; A8 generation') + p.add_argument('--routed-n32-library', type=Path, + help='Optional N32 expert decode library; adds resident N32 weight banks') + p.add_argument('--routed-n32-max-rows', type=int, choices=(8, 16), default=8, + help='Largest verification row count using the optional N32 expert library') + p.add_argument('--routed-n32-storage-library', type=Path, + help='Full N32 expert consumer; replace loaded N16 banks instead of retaining shadows') + p.add_argument('--cache-directory', type=Path, + default=Path(os.environ.get('XDG_CACHE_HOME', Path.home()/'.cache'))/'ornith-g256') + p.add_argument('--settings-output', type=Path) + p.add_argument('--dry-run', action='store_true', help='Print settings/environment without importing vLLM') + p.add_argument('--inspect-only', action='store_true', help='Parse installed vLLM arguments without loading a model') + return p + + +def make_settings(a): + context_limit = 262144 if a.mode == 'dflash' else 8192 + if not 1 <= a.max_seqs <= 8 or not 1 <= a.context <= context_limit: + raise ValueError(f'The {a.mode} runtime supports 1..8 sequences and 1..{context_limit} context') + if not 1024 <= a.port <= 65535 or not a.cache_gib > 0 or not a.served_name: + raise ValueError('Use a valid unprivileged port, positive cache size and model name') + if a.mode == 'dflash' and a.draft is None: + raise ValueError('--draft is required for DFlash2') + if a.mode != 'dflash' and a.draft is not None: + raise ValueError('--draft is only used with --mode dflash; MTP1 uses its target derivative') + if a.block_size is not None and a.block_size <= 0: + raise ValueError('--block-size must be positive') + if a.prefix_cache and (a.mode != 'dflash' or a.block_size not in (None, 1120)): + raise ValueError('--prefix-cache requires DFlash7 and --block-size1120 (or omitted)') + if a.fine_prefix_cache and not a.prefix_cache: + raise ValueError('--fine-prefix-cache requires --prefix-cache') + if a.prefix_cache and a.draft_tokens != 7: + raise ValueError('Prefix reuse currently requires7 draft tokens; test15 with prefix caching off') + if a.mode == 'dflash' and a.draft_tokens == 15 and a.max_seqs != 1: + raise ValueError('The experimental sixteen-token profile currently requires --max-seqs1; ' + 'concurrent verification needs its own expert precision/dispatch policy') + if a.routed_n32_max_rows != 8 and a.routed_n32_library is None: + raise ValueError('--routed-n32-max-rows requires --routed-n32-library') + if a.routed_n32_storage_library is not None and a.routed_n32_library is not None: + raise ValueError('Choose sole N32 storage or N32 decode shadows') + if a.iu4_prefill_library is not None and not ( + a.mode == 'dflash' and a.draft_tokens == 7 and a.max_seqs == 8 + and 65536 <= a.context <= 262144 and a.block_size in (None, 1120) + and a.prefix_cache and a.fine_prefix_cache + and a.compact_prefill and a.folded_decode and a.token_major_kv and a.routed_prefill_a4): + raise ValueError('--iu4-prefill-library requires agents64k: DFlash7, max8/64K, ' + 'page1120, fine prefix reuse and compact/folded token-major A4') + native = a.native_library_directory.expanduser().resolve() + settings = dict( + model=str(a.model.expanduser().resolve()), load_format='safetensors', dtype='bfloat16', + tensor_parallel_size=1, max_model_len=a.context, max_num_seqs=a.max_seqs, + max_num_batched_tokens=2048, gpu_memory_utilization=0.8, + kv_cache_memory_bytes=int(a.cache_gib*(1 << 30)), mamba_ssm_cache_dtype='float32', + enable_prefix_caching=a.prefix_cache, enable_chunked_prefill=True, enforce_eager=False, + async_scheduling=False, limit_mm_per_prompt={'image': int(a.enable_images), 'video': 0}, + seed=15035, disable_log_stats=False, quantization='ornith_g256', max_logprobs=248320, logprobs_mode='raw_logprobs', + worker_cls='ornith_g256.worker.OrnithG256Worker', + additional_config={'ornith_g256': dict( + {key: str(native/name) for key, name in LIBRARIES.items()}, + attention_mode='fast_fp32', dense_a8_max_rows=64, attention_query_tile=16)}, + compilation_config=dict(mode=3, cudagraph_mode='FULL_DECODE_ONLY', + cudagraph_capture_sizes=list(range(8, 65, 8))), + attention_config={'backend': 'CUSTOM'}, + ) + if a.context > 8192: + # DFlash's long-context path uses the installed ROCm backend. The + # retained fast native decoder only handles page1056 and <=8192; + # it is not used for the page1120 DFlash cache and needs no allocation. + settings['attention_config'] = {'backend': 'ROCM_ATTN'} + native_settings = settings['additional_config']['ornith_g256'] + native_settings['attention_mode'] = 'stock_rocm' + del native_settings['attention_library'] + if a.block_size is not None: + settings['block_size'] = a.block_size + if a.prefix_cache: + settings['block_size'] = 1120 + # The final Mamba boundary may not have DFlash's full lookahead block. + # Retain the preceding boundary and its draft window for cache lookup. + settings['prefix_cache_retention_interval'] = 1120 + if a.fine_prefix_cache: + settings['prefix_match_unit'] = 8 + settings['additional_config']['ornith_g256']['draft_full_retention'] = True + if a.compact_prefill: + if a.context <= 8192: + raise ValueError('--compact-prefill requires the long-context ROCm profile') + settings['additional_config']['ornith_g256']['compact_prefill'] = True + if a.folded_decode: + settings['additional_config']['ornith_g256']['folded_decode'] = True + if a.token_major_kv: + if not (a.compact_prefill and a.folded_decode): + raise ValueError('--token-major-kv requires --compact-prefill and --folded-decode') + settings['additional_config']['ornith_g256']['token_major_kv'] = True + if a.routed_prefill_a4: + settings['additional_config']['ornith_g256']['routed_prefill_activation_bits'] = 4 + if a.iu4_prefill_library is not None: + settings['additional_config']['ornith_g256']['iu4_prefill_library'] = str( + a.iu4_prefill_library.expanduser().resolve()) + if a.routed_n32_library is not None: + settings['additional_config']['ornith_g256'].update( + routed_decode_n32=True, + routed_n32_max_rows=a.routed_n32_max_rows, + routed_n32_library=str(a.routed_n32_library.expanduser().resolve())) + if a.routed_n32_storage_library is not None: + settings['additional_config']['ornith_g256'].update( + routed_storage_n32=True, + routed_library=str(a.routed_n32_storage_library.expanduser().resolve())) + if a.mode == 'dflash': + settings['speculative_config'] = dict( + method='dflash', model=str(a.draft.expanduser().resolve()), + num_speculative_tokens=a.draft_tokens, quantization=None, attention_backend='ROCM_ATTN') + block = a.draft_tokens + 1 + settings['compilation_config']['cudagraph_capture_sizes'] = list( + range(block, block * a.max_seqs + 1, block)) + if a.draft_tokens == 15: + if not (a.context > 8192 and a.compact_prefill and a.folded_decode + and a.token_major_kv): + raise ValueError('15 draft tokens require the long-context compact/folded token-major KV path') + settings['additional_config']['ornith_g256']['folded_decode_max_queries'] = 16 + settings['mamba_cache_mode'] = 'align' + elif a.mode == 'mtp1': + settings['speculative_config'] = dict( + method='mtp', num_speculative_tokens=1, moe_backend='triton', attention_backend='ROCM_ATTN') + settings['mamba_cache_mode'] = 'align' + settings['compilation_config']['cudagraph_capture_sizes'] = list(range(1, 17)) + else: + settings['compilation_config']['cudagraph_capture_sizes'] = list(range(1, 9)) + if a.enable_images: + if a.image_max_pixels < 65536: + raise ValueError('--image-max-pixels must be at least the processor minimum65536') + settings['mm_processor_kwargs'] = {'max_pixels': a.image_max_pixels} + settings['mm_encoder_attn_backend'] = 'TRITON_ATTN' + return settings + + +def environment(cache): + cache = str(cache.expanduser().resolve()) + return dict( + VLLM_PLUGINS='ornith_g256', OMP_NUM_THREADS='2', VLLM_MOE_SKIP_PADDING='1', + VLLM_ROCM_USE_AITER='0', VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS='0', + VLLM_DISABLE_SHARED_EXPERTS_STREAM='1', VLLM_USE_V2_MODEL_RUNNER='0', + VLLM_ENABLE_V1_MULTIPROCESSING='1', VLLM_ALLOW_INSECURE_SERIALIZATION='0', + VLLM_WORKER_MULTIPROC_METHOD='spawn', OPENBLAS_NUM_THREADS='2', + VLLM_NO_USAGE_STATS='1', DO_NOT_TRACK='1', + PYTHONDONTWRITEBYTECODE='1', HF_HUB_OFFLINE='1', TRANSFORMERS_OFFLINE='1', + XDG_CACHE_HOME=cache, AITER_JIT_DIR=cache+'/aiter', TRITON_CACHE_DIR=cache+'/triton', + TORCHINDUCTOR_CACHE_DIR=cache+'/inductor', VLLM_CACHE_ROOT=cache+'/vllm') + + +def main(argv=None): + a = parser().parse_args(argv) + settings = make_settings(a) + env = environment(a.cache_directory) + frontend = (dict(enable_auto_tool_choice=True, tool_call_parser='qwen3_xml', reasoning_parser='qwen3') + if a.enable_tools else {}) + if a.prefix_cache: + frontend['enable_prompt_tokens_details'] = True + if a.settings_output: + a.settings_output.parent.mkdir(parents=True, exist_ok=True) + a.settings_output.write_text(json.dumps(settings, indent=2, allow_nan=False)+'\n') + print(json.dumps(dict(settings=settings, environment=env, frontend=frontend, host=a.host, port=a.port, + served_name=a.served_name), indent=2, allow_nan=False), flush=True) + if a.dry_run: + return + for key in ('model',): + if not Path(settings[key]).is_dir(): + raise FileNotFoundError(settings[key]) + if a.draft and not a.draft.expanduser().is_dir(): + raise FileNotFoundError(a.draft) + for key in (*LIBRARIES, 'iu4_prefill_library'): + path = settings['additional_config']['ornith_g256'].get(key) + if path is None: + continue + if not Path(path).is_file(): + raise FileNotFoundError(path) + os.environ.update(env) + # Keep all vLLM imports after environment selection. The project's installed + # plugin entry point is also visible to worker processes via PYTHONPATH. + from vllm.entrypoints.serve.utils.api_utils import cli_env_setup + from vllm.entrypoints.launchers.cli_args import make_arg_parser, validate_parsed_serve_args + from vllm.utils.argparse_utils import FlexibleArgumentParser + from .dynamic_graphs import install_config_hook + install_config_hook() + from vllm.engine.arg_utils import AsyncEngineArgs + import torch + cli_env_setup() + server_parser = make_arg_parser(FlexibleArgumentParser(description='Ornith G256 serving')) + args = server_parser.parse_args([]) + for name, value in settings.items(): + if name not in AsyncEngineArgs.__dataclass_fields__ or not hasattr(args, name): + raise ValueError('Unsupported installed vLLM engine setting: '+name) + setattr(args, name, value) + args.host, args.port, args.served_model_name = a.host, a.port, [a.served_name] + args.disable_uvicorn_access_log = True + for name, value in frontend.items(): + if not hasattr(args, name): + if name == 'enable_prompt_tokens_details': + print('ORNITH_G256_PROMPT_TOKEN_DETAILS_UNAVAILABLE: use prefix-cache metrics', flush=True) + continue + raise ValueError('Unsupported installed vLLM frontend setting: '+name) + setattr(args, name, value) + validate_parsed_serve_args(args) + AsyncEngineArgs.from_cli_args(args) + if a.inspect_only: + if torch.cuda.is_initialized(): + raise RuntimeError('Argument inspection unexpectedly initialized CUDA') + print('ORNITH_G256_ARGUMENTS_OK', flush=True) + return + import uvloop + from vllm.entrypoints.launchers.api_server.entry import run_server + uvloop.run(run_server(args)) + + + + +# Isolated experiment: the original agents64k CLI establishes every base setting. +_original_make_settings = make_settings +def make_settings(a): + settings = _original_make_settings(a) + assert (a.mode == 'dflash' and a.draft_tokens == 7 and a.max_seqs == 8 + and a.prefix_cache and a.fine_prefix_cache and a.iu4_prefill_library) + settings['block_size'] = 1120 + settings['prefix_cache_retention_interval'] = 1120 + # Upstream Mamba alignment supports intermediate sub-block chunks. Keep + # Q1120 to retain the measured IU4 prefill path and current A4 PP shape. + settings['long_prefill_token_threshold'] = 1120 + settings['speculative_config'].update(num_speculative_tokens=15, + num_speculative_tokens_per_batch_size=[(1, 1, 15), (2, 8, 7)]) + settings['compilation_config']['cudagraph_capture_sizes'] = [16, 32, 48, 64] + native = settings['additional_config']['ornith_g256'] + native.update(dynamic_spec_profile=True, adaptive_c1_fallback=True, folded_decode_max_queries=16, + routed_storage_n32=True, + routed_library=str(Path(__file__).resolve().parents[2] / 'native' / 'libornith_routed_storage_n32.so')) + return settings + + +_shared_head_settings = make_settings +def make_settings(a): + s = _shared_head_settings(a) + s['additional_config']['ornith_g256']['head_library'] = str(Path(__file__).resolve().parents[2] / 'native' / 'libornith_head_i8_tile.so') + return s + +_fine_settings = make_settings +def make_settings(a): + s = _fine_settings(a) + s['prefix_match_unit'] = None + s['prefix_cache_retention_interval'] = 0 + s['additional_config']['ornith_g256']['draft_full_retention'] = False + return s + +if __name__ == '__main__': + main() diff --git a/bundle/plugin-site/ornith_g256/lifecycle.py b/bundle/plugin-site/ornith_g256/lifecycle.py new file mode 100644 index 0000000000000000000000000000000000000000..3742419cdc4cbdb838a21aa717d2e2d69fc368c5 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/lifecycle.py @@ -0,0 +1,210 @@ +"""Bounded worker-step error snapshots, separate from captured model work.""" +# Copyright 2026 Ciru. +from collections import deque +from dataclasses import dataclass +import threading + +import torch +from vllm.v1.outputs import AsyncModelRunnerOutput + +def _outside_capture(device): + if device.type == 'cuda': + with torch.cuda.device(device): + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError('Worker flag reset/check is forbidden during graph capture') + + +class OrnithWorkerFault(RuntimeError): + """A native error prevents this worker from delivering further outputs.""" + + +@dataclass +class _Slot: + host: object + event: object + + +@dataclass +class StepTicket: + sequence: int + name: str + slot: _Slot + queued: bool = False + completed: bool = False + flags: tuple[int, ...] | None = None + + +class CudaFlagTransport: + """One 160-byte D2H snapshot per 40-layer step, using preallocated storage.""" + def __init__(self, binding): + self.binding = binding + self.device = binding.backend.device + if self.device.type != 'cuda': + raise ValueError('Production worker flag transport requires ROCm GPU storage') + self.fallback_synchronizations = 0 + + def allocate_slot(self): + _outside_capture(self.device) + return _Slot(torch.empty(len(self.binding.layer_prefixes), dtype=torch.int32, + device='cpu', pin_memory=True), + torch.cuda.Event(blocking=True)) + + def queue(self, slot): + _outside_capture(self.device) + with torch.cuda.device(self.device): + slot.host.copy_(self.binding.error_flags, non_blocking=True) + slot.event.record(torch.cuda.current_stream(self.device)) + + def finish(self, slot): + _outside_capture(self.device) + with torch.cuda.device(self.device): + # Normal sampling already waited for its later D2H event. A direct + # return/startup/error path may need this one step-level wait. + if not slot.event.query(): + self.fallback_synchronizations += 1 + slot.event.synchronize() + return tuple(int(value) for value in slot.host.tolist()) + + +class WorkerStepLifecycle: + """Own snapshots across execute→sample and asynchronous output completion. + + Model calls stay serialized on one worker stream. Tickets permit bounded + overlap with CPU output handling, not concurrent use of the native arena. + Earlier tickets are checked before a later output can be delivered. + """ + def __init__(self, binding, max_inflight, *, _transport=None): + if type(max_inflight) is not int or not 1 <= max_inflight <= 64: + raise ValueError('Invalid bounded worker output capacity') + self.binding = binding + self.transport = _transport if _transport is not None else CudaFlagTransport(binding) + self._free = deque(self.transport.allocate_slot() for _ in range(max_inflight)) + self._pending = {} + self._lock = threading.RLock() + self._sequence = 0 + self._fault = None + self.resets = self.copies = self.completions = 0 + + def begin(self, name): + with self._lock: + self.raise_if_failed() + if not self._free: + raise RuntimeError('Ornith output snapshot pool exhausted before a new worker step') + slot = self._free.popleft() + ticket = StepTicket(self._sequence, str(name), slot) + self._sequence += 1 + try: + self.binding.reset_error_flags() + except BaseException as error: + self._free.appendleft(slot) + self.poison(error) + raise + self.resets += 1 + self._pending[ticket.sequence] = ticket + return ticket + + def raise_if_failed(self): + if self._fault is not None: + raise self._fault + + def poison(self, error): + """A failed GPU/worker boundary requires process replacement, not reuse.""" + with self._lock: + if self._fault is None: + self._fault = OrnithWorkerFault( + f'Ornith worker failed and cannot deliver more outputs: {type(error).__name__}: {error}') + self._fault.__cause__ = error + + def queue(self, ticket): + with self._lock: + if ticket.completed or ticket.sequence not in self._pending: + raise RuntimeError('Cannot snapshot a completed/unknown worker step') + if ticket.queued: + return + try: + self.transport.queue(ticket.slot) + except BaseException as error: + self.poison(error) + raise + ticket.queued = True + self.copies += 1 + + def complete(self, ticket): + with self._lock: + # Same-stream GPU ordering means completing a later output also + # admits all earlier snapshots. Never let a healthy later step + # conceal an earlier native error, even if consumers reorder calls. + for sequence in sorted(self._pending): + if sequence > ticket.sequence: + break + earlier = self._pending[sequence] + if not earlier.queued: + raise RuntimeError('Worker completion reached an unsnapshotted earlier step') + try: + flags = self.transport.finish(earlier.slot) + if len(flags) != len(self.binding.layer_prefixes): + raise RuntimeError('Worker snapshot width changed') + except BaseException as error: + self.poison(error) + raise + earlier.flags, earlier.completed = flags, True + if any(flags) and self._fault is None: + failed = {prefix: value for prefix, value in zip(self.binding.layer_prefixes, flags) if value} + self._fault = OrnithWorkerFault( + f'Ornith native worker step {earlier.sequence} ({earlier.name}) failed: {failed}') + self._free.append(earlier.slot) + del self._pending[sequence] + self.completions += 1 + self.raise_if_failed() + if not ticket.completed: + raise RuntimeError('Unknown or incomplete Ornith worker ticket') + return ticket.flags + + def abort(self, ticket, error): + """Capture flags on an exceptional worker boundary and preserve cause.""" + try: + if not ticket.queued: + self.queue(ticket) + self.complete(ticket) + except BaseException as flag_error: + if isinstance(flag_error, OrnithWorkerFault): + if flag_error is error: + raise + raise flag_error from error + error.add_note(f'Ornith flag capture also failed: {type(flag_error).__name__}: {flag_error}') + self.poison(error) + raise error + + def checked_output(self, output, ticket): + if isinstance(output, AsyncModelRunnerOutput): + return CheckedAsyncOutput(output, self, ticket) + self.complete(ticket) + return output + + def drain(self): + with self._lock: + tickets = list(self._pending.values()) + for ticket in tickets: + if not ticket.queued: + self.queue(ticket) + if tickets: + self.complete(tickets[-1]) + self.raise_if_failed() + + +class CheckedAsyncOutput(AsyncModelRunnerOutput): + """The existing executor calls get_output before delivering any tokens.""" + def __init__(self, delegate, lifecycle, ticket): + self._delegate, self._lifecycle, self._ticket = delegate, lifecycle, ticket + self._used = False + + def get_output(self): + if self._used: + raise RuntimeError('Async Ornith worker output may be consumed only once') + self._used = True + try: + output = self._delegate.get_output() + except BaseException as error: + self._lifecycle.abort(self._ticket, error) + self._lifecycle.complete(self._ticket) + return output diff --git a/bundle/plugin-site/ornith_g256/loader.py b/bundle/plugin-site/ornith_g256/loader.py new file mode 100644 index 0000000000000000000000000000000000000000..2bad6ec3f86660c4626c79c7c5e7843be9fbe0c8 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/loader.py @@ -0,0 +1,54 @@ +"""Strict direct packed-bank loader for pinned MoERunner delegation. + +Checkpoint keys use model.language_model.layers.N.mlp.experts.routed_experts.FIELD. +Qwen3.5's wrapper maps only the outer model prefix. MoERunner passes the +routed_experts.FIELD suffix to this bound RoutedExperts method unchanged. +This avoids upstream's per-expert/fused-weight matcher, which consumes and +silently drops unmatched direct parameter names. No installed class is patched. +""" +# Copyright 2026 Ciru. +from types import MethodType + +FIELDS=('w13_tilebank_codes','w13_tilebank_metadata', + 'w2_tilebank_codes','w2_tilebank_metadata') +PREFIX='routed_experts.' + + +def _load_direct(self,weights): + import torch + if getattr(self,'_ornith_weights_ready',False):raise ValueError('Tilebank weights already sealed') + for key,value in weights: + if not key.startswith(PREFIX) or key[len(PREFIX):] not in FIELDS: + raise ValueError('Unsupported tilebank checkpoint field '+key) + field=key[len(PREFIX):] + if field in self._tilebank_loaded_fields:raise ValueError('Duplicate tilebank field '+field) + parameter=self._parameters.get(field) + if parameter is None:raise ValueError('Missing registered tilebank parameter '+field) + if value.shape!=parameter.shape or value.dtype!=parameter.dtype: + raise ValueError('Tilebank shape/dtype mismatch for '+field) + if value.device.type=='meta' or parameter.device.type=='meta': + raise ValueError('Tilebank loader requires materialized tensors') + if field.endswith('_metadata'): + if value.device.type!='cpu' or value.dtype!=torch.uint32: + raise ValueError('Tilebank metadata admission expects CPU uint32 packed pairs') + pairs=value.contiguous().view(torch.float16).reshape(*value.shape,2) + if not bool(torch.isfinite(pairs).all()) or not bool((pairs[...,0]>=2**-14).all()): + raise ValueError('Tilebank metadata requires finite offset and normal-floor scale') + with torch.no_grad():parameter.copy_(value) + self._tilebank_loaded_fields.add(field) + # Relative to the MoERunner, not its RoutedExperts child. + yield PREFIX+field + + +def install_tilebank_loader(layer): + """Call once from the project quant method's create_weights on RoutedExperts.""" + if hasattr(layer,'_tilebank_loaded_fields'):raise ValueError('Tilebank loader already installed') + if any(name not in layer._parameters for name in FIELDS):raise ValueError('Create all four tilebank parameters first') + layer._tilebank_loaded_fields=set() + layer.load_weights=MethodType(_load_direct,layer) + + +def require_complete_tilebank_load(layer): + """Call during process_weights_after_loading, before runtime binding.""" + if getattr(layer,'_tilebank_loaded_fields',set())!=set(FIELDS): + raise ValueError('Incomplete tilebank packed-bank checkpoint coverage') diff --git a/bundle/plugin-site/ornith_g256/method.py b/bundle/plugin-site/ornith_g256/method.py new file mode 100644 index 0000000000000000000000000000000000000000..5767c3932888e84cd965839f473dd866fd9ac412 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/method.py @@ -0,0 +1,183 @@ +"""Direct packed loads and fresh-output linear/MoE implementations.""" +# Copyright 2026 Ciru. +import torch +from vllm.model_executor.layers.linear import LinearMethodBase +from vllm.model_executor.layers.fused_moe.config import FusedMoEQuantConfig, FusedMoEQuantDesc +from vllm.model_executor.layers.quantization.utils.quant_utils import GroupShape +from .moe_base import OrnithMoEMethodBase +from .loader import install_tilebank_loader, require_complete_tilebank_load + + +def register_packed(layer, specs): + layer._g256_loaded = set() + layer._g256_ready = False + for name, shape, dtype in specs: + param = torch.nn.Parameter(torch.empty(shape, dtype=dtype), requires_grad=False) + layer.register_parameter(name, param) + + def load(destination, value, *shard_args, name=name, param=param): + if layer._g256_ready or name in layer._g256_loaded or destination is not param: + raise ValueError(f"Duplicate or late packed load: {name}") + if value.shape != param.shape or value.dtype != param.dtype: + raise ValueError(f"Packed shape/dtype mismatch: {name}: {value.shape}/{value.dtype}") + # The fused qkvz checkpoint alias acquires a vLLM shard_id; its + # payload already contains all four projections, so no slicing. + with torch.no_grad(): param.copy_(value) + layer._g256_loaded.add(name) + param.weight_loader = load + + +class G256LinearMethod(LinearMethodBase): + fields = ("g256_codes", "g256_metadata") + + def __init__(self, prefix): + self.prefix, self.runtime, self.slot = prefix, None, None + + def create_weights(self, layer, input_size_per_partition, output_partition_sizes, + input_size, output_size, params_dtype, **attrs): + self.k, self.n = input_size, output_size + if (input_size_per_partition != input_size or sum(output_partition_sizes) != output_size + or getattr(layer, 'tp_size', 1) != 1 or params_dtype != torch.bfloat16 + or getattr(layer, 'has_bias', False) or not 0 < self.n <= 12288 + or self.n % 16 or not 0 < self.k <= 8192 or self.k % 256): + raise ValueError(f"Unsupported G256 linear geometry: {self.prefix}: N{self.n}/K{self.k}") + register_packed(layer, [(self.fields[0], (self.n//16, self.k//256, 32, 16), torch.uint32), + (self.fields[1], (self.n//16, self.k//256, 16), torch.uint32)]) + + def process_weights_after_loading(self, layer): + if layer._g256_loaded != set(self.fields): + raise ValueError(f"Incomplete packed weights: {self.prefix}") + layer._g256_ready = True + + def apply(self, layer, x, bias=None): + if self.runtime is None or not layer._g256_ready or bias is not None: + raise RuntimeError("G256 requires loaded weights and a bound worker") + shape = x.shape[:-1] + x = x.reshape(-1, self.k).contiguous() + out = torch.empty((x.shape[0], self.n), dtype=x.dtype, device=x.device) + from .native import dense_out + dense_out(x, layer.g256_codes, layer.g256_metadata, self.runtime.workspace, + out, self.runtime.slots[self.slot], self.runtime.capacity, + self.n, self.k, self.runtime.dense_geometry, self.runtime.dense_a8_max_rows) + return out.view(*shape, self.n) + + +class W8HeadMethod(G256LinearMethod): + fields = ("head_tilebank_codes", "head_scales") + + def create_weights(self, layer, input_size_per_partition, output_partition_sizes, + input_size, output_size, params_dtype, **attrs): + self.k, self.n = input_size, output_size + if (input_size_per_partition, input_size, output_size, params_dtype, + getattr(layer, 'tp_size', 1)) != (2048, 2048, 248320, torch.bfloat16, 1): + raise ValueError("W8 head requires the full TP1 Ornith vocabulary") + register_packed(layer, [(self.fields[0], (15520, 128, 4, 16), torch.uint32), + (self.fields[1], (248320, 16), torch.float16)]) + + def apply(self, layer, x, bias=None): + if self.runtime is None or not layer._g256_ready or bias is not None: + raise RuntimeError("W8 head requires loaded weights and a bound worker") + x = x.reshape(-1, self.k).contiguous() + out = torch.empty((x.shape[0], self.n), dtype=x.dtype, device=x.device) + from .native import head_out + # Startup dummy logits can have more rows than normal C1/C8 sampling. + for start in range(0, x.shape[0], 64): + head_out(x[start:start+64], layer.head_tilebank_codes, layer.head_scales, + self.runtime.workspace, out[start:start+64], self.runtime.slots[self.slot], + self.runtime.head_geometry) + return out + + +class G256DispatchConfig(FusedMoEQuantConfig): + def __init__(self): + super().__init__( + _a1=FusedMoEQuantDesc(dtype='ornith_s8_g256', shape=GroupShape(1, 256)), + _a2=FusedMoEQuantDesc(dtype='ornith_s8_g256', shape=GroupShape(1, 256)), + _w1=FusedMoEQuantDesc(dtype='ornith_affine_u4_g256', shape=GroupShape(1, 256)), + _w2=FusedMoEQuantDesc(dtype='ornith_affine_u4_g256', shape=GroupShape(1, 256)), + is_scale_swizzled=False) + + @property + def ocp_mx_scheme(self): return None + def config_name(self, dtype): + raise RuntimeError("G256 uses its own native dispatch") + + +class G256MoEMethod(OrnithMoEMethodBase): + def create_weights(self, layer, num_experts, hidden_size, + intermediate_size_per_partition, params_dtype, **attrs): + if (num_experts, hidden_size, intermediate_size_per_partition, params_dtype) != (256, 2048, 512, torch.bfloat16): + raise ValueError("G256 routed weights require E256/H2048/I512 BF16") + self.runtime, self.slot = None, None + layer._ornith_weights_ready = False + for proj, n, k in (("w13", 1024, 2048), ("w2", 2048, 512)): + for suffix, shape in (("codes", (256, n//16, k//256, 32, 16)), + ("metadata", (256, n//16, k//256, 16))): + layer.register_parameter(f"{proj}_tilebank_{suffix}", + torch.nn.Parameter(torch.empty(shape, dtype=torch.uint32), requires_grad=False)) + install_tilebank_loader(layer) + + def process_weights_after_loading(self, layer): + require_complete_tilebank_load(layer) + layer._ornith_weights_ready = True + + def prepare_n32_weights(self, layer, *, replace=False): + """Reorder at startup, optionally replacing the loaded N16 storage.""" + if not layer._ornith_weights_ready or self.runtime is not None: + raise RuntimeError('N32 shadows require loaded, unbound routed weights') + if getattr(layer, '_ornith_storage_n32', False): + raise RuntimeError('N32 expert storage is prepared once') + total = 0 + with torch.no_grad(): + for proj in ('w13', 'w2'): + for suffix in ('codes', 'metadata'): + name = f'{proj}_n32_{suffix}' + if hasattr(layer, name): + raise RuntimeError('N32 routed shadows are prepared once') + parent = getattr(layer, f'{proj}_tilebank_{suffix}') + e, nt, groups = parent.shape[:3] + if (parent.dtype != torch.uint32 or parent.device.type != 'cuda' + or not parent.is_contiguous() or nt % 2): + raise ValueError('N32 shadows require contiguous GPU U32 N16 banks') + if suffix == 'codes': + shadow = parent.reshape(e, nt//2, 2, groups, 32, 16).permute( + 0, 1, 3, 4, 2, 5).contiguous().reshape(e, nt//2, groups, 32, 32) + else: + shadow = parent.reshape(e, nt//2, 2, groups, 16).permute( + 0, 1, 3, 2, 4).contiguous().reshape(e, nt//2, groups, 32) + if replace: + # Preserve the registered parameter object/loader. Its + # payload changes layout once, before binding/capture. + # The source checkpoint remains in its load format. + parent.data = shadow + else: + layer.register_buffer(name, shadow, persistent=False) + total += shadow.numel()*shadow.element_size() + if replace: + layer._ornith_storage_n32 = True + return total + + def get_fused_moe_quant_config(self, layer): + if self.moe_quant_config is None: + self.moe_quant_config = G256DispatchConfig() + return self.moe_quant_config + + def apply(self, layer, x, topk_weights, topk_ids, shared_experts=None, shared_experts_input=None): + if self.runtime is None or not layer._ornith_weights_ready: + raise RuntimeError("G256 routed weights require a bound worker") + from .native import routed_out, routed_out_n32 + out = torch.empty_like(x) + if self.runtime.routed_decode_n32: + routed_out_n32(x, topk_weights, topk_ids, layer.w13_tilebank_codes, + layer.w13_tilebank_metadata, layer.w2_tilebank_codes, + layer.w2_tilebank_metadata, layer.w13_n32_codes, + layer.w13_n32_metadata, layer.w2_n32_codes, + layer.w2_n32_metadata, self.runtime.workspace, out, + self.runtime.slots[self.slot], self.runtime.capacity, + self.runtime.routed_n32_max_rows) + return out + routed_out(x, topk_weights, topk_ids, layer.w13_tilebank_codes, + layer.w13_tilebank_metadata, layer.w2_tilebank_codes, + layer.w2_tilebank_metadata, self.runtime.workspace, out, + self.runtime.slots[self.slot], self.runtime.capacity) + return out diff --git a/bundle/plugin-site/ornith_g256/moe_base.py b/bundle/plugin-site/ornith_g256/moe_base.py new file mode 100644 index 0000000000000000000000000000000000000000..6df214b7e205f299bfb4c78fe7ed4c4718fa4619 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/moe_base.py @@ -0,0 +1,67 @@ +"""Compact loader and explicitly configured grouped vLLM execution boundary.""" +# Copyright 2026 Ciru. +import torch +from vllm.model_executor.layers.fused_moe.activation import MoEActivation +from vllm.model_executor.layers.fused_moe.fused_moe_method_base import FusedMoEMethodBase + + +class NativeImplementationUnavailable(RuntimeError): + pass + + +class OrnithMoEMethodBase(FusedMoEMethodBase): + """Unchanged routing/geometry contract inherited by the G256 method.""" + def __init__(self, moe, quant_config, prefix): + super().__init__(moe) + self.quant_config, self.prefix = quant_config, prefix + self._grouped_runtime = None + self._grouped_layer_index = None + if (moe.num_experts, moe.num_local_experts, moe.num_logical_experts, + moe.experts_per_token, moe.hidden_dim, moe.intermediate_size) != (256, 256, 256, 8, 2048, 512): + raise ValueError("ornith_tilebank requires E256/top8/H2048/I512 without shared/redundant expert slots") + if moe.in_dtype != torch.bfloat16 or moe.activation != MoEActivation.SILU: + raise ValueError("ornith_tilebank requires BF16 inputs and ordinary SiLU gating") + if moe.has_bias or moe.is_lora_enabled or any( + getattr(moe, name, None) is not None for name in ("swiglu_limit", "swiglu_alpha", "swiglu_beta")): + raise ValueError("Bias, LoRA and modified SwiGLU are outside the initial Ornith contract") + self._validate_parallel(moe.moe_parallel_config) + if moe.aiter_fmoe_shared_expert_enabled or moe.rocm_aiter_fmoe_enabled: + raise ValueError("Disable AITER fused MoE/shared experts for Ornith IU4 bringup") + if moe.defer_moe_finalize or any(getattr(moe, name, None) is not None + for name in ('activation_situ_beta', 'activation_situ_linear_beta')): + raise ValueError('Deferred finalize and modified activations are unsupported') + + @staticmethod + def _validate_parallel(parallel): + if any(getattr(parallel, key) != 1 for key in ("tp_size", "ep_size", "dp_size", "pcp_size", "sp_size")): + raise ValueError("Initial ornith_tilebank adapter supports one rank, no sequence parallelism") + if parallel.use_ep or parallel.enable_eplb: + raise ValueError("Initial ornith_tilebank adapter does not support EP/EPLB") + + def maybe_roundup_sizes(self, hidden_size, intermediate_size_per_partition, + act_dtype, moe_parallel_config): + self._validate_parallel(moe_parallel_config) + if (hidden_size, intermediate_size_per_partition, act_dtype) != (2048, 512, torch.bfloat16): + raise ValueError("Unexpected Ornith shape/dtype; implicit padding is forbidden") + return hidden_size, intermediate_size_per_partition + + @property + def skip_forward_padding(self): + return True + + @property + def topk_indices_dtype(self): + return torch.int32 + + @staticmethod + def validate_layer_scope(layer): + if (layer.apply_router_weight_on_input or layer.use_grouped_topk + or layer.custom_routing_function is not None + or layer.num_expert_group is not None or layer.topk_group is not None + or layer.e_score_correction_bias is not None + or layer.scoring_func != 'softmax' or not layer.renormalize + or layer.routed_scaling_factor != 1.0): + raise ValueError('Ornith requires unchanged ordinary renormalized top8 routing') + + def apply_monolithic(self, layer, x, router_logits, input_ids=None): + raise NativeImplementationUnavailable("Ornith does not implement monolithic router execution") diff --git a/bundle/plugin-site/ornith_g256/native.py b/bundle/plugin-site/ornith_g256/native.py new file mode 100644 index 0000000000000000000000000000000000000000..834c9f490dadc6b2e451ab1ff631b8a53c4342fd --- /dev/null +++ b/bundle/plugin-site/ornith_g256/native.py @@ -0,0 +1,224 @@ +"""Current-stream native launches with explicit scratch/output mutations.""" +# Copyright 2026 Ciru. +import ctypes as C +from pathlib import Path +import torch + +_libs = {} +_PHASE_SEGMENTS = () + + +class DenseLayout(C.Structure): + _fields_ = [(n, C.c_size_t) for n in + ('workspace_bytes', 'transformed', 'low', 'high', 'scales', 'sums', 'projection')] + + +class RoutedLayout(C.Structure): + _fields_ = [(n, C.c_size_t) for n in + ('workspace_bytes', 'gate_x', 'gate_low', 'gate_high', 'gate_scales', + 'gate_sums', 'gate_y', 'middle', 'down_x', 'down_low', 'down_high', + 'down_scales', 'down_sums', 'route_out', 'final_f32')] + + +class HeadLayout(C.Structure): + _fields_ = [(n, C.c_size_t) for n in + ('workspace_bytes', 'transformed', 'activation_words', 'activation_scales', 'projection')] + + +def configure(settings, capacity, shapes): + routed_activation_bits = settings.get('routed_activation_bits', 8) + routed_prefill_activation_bits = settings.get('routed_prefill_activation_bits', 8) + routed_decode_n32 = settings.get('routed_decode_n32', False) + routed_storage_n32 = settings.get('routed_storage_n32', False) + if routed_storage_n32 and routed_decode_n32: + raise ValueError('Choose sole N32 storage or N32 decode shadows') + if settings.get('routed_n32_max_rows', 8) not in (8, 16): + raise ValueError('N32 routed decode crossover must be 8 or 16 rows') + if routed_activation_bits not in (4, 8) or routed_prefill_activation_bits not in (4, 8): + raise ValueError('G256 routed activation bits must be 4 or 8') + if routed_activation_bits == routed_prefill_activation_bits == 4: + raise ValueError('Choose either global routed A4 or prefill-only routed A4') + if routed_decode_n32 and routed_activation_bits != 8: + raise ValueError('N32 routed decode requires A8 verification') + routed_suffix = ('_a4' if routed_activation_bits == 4 else + '_a4_prefill' if routed_prefill_activation_bits == 4 else '') + signatures = { + 'dense': ('ornith_dense_g256', DenseLayout, [C.c_int]*4, + [C.c_void_p]*4+[C.c_size_t]+[C.c_void_p]*2+[C.c_int]*7+[C.c_void_p]), + 'routed': ('ornith_routed_direct', RoutedLayout, [C.c_int], + [C.c_void_p]*8+[C.c_size_t]+[C.c_void_p]*2+[C.c_int]*2+[C.c_void_p]), + 'head': ('ornith_head_i8_tile', HeadLayout, [C.c_int]*2, + [C.c_void_p]*4+[C.c_size_t]+[C.c_void_p]*2+[C.c_size_t]+[C.c_void_p]+[C.c_int]*4+[C.c_void_p]), + } + if routed_decode_n32: + if not settings.get('routed_n32_library'): + raise ValueError('routed_decode_n32 requires routed_n32_library') + # The isolated N32 library retains the routed layout/launch C ABI. + signatures['routed_n32'] = signatures['routed'] + for kind, (symbol, layout, query_args, launch_args) in signatures.items(): + path = str(Path(settings[kind+'_library']).resolve(strict=True)) + launch_name = symbol+'_launch'+(routed_suffix if kind == 'routed' else '') + if kind in _libs: + if _libs[kind][0] != path: raise RuntimeError("Cannot replace a native library in a live worker") + if _libs[kind][3].__name__ != launch_name: + raise RuntimeError('Cannot change activation precision in a live worker') + continue + lib = C.CDLL(path) + if kind == 'routed' and routed_storage_n32: + marker = lib.ornith_routed_storage_n32 + marker.argtypes, marker.restype = [], C.c_int + if marker() != 32: + raise ValueError('Sole N32 storage requires the full N32 consumer') + query, launch = getattr(lib, symbol+'_get_layout'), getattr(lib, launch_name) + query.argtypes, query.restype = query_args+[C.POINTER(layout)], C.c_int + launch.argtypes, launch.restype = launch_args, C.c_int + if kind == 'dense': + prepare = lib.ornith_dense_g256_prepare_bf16 + prepare.argtypes = [C.c_void_p]*6+[C.c_int]*3+[C.c_void_p] + prepare.restype = C.c_int + _libs[kind] = (path, lib, query, launch) + path, lib, query, launch = _libs['routed'] + for kind, symbol in [('routed_verify','ornith_routed_direct_launch'), + ('routed_prefill','ornith_routed_direct_launch_a4')]: + fn=getattr(lib,symbol) + fn.argtypes,fn.restype=launch.argtypes,launch.restype + _libs[kind]=(path,lib,query,fn) + sizes = [] + for n, k in shapes: + layout = DenseLayout() + check(_libs['dense'][2](capacity, n, k, 8, C.byref(layout)), 'dense layout') + sizes.append(layout.workspace_bytes) + sizes.append((capacity*k+n*k)*2) + layout = RoutedLayout() + check(_libs['routed'][2](capacity, C.byref(layout)), 'routed layout') + sizes.append(layout.workspace_bytes) + if routed_decode_n32: + n32_layout = RoutedLayout() + check(_libs['routed_n32'][2](capacity, C.byref(n32_layout)), 'N32 routed layout') + if bytes(n32_layout) != bytes(layout): + raise ValueError('N32 routed library must preserve the parent workspace layout') + layout = HeadLayout() + check(_libs['head'][2](64, 248320, C.byref(layout)), 'head layout') + sizes.append(layout.workspace_bytes) + if not sizes or min(sizes) <= 0: raise RuntimeError("Native workspace query failed") + return max(sizes) + + +def ptr(t): return C.c_void_p(t.data_ptr()) +def stream(x): return C.c_void_p(torch.cuda.current_stream(x.device).cuda_stream) +def check(status, operation): + if status: raise RuntimeError(f"{operation} failed with native status {status}") + + +def validate(x, out, capacity, k): + if (x.ndim != 2 or x.shape[1] != k or not 0 <= x.shape[0] <= capacity + or x.dtype != torch.bfloat16 or x.device.type != 'cuda' + or out.dtype != x.dtype or out.device != x.device + or not x.is_contiguous() or not out.is_contiguous()): + raise ValueError("Native G256 requires contiguous BF16 input/output within capacity") + + +@torch.library.custom_op('ornith_g256::dense_out', + mutates_args={'workspace', 'out', 'flags'}, device_types='cuda') +def dense_out(x: torch.Tensor, codes: torch.Tensor, metadata: torch.Tensor, + workspace: torch.Tensor, out: torch.Tensor, flags: torch.Tensor, + capacity: int, n: int, k: int, geometry: int, a8_max_rows: int) -> None: + if _PHASE_SEGMENTS and x.shape[0] >= _PHASE_SEGMENTS[-1][1]: + for start,end,prefill in _PHASE_SEGMENTS: + _dense_impl(x[start:end],codes,metadata,workspace,out[start:end],flags, + capacity,n,k,geometry,0 if prefill else capacity) + end=_PHASE_SEGMENTS[-1][1] + if end a8_max_rows: + # Reuse the arena for ephemeral BF16 transformed X and dequantized W. + # This A16 prefill path is intentionally distinct from A8 decode. + weight_offset = capacity*k*2 + transformed = workspace[:x.shape[0]*k*2].view(torch.bfloat16).view(x.shape[0], k) + weight = workspace[weight_offset:weight_offset+n*k*2].view(torch.bfloat16).view(n, k) + check(_libs['dense'][1].ornith_dense_g256_prepare_bf16( + ptr(x), ptr(codes), ptr(metadata), ptr(transformed), ptr(weight), ptr(flags), + x.shape[0], n, k, stream(x)), 'G256 BF16 prefill preparation') + torch.mm(transformed, weight.t(), out=out) + return + check(_libs['dense'][3](ptr(x), ptr(codes), ptr(metadata), ptr(workspace), workspace.numel(), + ptr(out), ptr(flags), x.shape[0], capacity, n, k, 8, 128, geometry, + stream(x)), 'G256 dense') + + +@torch.library.custom_op('ornith_g256::routed_out', + mutates_args={'workspace', 'out', 'flags'}, device_types='cuda') +def routed_out(x: torch.Tensor, routes: torch.Tensor, ids: torch.Tensor, + gate: torch.Tensor, gate_meta: torch.Tensor, down: torch.Tensor, down_meta: torch.Tensor, + workspace: torch.Tensor, out: torch.Tensor, flags: torch.Tensor, capacity: int) -> None: + if _PHASE_SEGMENTS and x.shape[0] >= _PHASE_SEGMENTS[-1][1]: + segments=list(_PHASE_SEGMENTS) + if segments[-1][1] None: + # Shape dispatch stays inside the opaque op; capture records the selected + # launch with persistent banks and the caller's existing scratch/stream. + if x.shape[0] <= n32_max_rows: + kind, bank = 'routed_n32', (gate_n32, gate_meta_n32, down_n32, down_meta_n32) + else: + kind, bank = 'routed', (gate, gate_meta, down, down_meta) + _routed_out(kind, x, routes, ids, *bank, workspace, out, flags, capacity) + + +@torch.library.custom_op('ornith_g256::head_out', + mutates_args={'workspace', 'out', 'flags'}, device_types='cuda') +def head_out(x: torch.Tensor, codes: torch.Tensor, scales: torch.Tensor, + workspace: torch.Tensor, out: torch.Tensor, flags: torch.Tensor, geometry: int) -> None: + validate(x, out, 64, 2048) + check(_libs['head'][3](ptr(x), ptr(codes), ptr(scales), ptr(workspace), workspace.numel(), + ptr(out), None, 0, ptr(flags), x.shape[0], 64, 248320, geometry, stream(x)), + 'W8 head') + + +@dense_out.register_fake +def _dense_fake(x, codes, metadata, workspace, out, flags, capacity, n, k, geometry, a8_max_rows): return None +@routed_out.register_fake +def _routed_fake(x, routes, ids, gate, gate_meta, down, down_meta, workspace, out, flags, capacity): return None +@routed_out_n32.register_fake +def _routed_n32_fake(x, routes, ids, gate, gate_meta, down, down_meta, + gate_n32, gate_meta_n32, down_n32, down_meta_n32, + workspace, out, flags, capacity, n32_max_rows=8): return None +@head_out.register_fake +def _head_fake(x, codes, scales, workspace, out, flags, geometry): return None diff --git a/bundle/plugin-site/ornith_g256/no_spec.py b/bundle/plugin-site/ornith_g256/no_spec.py new file mode 100644 index 0000000000000000000000000000000000000000..6d43588c750f3a29faa5b5c3f0750dbf4ed6a47a --- /dev/null +++ b/bundle/plugin-site/ornith_g256/no_spec.py @@ -0,0 +1,58 @@ +"""K0 skips query forwards but preserves DFlash target-context K/V for reentry.""" +# Copyright 2026 Ciru. +from functools import wraps + +_INSTALLED = False + + +def install(): + global _INSTALLED + if _INSTALLED: + return + import torch + from vllm.v1.spec_decode.dflash import DFlashProposer + original = DFlashProposer.propose + + @wraps(original) + def propose(self, num_speculative_tokens, *args, **kwargs): + if num_speculative_tokens != 0: + return original(self, num_speculative_tokens, *args, **kwargs) + if args: + raise ValueError('The isolated K0 context path expects runner keyword arguments') + # The six-layer query network is not needed to cache target context. + # Preserve the same combine -> accepted-context slot prep -> fused K/V + # insert path as upstream. Rejected verification tails get the existing + # negative slot mapping and are never inserted as accepted context. + self._last_draft_probs = None + hidden = self.model.combine_hidden_states(kwargs['target_hidden_states']) + # Use the already-supported Q8 input-preparation geometry. Query slots + # and masks are only scratch here: there is no query forward or sample. + # This avoids inventing a Q1 convolution path just to maintain context. + self.num_speculative_tokens = 7 + try: + self.set_inputs_first_pass( + target_token_ids=kwargs['target_token_ids'], + next_token_ids=kwargs['next_token_ids'], + target_positions=kwargs['target_positions'], + target_hidden_states=hidden, + token_indices_to_sample=kwargs['token_indices_to_sample'], + cad=kwargs['common_attn_metadata'], + num_rejected_tokens_gpu=kwargs.get('num_rejected_tokens_gpu'), + ) + count = self._dflash_num_context + self.model.precompute_and_store_context_kv( + self._dflash_hidden_states, + self._context_positions_buffer[:count], + self._context_slot_mapping_buffer[:count], + ) + finally: + self.num_speculative_tokens = 0 + count = getattr(self, '_ornith_context_only_calls', 0) + 1 + self._ornith_context_only_calls = count + if count <= 2: + print('ORNITH_K0_CONTEXT query forward skipped; target context KV maintained', flush=True) + ids = kwargs['next_token_ids'] + return torch.empty((ids.numel(), 0), dtype=torch.int64, device=ids.device) + + DFlashProposer.propose = propose + _INSTALLED = True diff --git a/bundle/plugin-site/ornith_g256/phase_dispatch.py b/bundle/plugin-site/ornith_g256/phase_dispatch.py new file mode 100644 index 0000000000000000000000000000000000000000..7560163280768ea603012f21475a33ad12b7f00a --- /dev/null +++ b/bundle/plugin-site/ornith_g256/phase_dispatch.py @@ -0,0 +1,45 @@ +"""Keep verification precision stable when new prompts join a target batch.""" +from functools import wraps +_installed=False + +def install(): + global _installed + if _installed:return + from vllm.v1.worker.gpu_model_runner import GPUModelRunner + from . import native + prepare_original=GPUModelRunner._prepare_inputs + execute_original=GPUModelRunner.execute_model + @wraps(prepare_original) + def prepare(self,scheduler_output,num_scheduled_tokens): + b=self.input_batch;n=b.num_reqs + phases=b.num_computed_tokens_cpu[:n] < b.num_prompt_tokens[:n] + counts=[int(x) for x in num_scheduled_tokens] + segments=[];requests=[];offset=0 + for req,(count,prefill) in enumerate(zip(counts,phases)): + phase=bool(prefill) + requests.append((offset,offset+count,phase, + int(b.num_computed_tokens_cpu[req])+count)) + if segments and segments[-1][2]==phase: + segments[-1]=(segments[-1][0],offset+count,phase) + else:segments.append((offset,offset+count,phase)) + offset+=count + native._PHASE_SEGMENTS=tuple(segments) if offset>64 and any(phases) and not all(phases) else () + # Same CPU metadata used by vLLM optimistic_seq_lens_cpu. Actual GPU + # seq_lens remain authoritative for attention; only prefill uses the + # exact CPU length for existing fixed-width IU4 eligibility. + native._ATTENTION_PHASE_REQUESTS=(tuple(requests) + if any(phases) and not all(phases) and not self.use_async_scheduling else ()) + if native._PHASE_SEGMENTS: + print('ORNITH_PHASE_DISPATCH '+str(native._PHASE_SEGMENTS),flush=True) + return prepare_original(self,scheduler_output,num_scheduled_tokens) + @wraps(execute_original) + def execute(self,*args,**kwargs): + native._PHASE_SEGMENTS=() + native._ATTENTION_PHASE_REQUESTS=() + try:return execute_original(self,*args,**kwargs) + finally: + native._PHASE_SEGMENTS=() + native._ATTENTION_PHASE_REQUESTS=() + GPUModelRunner._prepare_inputs=prepare + GPUModelRunner.execute_model=execute + _installed=True diff --git a/bundle/plugin-site/ornith_g256/prefill_draft.py b/bundle/plugin-site/ornith_g256/prefill_draft.py new file mode 100644 index 0000000000000000000000000000000000000000..f492b5bf9f8f1f38454026f8da7e8cfaf2875f89 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/prefill_draft.py @@ -0,0 +1,110 @@ +"""Avoid sampling DFlash candidates for scheduler-discarded prefill rows. + +Unfinished chunked-prefill requests have no valid sampled anchor token. Their +draft query rows are placeholders, not real decode predictions. The original +proposer still runs every forward/context-KV update; only the subsequent +candidate-head/selector call omits requests already marked for discard. +""" +# Copyright 2026 Ciru. +import dataclasses +import torch + + +_upstream_propose = None +_upstream_sample = None +_MISSING = object() +_MASK_ATTR = '_ornith_discard_prefill_requests' + + +def _propose_draft_token_ids(self, *args, **kwargs): + config = self.vllm_config + spec = config.speculative_config + if (getattr(config.model_config, 'quantization', None) != 'ornith_g256' + or not config.cache_config.enable_prefix_caching + or config.cache_config.mamba_cache_mode != 'align' + or spec is None or spec.method != 'dflash' or spec.num_speculative_tokens != 15 + or not getattr(self.drafter, 'is_dflash2', False) + or spec.disable_padded_drafter_batch): + return _upstream_propose(self, *args, **kwargs) + if config.scheduler_config.async_scheduling: + raise ValueError('Ornith prefill sampling scope requires synchronous scheduling') + # This is the same existing CPU mask used by _is_all_reqs_chunked_prefill; + # there is no device-to-host read or inference from hidden-state values. + mask = tuple(bool(value) for value in + self.discard_request_mask.np[:self.input_batch.num_reqs]) + prior = getattr(self.drafter, _MASK_ATTR, _MISSING) + setattr(self.drafter, _MASK_ATTR, mask) + try: + return _upstream_propose(self, *args, **kwargs) + finally: + if prior is _MISSING: + delattr(self.drafter, _MASK_ATTR) + else: + setattr(self.drafter, _MASK_ATTR, prior) + + +def _sample_draft_tokens(self, hidden_states, sampling_metadata): + discarded = getattr(self, _MASK_ATTR, None) + if not self.is_dflash2 or discarded is None or not any(discarded): + return _upstream_sample(self, hidden_states, sampling_metadata) + anchors = self._dflash_anchor_token_ids + if anchors is None: + raise ValueError('DFlash2 candidate sampling has no anchor metadata') + batch_size = anchors.numel() + steps = self.num_speculative_tokens + if (len(discarded) != batch_size or hidden_states.ndim != 2 + or hidden_states.shape[0] != batch_size * steps): + raise ValueError('Discard mask does not match DFlash2 request-major candidate rows') + active = [index for index, discard in enumerate(discarded) if not discard] + if not active: + # These IDs/probability rows are never accepted by the scheduler. Keep + # the original flattened B*K interface without invoking head/selector. + ids = torch.zeros(batch_size * steps, dtype=torch.long, device=hidden_states.device) + probs = (None if sampling_metadata.all_greedy else torch.zeros( + batch_size * steps, int(self.draft_model_config.hf_config.vocab_size), + dtype=torch.float32, device=hidden_states.device)) + return ids, probs + + rows = torch.tensor(active, dtype=torch.long, device=hidden_states.device) + compact_hidden = hidden_states.reshape(batch_size, steps, -1).index_select( + 0, rows).reshape(len(active) * steps, -1) + metadata = sampling_metadata + if sampling_metadata.temperature is not None: + # The DFlash2 sampler reads only all_greedy and per-request temperature; + # preserving all_greedy also preserves its probability-return contract. + metadata = dataclasses.replace(sampling_metadata, + temperature=sampling_metadata.temperature[:batch_size].index_select(0, rows)) + self._dflash_anchor_token_ids = anchors.index_select(0, rows) + try: + active_ids, active_probs = _upstream_sample(self, compact_hidden, metadata) + finally: + self._dflash_anchor_token_ids = anchors + if active_ids.numel() != len(active) * steps: + raise ValueError('Unexpected compact DFlash2 candidate-ID shape') + ids = active_ids.new_zeros((batch_size, steps)) + ids.index_copy_(0, rows, active_ids.reshape(len(active), steps)) + probs = None + if active_probs is not None: + if active_probs.ndim != 2 or active_probs.shape[0] != len(active) * steps: + raise ValueError('Unexpected compact DFlash2 probability shape') + probs = active_probs.new_zeros((batch_size, steps, active_probs.shape[-1])) + probs.index_copy_(0, rows, active_probs.reshape(len(active), steps, -1)) + probs = probs.flatten(0, 1).contiguous() + return ids.reshape(-1), probs + + +def install(): + """Install lazily from the validated Ornith prefix-cache annotation path.""" + from vllm.v1.worker.gpu_model_runner import GPUModelRunner + from vllm.v1.spec_decode.dflash import DFlashProposer + global _upstream_propose, _upstream_sample + current_propose = GPUModelRunner.propose_draft_token_ids + current_sample = DFlashProposer._sample_draft_tokens + if current_propose is _propose_draft_token_ids and current_sample is _sample_draft_tokens: + return + if _upstream_propose is None and _upstream_sample is None: + _upstream_propose, _upstream_sample = current_propose, current_sample + elif current_propose is not _upstream_propose or current_sample is not _upstream_sample: + raise RuntimeError('Another extension replaced DFlash proposal or sampling') + GPUModelRunner.propose_draft_token_ids = _propose_draft_token_ids + DFlashProposer._sample_draft_tokens = _sample_draft_tokens diff --git a/bundle/plugin-site/ornith_g256/prefix_cache.py b/bundle/plugin-site/ornith_g256/prefix_cache.py new file mode 100644 index 0000000000000000000000000000000000000000..12d34964c719cf94bb853ada654b33de78c9ff67 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/prefix_cache.py @@ -0,0 +1,279 @@ +"""Opt-in metadata bridge for the retained Ornith/DFlash2 prefix-cache layout. + +The installed planner has no marker for Qwen's ordinary sliding-window draft +layers. Its conservative fallback marks target Mamba groups as draft groups too. +This model-scoped bridge labels the exact, separate draft group. Fine prefix +reuse optionally retains draft KV under full-attention allocation while keeping +the drafter's sliding attention window and existing accepted-state copying. +""" +# Copyright 2026 Ciru. +import logging + + +logger = logging.getLogger(__name__) +EXPECTED_DRAFT_NAMES = frozenset( + f'model.layers.{index}.self_attn.attn' for index in range(40, 46) +) +_upstream_annotator = None +_upstream_mamba_split = None +_upstream_scheduler_init = None + + +def retain_draft_history(config): + return config.additional_config.get('ornith_g256', {}).get('draft_full_retention', False) + + +def retained_draft_specs(config, specs): + """Use the existing full-retention representation for six windowed layers. + + FullAttentionSpec.sliding_window explicitly preserves windowed computation + while retaining all physical pages. That permits existing fine-grained + full-attention/Mamba lookup and CoW without adding a new state cache. + """ + from .cache_full1120 import padded_specs + specs = padded_specs(config, specs) + if not retain_draft_history(config): + return specs + from dataclasses import fields + from vllm.v1.kv_cache_interface import AttentionSpec, FullAttentionSpec, SlidingWindowSpec + validate_prefix_cache_config(config) + result = dict(specs) + names = {name for name in specs if name.startswith('model.layers.')} + if names != EXPECTED_DRAFT_NAMES: + raise ValueError('Fine prefix reuse requires exactly the six DFlash2 KV layers') + for name in names: + spec = specs[name] + if (not isinstance(spec, SlidingWindowSpec) or spec.sliding_window != 4096 + or spec.extra_retained_tokens != 0): + raise ValueError('Expected original DFlash2 4096-token window') + result[name] = FullAttentionSpec( + **{field.name: getattr(spec, field.name) for field in fields(AttentionSpec)}, + sliding_window=spec.sliding_window) + return result + + +class _AttributeView: + """Read-through object view; overrides never mutate the wrapped objects.""" + def __init__(self, wrapped, **overrides): + object.__setattr__(self, '_wrapped', wrapped) + object.__setattr__(self, '_overrides', overrides) + + def __getattr__(self, name): + if name in self._overrides: + return self._overrides[name] + return getattr(self._wrapped, name) + + def __setattr__(self, name, value): + raise AttributeError('Alignment helper view is read-only') + + +def _mamba_block_aligned_split(self, request, num_new_tokens, + num_new_local_computed_tokens=0, + num_external_computed_tokens=0): + config = self.vllm_config + if (getattr(config.model_config, 'quantization', None) != 'ornith_g256' + or not self.cache_config.enable_prefix_caching): + return _upstream_mamba_split(self, request, num_new_tokens, + num_new_local_computed_tokens, num_external_computed_tokens) + spec = config.speculative_config + if (self.block_size != 1120 or self.cache_config.mamba_cache_mode != 'align' + or spec is None or spec.method != 'dflash' or spec.num_speculative_tokens != 15 + or self.scheduler_config.async_scheduling + or not self.scheduler_config.enable_chunked_prefill): + raise ValueError('Unexpected scheduler configuration for Ornith aligned prefix reuse') + # EngineCore resets cache_config.block_size to the minimum participating + # group size (560 for the draft), while self.block_size remains the resolved + # LCM (1120 for Mamba/target). The upstream helper reads the former. Give + # only this read-only helper the resolved alignment; keep the actual cache + # config, hash granularity, allocation and every other helper rule intact. + view = _AttributeView(self, cache_config=_AttributeView( + self.cache_config, block_size=self.block_size)) + return _upstream_mamba_split(view, request, num_new_tokens, + num_new_local_computed_tokens, num_external_computed_tokens) + + +def _install_mamba_alignment(): + # Cache-group annotation happens after model/config imports are complete. + # Importing the scheduler here avoids a plugin-registration import cycle. + from vllm.v1.core.sched.scheduler import Scheduler + global _upstream_mamba_split + current = Scheduler._mamba_block_aligned_split + if current is _mamba_block_aligned_split: + return + if _upstream_mamba_split is None: + _upstream_mamba_split = current + elif current is not _upstream_mamba_split: + raise RuntimeError('Another extension replaced Mamba prefill alignment') + Scheduler._mamba_block_aligned_split = _mamba_block_aligned_split + + +def _scheduler_init(self, *args, **kwargs): + # The upstream constructor creates the coordinator and all managers. No + # request has been admitted yet, so its derived lookup policy can be set + # coherently here while the model's VllmConfig is still directly available. + _upstream_scheduler_init(self, *args, **kwargs) + config = self.vllm_config + if (getattr(config.model_config, 'quantization', None) != 'ornith_g256' + or not config.cache_config.enable_prefix_caching): + return + if config.cache_config.block_size not in (560, 1120): + raise ValueError('Unexpected physical block size for Ornith DFlash prefix reuse') + validate_prefix_cache_config(_AttributeView(config, cache_config=_AttributeView( + config.cache_config, block_size=self.block_size))) + from vllm.v1.core.kv_cache_coordinator import HybridKVCacheCoordinator + from vllm.v1.kv_cache_interface import FullAttentionSpec, SlidingWindowSpec + + coordinator = self.kv_cache_manager.coordinator + groups = coordinator.kv_cache_config.kv_cache_groups + # Reuse the full model/layout validation, including all target names/types. + # Re-installing the hooks is idempotent; the config view exposes the resolved + # scheduler alignment instead of EngineCore's minimum physical group size. + validated_config = _AttributeView(config, cache_config=_AttributeView( + config.cache_config, block_size=self.block_size)) + specs = {name: group.kv_cache_spec for group in groups for name in group.layer_names} + annotate_draft_group(validated_config, specs, groups) + if (not isinstance(coordinator, HybridKVCacheCoordinator) + or coordinator.num_reprefillable_tokens != 0): + raise ValueError('Expected hybrid DFlash with no multi-module re-prefill tail') + draft_ids = [i for i, group in enumerate(groups) + if set(group.layer_names) == EXPECTED_DRAFT_NAMES] + if len(draft_ids) != 1 or coordinator.eagle_group_ids != set(draft_ids): + raise ValueError('Unexpected coordinator draft-group identity') + draft_id = draft_ids[0] + matching = [(i, group) for i, group in enumerate(coordinator.attention_groups) + if draft_id in group.group_ids] + expected_type = FullAttentionSpec if retain_draft_history(config) else SlidingWindowSpec + if (len(matching) != 1 or matching[0][1].group_ids != [draft_id] + or not isinstance(matching[0][1].spec, expected_type)): + raise ValueError('Expected one exclusive DFlash lookup group') + index, lookup_group = matching[0] + # DFlash context KV uses original target positions; query KV starts at the + # next accepted position. Unlike Eagle, no next-token embedding is written + # into the final accepted context slot. Remove the lookup margin AND drop + # together, and make cache retention use the same unshifted window policy. + coordinator.attention_groups[index] = lookup_group._replace(use_eagle=False) + coordinator.single_type_managers[draft_id].use_eagle = False + # Keep is_eagle_group/eagle_group_ids: they identify the draft and prevent + # the upstream conservative fallback from marking target groups as drafts. + logger.info('Ornith prefix reuse: DFlash context KV uses unshifted lookup; ' + 'disabled only its Eagle lookahead margin and last-block drop') + if not retain_draft_history(config): + # Isolated diagnostic: default GCD hashing otherwise auto-enables + # partial Mamba tails even with prefix_match_unit=None. + coordinator.enable_partial_hash_hits = False + self.mamba_partial_cache_hit = False + assert self.block_size == coordinator.scheduler_block_size == 1120 + assert coordinator._cache_hit_alignment_tokens == 1120 + assert coordinator._align_cacheable(3360) == 3360 + assert all(manager.scheduler_block_size == 1120 + for manager in coordinator.single_type_managers) + logger.info('Ornith full-boundary diagnostic: hash=%s scheduler=%s ' + 'hit_alignment=%s cacheable3360=%s partial_hits=%s ' + 'partial_prompt_stops=%s group_blocks=%s', + coordinator.hash_block_size, coordinator.scheduler_block_size, + coordinator._cache_hit_alignment_tokens, + coordinator._align_cacheable(3360), + coordinator.enable_partial_hash_hits, self.mamba_partial_cache_hit, + [manager.block_size for manager in coordinator.single_type_managers]) + if retain_draft_history(config): + if not (coordinator.enable_partial_hash_hits and self.mamba_partial_cache_hit + and coordinator.hash_block_size == 8): + raise ValueError('Fine prefix lookup and Mamba prompt-tail checkpoints were not enabled') + logger.info('Ornith fine prefix reuse: retained draft KV, unchanged 4096-token ' + 'attention window, match unit8, existing Mamba prompt-tail CoW') + + +def _install_dflash_lookup(): + from vllm.v1.core.sched.scheduler import Scheduler + global _upstream_scheduler_init + current = Scheduler.__init__ + if current is _scheduler_init: + return + if _upstream_scheduler_init is None: + _upstream_scheduler_init = current + elif current is not _upstream_scheduler_init: + raise RuntimeError('Another extension replaced scheduler construction') + Scheduler.__init__ = _scheduler_init + + +def validate_prefix_cache_config(config): + """Return whether this is the supported experiment, rejecting other opt-ins.""" + if not config.cache_config.enable_prefix_caching: + return False + spec = config.speculative_config + scheduler = config.scheduler_config + cache = config.cache_config + if (spec is None or spec.method != 'dflash' or spec.num_speculative_tokens != 15 + or cache.mamba_cache_mode != 'align' or cache.block_size != 1120 + or scheduler.async_scheduling or not scheduler.enable_chunked_prefill): + raise ValueError('Ornith prefix reuse requires DFlash15, Mamba align, KV block1120, ' + 'synchronous scheduling and chunked prefill') + if retain_draft_history(config): + if cache.prefix_match_unit != 8: + raise ValueError('Fine prefix reuse requires prefix_match_unit8') + elif cache.prefix_match_unit is not None: + raise ValueError('Finer prefix matching requires retained DFlash history') + return True + + +def annotate_draft_group(config, kv_cache_spec, groups): + """Validate source-derived names/types before marking the one draft group.""" + if (getattr(config.model_config, 'quantization', None) != 'ornith_g256' + or not config.cache_config.enable_prefix_caching): + return + validate_prefix_cache_config(config) + from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec, SlidingWindowSpec + + target = config.model_config + draft = config.speculative_config.draft_model_config + if (target.architecture != 'Qwen3_5MoeForConditionalGeneration' + or target.hf_text_config.num_hidden_layers != 40 + or draft.architecture != 'DFlash2DraftModel' + or draft.hf_text_config.num_hidden_layers != 6): + raise ValueError('Ornith prefix reuse only supports the retained40-layer target and6-layer DFlash2 draft') + actual_draft_names = {name for name in kv_cache_spec if name.startswith('model.layers.')} + if actual_draft_names != EXPECTED_DRAFT_NAMES: + raise ValueError('Unexpected Ornith draft KV layer names; refusing ambiguous prefix-cache annotation') + expected_type = FullAttentionSpec if retain_draft_history(config) else SlidingWindowSpec + if any(not isinstance(kv_cache_spec[name], expected_type) + or kv_cache_spec[name].sliding_window != 4096 for name in EXPECTED_DRAFT_NAMES): + raise ValueError('Expected six DFlash KV specs with the original4096-token window') + target_names = set(kv_cache_spec) - EXPECTED_DRAFT_NAMES + if (len(target_names) != 40 + or any(not name.startswith('language_model.model.layers.') for name in target_names) + or sum(isinstance(kv_cache_spec[name], MambaSpec) for name in target_names) != 30 + or sum(isinstance(kv_cache_spec[name], FullAttentionSpec) for name in target_names) != 10): + raise ValueError('Unexpected Ornith target KV layout; expected30GDN and10attention layers') + containing = [group for group in groups if EXPECTED_DRAFT_NAMES.intersection(group.layer_names)] + if (len(containing) != 1 or set(containing[0].layer_names) != EXPECTED_DRAFT_NAMES + or not isinstance(containing[0].kv_cache_spec, expected_type)): + raise ValueError('Ornith prefix reuse requires all six draft layers in one exclusive group') + if any(group.is_eagle_group for group in groups if group is not containing[0]): + raise ValueError('Upstream marked an Ornith target group as a drafter; refusing conflicting annotation') + containing[0].is_eagle_group = True + _install_mamba_alignment() + _install_dflash_lookup() + from .prefill_draft import install as install_prefill_sampling + install_prefill_sampling() + logger.info('Ornith prefix reuse: marked exactly six DFlash2 sliding-window layers as the draft group') + + +def _annotate_eagle_groups(vllm_config, kv_cache_spec, kv_cache_groups, use_deepseek_v4_fallback=False): + # Preserve every installed upstream annotation rule before applying ours. + _upstream_annotator(vllm_config, kv_cache_spec, kv_cache_groups, + use_deepseek_v4_fallback=use_deepseek_v4_fallback) + annotate_draft_group(vllm_config, kv_cache_spec, kv_cache_groups) + + +def install(): + """Patch the planner's module-global hook once, without installed-file edits.""" + from vllm.v1.core import kv_cache_utils + global _upstream_annotator + current = kv_cache_utils._annotate_eagle_groups + if current is _annotate_eagle_groups: + return + if _upstream_annotator is None: + _upstream_annotator = current + elif current is not _upstream_annotator: + raise RuntimeError('Another extension replaced KV draft-group annotation') + kv_cache_utils._annotate_eagle_groups = _annotate_eagle_groups diff --git a/bundle/plugin-site/ornith_g256/runtime.py b/bundle/plugin-site/ornith_g256/runtime.py new file mode 100644 index 0000000000000000000000000000000000000000..c5f3a8326e6a5fb534d2299c479f229156f17920 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/runtime.py @@ -0,0 +1,97 @@ +"""One serialized native arena, bound before memory profiling and capture.""" +# Copyright 2026 Ciru. +from types import SimpleNamespace +import torch +import vllm.envs as envs +from vllm.logger import init_logger +from .method import G256LinearMethod, G256MoEMethod, W8HeadMethod +from .native import configure + +logger = init_logger('vllm.ornith_g256.runtime') + + +class Runtime: + def __init__(self, model, settings, capacity, device, *, max_model_len=4096, max_num_seqs=8): + if not envs.VLLM_DISABLE_SHARED_EXPERTS_STREAM: + raise ValueError("Set VLLM_DISABLE_SHARED_EXPERTS_STREAM=1 for the shared arena") + self.backend = SimpleNamespace(device=torch.device(device)) + self.capacity = capacity + self.routed_decode_n32 = settings.get('routed_decode_n32', False) + self.routed_storage_n32 = settings.get('routed_storage_n32', False) + self.routed_n32_max_rows = settings.get('routed_n32_max_rows', 8) + self.routed_n32_bytes = 0 + self.dense_geometry = settings.get('dense_geometry', 2) + self.dense_a8_max_rows = settings.get('dense_a8_max_rows', 32) + if self.dense_a8_max_rows not in (32, 64): + raise ValueError('G256 dense A8 crossover must be32 or64 rows') + self.head_geometry = settings.get('head_geometry', 2) + if self.dense_geometry not in (0, 1, 2) or self.head_geometry not in (0, 1, 2): + raise ValueError("Native geometry must be 0, 1, or 2") + layers = [(name, layer) for name, layer in model.named_modules() + if isinstance(getattr(layer, 'quant_method', None), (G256LinearMethod, G256MoEMethod))] + routed = [layer for _, layer in layers if isinstance(layer.quant_method, G256MoEMethod)] + heads = [layer for _, layer in layers if isinstance(layer.quant_method, W8HeadMethod)] + dense = [layer for _, layer in layers if type(layer.quant_method) is G256LinearMethod] + if (len(routed), len(dense), len(heads)) != (40, 160, 1): + raise ValueError(f"Expected 40 routed/160 dense/1 head, got {len(routed)}/{len(dense)}/{len(heads)}") + shapes = sorted({(layer.quant_method.n, layer.quant_method.k) for layer in dense}) + size = configure(settings, capacity, shapes) + if self.routed_storage_n32: + repacked = sum(layer.quant_method.prepare_n32_weights(layer, replace=True) + for layer in routed) + logger.info('Ornith sole N32 expert storage ready: %d layers, %d bytes repacked; ' + 'no N16 shadows retained', len(routed), repacked) + if self.routed_decode_n32: + for layer in routed: + self.routed_n32_bytes += layer.quant_method.prepare_n32_weights(layer) + logger.info('Ornith N32 decode shadows ready: %d layers, %d bytes; T<=%d', + len(routed), self.routed_n32_bytes, self.routed_n32_max_rows) + attention = [] + self.attention_backend = self.attention_buffers = None + self.attention_verify_backend = self.attention_verify_buffers = None + if settings.get('attention_mode', 'column') == 'fast_fp32': + from .attention_fast import NativeFastFP32, OrnithG256FastAttentionImpl + attention = [(name, layer) for name, layer in model.named_modules() + if isinstance(getattr(layer, 'impl', None), OrnithG256FastAttentionImpl)] + if len(attention) != 10: + raise ValueError(f'Expected ten target fast-attention layers, got {len(attention)}') + self.attention_backend = NativeFastFP32(settings['attention_library'], + Ccap=max_num_seqs, Lcap=min(max_model_len, 8192), device=device) + size = max(size, self.attention_backend.arena_bytes) + if settings.get('attention_verify_library'): + from .attention_verify import NativeVerifyFP32 + self.attention_verify_backend = NativeVerifyFP32(settings['attention_verify_library'], + Ccap=max_num_seqs * 8, Lcap=min(max_model_len, 8192), device=device) + size = max(size, self.attention_verify_backend.arena_bytes) + self.workspace = torch.empty(size, dtype=torch.uint8, device=device) + self.error_flags = torch.zeros(len(layers)+len(attention), dtype=torch.int32, device=device) + self.slots = tuple(self.error_flags[i:i+1] for i in range(len(layers)+len(attention))) + self.layer_prefixes = tuple(name for name, _ in layers+attention) + self.layers = tuple(layer for _, layer in layers) + for slot, layer in enumerate(self.layers): + method = layer.quant_method + if method.runtime is not None: raise RuntimeError("Native layer already bound") + method.runtime, method.slot = self, slot + if self.attention_backend is not None: + self.attention_buffers = self.attention_backend.bind_arena(self.workspace) + if self.attention_verify_backend is not None: + self.attention_verify_buffers = self.attention_verify_backend.bind_arena(self.workspace) + for slot, (name, layer) in enumerate(attention, start=len(layers)): + layer.impl.bind_native(self.attention_backend, self.attention_buffers, + self.error_flags, slot, name) + if self.attention_verify_backend is not None: + layer.impl.bind_verify(self.attention_verify_backend, self.attention_verify_buffers) + self.shapes = shapes + + def validate_identity(self): pass + + def reset_error_flags(self): + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("Reset native error flags outside graph capture") + self.error_flags.zero_() + + def check_error_flags(self): + flags = self.error_flags.cpu().tolist() + if any(flags): + raise RuntimeError(str({name: code for name, code in zip(self.layer_prefixes, flags) if code})) + return flags diff --git a/bundle/plugin-site/ornith_g256/worker.py b/bundle/plugin-site/ornith_g256/worker.py new file mode 100644 index 0000000000000000000000000000000000000000..e9f895bf9543443c80917c00840d85762e517e7b --- /dev/null +++ b/bundle/plugin-site/ornith_g256/worker.py @@ -0,0 +1,231 @@ +"""Own configurable worker; reuse the existing step-level flag transport.""" +# Copyright 2026 Ciru. +import torch +from vllm.v1.worker.gpu_worker import Worker as VllmGPUWorker +from .worker_base import OrnithWorkerBase +from .lifecycle import WorkerStepLifecycle +from .config import OrnithG256Config +from .runtime import Runtime +from .prefix_cache import validate_prefix_cache_config + + +class OrnithG256Worker(OrnithWorkerBase): + def __init__(self, vllm_config, local_rank, rank, distributed_init_method, is_driver_worker=False): + from .no_spec import install as install_no_spec + install_no_spec() + c = vllm_config + self.settings = dict(c.additional_config.get('ornith_g256', {})) + for key in ('dense_library', 'routed_library', 'head_library'): + if not self.settings.get(key): raise ValueError(f"additional_config.ornith_g256.{key} is required") + tile = self.settings.get('attention_query_tile', 32) + if tile == 16: + from .attention_tile import install as install_attention_tile + install_attention_tile(compact_prefill=self.settings.get("compact_prefill", False), + folded_decode=self.settings.get("folded_decode", False), + folded_decode_max_queries=self.settings.get("folded_decode_max_queries", 8)) + elif tile != 32: + raise ValueError('Attention query tile must be16 or32') + if self.settings.get('token_major_kv', False): + if (tile != 16 or not self.settings.get('compact_prefill', False) + or not self.settings.get('folded_decode', False) + or c.cache_config.block_size not in (560, 1120)): + raise ValueError('Token-major KV requires compact/folded target attention and page1120/2240') + from .attention_storage import install as install_attention_storage + install_attention_storage() + if __import__('os').environ.get('ORNITH_PERSISTENT_IU4_LIBRARY'): + from .attention_iu4_persistent import install as install_persistent_iu4 + install_persistent_iu4() + spec = c.speculative_config + context = c.model_config.max_model_len + if (self.settings.get('iu4_prefill_library') + and not self.settings.get('dynamic_spec_profile')): + graph = c.compilation_config + graph_mode = getattr(graph.cudagraph_mode, 'name', graph.cudagraph_mode) + capture_sizes = graph.cudagraph_capture_sizes + if (spec is None or spec.method != 'dflash' or spec.num_speculative_tokens != 7 + or not 65536 <= context <= 262144 or c.scheduler_config.max_num_seqs != 8 + or c.scheduler_config.async_scheduling + or c.cache_config.block_size != 1120 + or not c.cache_config.enable_prefix_caching + or tile != 16 or self.settings.get('attention_mode') != 'stock_rocm' + or not all(self.settings.get(key, False) for key in + ('compact_prefill', 'folded_decode', 'token_major_kv', 'draft_full_retention')) + or self.settings.get('routed_prefill_activation_bits') != 4 + or graph_mode != 'FULL_DECODE_ONLY' or not capture_sizes + or max(capture_sizes) > 64): + raise ValueError('IU4 prefill requires the synchronous agents64k profile: ' + 'DFlash7, max8/64K, page1120, fine prefix reuse, compact/folded ' + 'token-major A4 and FULL_DECODE_ONLY graphs up to64 rows') + if context > 8192 and ( + context > 262144 or spec is None or spec.method != 'dflash' + or self.settings.get('attention_mode') != 'stock_rocm'): + raise ValueError('Context above8192 requires DFlash with stock_rocm attention, up to262144') + self._spec_enabled = spec is not None + validate_prefix_cache_config(c) + if spec is not None: + if (spec.method not in ('mtp', 'dflash') + or (spec.method == 'mtp' and spec.num_speculative_tokens != 1) + or (spec.method == 'dflash' and spec.num_speculative_tokens not in (3, 7, 15)) + or c.cache_config.mamba_cache_mode not in ('align', 'none')): + raise ValueError("G256 speculation supports MTP1 or DFlash block4/8/16; prefix reuse is restricted to DFlash7") + from .gdn_spec import install + install() + if spec.method == 'dflash': + from .dflash_spec import install as install_dflash + install_dflash() + p = c.parallel_config + if (not isinstance(c.quant_config, OrnithG256Config) or c.model_config.dtype != torch.bfloat16 + or any(getattr(p, key) != 1 for key in + ('tensor_parallel_size', 'pipeline_parallel_size', 'data_parallel_size', 'prefill_context_parallel_size')) + or p.enable_expert_parallel or p.enable_dbo or p.use_sequence_parallel_moe + or c.use_v2_model_runner or c.scheduler_config.async_scheduling + or c.scheduler_config.max_num_seqs > 8 or c.scheduler_config.max_num_batched_tokens > 2048 + or c.lora_config is not None + or c.model_config.runner_type != 'generate' or c.model_config.enable_sleep_mode): + raise ValueError("G256 prototype requires BF16, TP1 V1 generation, synchronous scheduling, <=8 sequences/2048 tokens") + if self.settings.get('dynamic_spec_profile'): + from .dynamic_graphs import install + install() + self._ornith_binding = self._ornith_lifecycle = self._ornith_pending = None + VllmGPUWorker.__init__(self, vllm_config=vllm_config, local_rank=local_rank, rank=rank, + distributed_init_method=distributed_init_method, is_driver_worker=is_driver_worker) + + def load_model(self, *, load_dummy_weights=False): + from .dense_source import install as install_dense_source + install_dense_source() + from .dense_n32 import install as install_dense_n32 + install_dense_n32() + from .gdn_compact import install + install() + from .phase_dispatch import install + install() + if load_dummy_weights or self._ornith_binding is not None: + raise RuntimeError("G256 prototype loads the complete checkpoint once") + from .dflash_conv_boundary import install as install_conv_boundaries + install_conv_boundaries() + VllmGPUWorker.load_model(self, load_dummy_weights=False) + spec = self.vllm_config.speculative_config + if spec is not None and spec.method == 'dflash': + # Installed AOT selector can specialize on batch1 and then + # reject real batch4/8. Keep only this small selector eager; + # target verification and the draft transformer retain graphs. + draft = self.model_runner.drafter.model + if hasattr(draft, 'unwrap'): + draft = draft.unwrap() + draft.model.candidate_selector.do_not_compile = True + model = self.model_runner.get_model() + model.eval() + before = torch.cuda.memory_allocated(self.device) + self._ornith_binding = Runtime(model, self.settings, + self.vllm_config.scheduler_config.max_num_batched_tokens, self.device, + max_model_len=self.vllm_config.model_config.max_model_len, + max_num_seqs=self.vllm_config.scheduler_config.max_num_seqs) + if self.settings.get('compact_prefill', False): + from .attention_compact import configure + configure(max_num_seqs=self.vllm_config.scheduler_config.max_num_seqs, + max_model_len=self.vllm_config.model_config.max_model_len, + max_num_batched_tokens=self.vllm_config.scheduler_config.max_num_batched_tokens, + device=self.device, + iu4_prefill_library=self.settings.get('iu4_prefill_library')) + self._ornith_resident_bytes = torch.cuda.memory_allocated(self.device)-before + self.model_runner.model_memory_usage += self._ornith_resident_bytes + self._ornith_lifecycle = WorkerStepLifecycle( + self._ornith_binding, max_inflight=self.vllm_config.max_concurrent_batches+1) + + def get_kv_cache_spec(self): + from .prefix_cache import retained_draft_specs + return retained_draft_specs(self.vllm_config, super().get_kv_cache_spec()) + + def execute_model(self, scheduler_output): + # In this installed V1 runner normal sample logits run in execute_model. + # Prompt logits run later in sample_tokens, after the inherited snapshot. + for request in scheduler_output.scheduled_new_reqs: + sampling = request.sampling_params + if sampling is not None and sampling.prompt_logprobs is not None: + raise ValueError("ornith_g256 prototype does not support prompt_logprobs") + if not self._spec_enabled: + return super().execute_model(scheduler_output) + lifecycle = self._require_lifecycle() + if self._ornith_pending is not None: + raise RuntimeError('sample_tokens must finish the preceding speculative step') + ticket = lifecycle.begin('execute-model-and-draft') + try: + output = VllmGPUWorker.execute_model(self, scheduler_output) + except Exception as error: + lifecycle.abort(ticket, error) + if output is None: + # Proposal generation, including the shared native W8 head, runs + # inside sample_tokens. Snapshot only after those launches. + self._ornith_pending = ticket + return None + lifecycle.queue(ticket) + return lifecycle.checked_output(output, ticket) + + def compile_or_warm_up_model(self): + if __import__('os').environ.get('ORNITH_PERSISTENT_IU4_LIBRARY'): + from .attention_iu4_persistent import bind + added = bind(self.model_runner) + self._ornith_resident_bytes += added + self.model_runner.model_memory_usage += added + if not self._spec_enabled: + return super().compile_or_warm_up_model() + # Exploratory capture uses vLLM's synthetic GDN inputs/state. Their + # output values are discarded; unlike real requests they are not a + # finite-input contract. Preserve other faults and log the observed + # numerical flags. Every real execute/sample still resets and checks + # the complete flag bank through the ordinary lifecycle. + lifecycle = self._require_lifecycle() + ticket = lifecycle.begin('speculative-dummy-warmup') + try: + result = VllmGPUWorker.compile_or_warm_up_model(self) + flags = self._ornith_binding.error_flags.cpu().tolist() + observed = {} + from .method import G256LinearMethod, G256MoEMethod + for i, (name, code) in enumerate(zip(self._ornith_binding.layer_prefixes, flags)): + method = (self._ornith_binding.layers[i].quant_method + if i < len(self._ornith_binding.layers) else None) + # Larger DFlash capture batches use the BF16 prefill path, + # which propagates undefined dummy GDN values downstream. + # Dense bits 1/2/4/8 and routed bits 1/2/8/16/32 are + # numerical; routed bit 4 (invalid expert ID) stays fatal. + numeric_mask = (15 if type(method) is G256LinearMethod else + 59 if isinstance(method, G256MoEMethod) else 0) + if code and code & ~numeric_mask == 0: + observed[name] = code + self._ornith_binding.error_flags[i] = 0 + if observed: + from vllm.logger import init_logger + init_logger(__name__).info( + 'Exploratory synthetic warmup numerical flags (real-request checks remain enabled): %s', observed) + lifecycle.queue(ticket) + lifecycle.complete(ticket) + return result + except Exception as error: + lifecycle.abort(ticket, error) + + def sample_tokens(self, grammar_output): + if not self._spec_enabled: + return super().sample_tokens(grammar_output) + lifecycle = self._require_lifecycle() + ticket = self._ornith_pending + if ticket is None: + return VllmGPUWorker.sample_tokens(self, grammar_output) + self._ornith_pending = None + try: + output = VllmGPUWorker.sample_tokens(self, grammar_output) + except Exception as error: + lifecycle.abort(ticket, error) + lifecycle.queue(ticket) + return lifecycle.checked_output(output, ticket) + + def ornith_persistent_snapshot(self): + from .attention_iu4_persistent import snapshot + return snapshot() + + def ornith_served_snapshot(self): + self._ornith_lifecycle.drain() + binding = self._ornith_binding + return dict(flags=binding.check_error_flags(), arena_bytes=binding.workspace.numel(), + dense_shapes=binding.shapes, native_layers=len(binding.layer_prefixes), + resets=self._ornith_lifecycle.resets, copies=self._ornith_lifecycle.copies, + completions=self._ornith_lifecycle.completions) diff --git a/bundle/plugin-site/ornith_g256/worker_base.py b/bundle/plugin-site/ornith_g256/worker_base.py new file mode 100644 index 0000000000000000000000000000000000000000..631b257056d278375dda0f06a7c3c58573797f97 --- /dev/null +++ b/bundle/plugin-site/ornith_g256/worker_base.py @@ -0,0 +1,111 @@ +"""Project-owned worker lifecycle, retained from Ciru's tilebank adapter. + +Only inherited methods used by the G256 worker are included. Model admission +and native binding belong to OrnithG256Worker. +""" +# Copyright 2026 Ciru. +from vllm.v1.worker.gpu_worker import Worker as VllmGPUWorker + + +class OrnithWorkerBase(VllmGPUWorker): + def _require_lifecycle(self): + if self._ornith_lifecycle is None: + raise RuntimeError('Ornith worker model must load and bind before execution') + self._ornith_lifecycle.raise_if_failed() + return self._ornith_lifecycle + + def _startup_region(self, name, operation): + lifecycle = self._require_lifecycle() + if self._ornith_pending is not None: + raise RuntimeError('Startup/profile cannot overlap an unfinished worker step') + ticket = lifecycle.begin(name) + try: + result = operation() + except Exception as error: + lifecycle.abort(ticket, error) + lifecycle.queue(ticket) + lifecycle.complete(ticket) + return result + + def determine_available_memory(self): + return self._startup_region('profile-memory', super().determine_available_memory) + + def compile_or_warm_up_model(self): + return self._startup_region('warmup-and-capture', super().compile_or_warm_up_model) + + def execute_model(self, scheduler_output): + lifecycle = self._require_lifecycle() + if self._ornith_pending is not None: + raise RuntimeError('sample_tokens must complete the preceding Ornith execute_model step') + ticket = lifecycle.begin('execute-model') + try: + output = super().execute_model(scheduler_output) + except Exception as error: + lifecycle.abort(ticket, error) + # This D2H precedes the ordinary sampler's CPU completion event. No + # device/error data is read here, and flags are not reset by sampling. + lifecycle.queue(ticket) + if output is None: + self._ornith_pending = ticket + return None + return lifecycle.checked_output(output, ticket) + + def sample_tokens(self, grammar_output): + lifecycle = self._require_lifecycle() + ticket = self._ornith_pending + if ticket is None: + return super().sample_tokens(grammar_output) + self._ornith_pending = None + try: + output = super().sample_tokens(grammar_output) + except Exception as error: + lifecycle.abort(ticket, error) + return lifecycle.checked_output(output, ticket) + + def sleep(self, level=1): + raise RuntimeError('Ornith worker sleep/reload is unsupported while native graphs own the arena') + + def wake_up(self, tags=None): + raise RuntimeError('Ornith worker wake-up requires a fresh worker/model load') + + def reload_weights(self, *args, **kwargs): + raise RuntimeError('Reloading sealed Ornith weights requires a fresh worker') + + def update_config(self, overrides): + raise RuntimeError('Changing the sealed Ornith runtime requires a fresh worker') + + def update_max_model_len(self, max_model_len): + raise RuntimeError('Changing Ornith capacity requires a fresh worker') + + def start_weight_update(self): + raise RuntimeError('Ornith native weight updates require a fresh worker') + + def start_draft_weight_update(self): + raise RuntimeError('Ornith draft weight updates are not admitted') + + def update_weights(self, update_info): + raise RuntimeError('Ornith native weight updates require a fresh worker') + + def finish_weight_update(self): + raise RuntimeError('Ornith native weight updates require a fresh worker') + + def shutdown(self): + error = None + try: + if self._ornith_lifecycle is not None: + self._ornith_lifecycle.drain() + except Exception as caught: + error = caught + try: + # The pinned parent first synchronizes and destroys captured model + # graphs. The worker binding still strongly owns arena/library here. + super().shutdown() + except BaseException: + # If parent graph destruction fails, retain native storage/library + # until process teardown; releasing it here could dangle graph args. + raise + self._ornith_pending = self._ornith_lifecycle = self._ornith_binding = None + if hasattr(self.vllm_config, '_ornith_grouped_runtime'): + delattr(self.vllm_config, '_ornith_grouped_runtime') + if error is not None: + raise error diff --git a/bundle/serve.sh b/bundle/serve.sh new file mode 100644 index 0000000000000000000000000000000000000000..02a7efabec6d5c30761170b5c65d715f294bc6f8 --- /dev/null +++ b/bundle/serve.sh @@ -0,0 +1,18 @@ +#!/usr/bin/env bash +# Copyright 2026 Ciru. Ornith1.5 Ciru Halo Agent production profile. +set -euo pipefail +bundle_root=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +source "$bundle_root/paths.env" +export ORNITH_C1_POLICY=auto +export ORNITH_PERSISTENT_IU4_LIBRARY="$bundle_root/native/libornith_persistent_iu4.so" +export ORNITH_PERSISTENT_IU4_MIN_SEQ=32768 +export CC="${CC:-$(command -v gcc)}" CXX="${CXX:-$(command -v g++)}" +exec bash "$bundle_root/packaging/serve.sh" \ + --runtime-root "$ORNITH_RUNTIME_ROOT" --plugin-site "$bundle_root/plugin-site" \ + --model "$ORNITH_MODEL" --draft "$ORNITH_DRAFT" \ + --native-library-directory "$bundle_root/native" --cache-directory "$bundle_root/cache" \ + --context 262144 --cache-gib 44 --max-seqs 8 --block-size 1120 \ + --enable-tools --compact-prefill --folded-decode --token-major-kv --routed-prefill-a4 \ + --draft-tokens 7 --prefix-cache --fine-prefix-cache \ + --iu4-prefill-library "$bundle_root/native/libornith_attention_iu4.so" \ + --served-name ciru-halo-agent "$@" diff --git a/runtime/INSTALL-ORNITH-RUNTIME.sh b/runtime/INSTALL-ORNITH-RUNTIME.sh new file mode 100644 index 0000000000000000000000000000000000000000..20d55270379bcbaccd573c4d6b0697df65ea0bf6 --- /dev/null +++ b/runtime/INSTALL-ORNITH-RUNTIME.sh @@ -0,0 +1,18 @@ +#!/usr/bin/env bash +set -euo pipefail +task_packages=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +task_root=${1:?usage: INSTALL-ORNITH-RUNTIME.sh NEW_RUNTIME_ROOT} +if [[ -e "$task_root" ]]; then + printf 'Refusing to overwrite existing runtime: %s\n' "$task_root" >&2 + exit 2 +fi +mkdir -p "$task_root" +task_root=$(cd -- "$task_root" && pwd) +# Reuse the qualified wheel installer without modifying its package pins. +export CIRU_RUNTIME_ROOT="$task_root/sources" +export CIRU_VLLM_VENV="$task_root/venv" +bash "$task_packages/INSTALL-RUNTIME.sh" +ln -s sources/vllm-glm53-strix "$task_root/vllm" +ln -s sources/aiter-gfx1151 "$task_root/aiter" +ln -s sources/vllm-glm53-strix/runtime-env.sh "$task_root/runtime-env.sh" +printf 'Ornith runtime prepared. Set ORNITH_RUNTIME_ROOT=%s\n' "$task_root" diff --git a/runtime/INSTALL-RUNTIME.sh b/runtime/INSTALL-RUNTIME.sh new file mode 100644 index 0000000000000000000000000000000000000000..4f1a5c5ade5006887d666941e80b1929dae5bde1 --- /dev/null +++ b/runtime/INSTALL-RUNTIME.sh @@ -0,0 +1,103 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +task_package_dir="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" +task_runtime_root="${CIRU_RUNTIME_ROOT:-/srv/llm/runtime}" +task_venv="${CIRU_VLLM_VENV:-/srv/llm/venvs/vllm-rocm10-gfx1151}" +task_vllm_source="$task_runtime_root/vllm-glm53-strix" +task_aiter_source="$task_runtime_root/aiter-gfx1151" +task_host_library_file="$task_vllm_source/ciru-host-library-path.txt" +task_host_library_path="${CIRU_HOST_LIBRARY_PATH:-}" + +if [[ "$task_host_library_path" == *$'\n'* || "$task_host_library_path" == *$'\r'* ]]; then + printf '%s\n' 'CIRU_HOST_LIBRARY_PATH must be a single colon-separated line.' >&2 + exit 2 +fi + +for task_command in uv tar install find; do + command -v "$task_command" >/dev/null || { + printf 'required command is missing: %s\n' "$task_command" >&2 + exit 2 + } +done + +for task_command in gcc g++ git pkg-config xxd; do + command -v "$task_command" >/dev/null || { + printf 'AITER JIT build command is missing: %s\n' "$task_command" >&2 + printf '%s\n' 'On NixOS, rerun inside a development shell providing gcc, pkg-config, xxd, libdrm, elfutils, numactl, openssl, and git.' >&2 + exit 2 + } +done + +if [[ ! -e /lib64/ld-linux-x86-64.so.2 ]]; then + printf '%s\n' 'The standard Linux ELF interpreter is unavailable.' >&2 + printf '%s\n' 'On NixOS, enable programs.nix-ld before installing the portable Python/ROCm wheel stack.' >&2 + exit 2 +fi + +mapfile -t task_vllm_archives < <(find "$task_package_dir" -maxdepth 1 -type f -name 'ciru-halo-agent-vllm-source.tar.gz' -print) +mapfile -t task_aiter_archives < <(find "$task_package_dir" -maxdepth 1 -type f -name 'ciru-halo-agent-aiter-source.tar.gz' -print) +mapfile -t task_vllm_wheels < <(find "$task_package_dir/wheels" -maxdepth 1 -type f -name 'vllm-*.whl' -print) +mapfile -t task_aiter_wheels < <(find "$task_package_dir/wheels" -maxdepth 1 -type f -name 'amd_aiter-*.whl' -print) + +test "${#task_vllm_archives[@]}" -eq 1 +test "${#task_aiter_archives[@]}" -eq 1 +test "${#task_vllm_wheels[@]}" -eq 1 +test "${#task_aiter_wheels[@]}" -eq 1 +test -f "$task_package_dir/requirements-runtime.lock" +test -f "$task_package_dir/runtime-env.sh" + +for task_target in "$task_vllm_source" "$task_aiter_source" "$task_venv"; do + if [[ -e "$task_target" ]]; then + printf 'refusing to overwrite existing path: %s\n' "$task_target" >&2 + exit 2 + fi +done + +mkdir -p "$task_runtime_root" "$(dirname -- "$task_venv")" +tar -xzf "${task_vllm_archives[0]}" -C "$task_runtime_root" +tar -xzf "${task_aiter_archives[0]}" -C "$task_runtime_root" +install -m 0644 "$task_package_dir/runtime-env.sh" "$task_vllm_source/runtime-env.sh" + +printf '%s\n' "$task_host_library_path" > "$task_host_library_file" +chmod 0644 "$task_host_library_file" + +uv python install 3.14.3 +uv venv --python 3.14.3 "$task_venv" + +task_amd_index=https://stable.repo.amd.com/rocm/whl-next/ +uv pip install --python "$task_venv/bin/python" \ + --index-url "$task_amd_index" \ + 'rocm[libraries,devel,device-gfx1151]==10.0.0' \ + 'torch[device-gfx1151]==2.13.0+rocm10.0.0' \ + 'torchvision[device-gfx1151]==0.28.0+rocm10.0.0' \ + 'torchaudio==2.11.0.2+rocm10.0.0' + +# Materialize the SDK compiler/library layout used by the fused ROCm kernels. +"$task_venv/bin/python" -m rocm_sdk init + +uv pip install --python "$task_venv/bin/python" \ + --extra-index-url "$task_amd_index" \ + --index-strategy unsafe-best-match \ + -r "$task_package_dir/requirements-runtime.lock" + +uv pip install --python "$task_venv/bin/python" --no-deps \ + "${task_aiter_wheels[0]}" "${task_vllm_wheels[0]}" + +task_site="$($task_venv/bin/python -c 'import site; print(site.getsitepackages()[0])')" +uv pip install --python "$task_venv/bin/python" \ + "$task_site/_rocm_sdk_core/share/amd_smi" + +if [[ -f "$task_package_dir/aiter-jit-gfx1151/module_aiter_core.so" ]]; then + install -D -m 0755 \ + "$task_package_dir/aiter-jit-gfx1151/module_aiter_core.so" \ + "$task_venv/var/aiter-jit-gfx1151/module_aiter_core.so" +fi + +printf '%s\n' 'Runtime engine installed.' +printf 'VLLM_SOURCE=%s\n' "$task_vllm_source" +printf 'AITER_SOURCE=%s\n' "$task_aiter_source" +printf 'VLLM_VENV=%s\n' "$task_venv" +printf 'VLLM_RUNTIME_ENV=%s\n' "$task_vllm_source/runtime-env.sh" +printf 'CIRU_HOST_LIBRARY_PATH_FILE=%s\n' "$task_host_library_file" +printf '%s\n' 'Source runtime-env.sh from the node launcher before starting vLLM.' diff --git a/runtime/README.md b/runtime/README.md new file mode 100644 index 0000000000000000000000000000000000000000..7038a27addb6bb0591b295ff71e61f6d8ba7a896 --- /dev/null +++ b/runtime/README.md @@ -0,0 +1,26 @@ +# Ornith native runtime package + +This package reuses the retained GLM runtime engine installer and exact local +vLLM/AITER wheels. It includes their source archives and retained licenses; +it contains no GLM or Ornith model weights, NHI services or model selectors. + +On a compatible Linux host with uv, GCC/G++, git, pkg-config and xxd installed: + + bash INSTALL-ORNITH-RUNTIME.sh /path/to/new/ornith-runtime + +The installer downloads pinned Python/AMD/runtime dependencies. No installation +or download occurs while exporting this package. The existing wheel package +documents Ubuntu24.04+ or an equivalent ABI; other distributions remain subject +to actual installation validation. On NixOS, use the included shell.nix and +nixos-module.nix prerequisites to supply nix-ld and the host library path. + +The wrapper adapts the existing install to ORNITH_RUNTIME_ROOT. It never starts +a model or changes services. Point the Ornith model bundle at that runtime root. +The Ornith plugin, weights, draft and seven model-specific native libraries are +separate release assets. Existing native RUNPATHs are resolved through the +runtime's LD_LIBRARY_PATH; their presence alone is not a deployment failure. + +UPSTREAM-RUNTIME.md preserves the runtime's source identities, package contract +and license provenance. Its GLM model/workload instructions are historical +context and do not configure Ornith. The source archives retain vLLM Apache-2.0, +AITER MIT and Composable Kernel MIT notices. Retain downloaded component notices. diff --git a/runtime/UPSTREAM-RUNTIME.md b/runtime/UPSTREAM-RUNTIME.md new file mode 100644 index 0000000000000000000000000000000000000000..07d3b09109f218ff59f76fb5868b192d531e1866 --- /dev/null +++ b/runtime/UPSTREAM-RUNTIME.md @@ -0,0 +1,271 @@ +# GLM5.3 Flash CIRU STRIX IU4 runtime packages + +This directory is the self-contained runtime-engine deliverable. +The 128K default (131,200 tokens, 12 GiB KV and 8 GiB host staging per rank) passed the +bounded 131,000-token generation/restart/reuse gate on packaged dev9. +The earlier minimal packaged 64K NHI functional gate passed, including +useful disk-prefix reuse on dev8. This is not a new full +performance or quality qualification; see the +[versioned gate report](../../benchmarks/packaged-nhi-gate.md). +It does not depend on a GitHub repository. `INSTALL-RUNTIME.sh` installs the +documented layout: + +- `/srv/llm/runtime/vllm-glm53-strix` +- `/srv/llm/runtime/aiter-gfx1151` +- `/srv/llm/venvs/vllm-rocm10-gfx1151` +- `/srv/llm/runtime/vllm-glm53-strix/runtime-env.sh` + +The same runtime environment supplies FastAPI, httpx, and Uvicorn for the +Apache-2.0 mirrored generation frontend under +`../generation-frontend/`. External-launcher TP2 generation must enter through +that frontend (port 8083 by default), which submits an identical request to +both rank-local port 8100 APIs, returns rank 0, and drains rank 1. The frontend +may run on either node and has no CiruStrixLink application dependency. + +The current candidate identities are: + +- vLLM source base: `9255fd9fb9fedf4b29d574a8d8bb21d93892cc98`, plus the + uncommitted Python-only store/load group projection, Mamba cached-state + indexing, and paired prefix admission fixes; this is not a new dev9 commit. +- vLLM runtime version: `0.1.0rc2.dev9+g9255fd9fb9.rocm100` +- vLLM wheel: + `vllm-0.1.0rc2.dev9+g9255fd9fb9.rocm100-cp314-cp314-linux_x86_64.whl` +- Patched vLLM source archive: + `GLM5.3-Flash-CIRU-STRIX-IU4-vllm-source-v0.1.0-rc2-dev9.tar.gz` +- AITER: `ec6b1a5d0bdbc9d43f9375dcc1516f87efac1c57` + (identity-only child of `2fe3fff6adb4ecbdc0aa2b9c2ce9618b5e6b3046`) +- AITER Composable Kernel: `fdf4bb7fcc984811cef48ce817d89aac064b984a` +- Python: 3.14.3 +- ROCm: 10.0.0, gfx1151 +- PyTorch: 2.13.0+rocm10.0.0 +- TileLang: 0.1.10, with Apache TVM FFI 0.1.10 (required for fused GLM mHC) + +TileLang is a production dependency, not optional profiling tooling. Omitting it +silently sends this model's mHC operations through vLLM's unfused Torch fallback. +The completed requirements repair pins TileLang and its prerequisites. With +NHI, DFlash2 k7, the 64K profile, prefix caching enabled, and the default +2,304-token batch budget, the packaged 2,048-prompt / 128-output check measured +402.455 prompt tokens/s, 5.089 s TTFT, 23.686 decode tokens/s (excluding prefill), +and 56.593% draft acceptance. The production pair was restored to that +configuration for that check; these are recorded 64K results, not qualification +of the newly requested 128K default. + +A matched cache-off comparison used a 20,499-token prompt plus 128 generated +tokens: the 2,304-token budget gave 395.305 prompt tokens/s and 51.856 s TTFT; +8,192 gave 377.484 prompt tokens/s and 54.304 s TTFT. The larger budget did not +help this workload and is not endorsed as the production default. Earlier +cache-off recovery checks with the 8,192 budget measured 425 prompt tokens/s +at 2,048 prompt tokens and 391 at 8,119 prompt tokens. + +A 20,480-token batch budget did not pass startup with the 64K/6 GiB KV profile: +the engine required 10.35 GiB of KV allocation, so there is no speed result +for that budget. The disk-prefix-qualified default remains 2,304; +`MAX_NUM_BATCHED_TOKENS` permits explicit experimental overrides. These are +bounded performance checks, not a full context sweep or broad acceptance or +quality requalification. The separate bounded 128K generation/restart/reuse +gate passed; 64K remains the lower-memory fallback and 256K is experimental. + +The model launcher defaults to `CONTEXT_PROFILE=128k`. Fresh user/NHI service +installations select context profile 2; existing selections are preserved. +Use `glm53-context 2` or `sudo glm53-nhi-context USER 2` on both nodes to move +an existing installation to that default for its next paired start. Profile 1 +retains 64K/6 GiB KV, and profile 3 is unvalidated 256K/24 GiB KV; all use the +8 GiB host tier and the default 2,304-token batch budget. + +The production launcher defaults `VLLM_NHI_TIMEOUT_MS` to 30,000 so a rank's +first-use JIT compilation does not trigger the former one-second peer timeout. +This is a failure deadline, not an added delay: successful exchanges return +immediately. + +The vLLM and AITER wheels are built on Ubuntu 24.04 (glibc 2.39, GCC 13) and +are tagged `cp314-cp314-linux_x86_64`. Their compiled extensions are qualified +without Nix-store, `/srv/llm/work`, or other host-specific RPATH/RUNPATH +entries. They are not manylinux wheels and should be treated as Ubuntu 24.04+ +or equivalently compatible x86-64 artifacts. + +## NixOS installation + +Import the included module from your NixOS configuration, then rebuild the +host once: + +```nix +{ + imports = [ /path/to/runtime/packages/nixos-module.nix ]; +} +``` + +```bash +sudo nixos-rebuild switch +``` + +The module enables `programs.nix-ld` for the portable Python/ROCm wheels and +publishes `CIRU_HOST_LIBRARY_PATH` for login sessions. Enter the included +development shell and run the installer from this directory: + +```bash +cd /path/to/GLM5.3-Flash-CIRU-STRIX-IU4/runtime/packages +nix-shell ./shell.nix +bash ./INSTALL-RUNTIME.sh +exit +``` + +`shell.nix` provides `uv`, GCC, CMake, Ninja, Make, Git, `pkg-config`, `xxd`, +and the standard archive/file commands. It also exports the same +`CIRU_HOST_LIBRARY_PATH` used by `runtime-env.sh`. No Nix-store path is +embedded in the packaged wheels or shared objects. + +At install time, `INSTALL-RUNTIME.sh` writes the shell's evaluated path as one +inert line in +`/srv/llm/runtime/vllm-glm53-strix/ciru-host-library-path.txt`. +`runtime-env.sh` reads that data file only when `CIRU_HOST_LIBRARY_PATH` is +absent; it never sources the file as shell code. User and privileged NHI +services therefore do not depend on a login manager or `nix-shell` +environment inheriting the variable. The generated file is host-specific +install state and is not embedded in any packaged wheel or shared object. + +After installation, the release root's `run-node.sh` automatically enters this +packaged `shell.nix` when `nix-shell` is available. This supplies Triton with +the NixOS compiler wrappers and host development headers it needs during +initialization. Other Linux hosts launch directly with their installed +toolchain. + +The source archives have no `.git` directory and retain the upstream license +and notice files. The dev9 vLLM archive retains dev8's two matching +projections: `_build_partial_tail_store_jobs` remaps physical KV +group IDs before aligned/partial stores, and `update_state_after_alloc` +selects the same transferable groups from physical block allocations before +loading. The second change fixes dev7's load-allocation assertion after a +successful cache lookup. Dev9 additionally indexes cached recurrent states +using the Mamba group's 2,304-token block size, not the global 64-token block +size (source slot 55 rather than 2,015 in the observed failure), and coordinates +capacity-bounded prefix admission across the mirrored external-launcher ranks. +The final positive cached-read gate passed: both ranks admitted 129,024 tokens +and loaded 1,600,792,576 cache bytes onto GPU. The earlier repair-write request +completed without GPU cache hits and was not counted as the read gate. +Native extensions are unchanged. +The archive's `ciru-release/DEV9_CACHE_PATCH.md` records the Python delta. +The candidate launcher caps prefill at 2,304 batched tokens to materialize the +KDA checkpoint. The AITER archive contains +the exact Composable Kernel submodule contents. The separately packaged +seven-kernel source archive excludes prior build output and shared objects. + +`requirements-runtime.lock` pins only the direct production requirements for +this local HTTP serving profile, mirrored frontend, and AITER JIT. The +installer resolves their transitive dependencies. ROCm/PyTorch is installed +as a separate exact stack; +cloud SDKs, telemetry exporters, profiling/benchmark tools, docs/tests, +training packages, and unused alternate kernel/model-streaming stacks are not +part of the public runtime resolution. + +The live isolated install identified and fixed two omissions in the original +package: `runtime-env.sh` now accepts ROCm 10's `_rocm_sdk_core` plus +`_rocm_sdk_libraries` wheel layout, and the production lock includes +`model-hosting-container-standards==0.1.16` as required by vLLM. + +The installer also runs `python -m rocm_sdk init` after installing the AMD stack. +This expands the already-installed development bundle and its device links into +`_rocm_sdk_devel`, providing the compiler/library layout used by the fused +kernels. It is installation work, not a model startup or a per-request check. + +## Privileged NHI runtime staging + +The `CAP_SYS_RAWIO`-bearing system service requires a separate root-owned +runtime, conventionally `/opt/ciru/glm53-iu4`. Copy the venv, vLLM source, +AITER source, runtime environment, and gfx1151 libraries there before using +`install-nhi-system-service.sh`. Copy the complete CPython 3.14.3 runtime too: +copying a UV-created venv alone leaves its interpreter symlink pointing into +the original user's home directory. + +With CPython staged at `/opt/ciru/glm53-iu4/python`, the required relationships +are: + +- `venv/bin/python` resolves to + `/opt/ciru/glm53-iu4/python/bin/python3.14`; the `python3` and `python3.14` + aliases must resolve into that same root-owned runtime. +- `venv/pyvenv.cfg` contains + `home = /opt/ciru/glm53-iu4/python/bin`. +- The full staged tree and its symlinks are owned by `root:root`, with no + group/other write permission. Directories must be traversable and files + readable by the service user; preserve executable bits on programs. The + staging repair uses `chmod -R u=rwX,go=rX` on this exact runtime tree, not on + model data or user home directories. + +Do not hardlink a user-owned source tree into this privileged tree: changing +ownership would also change the source inode. Use copies or reflinks. The +packaged launcher invokes the staged Python as +`python -m torch.distributed.run`, avoiding stale copied `torchrun` shebangs. +This runtime relocation is a separate operator step; the public NHI service +installer validates the paths but does not automate the Python copy/relink. + +`runtime-env.sh` defaults `AITER_JIT_DIR` to the service user's writable +`${XDG_CACHE_HOME:-$HOME/.cache}/GLM5.3-Flash-CIRU-STRIX-IU4/aiter` directory. +When that directory lacks `module_aiter_core.so`, it seeds the module from +`$VLLM_VENV/var/aiter-jit-gfx1151/module_aiter_core.so` if the packaged module +is present. Existing JIT modules are never overwritten. This keeps the +privileged runtime tree immutable while avoiding an unnecessary AITER core +C++ rebuild on the first use of a new cache. An explicit `AITER_JIT_DIR` +continues to override the default. + +### NixOS host compiler for the NHI service + +The capability-bearing system service does not enter the user launcher's +`nix-shell`. Its host-side Triton launcher builds require the Nix compiler +wrappers and their standard headers; bare ROCm clang is not a replacement. +Add this Nix-managed systemd drop-in to the host configuration on both nodes, +then apply the NixOS configuration before starting the pair: + +```nix +{ pkgs, ... }: { + environment.etc."systemd/system.attached/GLM5.3-Flash-CIRU-STRIX-IU4-nhi@.service.d/10-nixos-toolchain.conf".text = '' + [Service] + Environment="CC=${pkgs.stdenv.cc}/bin/cc" + Environment="CXX=${pkgs.stdenv.cc}/bin/c++" + ''; +} +``` + +`runtime-env.sh` preserves these explicit `CC`/`CXX` values. The drop-in is +host configuration, not an embedded build-host path in the public wheel. + +### Host-staging scratch versus persistent prefix storage + +The packaged NHI unit sets +`TemporaryFileSystem=/dev/shm:rw,nosuid,nodev,mode=1777,size=16G`. Each service +therefore owns an isolated tmpfs for `vllm_offload_.mmap` and related +runtime scratch; that mount disappears when the service stops. The 16 GiB +value is a limit, not a preallocated reservation. This avoids a stopped +instance's 8 GiB staging file remaining alongside the next instance's file. + +The disk prefix cache under `PREFIX_CACHE_ROOT` is separate local-NVMe data +and survives this scratch teardown. The 8 GiB tier passed earlier 64K restart/reuse: +dev8 admitted 2,304 cached tokens and loaded 127,419,136 bytes to GPU on each +rank after reading the compatible dev7-written cache. No allocation failures +were observed. The one-token replay is a cache-functional check, not a TG or +quality benchmark, and does not establish a guaranteed free-memory reserve. + +The later 128K dev9 gate replayed the same 131,000-token request after restart +and admitted 129,024 cached tokens on both ranks, with 7.568 s TTFT and actual +GPU loading. The 8 GiB setting still limits RAM staging only: the filesystem +tier has no automatic eviction or byte quota, and the long-prompt test required +substantial NVMe headroom. See the [gate report](../../benchmarks/packaged-nhi-gate.md) +and [storage guidance](../../README.md#persistent-disk-prefix-cache). + +## Kernel source licensing + +`GLM5.3-Flash-CIRU-STRIX-IU4-kernels-source-v0.1.0-rc1.tar.gz` is the clean +source snapshot for the seven model-specific native libraries. The release +owner approved Apache-2.0 for the six Ciru-authored groups: `dense-kda`, +`m4-residual`, `m8-align`, `resident-g128`, `top8-epilogue`, and `iu4-m1`. +Their sources and Ciru-authored builders carry SPDX/copyright headers, and the +archive includes the full Apache-2.0 `LICENSE` plus a concise provenance +`NOTICE`. A provenance audit found no outside permission blocker in those +groups. + +The `nhi-m8-bf16` sources are not relicensed. The bridge retains its MIT SPDX +identifiers, and the bundled USB4STREAM Linux UAPI header retains +`GPL-2.0 WITH Linux-syscall-note`; those files are byte-identical to the +previous qualified source snapshot. + +This runtime-engine bundle alone is not the complete model payload. The two +rank payloads and seven qualified model-specific shared libraries must also be +present before an end-to-end model launch is possible. diff --git a/runtime/aiter-jit-gfx1151/module_aiter_core.so b/runtime/aiter-jit-gfx1151/module_aiter_core.so new file mode 100644 index 0000000000000000000000000000000000000000..627cdbe891935f4444faa3900a2fef7fdbf93969 --- /dev/null +++ b/runtime/aiter-jit-gfx1151/module_aiter_core.so @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3580569033331d904b3ca2b899e3bea00324d098dc5a28c3c8f76a124f24f846 +size 567024 diff --git a/runtime/ciru-halo-agent-aiter-source.tar.gz b/runtime/ciru-halo-agent-aiter-source.tar.gz new file mode 100644 index 0000000000000000000000000000000000000000..b70d0c79213ac02fbaebe64a95066bef104716a3 --- /dev/null +++ b/runtime/ciru-halo-agent-aiter-source.tar.gz @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:986e7d734d4c32e70788ad0052a4c48196487b22db1bb9de1b220e555bb3de99 +size 55734195 diff --git a/runtime/ciru-halo-agent-vllm-source.tar.gz b/runtime/ciru-halo-agent-vllm-source.tar.gz new file mode 100644 index 0000000000000000000000000000000000000000..1fdd7cb4a1a86f405a63916511b21fd7c185270d --- /dev/null +++ b/runtime/ciru-halo-agent-vllm-source.tar.gz @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c9b6c30d9690591a0e01ff05fcbf2b3456e9207b92bf4268a4c62a294da965bb +size 40520111 diff --git a/runtime/nixos-module.nix b/runtime/nixos-module.nix new file mode 100644 index 0000000000000000000000000000000000000000..ba06133c7af35f5155dc87e00f24b1e68a75c241 --- /dev/null +++ b/runtime/nixos-module.nix @@ -0,0 +1,22 @@ +{ lib, pkgs, ... }: + +let + ciruHostLibraries = with pkgs; [ + stdenv.cc.cc.lib + zlib + vulkan-loader + libdrm + elfutils + numactl + openssl + ]; +in +{ + programs.nix-ld = { + enable = true; + libraries = ciruHostLibraries; + }; + + environment.sessionVariables.CIRU_HOST_LIBRARY_PATH = + lib.makeLibraryPath ciruHostLibraries; +} diff --git a/runtime/requirements-runtime.lock b/runtime/requirements-runtime.lock new file mode 100644 index 0000000000000000000000000000000000000000..c7e6970b30b446e4ea9b007aa2c0ad150aeb017c --- /dev/null +++ b/runtime/requirements-runtime.lock @@ -0,0 +1,90 @@ +# GLM5.3 Flash CIRU STRIX IU4 production runtime requirements. +# +# This is the qualified direct production resolution for the local HTTP model +# server and AITER JIT. Transitive dependencies are resolved by uv. The ROCm, +# PyTorch, torchvision, and torchaudio stack is installed separately by +# INSTALL-RUNTIME.sh from AMD's ROCm 10 index. +# +# Deliberately excluded: cloud storage/provider SDKs, telemetry exporters, +# profiling/benchmark packages, docs/test tooling, training/fine-tuning tools, +# and optional alternate kernel/model-streaming stacks not used by this model. + +# AITER build/JIT runtime +einops==0.8.2 +flydsl==0.1.4 +ninja==1.13.0 +packaging==26.3 +pandas==3.0.5 +psutil==7.2.2 +pybind11==3.1.0 + +# Required GLM mHC fusion; without TileLang, vLLM silently uses unfused Torch. +tilelang==0.1.10 +apache-tvm-ffi==0.1.10 +ml-dtypes==0.6.0 +z3-solver==4.15.4.0 + +# vLLM model/tokenizer/runtime core +blake3==1.0.9 +cachetools==7.1.7 +cloudpickle==3.1.2 +compressed-tensors==0.17.0 +depyf==0.20.0 +filelock==3.32.3 +huggingface-hub==1.29.0 +numba==0.65.0 +numpy==2.4.6 +opencv-python-headless==5.0.0.93 +pillow==12.3.0 +protobuf==6.33.6 +py-cpuinfo==9.0.0 +regex==2026.7.19 +requests==2.34.2 +safetensors==0.8.0 +sentencepiece==0.2.2 +setuptools==79.0.1 +six==1.17.0 +tokenizers==0.23.1 +tqdm==4.70.0 +transformers==5.16.1 +typing-extensions==4.16.0 + +# HTTP/OpenAI-compatible serving +aiohttp==3.14.3 +email-validator==2.3.0 +fastapi==0.136.3 +httptools==0.8.0 +httpx==0.28.1 +model-hosting-container-standards==0.1.16 +openai==3.6.0 +prometheus-client==0.26.0 +prometheus-fastapi-instrumentator==8.1.0 +pydantic==2.13.5 +python-dotenv==1.2.3 +python-json-logger==4.2.0 +python-multipart==0.0.32 +pyyaml==6.0.3 +setproctitle==1.3.7 +starlette==1.6.0 +uvicorn==0.52.4 +uvloop==0.22.1 +watchfiles==1.2.0 +websockets==17.1 + +# Chat templates, structured output, and tool-call parsing +cbor2==6.1.4 +ijson==3.5.1 +jsonschema==4.26.0 +lark==1.2.2 +llguidance==1.7.6 +lm-format-enforcer==0.11.3 +mcp==2.1.1 +mistral-common==1.11.7 +msgspec==0.21.1 +openai-harmony==0.0.8 +outlines-core==0.2.14 +partial-json-parser==0.2.1.1.post7 +pybase64==1.5.0 +pyzmq==27.2.0 +tiktoken==0.14.0 +xgrammar==0.2.3 diff --git a/runtime/runtime-env.sh b/runtime/runtime-env.sh new file mode 100644 index 0000000000000000000000000000000000000000..b855601015fec11c1614d9f8e553adda96e6d236 --- /dev/null +++ b/runtime/runtime-env.sh @@ -0,0 +1,124 @@ +#!/usr/bin/env bash +# Source-only environment for the wheel-installed GLM5.3 Strix runtime. + +if [[ "${BASH_SOURCE[0]}" == "$0" ]]; then + printf 'source this file; do not execute it directly\n' >&2 + exit 2 +fi + +: "${VLLM_SOURCE:?set VLLM_SOURCE to /srv/llm/runtime/vllm-glm53-strix}" +: "${VLLM_VENV:?set VLLM_VENV to /srv/llm/venvs/vllm-rocm10-gfx1151}" +: "${AITER_SOURCE:?set AITER_SOURCE to /srv/llm/runtime/aiter-gfx1151}" + +python_bin="$VLLM_VENV/bin/python" +[[ -x "$python_bin" ]] || { + printf 'Python is missing from VLLM_VENV: %s\n' "$python_bin" >&2 + return 2 +} + +site_packages="$($python_bin - <<'PY' +import site +print(site.getsitepackages()[0]) +PY +)" +rocm_core="${ROCM_CORE_ROOT:-$site_packages/_rocm_sdk_core}" +rocm_devel="${ROCM_DEVEL_ROOT:-$site_packages/_rocm_sdk_devel}" +rocm_libraries="${ROCM_LIBRARIES_ROOT:-$site_packages/_rocm_sdk_libraries}" +torch_lib="$site_packages/torch/lib" + +# ROCm 10 wheels may place the compiler and headers in _rocm_sdk_core +# instead of the older _rocm_sdk_devel split. +if [[ -x "$rocm_devel/bin/hipcc" ]]; then + rocm_tool_root="$rocm_devel" +else + rocm_tool_root="$rocm_core" +fi +if [[ -x "$rocm_tool_root/lib/llvm/bin/clang" ]]; then + rocm_llvm_bin="$rocm_tool_root/lib/llvm/bin" +else + rocm_llvm_bin="$rocm_tool_root/llvm/bin" +fi + +for required in \ + "$VLLM_SOURCE/ciru-release/SOURCE_FILES.txt" \ + "$AITER_SOURCE/ciru-release/SOURCE_FILES.txt" \ + "$site_packages/vllm/__init__.py" \ + "$site_packages/aiter/__init__.py" \ + "$rocm_tool_root/bin/hipcc" \ + "$rocm_llvm_bin/clang" \ + "$rocm_llvm_bin/clang++" \ + "$rocm_tool_root/lib" \ + "$rocm_core/lib/llvm/amdgcn/bitcode" \ + "$rocm_core/share/amd_smi" \ + "$torch_lib"; do + [[ -e "$required" ]] || { + printf 'required runtime path is missing: %s\n' "$required" >&2 + return 2 + } +done + +host_library_path="${CIRU_HOST_LIBRARY_PATH:-}" +if [[ -z "${CIRU_HOST_LIBRARY_PATH+x}" ]]; then + host_library_file="$VLLM_SOURCE/ciru-host-library-path.txt" + if [[ -r "$host_library_file" ]]; then + IFS= read -r host_library_path < "$host_library_file" || host_library_path="" + export CIRU_HOST_LIBRARY_PATH="$host_library_path" + fi +fi +rocm_library_path="$rocm_tool_root/lib:$rocm_tool_root/lib/rocm_sysdeps/lib" +rocm_cmake_path="$rocm_tool_root/lib/cmake" +if [[ -d "$rocm_libraries/lib" ]]; then + rocm_library_path="$rocm_library_path:$rocm_libraries/lib:$rocm_libraries/lib/rocm_sysdeps/lib" + rocm_cmake_path="$rocm_cmake_path:$rocm_libraries/lib/cmake" +fi +export PATH="$rocm_tool_root/bin:$rocm_llvm_bin:$VLLM_VENV/bin:$PATH" +# Match the working SDK layout: host C++ runtime, expanded ROCm SDK, then Torch. +export LD_LIBRARY_PATH="${host_library_path:+$host_library_path:}$rocm_library_path:$torch_lib${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}" +if [[ -z "${CC:-}" ]]; then + if command -v cc >/dev/null 2>&1; then + export CC="$(command -v cc)" + else + export CC="$rocm_llvm_bin/clang" + fi +fi +if [[ -z "${CXX:-}" ]]; then + if command -v c++ >/dev/null 2>&1; then + export CXX="$(command -v c++)" + else + export CXX="$rocm_llvm_bin/clang++" + fi +fi +# The extracted source trees are provenance/build inputs. Do not add them to +# PYTHONPATH: doing so would shadow the compiled packages installed by wheel. +export PYTHONPATH="$rocm_core/share/amd_smi${PYTHONPATH:+:$PYTHONPATH}" +export HIP_DEVICE_LIB_PATH="$rocm_core/lib/llvm/amdgcn/bitcode" +export ROCM_PATH="$rocm_tool_root" +export ROCM_HOME="$rocm_tool_root" +export HIP_PATH="$rocm_tool_root" +export CMAKE_PREFIX_PATH="$rocm_cmake_path:$site_packages/torch/share/cmake${CMAKE_PREFIX_PATH:+:$CMAKE_PREFIX_PATH}" +export AITER_JIT_DIR="${AITER_JIT_DIR:-${XDG_CACHE_HOME:-$HOME/.cache}/GLM5.3-Flash-CIRU-STRIX-IU4/aiter}" +install -d "$AITER_JIT_DIR" || return 2 +aiter_packaged_module="$VLLM_VENV/var/aiter-jit-gfx1151/module_aiter_core.so" +if [[ -f "$aiter_packaged_module" && ! -e "$AITER_JIT_DIR/module_aiter_core.so" ]]; then + cp -n "$aiter_packaged_module" "$AITER_JIT_DIR/module_aiter_core.so" || return 2 +fi + +export HIP_VISIBLE_DEVICES="${HIP_VISIBLE_DEVICES:-0}" +export ROCR_VISIBLE_DEVICES="${ROCR_VISIBLE_DEVICES:-0}" +export HIP_FORCE_DEV_KERNARG=1 +export PYTORCH_ROCM_ARCH=gfx1151 +export GPU_ARCHS=gfx1151 +export VLLM_TARGET_DEVICE=rocm +export VLLM_ROCM_USE_AITER="${VLLM_ROCM_USE_AITER:-1}" +export VLLM_ROCM_USE_AITER_MOE="${VLLM_ROCM_USE_AITER_MOE:-0}" +export VLLM_ROCM_USE_SKINNY_GEMM="${VLLM_ROCM_USE_SKINNY_GEMM:-0}" +export VLLM_ROCM_MOE_M1_DIRECT_ROUTE="${VLLM_ROCM_MOE_M1_DIRECT_ROUTE:-1}" +export FLASH_ATTENTION_TRITON_AMD_ENABLE=TRUE +export VLLM_ENABLE_V1_MULTIPROCESSING="${VLLM_ENABLE_V1_MULTIPROCESSING:-0}" +export HF_HUB_OFFLINE="${HF_HUB_OFFLINE:-1}" +export TRANSFORMERS_OFFLINE="${TRANSFORMERS_OFFLINE:-1}" +export PYTHONHASHSEED="${PYTHONHASHSEED:-1}" + +unset CUDA_VISIBLE_DEVICES +unset host_library_file host_library_path rocm_library_path rocm_cmake_path +unset aiter_packaged_module diff --git a/runtime/shell.nix b/runtime/shell.nix new file mode 100644 index 0000000000000000000000000000000000000000..e9e725da363335c9bf4c5dda589e2bdd411e0bbb --- /dev/null +++ b/runtime/shell.nix @@ -0,0 +1,30 @@ +{ pkgs ? import { } }: + +let + hostLibraries = with pkgs; [ + stdenv.cc.cc.lib + zlib + vulkan-loader + libdrm + elfutils + numactl + openssl + ]; +in +pkgs.mkShell { + packages = with pkgs; [ + uv + gcc + gnumake + cmake + ninja + git + pkg-config + unixtools.xxd + gnutar + coreutils + findutils + ]; + + CIRU_HOST_LIBRARY_PATH = pkgs.lib.makeLibraryPath hostLibraries; +} diff --git a/runtime/wheels/amd_aiter-0.1.0rc1-cp314-cp314-linux_x86_64.whl b/runtime/wheels/amd_aiter-0.1.0rc1-cp314-cp314-linux_x86_64.whl new file mode 100644 index 0000000000000000000000000000000000000000..1a820f3c3ae1d961673fb1bf901beb6fa1b34c49 --- /dev/null +++ b/runtime/wheels/amd_aiter-0.1.0rc1-cp314-cp314-linux_x86_64.whl @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e392272a101a5f11809818df222ccfe68892c4d541a83eefdfd3e6e3e08d7cb1 +size 67412360 diff --git a/runtime/wheels/vllm-0.1.0rc2.dev9+g9255fd9fb9.rocm100-cp314-cp314-linux_x86_64.whl b/runtime/wheels/vllm-0.1.0rc2.dev9+g9255fd9fb9.rocm100-cp314-cp314-linux_x86_64.whl new file mode 100644 index 0000000000000000000000000000000000000000..7c58f46b3d57fe05d5ce63d9083acaa81a9f634f --- /dev/null +++ b/runtime/wheels/vllm-0.1.0rc2.dev9+g9255fd9fb9.rocm100-cp314-cp314-linux_x86_64.whl @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ba5b1c0f957c41cf0d4d0af36df59cbe25ae5d0640766baea6d7ce32506d2317 +size 44318322