commit 60bf1b7b02896c1523c3cb895dacb7553d4377c8 Author: ent Date: Wed Jun 24 13:02:35 2026 +1000 qwen3.5-122B-A10B on DGX Spark: vLLM + DFlash + dense-bandwidth stack, one-shot installer diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..075bdad --- /dev/null +++ b/.gitignore @@ -0,0 +1,20 @@ +.DS_Store +*.swp +*.swo +*~ +.idea/ +.vscode/ +__pycache__/ +*.pyc + +# Logs / local results +*.log +/local/ +/results/ + +# Never commit weights or checkpoints +*.safetensors +*.gguf +*.bin +/hybrid-ckpt/ +*-hybrid-int4-fp8/ diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..cee33fd --- /dev/null +++ b/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 qwen3.5-122B-A10B-on-spark contributors + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md new file mode 100644 index 0000000..2809c11 --- /dev/null +++ b/README.md @@ -0,0 +1,219 @@ +# qwen3.5-122B-A10B-on-spark + +[`Qwen3.5-122B-A10B`](https://huggingface.co/Intel/Qwen3.5-122B-A10B-int4-AutoRound) +(hybrid GDN + mamba + 128-expert MoE, ~10B active) running on a single +**NVIDIA DGX Spark** (GB10 / SM121, 128 GiB unified) under **vLLM**, with +**[DFlash](https://modal.com/blog/spec-is-all-u-need) block-diffusion speculative +decode** and an optional **dense-bandwidth patch stack** — measured end-to-end, +with a per-token bandwidth model that explains every number. + +**Status:** Working end-to-end, one-shot install. On real Hermes-agent +tool-call turns, **DFlash decode reaches a median ~81 tok/s on GB10** — +**~2× the native MTP-2 head (~40 tok/s)** on the same workload, and above the +best previously published number for this model on Spark (albond's fully-patched +MTP stack, 51.58 tok/s end-to-end). DFlash's acceptance is task-dependent (it +block-drafts 12 tokens in one parallel forward), so the win is largest on +structured/tool-call/code traffic and collapses to parity on open-ended prose. + +A separate **dense-bandwidth stack** (hybrid INT4+FP8 shared experts + int8 +lm-head) adds **+28 % to no-spec / base decode** (28.2 → 36.0 tok/s) but, by the +amortization law below, washes out to ~null on high-acceptance agent traffic — +so it's a lever for *base / low-acceptance* serving, not for the agent path. + +- **Engine:** [`vLLM`](https://github.com/vllm-project/vllm) 0.23, sm121 build with the DFlash PRs, via the prebuilt image `ghcr.io/aeon-7/aeon-vllm-ultimate:2026-06-18-v0.23.0-dflashfix`. No host build — the four runtime patches in [`runtime/`](runtime/) are applied at serve time. +- **Target:** [`Intel/Qwen3.5-122B-A10B-int4-AutoRound`](https://huggingface.co/Intel/Qwen3.5-122B-A10B-int4-AutoRound) — INT4 (AutoRound/GPTQ) routed experts + attention, BF16 shared experts/embeddings/head, ~62 GiB. (Safetensors, *not* GGUF — vLLM serves HF checkpoints directly.) +- **Drafter:** [`z-lab/Qwen3.5-122B-A10B-DFlash`](https://huggingface.co/z-lab/Qwen3.5-122B-A10B-DFlash) — 0.8B / 6-layer non-causal block-diffusion drafter (block 16), shares the target's `embed_tokens` + `lm_head`, ~1.6 GiB. +- **Hardware:** NVIDIA DGX Spark, GB10, SM121, 128 GiB LPDDR5X unified, ~273 GB/s. + +## Quick start + +On a DGX Spark with Docker + the NVIDIA container runtime: + +```bash +curl -sSL https://raw.githubusercontent.com/Entrpi/qwen3.5-122B-A10B-on-spark/main/install.sh | bash -s -- --start +``` + +That one command: + +1. Verifies the host (aarch64, GB10/SM121, Docker GPU access, free disk). +2. Pulls the sm121 vLLM image (~40 GiB, one-time). +3. Downloads the INT4 target (~62 GiB) + DFlash drafter (~1.6 GiB) into the HF cache. +4. Starts the `dflash` profile on `:8000`, waits until READY, and runs the + "capital of France" smoke test (asserts "Paris"). + +**Already have the model?** Skip the 62 GiB download: + +```bash +# point at a checkpoint dir you already have (mounted read-only at /model): +./install.sh --start --model-dir /path/to/Qwen3.5-122B-A10B-int4-AutoRound +# or reuse an existing HF cache (download becomes a no-op if already present): +./install.sh --start --hf-home /mnt/big/hf +``` + +Preview without running: `... | bash -s -- --help`. + +## Hardware requirements + +| | | +|---|---| +| Validated on | NVIDIA DGX Spark (GB10, SM121, 128 GiB unified) | +| Likely to work | other Blackwell with `--force` (untested) | +| Runtime | Docker + NVIDIA container runtime (`docker run --gpus all`) | +| Disk | ≥ 75 GiB free (image + weights); ≥ 150 GiB if `--build-hybrid` | +| OS | aarch64 Linux (Grace) | +| Memory | 128 GiB unified is enough for the model + DFlash drafter + KV @ 16k | + +GB10 is detected via `nvidia-smi --query-gpu=compute_cap` returning `12.1`; +anything else needs `--force`. + +## What you get — profiles + +Pick with `--profile`: + +| Profile | Stack | Best for | Measured | +|---|---|---|---| +| **`dflash`** *(default)* | INT4 + DFlash n=12 | agents / tool-calls / code | **~81 tok/s** Hermes · 53.7 albond-bench | +| `dense` | hybrid INT4+FP8 + int8 lm-head + DFlash n=12 | base / low-accept serving | 36.0 base (+28%) · 59.0 albond-bench | +| `base` | plain INT4, no spec | airtight baseline | 28.2 tok/s c=1 | +| `mtp` | INT4 + native MTP-2 head | comparison | ~40 tok/s Hermes | + +The server is OpenAI-compatible (`/v1/chat/completions` with tool calls + SSE, +`/v1/completions`, `/v1/models`) and serves under the model name `qwen`. + +## Benchmarks + +All single-stream (c=1), temperature 0, GB10. "Hermes" = regenerating the next +assistant turn over 10 real conversations from a live agent's `state.db` (73 % +tool-calls); "albond-bench" = albond's own end-to-end harness (completion_tokens +/ total wallclock incl. prefill, 5 prompts, run-1 discarded — directly +comparable to his published 51.58). + +### DFlash vs MTP, same harness, unpatched + +| Workload (accept len) | base no-spec | MTP-2 | **DFlash n=12** | +|---|---|---|---| +| Prose (~2.3) | 28.2 | 33.7 | 33.2 *(use n=4)* | +| Code (~5.4) | 28.2 | 40.5 | **54.5** | +| Counting (~11) | 28.2 | 43.7 *(MTP caps at acc 3)* | **124.5** | +| **Hermes, real turns (8.3)** | — | **39.9** | **~81** | +| albond-bench e2e (6.5) | — | — | **53.7** | + +MTP-2 drafts 2 tokens *sequentially* (acceptance caps at ~3); DFlash block-drafts +12 in **one parallel forward**, so on predictable/agent traffic it accepts 5–11 +and pulls ~2× ahead. They tie only on low-acceptance prose. **53.7 unpatched +already clears albond's fully-patched MTP (51.58)** under his own method. + +### The dense-bandwidth stack (`dense` profile) + +Two independent always-on levers, ported to vLLM 0.23 as runtime patches: +hybrid INT4+FP8 (BF16 shared experts → calibrated FP8) and int8 lm-head (the +248 320-row vocab projection → int8 w8a16 GEMV, ~2× the bf16 read). + +| Config | base (acc 1) | DFlash spec, albond-bench (acc 6.4) | Hermes (acc 8.3) | +|---|---|---|---| +| INT4 baseline | 28.2 | 53.7 | ~81 | +| + hybrid-FP8 | 30.4 (+7.8%) | 57.0 (+6.1%) | ~80 | +| + int8 lm-head | 32.7 (+16%) | — | — | +| **+ both** | **36.0 (+28%)** | **59.0 (+10%)** | ~80–87 *(noise)* | + +## The amortization law + +The dense levers cut **always-on** weight reads (shared experts + lm-head, read +every token). Under speculative decode the verify forward reads those weights +**once and amortizes them across the accepted block**, so the gain shrinks as +acceptance rises — monotonically, across the whole curve: + +``` +dense stack uplift: +28% (base, accept 1) + → +10% (albond-bench, accept ~6.4) + → ~0% (Hermes, accept ~8.3) +``` + +**Consequence:** for the **agent path (`dflash`)**, DFlash's own high acceptance +already saturates the dense levers — its remaining bottleneck is *routed-expert* +verify-batch reads, which no dense-weight quant touches. For **base / low-accept +serving (`dense`)**, the stack is a real +28 %. See [`docs/FINDINGS.md`](docs/FINDINGS.md). + +## Under the hood: the four runtime patches + +vLLM is unmodified on disk; [`runtime/serve.sh`](runtime/serve.sh) edits the +installed package in-place before `vllm serve` (idempotent, sentinel-guarded): + +| Patch | What it does | Needed by | +|---|---|---| +| [`patch_unify2.py`](runtime/patch_unify2.py) | scale-block KV-cache **unify** so the hybrid GDN+mamba target absorbs the drafter's attention spec (the original assert can't); + `--no-enable-prefix-caching` routes to the no-hash-assert coordinator | **DFlash** (any spec profile) | +| [`patch_inc_hybrid.py`](runtime/patch_inc_hybrid.py) | adds an `INCConfig.maybe_update_config` override that detects FP8 dense layers in the hybrid checkpoint and dispatches `Fp8LinearMethod` for `shared_expert` | `dense` | +| [`patch_int8_lmhead_v3.py`](runtime/patch_int8_lmhead_v3.py) | replaces the lm-head matmul in `_get_logits` with a batched int8 w8a16 Triton GEMV (keeps bf16 weight for the shared drafter) | `dense` | +| [`patch_fla_shmem.py`](runtime/patch_fla_shmem.py) | lets the FLA GDN chunk kernels use big tiles on sm121's 99 KiB shmem (prefill/TTFT only; harmless) | always (free) | + +Why DFlash needs the unify patch at all, why the drafter must run `FLASH_ATTN` +(non-causal), and the full vLLM-vs-SGLANG dead-end history are in +[`docs/FINDINGS.md`](docs/FINDINGS.md). + +## Repo layout + +``` +install.sh One-shot installer (curl | bash | --help) +runtime/ Mounted read-only at /host inside the container: + serve.sh vLLM serve wrapper (applies the patches, then serves) + patch_unify2.py DFlash KV-unify fix + patch_inc_hybrid.py hybrid INT4+FP8 dispatch + patch_int8_lmhead_v3.py int8 lm-head GEMV + patch_fla_shmem.py FLA sm121 big-tile (prefill) + mtp_serve.sh MTP-2 comparison serve +scripts/ Host-side helpers: + monitor.sh Container-startup monitor with OOM auto-kill guard + bench_decode.py Decode-only tok/s (excludes TTFT) + bench_albond.py albond's e2e method (comparable to his 51.58) + hermes_bench.py Real agent turns from ~/.hermes/state.db + run_bank.sh prose/code/counting/hermes bank on any server +tools/ + build-hybrid-checkpoint.py Build the hybrid INT4+FP8 ckpt (for --build-hybrid) + inspect_ckpt.py Which layers are INT4 vs BF16 vs FP8 + validate_*.py Standalone correctness checks for the patches +docs/ + FINDINGS.md The full investigation, methodology, and the + amortization-law derivation +``` + +## Reproducing + +```bash +# default agent path (DFlash) + smoke test: +./install.sh --start + +# the dense stack (build the hybrid ckpt once, ~20 min, then serve): +./install.sh --build-hybrid +./install.sh --start --profile dense + +# benches (run on the host against the server; need: pip install requests): +python3 scripts/bench_decode.py --base-url http://127.0.0.1:8000 --model qwen \ + --prompt "Write a detailed essay about the history of tea." +python3 scripts/bench_albond.py http://127.0.0.1:8000 "dflash" # e2e, vs 51.58 +python3 scripts/hermes_bench.py --base-url http://127.0.0.1:8000 # real agent turns + +# MTP comparison: +./install.sh --start --profile mtp +``` + +## How this fits with related work + +| Piece | Role | +|---|---| +| [`vLLM`](https://github.com/vllm-project/vllm) | the inference engine; this repo serves Qwen3.5 + DFlash on it, unmodified-on-disk | +| [`Intel/...int4-AutoRound`](https://huggingface.co/Intel/Qwen3.5-122B-A10B-int4-AutoRound) · [`z-lab/...DFlash`](https://huggingface.co/z-lab/Qwen3.5-122B-A10B-DFlash) | the target + drafter weights | +| [`albond/DGX_Spark_Qwen3.5-122B-A10B-AR-INT4`](https://github.com/albond/DGX_Spark_Qwen3.5-122B-A10B-AR-INT4) | the MTP + hybrid-FP8 + int8-lmhead recipe we benchmarked against and ported the dense levers from | +| [`Entrpi/ds4-on-spark`](https://github.com/Entrpi/ds4-on-spark) | sibling repo, same hardware, different model (DeepSeek-V4-Flash via ds4) | +| [Modal: *Speculative decoding is all you need*](https://modal.com/blog/spec-is-all-u-need) | the DFlash block-diffusion drafter and the task-dependent-acceptance framing | + +## Acknowledgements + +- [`z-lab`](https://huggingface.co/z-lab) / [Modal](https://modal.com/blog/spec-is-all-u-need) — the DFlash drafter and block-diffusion speculative decode. +- [`Intel/AutoRound`](https://huggingface.co/Intel) — the INT4 target quantization. +- [`vLLM`](https://github.com/vllm-project/vllm) and the AEON sm121 image maintainers — the engine and the DFlash-enabled GB10 build. +- [`albond`](https://github.com/albond/DGX_Spark_Qwen3.5-122B-A10B-AR-INT4) — the MTP/hybrid-FP8/int8-lmhead recipe and the end-to-end benchmark methodology. + +## License + +MIT — see [LICENSE](LICENSE). The patches are original; vendored third-party +files (`tools/build-hybrid-checkpoint.py`) retain their upstream attribution. diff --git a/docs/FINDINGS.md b/docs/FINDINGS.md new file mode 100644 index 0000000..5b09b0f --- /dev/null +++ b/docs/FINDINGS.md @@ -0,0 +1,133 @@ +# FINDINGS — DFlash + dense levers for Qwen3.5-122B-A10B on DGX Spark + +Single-stream (c=1) decode of `Qwen3.5-122B-A10B` (hybrid GDN + mamba + 128-expert +MoE, ~10B active) on GB10 / SM121, 128 GiB unified, ~273 GB/s. The agent this +backs (Hermes) is ~73 % tool-calls. All numbers temperature 0. + +## 1. Getting DFlash to run on the hybrid 122B in vLLM + +The DFlash drafter is **non-causal** (it block-drafts 16 tokens in one parallel +forward) — only the `FLASH_ATTN` (FA2) backend supports non-causal attention. +But the hybrid GDN+mamba+MoE target's KV-cache page geometry won't *unify* with +the drafter's attention spec: + +- vLLM auto-aligns the hybrid (attention block 2240, mamba page padded +0.54 % to + match), so `max_page_size` is a *padded* value. `unify_kv_cache_spec_page_size` + scales the drafter's attention block by `ratio` and then asserts + `page == max` — which fails, because `page_size_bytes` ignores `block_size` + once `page_size_padded` is set. +- **Fix** ([`patch_unify2.py`](../runtime/patch_unify2.py)): keep the *scaled* + `block_size` **and** pad the <1 % remainder (mirrors vLLM's own + `HiddenStateCacheSpec` handling). The earlier "pad but keep block_size=16" + patch mis-strided the drafter KV → acceptance collapsed to 1.47 (a real bug, + not a quant mismatch). +- `--mamba-block-size 256` (from other Spark recipes) **breaks** the 122B: it + makes the mamba group block ≠ cache block, tripping the coordinator hash + assert. Omit it. +- `--no-enable-prefix-caching` routes to `KVCacheCoordinatorNoPrefixCache`, which + has no line-504 hash assert. Prefix caching is irrelevant at c=1 anyway. + +Working stack: `patch_unify2` + prefix-off + INT4 (bf16 KV) target + drafter +pinned to `FLASH_ATTN`. **No FA4 shim needed** — vLLM gates FA4 to cap families +90/100/110 (excludes 120), so the drafter runs FA2. (The whole fa4-sm120 saga is +SGLang-only; SGLang's DFlash works too but its sm121 base decode is ~2× slower +than vLLM's, so it loses on absolute throughput.) + +## 2. DFlash vs MTP — acceptance is task-dependent + +MTP-2 (the native head) drafts 2 tokens **sequentially**, so acceptance caps at +~3. DFlash block-drafts 12 in **one parallel forward**, so its acceptance fills +the block on predictable traffic. Same harness, unpatched, flash_attn: + +| Workload | accept (MTP-2 / DFlash) | tok/s (MTP-2 / DFlash) | +|---|---|---| +| Prose | 2.24 / 2.3 | 33.7 / 33.2 *(tie; use DFlash n=4)* | +| Code | 2.77 / 5.4 | 40.5 / 54.5 | +| Counting | **3.00 (maxed)** / 11 | 43.7 / 124.5 | +| Hermes (real) | 2.88 / **8.66** | 39.9 / **~81** | + +The "DFlash caps at 2.3 / 33 tok/s" story was a **prose-benchmark artifact**. +On agent/code traffic DFlash pulls ~2× ahead because MTP is acceptance-saturated. +`n` (num_speculative_tokens) is task-dependent: prose → 4, agent/code → 12+. + +## 3. Methodology — two non-comparable harnesses + +- `bench_decode.py` = **decode-only** tok/s (excludes TTFT). Good for c=1 kernel + comparisons. This is the ~81 Hermes number. +- `bench_albond.py` = **end-to-end** (completion_tokens / total wallclock incl. + prefill, non-streaming, 5 prompts, run-1 discarded). This reproduces albond's + own method and is directly comparable to his published **51.58**. + +Apples-to-apples (albond's method): **DFlash n=12 unpatched = 53.7 tok/s +cross-prompt mean — already above his fully-patched MTP stack (51.58).** + +## 4. The dense-bandwidth levers and the amortization law + +albond's non-MTP wins are *always-on* bandwidth cuts. We ported the two that +transfer to vLLM 0.23 + DFlash: + +- **hybrid INT4+FP8** ([`patch_inc_hybrid.py`](../runtime/patch_inc_hybrid.py)): + the Intel base already stores **attention as INT4** (0.5 B/param, *better* than + albond's FP8 attention), so the only thing to gain is the BF16 **shared + experts** → calibrated FP8 (144 layers, 0.48 GB saved). The dispatch patch adds + an `INCConfig.maybe_update_config` override (AEON 0.23's hook signature takes + `hf_config=`, unlike albond's 0.19) that detects FP8 dense layers and + dispatches `Fp8LinearMethod` for them. +- **int8 lm-head** ([`patch_int8_lmhead_v3.py`](../runtime/patch_int8_lmhead_v3.py)): + the 248 320-row vocab projection is the single largest dense read (1.5 GB BF16, + *every token*). A batched int8 w8a16 Triton GEMV reads it at ~227 GB/s (vs bf16 + ~6.5–8.8 ms) — **~2× faster, argmax-exact**. Prior ports failed not on the + kernel but on **integration**: zeroing the lm-head weight corrupted the + *drafter-shared* head (garbage), and a per-row loop for B>4 was slower under + spec. v3 uses one batched kernel and **keeps** the bf16 weight. + +Why the denominator matters: the 0.48 GB shared-expert saving is **0.7 % of the +71 GB on disk** but **~8 % of the ~6 GB *active per-token* footprint** (the disk +is mostly sparse routed experts). Shared experts and the lm-head are **dense — +read every token** — so at base decode the savings land in full: + +| Config | base (acc 1) | albond-bench (acc 6.4) | Hermes (acc 8.3) | +|---|---|---|---| +| INT4 baseline | 28.2 | 53.7 | ~81 | +| + hybrid-FP8 | 30.4 (+7.8%) | 57.0 (+6.1%) | ~80 | +| + int8 lm-head | 32.7 (+16%) | — | — | +| **+ both** | **36.0 (+28%)** | **59.0 (+10%)** | ~80–87 (noise) | + +The levers compose additively (step savings 2.6 + 4.9 ≈ 7.7 ms). But the uplift +**decays monotonically with acceptance**: + +> Under speculative decode the verify forward processes ~`accept` positions and +> reads each dense weight **once**, amortized across them. So a dense-weight cut +> that is +X % at base is ~+X/accept % under spec. + +``` ++28% base (accept 1) → +10% albond-bench (accept 6.4) → ~0% Hermes (accept 8.3) +``` + +**Consequences** + +- For the **agent path** (`dflash`), DFlash's own high acceptance already + amortizes the dense levers to ~null. Its remaining bottleneck is **routed-expert + verify-batch reads** (each of ~13 verify positions routes to different experts) + — untouched by any dense-weight quant. To push Hermes further you must attack + *that*: a smaller/faster drafter, lower `n` at equal acceptance, or sub-INT4 + routed experts. +- For **base / low-acceptance** serving (`dense`), the stack is a real **+28 %** + (36 tok/s) and is the recommended config there. + +## 5. Things that did NOT help c=1 decode + +- **FLASHINFER target backend** — null both short-context and Hermes (attention + isn't the bottleneck on this GDN/mamba-heavy MoE; most layers are linear + attention). albond's "+16 %" was on his dense-attention MTP path. +- **b12x / native FP4 MoE** — null at c=1 (a throughput/concurrency lever, not a + latency one; at batch 1 the active-expert GEMM is tiny). +- **FLA sm121 big-tile shmem fix** — real bug, but prefill/TTFT only; c=1 decode + uses the GDN *recurrent* path, a different kernel. Kept (free TTFT win). +- **PR#38325 swapAB FP8 GEMM** — marginal (+0.76 %), only with the FP8 checkpoint. + +## Production recommendation + +Ship **`dflash`** for the agent (DFlash unpatched, ~81 tok/s, ~2× MTP, > albond's +patched 51.58). Reserve **`dense`** for base / low-acceptance serving (+28 %). +The dense patches are upside there, not a requirement for the agent to win. diff --git a/install.sh b/install.sh new file mode 100755 index 0000000..94125c1 --- /dev/null +++ b/install.sh @@ -0,0 +1,319 @@ +#!/usr/bin/env bash +# install.sh — Qwen3.5-122B-A10B + DFlash speculative decode on NVIDIA DGX Spark +# (GB10 / SM121, 128 GiB unified), via vLLM in Docker. +# +# curl -sSL https://raw.githubusercontent.com/Entrpi/qwen3.5-122B-A10B-on-spark/main/install.sh | bash +# curl -sSL https://raw.githubusercontent.com/Entrpi/qwen3.5-122B-A10B-on-spark/main/install.sh | bash -s -- --help +# +# What this does (every step idempotent — safe to re-run): +# +# 1. Verifies the host is a DGX Spark / GB10 (SM121) with Docker + the NVIDIA +# container runtime, and enough free disk for the chosen profile. +# 2. Pulls the prebuilt sm121 vLLM image (DFlash-enabled, vLLM 0.23). +# 3. Downloads the INT4 target + DFlash drafter from Hugging Face into the HF +# cache — OR reuses a checkpoint you already have (--model-dir / --hf-home). +# 4. (optional) Builds the hybrid INT4+FP8 checkpoint for the "dense" profile. +# 5. Starts the vLLM server for the chosen --profile, waits until READY, and +# runs the "capital of France" smoke test (expects "Paris"). +# +# The script makes NO changes outside: +# - the Docker image cache (the pulled image) +# - $HF_HOME (default ~/.cache/huggingface) +# - $HYBRID_DIR (only with --build-hybrid) +# - the running container named $NAME (only with --start / smoke) +# +# This repo provides the install + serve + patch + benchmark layer ON TOP of: +# - Intel/Qwen3.5-122B-A10B-int4-AutoRound (target weights) +# - z-lab/Qwen3.5-122B-A10B-DFlash (block-diffusion drafter) +# - ghcr.io/aeon-7/aeon-vllm-ultimate (sm121 DFlash-enabled vLLM) +# +# License: MIT. Source: https://github.com/Entrpi/qwen3.5-122B-A10B-on-spark + +set -euo pipefail + +# ============================================================================ +# 0. defaults + flag parsing +# ============================================================================ + +# Prebuilt sm121 vLLM image with the DFlash PRs + the .pth that auto-applies our +# KV-unify patch is NOT baked in — we apply patches at serve time from runtime/. +IMAGE="${QWEN_IMAGE:-ghcr.io/aeon-7/aeon-vllm-ultimate:2026-06-18-v0.23.0-dflashfix}" + +TARGET_REPO="${TARGET_REPO:-Intel/Qwen3.5-122B-A10B-int4-AutoRound}" # INT4 target (~62 GiB) +DRAFT_REPO="${DRAFT_REPO:-z-lab/Qwen3.5-122B-A10B-DFlash}" # 0.8B drafter (~1.6 GiB) +FP8_REPO="${FP8_REPO:-Qwen/Qwen3.5-122B-A10B-FP8}" # FP8 donor for --build-hybrid + +HF_HOME="${HF_HOME:-$HOME/.cache/huggingface}" +HYBRID_DIR="${HYBRID_DIR:-$HOME/qwen3.5-122b-hybrid-int4-fp8}" +MODEL_DIR="" # --model-dir: a pre-downloaded INT4 checkpoint dir + +# This repo's own dir (works for `curl | bash` too: falls back to a clone). +REPO_DIR="${REPO_DIR:-$(cd "$(dirname "${BASH_SOURCE[0]:-$0}")" 2>/dev/null && pwd || echo "$HOME/code/qwen3.5-122B-A10B-on-spark")}" +REPO_URL="${REPO_URL:-https://github.com/Entrpi/qwen3.5-122B-A10B-on-spark.git}" + +NAME="${NAME:-qwen-spark}" +PROFILE="dflash" # dflash | dense | base | mtp +NSPEC="" # override num_speculative_tokens (default per profile) +PORT="${PORT:-8000}" +CTX="${CTX:-16384}" +GPU_MEM="${GPU_MEM:-0.8}" +BACKEND="${BACKEND:-flash_attn}" + +FORCE_HW=0 +SKIP_PULL=0 +SKIP_DOWNLOAD=0 +BUILD_HYBRID=0 +START_SERVER=0 +SKIP_SMOKE=0 + +usage() { + cat <&2; usage; exit 2 ;; + esac +done + +case "$PROFILE" in dflash|dense|base|mtp) ;; *) echo "Bad --profile: $PROFILE" >&2; exit 2 ;; esac + +c_red() { printf '\033[31m%s\033[0m' "$*"; } +c_green() { printf '\033[32m%s\033[0m' "$*"; } +c_yellow(){ printf '\033[33m%s\033[0m' "$*"; } +log() { printf '%s %s\n' "[$(date +%H:%M:%S)]" "$*"; } +die() { printf '\n%s %s\n' "$(c_red FATAL:)" "$*" >&2; exit 1; } +warn(){ printf '%s %s\n' "$(c_yellow WARN:)" "$*" >&2; } +ok() { printf '%s %s\n' "$(c_green OK:)" "$*"; } + +# ============================================================================ +# 1. host verification +# ============================================================================ + +verify_host() { + log "Verifying host..." + local m; m=$(uname -m) + if [[ "$m" != "aarch64" ]] && [[ "$FORCE_HW" -eq 0 ]]; then + die "Expected aarch64 (Grace+Blackwell); got $m. Pass --force to skip." + fi + command -v docker >/dev/null 2>&1 || die "docker not found. Install Docker + the NVIDIA container runtime." + command -v nvidia-smi >/dev/null 2>&1 || die "nvidia-smi not found. Need the NVIDIA driver." + local gpu; gpu=$(nvidia-smi --query-gpu=name,compute_cap --format=csv,noheader 2>/dev/null || true) + [[ -n "$gpu" ]] || die "nvidia-smi failed to enumerate GPUs." + log "GPU: $gpu" + if ! echo "$gpu" | grep -qE '12\.1|GB10|Spark'; then + [[ "$FORCE_HW" -eq 1 ]] || die "Not detecting GB10/SM12.1. Pass --force (and maybe --backend) to proceed." + warn "Host is not GB10/SM121; proceeding under --force (untested)." + fi + # Docker can see the GPU? + if ! docker info 2>/dev/null | grep -qiE 'nvidia|Default Runtime: nvidia' \ + && ! docker run --rm --gpus all "$IMAGE" true 2>/dev/null; then + warn "Could not confirm Docker GPU access (nvidia-container-toolkit). 'docker run --gpus all' must work." + fi + # Disk + local need=75; [[ "$BUILD_HYBRID" -eq 1 ]] && need=150 + local free; free=$(df -BG "$HOME" | awk 'NR==2{gsub("G","",$4);print $4}') + if (( free < need )) && [[ "$SKIP_DOWNLOAD" -eq 0 ]] && [[ -z "$MODEL_DIR" ]]; then + die "Need >= ${need} GiB free under $HOME for profile '$PROFILE'; have ${free} GiB. Use --model-dir / --no-download, or free space." + fi + ok "Host checks passed (free ${free} GiB)." +} + +# ============================================================================ +# 2. pull image +# ============================================================================ + +pull_image() { + if [[ "$SKIP_PULL" -eq 1 ]]; then log "Skipping docker pull (--no-pull)."; return; fi + if docker image inspect "$IMAGE" >/dev/null 2>&1; then ok "Image present: $IMAGE"; return; fi + log "Pulling $IMAGE (~40 GiB, one-time) ..." + docker pull "$IMAGE" + ok "Image pulled." +} + +# ============================================================================ +# 3. download models (idempotent — snapshot_download no-ops if cached) +# ============================================================================ + +hf_get() { # repo -> populate HF cache via the image's huggingface_hub + local repo="$1" + docker run --rm --net=host -e HF_HOME=/hf ${HF_TOKEN:+-e HF_TOKEN="$HF_TOKEN"} \ + -v "$HF_HOME:/hf" --entrypoint python3 "$IMAGE" \ + -c "from huggingface_hub import snapshot_download as s; s('$repo')" +} + +download_models() { + if [[ "$SKIP_DOWNLOAD" -eq 1 ]]; then log "Skipping HF download (--no-download)."; return; fi + mkdir -p "$HF_HOME" + if [[ -n "$MODEL_DIR" ]]; then + [[ -f "$MODEL_DIR/config.json" ]] || die "--model-dir $MODEL_DIR has no config.json" + log "Using pre-downloaded target at $MODEL_DIR (skipping target download)." + else + log "Fetching target $TARGET_REPO into $HF_HOME ..." + hf_get "$TARGET_REPO" + fi + log "Fetching drafter $DRAFT_REPO ..." + hf_get "$DRAFT_REPO" + ok "Models ready." +} + +# ============================================================================ +# 4. optional: build the hybrid INT4+FP8 checkpoint +# ============================================================================ + +build_hybrid() { + [[ "$BUILD_HYBRID" -eq 1 ]] || return + if [[ -f "$HYBRID_DIR/model.safetensors.index.json" ]]; then ok "Hybrid ckpt present: $HYBRID_DIR"; return; fi + local gptq="$MODEL_DIR" + if [[ -z "$gptq" ]]; then + gptq=$(docker run --rm -v "$HF_HOME:/hf" -e HF_HOME=/hf --entrypoint python3 "$IMAGE" \ + -c "from huggingface_hub import snapshot_download as s; print(s('$TARGET_REPO'))" | tail -1) + gptq="/hf-snap" # mount the cache; resolve inside the container below + fi + mkdir -p "$HYBRID_DIR" + log "Building hybrid INT4+FP8 checkpoint -> $HYBRID_DIR (~20 min) ..." + docker run --rm --net=host -e HF_HOME=/hf ${HF_TOKEN:+-e HF_TOKEN="$HF_TOKEN"} \ + -v "$HF_HOME:/hf" -v "$HYBRID_DIR:/out" -v "$REPO_DIR/tools:/tools:ro" \ + ${MODEL_DIR:+-v "$MODEL_DIR:/gptq:ro"} \ + --entrypoint bash "$IMAGE" -c ' + set -e + GPTQ="'"${MODEL_DIR:+/gptq}"'" + if [ -z "$GPTQ" ]; then + GPTQ=$(python3 -c "from huggingface_hub import snapshot_download as s; print(s(\"'"$TARGET_REPO"'\"))") + fi + python3 /tools/build-hybrid-checkpoint.py --gptq-dir "$GPTQ" \ + --fp8-repo "'"$FP8_REPO"'" --output /out --force + rm -rf /out/.fp8_cache' + ok "Hybrid checkpoint built: $HYBRID_DIR" +} + +# ============================================================================ +# 5. start server (+ smoke test) +# ============================================================================ + +ensure_runtime() { # make sure runtime/ (serve wrapper + patches) is on disk + if [[ -f "$REPO_DIR/runtime/serve.sh" ]]; then return; fi + log "runtime/ not found next to install.sh — cloning repo to $HOME/code/qwen3.5-122B-A10B-on-spark" + REPO_DIR="$HOME/code/qwen3.5-122B-A10B-on-spark" + [[ -d "$REPO_DIR/.git" ]] || git clone --depth 1 "$REPO_URL" "$REPO_DIR" + [[ -f "$REPO_DIR/runtime/serve.sh" ]] || die "runtime/serve.sh still missing after clone." +} + +start_server() { + [[ "$START_SERVER" -eq 1 ]] || { log "Setup complete. Re-run with --start to launch the server."; return; } + ensure_runtime + + # profile -> serve args + env + mounts + local nspec model_env=() mounts=() serve_args + case "$PROFILE" in + dflash) nspec="${NSPEC:-12}"; serve_args="$nspec $BACKEND" ;; + dense) nspec="${NSPEC:-12}"; serve_args="$nspec $BACKEND" + [[ -f "$HYBRID_DIR/model.safetensors.index.json" ]] || die "dense profile needs the hybrid ckpt — run with --build-hybrid first." + model_env=(-e MODEL=/model -e INC_HYBRID=1 -e INT8_LMHEAD_V3=1) + mounts=(-v "$HYBRID_DIR:/model:ro") ;; + base) nspec="${NSPEC:-0}"; serve_args="$nspec $BACKEND" ;; + mtp) nspec="${NSPEC:-2}"; serve_args="$nspec $BACKEND" ;; + esac + if [[ -n "$MODEL_DIR" && "$PROFILE" != "dense" ]]; then + model_env=(-e MODEL=/model); mounts=(-v "$MODEL_DIR:/model:ro") + fi + local wrapper="/host/serve.sh"; [[ "$PROFILE" == "mtp" ]] && wrapper="/host/mtp_serve.sh" + + log "Starting profile=$PROFILE (nspec=$nspec, ctx=$CTX, gpu-mem=$GPU_MEM) as container '$NAME' ..." + docker rm -f "$NAME" >/dev/null 2>&1 || true + # shellcheck disable=SC2086 + docker run -d --name "$NAME" --gpus all --net=host --ipc=host --ulimit memlock=-1:-1 \ + -e HF_HOME=/hf -e MAX_MODEL_LEN="$CTX" -e GPU_MEM="$GPU_MEM" ${HF_TOKEN:+-e HF_TOKEN="$HF_TOKEN"} \ + "${model_env[@]}" \ + -v "$HF_HOME:/hf" -v "$REPO_DIR/runtime:/host:ro" "${mounts[@]}" \ + --entrypoint bash "$IMAGE" "$wrapper" $serve_args >/dev/null + log "Container started. Model load + compile is ~8-12 min. Tail: docker logs -f $NAME" + + log "Waiting for http://127.0.0.1:$PORT/health ..." + local i + for i in $(seq 1 180); do + if ! docker ps --format '{{.Names}}' | grep -q "^$NAME$"; then + docker logs "$NAME" 2>&1 | tail -30; die "Container exited during load. See log above." + fi + if curl -sf "http://127.0.0.1:$PORT/health" >/dev/null 2>&1; then + ok "Server READY on http://127.0.0.1:$PORT" + break + fi + sleep 5 + done + curl -sf "http://127.0.0.1:$PORT/health" >/dev/null 2>&1 || die "Server not ready within ~15 min. docker logs $NAME" + + [[ "$SKIP_SMOKE" -eq 1 ]] && { log "Skipping smoke test (--no-smoke)."; return; } + log "Smoke test: 'capital of France' ..." + local out + out=$(curl -s "http://127.0.0.1:$PORT/v1/completions" -H 'Content-Type: application/json' \ + -d "{\"model\":\"qwen\",\"prompt\":\"What is the capital of France? Answer in one word.\",\"max_tokens\":8,\"temperature\":0}" \ + | python3 -c "import sys,json;print(json.load(sys.stdin)['choices'][0]['text'])" 2>/dev/null || true) + echo " -> $out" + echo "$out" | grep -qi paris && ok "Smoke test PASSED — 'Paris'." || die "Smoke test FAILED — 'Paris' not in output." +} + +# ============================================================================ +# main +# ============================================================================ + +verify_host +pull_image +download_models +build_hybrid +start_server + +echo +ok "Done (profile=$PROFILE)." +echo " Server: http://127.0.0.1:$PORT/v1 (model name: qwen)" +echo " Logs: docker logs -f $NAME" +echo " Stop: docker rm -f $NAME" +echo " Bench: python3 scripts/bench_decode.py --base-url http://127.0.0.1:$PORT --model qwen --prompt 'Write an essay about tea.'" +echo " Agent: python3 scripts/hermes_bench.py --base-url http://127.0.0.1:$PORT # real tool-call turns" diff --git a/runtime/mtp_serve.sh b/runtime/mtp_serve.sh new file mode 100755 index 0000000..d971725 --- /dev/null +++ b/runtime/mtp_serve.sh @@ -0,0 +1,25 @@ +#!/bin/bash +# mtp_serve.sh — native qwen3_5 MTP-N head for the comparison (`--profile mtp`). +# The MTP head (1 layer, reuses target KV/embed/lm_head) is in the Intel +# checkpoint (mtp.layers.0) — no separate drafter, no unify patch needed. +# $1 = num_speculative_tokens (default 2 = the "MTP-2" recipe); $2 = backend. +set -euo pipefail +NSPEC="${1:-2}" +BACKEND="${2:-flash_attn}" +MODEL="${MODEL:-Intel/Qwen3.5-122B-A10B-int4-AutoRound}" +MAX_MODEL_LEN="${MAX_MODEL_LEN:-16384}" +GPU_MEM="${GPU_MEM:-0.8}" +PORT="${PORT:-8000}" +echo "[mtp] qwen3_5_mtp — backend=$BACKEND, num_speculative_tokens=$NSPEC, model=$MODEL" +exec vllm serve "$MODEL" \ + --served-model-name qwen \ + --host 0.0.0.0 --port "$PORT" \ + --max-model-len "$MAX_MODEL_LEN" \ + --max-num-seqs 16 \ + --max-num-batched-tokens "$MAX_MODEL_LEN" \ + --gpu-memory-utilization "$GPU_MEM" \ + --no-enable-prefix-caching \ + --enable-chunked-prefill \ + --trust-remote-code \ + --attention-backend "$BACKEND" \ + --speculative-config "{\"method\":\"qwen3_5_mtp\",\"num_speculative_tokens\":$NSPEC,\"model\":\"$MODEL\"}" diff --git a/runtime/patch_fla_shmem.py b/runtime/patch_fla_shmem.py new file mode 100644 index 0000000..9c625be --- /dev/null +++ b/runtime/patch_fla_shmem.py @@ -0,0 +1,32 @@ +#!/usr/bin/env python3 +"""spark-fla-shmem: let the FLA (flash-linear-attention / GDN) Triton kernels use their +BIG tiles on sm121 (GB10 / DGX Spark). + +The FLA Backend gate (vllm/model_executor/layers/fla/ops/utils.py, identical in SGLang) +picks big tiles only if device_max_shared_mem >= Backend.DEFAULT (102400 = 100 KiB): + cumsum.py : BS_LIST = [32,64] if check_shared_mem() else [16,32] + chunk_o.py: BKV_LIST = [64,128] if check_shared_mem() else [32,64] +sm121 reports max_shared_mem = 101376 (99 KiB), JUST below 102400 -> check returns False +-> small tiles -> slower GDN/linear-attention (the per-token hot path). But 101376 is +EXACTLY Backend.ADA, and RTX 4090 (ADA, same 99 KiB) runs the big tiles fine, so they +provably fit in 99 KiB. Lower DEFAULT to 101376 so sm121 (and ADA) pass the gate. +""" +import pathlib, sys + +P = pathlib.Path( + "/usr/local/lib/python3.12/site-packages/vllm/model_executor/layers/fla/ops/utils.py" +) +src = P.read_text() +if "spark-fla-shmem" in src: + print("[patch_fla_shmem] already patched", flush=True) + sys.exit(0) + +OLD = " DEFAULT = 102400 # Default" +NEW = " DEFAULT = 101376 # spark-fla-shmem: 102400->101376 so sm121 GB10 (99 KiB) uses big GDN tiles (already fit on ADA's identical 99 KiB)" + +if OLD not in src: + print("[patch_fla_shmem] ERROR: target line not found — FLA source differs", flush=True) + sys.exit(1) + +P.write_text(src.replace(OLD, NEW)) +print("[patch_fla_shmem] patched OK (DEFAULT 102400 -> 101376; big GDN tiles on sm121)", flush=True) diff --git a/runtime/patch_inc_hybrid.py b/runtime/patch_inc_hybrid.py new file mode 100644 index 0000000..7e8264b --- /dev/null +++ b/runtime/patch_inc_hybrid.py @@ -0,0 +1,163 @@ +#!/usr/bin/env python3 +"""spark-dflash-hybrid-fp8: port of albond's hybrid INT4+FP8 dispatch +(patches/01-hybrid-int4-fp8/inc.py.patch) onto AEON 0.23's vllm INCConfig. + +Adds, to vllm/model_executor/layers/quantization/inc.py: + * INCConfig.fp8_config / fp8_layers fields + * maybe_update_config OVERRIDE — AEON 0.23 signature (model_name, hf_config=None, + revision=None); the base hook is already CALLED from config/vllm.py:634, so we + only supply the override. Scans the checkpoint's safetensors metadata for + float8_e4m3fn weights that have a .weight_scale_inv, builds an Fp8Config, and + records those layer prefixes. + * _is_layer_fp8 — exact + fused + substring match against fp8_layers + * FP8 dispatch at BOTH dense short-circuits: get_quant_method's extra_config + override AND the apply_*_quant_layer not-quantized blocks. + +Idempotent; sentinel 'spark-dflash-hybrid-fp8'. Mirrors patch_unify2.py's style. +""" +import sys + +import vllm.model_executor.layers.quantization.inc as inc_mod + +path = inc_mod.__file__ +src = open(path).read() +SENT = "spark-dflash-hybrid-fp8" +if SENT in src: + print(f"[patch_inc_hybrid] already applied: {path}") + sys.exit(0) + +# 1. __init__ fields (anchor unique: only INCConfig has pack_factor = Fraction(32,..)) +a1 = " self.pack_factor = Fraction(32, weight_bits)\n" +b1 = a1 + ( + " # spark-dflash-hybrid-fp8: populated by maybe_update_config\n" + " self.fp8_config = None\n" + " self.fp8_layers = set()\n" +) +assert src.count(a1) == 1, f"anchor1 count={src.count(a1)}" +src = src.replace(a1, b1) + +# 2. apply_vllm_mapper: remap fp8_layers (HF->vLLM names) after extra_config remap +a2 = ( + " if self.extra_config is not None:\n" + " self.extra_config = hf_to_vllm_mapper.apply_dict(self.extra_config)\n" +) +b2 = a2 + ( + " if self.fp8_layers: # spark-dflash-hybrid-fp8\n" + " self.fp8_layers = set(\n" + " hf_to_vllm_mapper.apply_list(list(self.fp8_layers))\n" + " )\n" +) +assert src.count(a2) == 1, f"anchor2 count={src.count(a2)}" +src = src.replace(a2, b2) + +# 3. insert maybe_update_config + _is_layer_fp8 before apply_awq_quant_layer +a3 = ' def apply_awq_quant_layer(self, layer, prefix: str, backend: str = "auto"):\n' +methods = ''' def maybe_update_config( # spark-dflash-hybrid-fp8 + self, + model_name: str, + hf_config=None, + revision: str | None = None, + ): + """Detect FP8 dense layers in a hybrid INT4+FP8 checkpoint.""" + import torch as _torch + from safetensors.torch import _TYPES as _SF + from vllm.transformers_utils.config import get_safetensors_params_metadata + from vllm.model_executor.layers.quantization.fp8 import Fp8Config + metadata = get_safetensors_params_metadata(model_name, revision=revision) + fp8_weights = {} + for pn, info in metadata.items(): + ds = info.get("dtype", None) + if ds is None: + continue + if _SF.get(ds) == _torch.float8_e4m3fn and pn.endswith(".weight"): + sn = pn.replace(".weight", ".weight_scale_inv") + if sn in metadata: + fp8_weights[pn] = info + if not fp8_weights: + logger.info("spark-dflash-hybrid-fp8: no FP8 dense layers detected") + return + block_size = None + for pn, info in fp8_weights.items(): + sn = pn.replace(".weight", ".weight_scale_inv") + ws = info.get("shape", []) + ss = metadata[sn].get("shape", []) + if len(ws) == 2 and len(ss) == 2: + block_size = [ws[0] // ss[0], ws[1] // ss[1]] + break + if block_size is None: + return + self.fp8_config = Fp8Config( + is_checkpoint_fp8_serialized=True, + activation_scheme="dynamic", + weight_block_size=block_size, + ) + self.fp8_layers = {n.rsplit(".weight", 1)[0] for n in fp8_weights} + _sample = sorted(self.fp8_layers)[:3] + logger.info( + "spark-dflash-hybrid-fp8: detected %d FP8 dense layers " + "(block_size=%s) e.g. %s", + len(self.fp8_layers), block_size, _sample, + ) + + def _is_layer_fp8(self, prefix: str) -> bool: # spark-dflash-hybrid-fp8 + if not self.fp8_layers: + return False + if prefix in self.fp8_layers: + return True + fused = getattr(self, "packed_modules_mapping", {}) + proj = prefix.split(".")[-1] + if proj in fused: + shards = [prefix.replace(proj, s) for s in fused[proj]] + return all( + any(fl in sp for fl in self.fp8_layers) for sp in shards + ) + return any(fl in prefix for fl in self.fp8_layers) + +''' +assert src.count(a3) == 1, f"anchor3 count={src.count(a3)}" +src = src.replace(a3, methods + a3) + +# 4. FP8 dispatch in the not-quantized blocks (awq/gptq/xpu/cpu are byte-identical; +# guard is inert unless fp8_config set, so patching all is safe) +a4 = ( + " if not self.check_quantized(weight_bits):\n" + " if isinstance(layer, (LinearBase, ParallelLMHead)):\n" + " return UnquantizedLinearMethod()\n" + " else:\n" + " return None\n" +) +b4 = ( + " if not self.check_quantized(weight_bits):\n" + " if self.fp8_config and self._is_layer_fp8(prefix): # spark-dflash-hybrid-fp8\n" + " from vllm.model_executor.layers.quantization.fp8 import (\n" + " Fp8LinearMethod,\n" + " )\n" + " return Fp8LinearMethod(self.fp8_config)\n" + " if isinstance(layer, (LinearBase, ParallelLMHead)):\n" + " return UnquantizedLinearMethod()\n" + " else:\n" + " return None\n" +) +n4 = src.count(a4) +assert n4 >= 2, f"anchor4 count={n4}" +src = src.replace(a4, b4) + +# 5. FP8 dispatch in get_quant_method's extra_config (bits>=16) override +a5 = ( + ' ) and self.extra_config[layer_name].get("bits", 16) >= 16:\n' + " return UnquantizedLinearMethod()\n" +) +b5 = ( + ' ) and self.extra_config[layer_name].get("bits", 16) >= 16:\n' + " if self.fp8_config and self._is_layer_fp8(prefix): # spark-dflash-hybrid-fp8\n" + " from vllm.model_executor.layers.quantization.fp8 import (\n" + " Fp8LinearMethod,\n" + " )\n" + " return Fp8LinearMethod(self.fp8_config)\n" + " return UnquantizedLinearMethod()\n" +) +assert src.count(a5) == 1, f"anchor5 count={src.count(a5)}" +src = src.replace(a5, b5) + +open(path, "w").write(src) +print(f"[patch_inc_hybrid] applied {SENT} to {path} (not-quant blocks x{n4})") diff --git a/runtime/patch_int8_lmhead_v3.py b/runtime/patch_int8_lmhead_v3.py new file mode 100644 index 0000000..cdafabd --- /dev/null +++ b/runtime/patch_int8_lmhead_v3.py @@ -0,0 +1,124 @@ +#!/usr/bin/env python3 +"""INT8 W8A16 lm-head v3 for AEON vLLM 0.23 (sm121). Replaces ONLY the +`lm_head.quant_method.apply(...)` call inside LogitsProcessor._get_logits with a +batched int8 GEMV (one kernel launch for any batch), leaving the existing TP +gather + org_vocab_size trim untouched. + +Fixes vs the broken v2 port: + * KEEPS the bf16 lm_head weight (DFlash drafter shares it) — does NOT zero it. + Trades the memory saving for correctness; the speed win is the int8 read in + _get_logits, independent of keeping bf16 around. + * Single BATCHED kernel (dot-based, pad B->16) for ALL B — no per-row Python + loop (the v2 B>4 loop was what made spec decode SLOWER). + * Fixed proven config (N128/K128/w4/s3, ~227 GB/s, argmax-exact vs bf16 on the + real [248320,3072] shape) — no autotune (avoids sm121 bad-config miscompiles). + +Sentinel DGX_SPARK_INT8_LMHEAD_V3. Verified standalone: int8 3.35ms vs bf16 8.8ms +(B=1) / 6.5ms (B=13); maxerr 3e-4 < quant floor, argmax 100%. +""" +import os +import sys + +TARGET = "/usr/local/lib/python3.12/site-packages/vllm/model_executor/layers/logits_processor.py" + +ANCHOR = ( + " # Get the logits for the next tokens.\n" + " logits = lm_head.quant_method.apply(lm_head, hidden_states, bias=embedding_bias)\n" +) +REPLACE = ( + " # DGX_SPARK_INT8_LMHEAD_V3: int8 w8a16 GEMV for the huge vocab projection\n" + " logits = _spark_int8_lmhead_apply(self, lm_head, hidden_states, embedding_bias)\n" +) + +MODULE_CODE = ''' + +# ===================== DGX_SPARK_INT8_LMHEAD_V3 ===================== +import triton as _spark_triton +import triton.language as _spark_tl + + +@_spark_triton.jit +def _spark_k_int8(x_ptr, w_ptr, s_ptr, o_ptr, B, N, K, + sxb, sxk, swn, swk, sob, son, + BLOCK_B: _spark_tl.constexpr, BLOCK_N: _spark_tl.constexpr, + BLOCK_K: _spark_tl.constexpr): + pid_n = _spark_tl.program_id(0) + offs_b = _spark_tl.arange(0, BLOCK_B) + offs_n = pid_n * BLOCK_N + _spark_tl.arange(0, BLOCK_N) + offs_k = _spark_tl.arange(0, BLOCK_K) + x_ptrs = x_ptr + offs_b[:, None] * sxb + offs_k[None, :] * sxk + w_ptrs = w_ptr + offs_n[:, None] * swn + offs_k[None, :] * swk + acc = _spark_tl.zeros((BLOCK_B, BLOCK_N), dtype=_spark_tl.float32) + for k in range(0, K, BLOCK_K): + km = (offs_k[None, :] + k) < K + x = _spark_tl.load(x_ptrs, mask=(offs_b[:, None] < B) & km, other=0.0).to(_spark_tl.float16) + w = _spark_tl.load(w_ptrs, mask=(offs_n[:, None] < N) & km, other=0).to(_spark_tl.float16) + acc += _spark_tl.dot(x, w.T) + x_ptrs += BLOCK_K * sxk + w_ptrs += BLOCK_K * swk + s = _spark_tl.load(s_ptr + offs_n, mask=offs_n < N, other=0.0).to(_spark_tl.float32) + acc = acc * s[None, :] + o_ptrs = o_ptr + offs_b[:, None] * sob + offs_n[None, :] * son + _spark_tl.store(o_ptrs, acc, mask=(offs_b[:, None] < B) & (offs_n[None, :] < N)) + + +def _spark_int8_gemm(hidden, w_int8, w_scale): + import torch + N, K = w_int8.shape + x = hidden.reshape(-1, K) + B = x.shape[0] + BLOCK_B = max(16, _spark_triton.next_power_of_2(B)) + out = torch.empty(B, N, dtype=torch.float32, device=x.device) + xf = x.to(torch.float16) + grid = ((N + 127) // 128,) + _spark_k_int8[grid](xf, w_int8, w_scale, out, B, N, K, + xf.stride(0), xf.stride(1), w_int8.stride(0), w_int8.stride(1), + out.stride(0), out.stride(1), + BLOCK_B=BLOCK_B, BLOCK_N=128, BLOCK_K=128, + num_warps=4, num_stages=3) + return out.reshape(hidden.shape[:-1] + (N,)) + + +def _spark_int8_lmhead_apply(self, lm_head, hidden_states, embedding_bias): + import sys + import torch + if not getattr(lm_head, "_spark_int8_ready", None) is True and \\ + not getattr(lm_head, "_spark_int8_disabled", False): + w = getattr(lm_head, "weight", None) + if (w is not None and w.dtype in (torch.bfloat16, torch.float16) + and w.dim() == 2 and w.shape[0] > 100000): + with torch.no_grad(): + scales = (w.float().abs().amax(dim=1) / 127.0).clamp(min=1e-12) + w_int8 = (w.float() / scales.unsqueeze(1)).round().clamp(-127, 127).to(torch.int8) + lm_head._spark_w_int8 = w_int8.contiguous() + lm_head._spark_w_scale = scales.to(torch.float16) + lm_head._spark_int8_ready = True + print("DGX_SPARK_INT8_LMHEAD_V3: lm_head -> int8 (%s), bf16 kept for shared drafter" + % (list(w_int8.shape),), file=sys.stderr, flush=True) + else: + lm_head._spark_int8_disabled = True + if getattr(lm_head, "_spark_int8_ready", False) and embedding_bias is None: + return _spark_int8_gemm(hidden_states, lm_head._spark_w_int8, lm_head._spark_w_scale) + return lm_head.quant_method.apply(lm_head, hidden_states, bias=embedding_bias) +# =================== end DGX_SPARK_INT8_LMHEAD_V3 =================== +''' + + +def main(): + if not os.path.exists(TARGET): + print(f"FAIL: {TARGET} not found"); sys.exit(1) + src = open(TARGET).read() + if "DGX_SPARK_INT8_LMHEAD_V3" in src: + print("SKIP: int8 lm-head v3 already applied"); return + if ANCHOR not in src: + print("FAIL: _get_logits apply-anchor not found"); sys.exit(1) + if src.count(ANCHOR) != 1: + print(f"FAIL: anchor count={src.count(ANCHOR)} (expected 1)"); sys.exit(1) + src = src.replace(ANCHOR, REPLACE) + src = src + MODULE_CODE + open(TARGET, "w").write(src) + print("OK: int8 lm-head v3 applied") + + +if __name__ == "__main__": + main() diff --git a/runtime/patch_unify2.py b/runtime/patch_unify2.py new file mode 100644 index 0000000..f3a7b3c --- /dev/null +++ b/runtime/patch_unify2.py @@ -0,0 +1,55 @@ +#!/usr/bin/env python3 +"""spark-dflash-unify (CORRECTED): let the DFlash drafter's attention KV spec unify +with the hybrid GDN+mamba target's *padded* page size WITHOUT corrupting the drafter. + +Root cause (AEON 0.23.0+aeon.sm121a.dflash): with --mamba-block-size set, vLLM's own +hybrid alignment makes target mamba page == target attention page by PADDING the mamba +page (e.g. "+0.54%"). That padded value becomes max_page_size. The drafter's attention +page is smaller; unify_kv_cache_spec_page_size scales its block_size by +ratio = max_page_size // layer_page_size, but because max_page_size is a *padded* (non +block-size-linear) number, the scaled page lands just under it and +`assert new_spec.page_size_bytes == max_page_size` fires. + +The OLD spark-pad-unify patch "fixed" this by `replace(layer_spec, +page_size_padded=max_page_size)` -- but it DROPPED the block_size scaling, leaving the +drafter at block_size=16 behind a max-sized physical page. That mis-strided the drafter +KV and pinned mean accept length at ~1.47 (garbage drafts), so DFlash net-lost. + +This corrected patch mirrors what get_kv_cache_groups already does for +HiddenStateCacheSpec layers: keep the SCALED block_size AND pad the <1% remainder. +Strides stay correct (block_size is properly scaled); only the page tail is padded. +""" +import pathlib, sys + +P = pathlib.Path("/usr/local/lib/python3.12/site-packages/vllm/v1/core/kv_cache_utils.py") +src = P.read_text() + +if "spark-dflash-unify" in src: + print("[patch_unify2] already patched", flush=True) + sys.exit(0) + +OLD = ''' new_spec = replace(layer_spec, block_size=new_block_size) + assert new_spec.page_size_bytes == max_page_size + new_kv_cache_spec[layer_name] = new_spec''' + +NEW = ''' new_spec = replace(layer_spec, block_size=new_block_size) + if new_spec.page_size_bytes != max_page_size: + # spark-dflash-unify: max_page_size is a *padded* hybrid page; the + # scaled attention page lands just under it. Pad the remainder while + # KEEPING the scaled block_size, exactly as get_kv_cache_groups does + # for HiddenStateCacheSpec. (The old patch dropped the scaling -> + # mis-strided the DFlash drafter -> accept len stuck ~1.47.) + new_spec = replace( + layer_spec, + block_size=new_block_size, + page_size_padded=max_page_size, + ) + assert new_spec.page_size_bytes == max_page_size + new_kv_cache_spec[layer_name] = new_spec''' + +if OLD not in src: + print("[patch_unify2] ERROR: target block not found — vLLM source differs", flush=True) + sys.exit(1) + +P.write_text(src.replace(OLD, NEW)) +print("[patch_unify2] patched OK (scaled-block + pad-remainder)", flush=True) diff --git a/runtime/serve.sh b/runtime/serve.sh new file mode 100755 index 0000000..3dad0d7 --- /dev/null +++ b/runtime/serve.sh @@ -0,0 +1,59 @@ +#!/bin/bash +# serve.sh — runs INSIDE the sm121 vLLM container (mounted at /host). Applies the +# runtime monkeypatches, then `vllm serve`s the Qwen3.5-122B-A10B INT4 target with +# the DFlash drafter. Driven by install.sh; can also be run by hand. +# +# args: $1 = num_speculative_tokens (0 = no-spec baseline) +# $2 = target attention backend (flash_attn | FLASHINFER) +# env: MODEL target path/repo (default Intel INT4; /model for hybrid) +# INC_HYBRID=1 apply the hybrid INT4+FP8 dense-expert dispatch patch +# INT8_LMHEAD_V3=1 apply the int8 lm-head GEMV patch +# MAX_MODEL_LEN GPU_MEM PORT +# +# Stack rationale: the DFlash drafter is non-causal -> needs FLASH_ATTN (FA2). The +# hybrid GDN+mamba+MoE target's KV page geometry won't absorb the drafter's +# attention spec without patch_unify2 (scale-block unify) + prefix-caching OFF +# (NoPrefixCache coordinator, dodges the hash assert). See docs/FINDINGS.md. +set -euo pipefail +NSPEC="${1:-12}" +BACKEND="${2:-flash_attn}" +MODEL="${MODEL:-Intel/Qwen3.5-122B-A10B-int4-AutoRound}" +DRAFT="${DRAFT:-z-lab/Qwen3.5-122B-A10B-DFlash}" +MAX_MODEL_LEN="${MAX_MODEL_LEN:-16384}" +GPU_MEM="${GPU_MEM:-0.8}" +PORT="${PORT:-8000}" + +# FLA sm121 big-tile shmem fix (prefill/TTFT only on sm121; harmless, free). +echo "[serve] FLA sm121 big-tile shmem patch" +python3 /host/patch_fla_shmem.py || true + +if [ "${INC_HYBRID:-0}" = "1" ]; then + echo "[serve] hybrid INT4+FP8 dispatch patch (inc.py)" + python3 /host/patch_inc_hybrid.py +fi +if [ "${INT8_LMHEAD_V3:-0}" = "1" ]; then + echo "[serve] int8 lm-head v3 patch (batched w8a16 GEMV)" + python3 /host/patch_int8_lmhead_v3.py +fi + +if [ "$NSPEC" = "0" ]; then + SPEC_ARG=() + echo "[serve] NO-SPEC baseline (identical flags, prefix-off)" +else + SPEC_ARG=(--speculative-config "{\"method\":\"dflash\",\"model\":\"$DRAFT\",\"num_speculative_tokens\":$NSPEC,\"attention_backend\":\"FLASH_ATTN\"}") + echo "[serve] DFlash n=$NSPEC, target-backend=$BACKEND, drafter=FLASH_ATTN, model=$MODEL" +fi +python3 /host/patch_unify2.py || { [ "$NSPEC" = "0" ] && true; } + +exec vllm serve "$MODEL" \ + --served-model-name qwen \ + --host 0.0.0.0 --port "$PORT" \ + --max-model-len "$MAX_MODEL_LEN" \ + --max-num-seqs 16 \ + --max-num-batched-tokens "$MAX_MODEL_LEN" \ + --gpu-memory-utilization "$GPU_MEM" \ + --no-enable-prefix-caching \ + --enable-chunked-prefill \ + --trust-remote-code \ + --attention-backend "$BACKEND" \ + "${SPEC_ARG[@]}" diff --git a/scripts/bench_albond.py b/scripts/bench_albond.py new file mode 100644 index 0000000..8644275 --- /dev/null +++ b/scripts/bench_albond.py @@ -0,0 +1,77 @@ +#!/usr/bin/env python3 +"""Faithful reproduction of albond's bench_qwen35.sh methodology so our DFlash/MTP +numbers are directly comparable to his reported 51.58 tok/s. + +His method (verbatim): non-streaming /v1/chat/completions, time the WHOLE request, +tok/s = completion_tokens / wall_time (INCLUDES prefill + overhead = END-TO-END). +5 prompts (Q&A 256 / Code 512 / JSON 1024 / Math 64 / LongCode 2048), temp 0, run 1 +discarded as JIT warmup. We add: decode-only tok/s isn't measured here on purpose +(his isn't either) + an aggregate spec-accept-len from /metrics deltas for context. +""" +import json, sys, time, urllib.request, statistics + +BASE = sys.argv[1] if len(sys.argv) > 1 else "http://127.0.0.1:8000" +MODEL = "qwen" +PROMPTS = [ + ("Q&A", "What are the main differences between TCP and UDP? Be concise.", 256), + ("Code", "Write a Python function that implements binary search on a sorted list. Include type hints and docstring.", 512), + ("JSON", "Generate a JSON array of 10 fictional employees with fields: name, age, department, salary, email, skills (array of 3). Output ONLY valid JSON, no explanation.", 1024), + ("Math", "What is 7823 * 4519? Show only the answer.", 64), + ("LongCode", "Write a complete Python implementation of a red-black tree with insert, delete, search, and in-order traversal. Include all rotation methods.", 2048), +] + + +def scrape(): + acc = dr = 0.0 + try: + txt = urllib.request.urlopen(BASE + "/metrics", timeout=10).read().decode() + for ln in txt.splitlines(): + if ln.startswith("#") or not ln.split(): + continue + v = float(ln.split()[-1]) + if "spec_decode_num_accepted_tokens_total" in ln: + acc += v + elif "spec_decode_num_drafts_total" in ln: + dr += v + except Exception: + pass + return acc, dr + + +def chat_e2e(prompt, max_tokens): + body = json.dumps({"model": MODEL, "messages": [{"role": "user", "content": prompt}], + "max_tokens": max_tokens, "temperature": 0.0}).encode() + req = urllib.request.Request(BASE + "/v1/chat/completions", data=body, + headers={"Content-Type": "application/json", "Authorization": "Bearer x"}) + t0 = time.perf_counter() + r = json.loads(urllib.request.urlopen(req, timeout=600).read()) + elapsed = time.perf_counter() - t0 + ct = r["usage"]["completion_tokens"] + return ct, elapsed, ct / elapsed if elapsed > 0 else 0.0 + + +def main(): + label = sys.argv[2] if len(sys.argv) > 2 else "server" + print(f"=== albond-method e2e bench :: {label} ===") + a0, d0 = scrape() + run2 = {} + for run in (1, 2): + tag = "WARMUP(discard)" if run == 1 else "RUN2" + for name, prompt, mt in PROMPTS: + try: + ct, el, tps = chat_e2e(prompt, mt) + except Exception as e: + print(f" [{name}] FAILED: {type(e).__name__}: {e}"); continue + if run == 2: + run2[name] = tps + print(f" {tag:16s} [{name:8s}] {ct:4d} tok in {el:6.2f}s = {tps:6.1f} tok/s (e2e)") + a1, d1 = scrape() + print(f"\n{label} RUN2 e2e tok/s: " + " ".join(f"{k}={v:.1f}" for k, v in run2.items())) + if run2: + print(f" cross-prompt mean (RUN2) = {statistics.mean(run2.values()):.1f} tok/s (albond reports 51.58)") + if d1 - d0 > 0: + print(f" aggregate spec accept length over bench = {1 + (a1-a0)/(d1-d0):.2f}") + + +if __name__ == "__main__": + main() diff --git a/scripts/bench_decode.py b/scripts/bench_decode.py new file mode 100644 index 0000000..64b8591 --- /dev/null +++ b/scripts/bench_decode.py @@ -0,0 +1,138 @@ +#!/usr/bin/env python3 +"""Single-stream (c=1) decode-rate benchmark for an OpenAI-compatible vLLM server. + +Measures *decode* tok/s (excludes TTFT/prefill) over N sequential requests, and +optionally scrapes vLLM /metrics for speculative-decode acceptance length. + +stdlib only — runs anywhere with python3, no pip installs. + + python3 bench_decode.py --base-url http://127.0.0.1:8000 \ + --model Intel/Qwen3.5-122B-A10B-int4-AutoRound \ + --max-tokens 256 --runs 5 --label "int4 baseline" +""" +import argparse, json, statistics, sys, time, urllib.request, urllib.error + +PROMPT = ("You are a careful writer. Write a detailed, continuous explanation of how " + "a modern mixture-of-experts transformer performs autoregressive decoding, " + "covering routing, KV cache, and memory bandwidth. Begin now:\n\n") + + +def post_stream(base_url, model, prompt, max_tokens, timeout): + """Stream /v1/completions; return (completion_tokens, t_first, t_last).""" + body = json.dumps({ + "model": model, "prompt": prompt, "max_tokens": max_tokens, + "temperature": 0.0, "stream": True, + "stream_options": {"include_usage": True}, + # force the full token budget so we measure steady-state decode + "ignore_eos": True, "min_tokens": max_tokens, + }).encode() + req = urllib.request.Request(base_url.rstrip("/") + "/v1/completions", + data=body, headers={"Content-Type": "application/json", + "Authorization": "Bearer x"}) + t_first = t_last = None + completion_tokens = 0 + chunks = 0 + with urllib.request.urlopen(req, timeout=timeout) as r: + for raw in r: + line = raw.decode("utf-8", "replace").strip() + if not line.startswith("data:"): + continue + data = line[5:].strip() + if data == "[DONE]": + break + try: + obj = json.loads(data) + except json.JSONDecodeError: + continue + usage = obj.get("usage") + if usage and usage.get("completion_tokens"): + completion_tokens = usage["completion_tokens"] + choices = obj.get("choices") or [] + if choices and choices[0].get("text"): + now = time.perf_counter() + if t_first is None: + t_first = now + t_last = now + chunks += 1 + if completion_tokens == 0: + completion_tokens = chunks # fallback: 1 chunk ~= 1 token + return completion_tokens, t_first, t_last + + +def scrape_metrics(base_url, timeout=10): + """Return dict of spec-decode counters from vLLM /metrics, if present.""" + keys = ("num_accepted_tokens", "num_draft_tokens", "num_drafts", + "accepted_tokens", "draft_tokens") + out = {} + try: + with urllib.request.urlopen(base_url.rstrip("/") + "/metrics", timeout=timeout) as r: + for line in r.read().decode("utf-8", "replace").splitlines(): + if line.startswith("#"): + continue + if "spec_decode" in line or "speculat" in line: + name, _, val = line.partition(" ") + try: + out[name] = out.get(name, 0.0) + float(val) + except ValueError: + pass + except (urllib.error.URLError, OSError): + pass + return out + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--base-url", default="http://127.0.0.1:8000") + ap.add_argument("--model", required=True) + ap.add_argument("--max-tokens", type=int, default=256) + ap.add_argument("--runs", type=int, default=5) + ap.add_argument("--warmup", type=int, default=1) + ap.add_argument("--timeout", type=float, default=600) + ap.add_argument("--label", default="") + ap.add_argument("--prompt", default=PROMPT) + args = ap.parse_args() + + m_before = scrape_metrics(args.base_url) + for _ in range(args.warmup): + try: + post_stream(args.base_url, args.model, args.prompt, 32, args.timeout) + except Exception as e: + print(f"warmup failed: {e}", file=sys.stderr); sys.exit(2) + + decode_tps, ttfts, e2e_tps = [], [], [] + for i in range(args.runs): + t0 = time.perf_counter() + toks, tf, tl = post_stream(args.base_url, args.model, args.prompt, args.max_tokens, args.timeout) + t1 = time.perf_counter() + if not toks or tf is None or tl is None or tl <= tf: + print(f" run {i}: degenerate (toks={toks})", file=sys.stderr); continue + dec = (toks - 1) / (tl - tf) + decode_tps.append(dec); ttfts.append((tf - t0) * 1000); e2e_tps.append(toks / (t1 - t0)) + print(f" run {i}: {toks} tok decode={dec:6.1f} tok/s ttft={ (tf-t0)*1000:6.0f} ms") + m_after = scrape_metrics(args.base_url) + + if not decode_tps: + print("no successful runs", file=sys.stderr); sys.exit(1) + print(f"\n=== {args.label or args.model} ===") + print(f"decode tok/s : median {statistics.median(decode_tps):.1f} " + f"mean {statistics.mean(decode_tps):.1f} min {min(decode_tps):.1f} max {max(decode_tps):.1f}") + print(f"ttft ms : median {statistics.median(ttfts):.0f}") + print(f"e2e tok/s : median {statistics.median(e2e_tps):.1f}") + + # spec-decode acceptance length, if counters moved + def delta(k): + return m_after.get(k, 0.0) - m_before.get(k, 0.0) + acc = next((delta(k) for k in m_after if "accepted" in k), 0.0) + drafts = next((delta(k) for k in m_after if "num_drafts" in k or ("draft" in k and "tokens" not in k)), 0.0) + dtoks = next((delta(k) for k in m_after if "draft_tokens" in k or "num_draft_tokens" in k), 0.0) + if acc or dtoks: + al = (acc / drafts) if drafts else float("nan") + rate = (acc / dtoks) if dtoks else float("nan") + print(f"spec accept : +{acc:.0f} accepted, +{dtoks:.0f} drafted, " + f"mean accept len ~{al:.2f}, accept rate ~{rate:.1%}") + else: + print("spec accept : (no spec-decode counters — baseline/no drafter)") + + +if __name__ == "__main__": + main() diff --git a/scripts/hermes_bench.py b/scripts/hermes_bench.py new file mode 100644 index 0000000..481b910 --- /dev/null +++ b/scripts/hermes_bench.py @@ -0,0 +1,168 @@ +#!/usr/bin/env python3 +"""Real-world DFlash/MTP acceptance bench on the actual Hermes agent workload. + +Reconstructs real conversation contexts from ~/.hermes/state.db and has the loaded +server regenerate the *next assistant turn* (so the model that produced the history is +irrelevant — only the realistic context matters). Measures mean acceptance length +(authoritative vLLM /metrics deltas) + decode tok/s, aggregated over N real turns. +Privacy: runs entirely on the box; prints only metrics, not conversation content. +""" +import argparse, json, sqlite3, sys, time, urllib.request, urllib.error + +DB = "/home/ent/.hermes/state.db" + + +def scrape(base): + out = {"acc": 0.0, "drafts": 0.0, "dtoks": 0.0} + try: + txt = urllib.request.urlopen(base.rstrip("/") + "/metrics", timeout=10).read().decode("utf-8", "replace") + except Exception: + return out + for ln in txt.splitlines(): + if ln.startswith("#"): + continue + v = ln.split()[-1] if ln.split() else "0" + try: + val = float(v) + except ValueError: + continue + if "spec_decode_num_accepted_tokens_total" in ln: + out["acc"] += val + elif "spec_decode_num_drafts_total" in ln: + out["drafts"] += val + elif "spec_decode_num_draft_tokens_total" in ln: + out["dtoks"] += val + return out + + +def build_messages(cur, session_id, system_prompt, char_budget=24000): + rows = list(cur.execute( + "SELECT role, content, tool_calls, tool_call_id, tool_name FROM messages " + "WHERE session_id=? ORDER BY id", (session_id,))) + # find the LAST assistant turn -> generate it; prompt = everything before it + last_asst = None + for i, r in enumerate(rows): + if r[0] == "assistant": + last_asst = i + if last_asst is None or last_asst == 0: + return None + pre = rows[:last_asst] + msgs = [] + for role, content, tool_calls, tool_call_id, tool_name in pre: + content = content or "" + if role == "user": + msgs.append({"role": "user", "content": content}) + elif role == "assistant": + m = {"role": "assistant", "content": content} + if tool_calls: + try: + tc = json.loads(tool_calls) + if isinstance(tc, list) and tc: + m["tool_calls"] = tc + if not content: + m["content"] = "" + except Exception: + pass + msgs.append(m) + elif role == "tool": + msgs.append({"role": "tool", "content": content, + "tool_call_id": tool_call_id or "call_0"}) + # skip session_meta + if not msgs: + return None + # truncate oldest non-system messages to fit budget + sys_msg = [{"role": "system", "content": system_prompt}] if system_prompt else [] + while sum(len(json.dumps(m)) for m in msgs) > char_budget and len(msgs) > 1: + msgs.pop(0) + return sys_msg + msgs + + +def stream_chat(base, model, messages, max_tokens, timeout): + body = json.dumps({"model": model, "messages": messages, "max_tokens": max_tokens, + "temperature": 0.0, "stream": True, + "stream_options": {"include_usage": True}}).encode() + req = urllib.request.Request(base.rstrip("/") + "/v1/chat/completions", data=body, + headers={"Content-Type": "application/json", "Authorization": "Bearer x"}) + t_first = t_last = None + toks = 0 + with urllib.request.urlopen(req, timeout=timeout) as r: + for raw in r: + ln = raw.decode("utf-8", "replace").strip() + if not ln.startswith("data:"): + continue + d = ln[5:].strip() + if d == "[DONE]": + break + try: + o = json.loads(d) + except json.JSONDecodeError: + continue + u = o.get("usage") + if u and u.get("completion_tokens"): + toks = u["completion_tokens"] + ch = o.get("choices") or [] + if ch and (ch[0].get("delta") or {}).get("content"): + now = time.perf_counter() + if t_first is None: + t_first = now + t_last = now + return toks, t_first, t_last + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--base-url", default="http://127.0.0.1:8000") + ap.add_argument("--model", default="qwen") + ap.add_argument("--max-tokens", type=int, default=200) + ap.add_argument("--n-samples", type=int, default=10) + ap.add_argument("--min-msgs", type=int, default=6) + ap.add_argument("--timeout", type=float, default=300) + ap.add_argument("--label", default="HERMES real") + args = ap.parse_args() + + db = sqlite3.connect(DB) + cur = db.cursor() + sess = list(cur.execute( + "SELECT s.id, s.system_prompt FROM sessions s " + "WHERE s.message_count >= ? ORDER BY s.started_at DESC LIMIT 40", (args.min_msgs,))) + + m0 = scrape(args.base_url) + decode_tps, used = [], 0 + for sid, sysp in sess: + if used >= args.n_samples: + break + try: + msgs = build_messages(cur, sid, sysp) + except Exception: + continue + if not msgs: + continue + try: + toks, tf, tl = stream_chat(args.base_url, args.model, msgs, args.max_tokens, args.timeout) + except (urllib.error.HTTPError, urllib.error.URLError, OSError) as e: + print(f" skip {sid[:24]}: {type(e).__name__}", file=sys.stderr) + continue + if toks and tf and tl and tl > tf: + dec = (toks - 1) / (tl - tf) + decode_tps.append(dec) + used += 1 + print(f" sample {used}: {toks} tok decode={dec:5.1f} tok/s ({sid[:28]})") + m1 = scrape(args.base_url) + + dacc = m1["acc"] - m0["acc"] + ddr = m1["drafts"] - m0["drafts"] + ddt = m1["dtoks"] - m0["dtoks"] + import statistics + print(f"\n=== {args.label} (n={used} real turns) ===") + if decode_tps: + print(f"decode tok/s : median {statistics.median(decode_tps):.1f} mean {statistics.mean(decode_tps):.1f} " + f"min {min(decode_tps):.1f} max {max(decode_tps):.1f}") + if ddr > 0: + print(f"accept length: {1 + dacc/ddr:.2f} (accepted {dacc:.0f} / drafts {ddr:.0f}; " + f"draft tokens {ddt:.0f}; accept rate {dacc/ddt:.1%})") + else: + print("accept length: (no draft activity captured)") + + +if __name__ == "__main__": + main() diff --git a/scripts/monitor.sh b/scripts/monitor.sh new file mode 100755 index 0000000..c7e1f68 --- /dev/null +++ b/scripts/monitor.sh @@ -0,0 +1,38 @@ +#!/usr/bin/env bash +# Monitor SGLang container startup with an OOM auto-kill guard. +# Breaks on: READY, ERROR (traceback), OOM-GUARD (avail mem too low), or timeout. +set -u +NAME="${1:-sglang-dflash}" +MAX_ITERS="${2:-60}" # 60 * 20s = 20 min +FLOOR_MB="${3:-4096}" # kill if available memory drops below this + +for i in $(seq 1 "$MAX_ITERS"); do + if ! docker ps --filter "name=$NAME" --format '{{.Names}}' | grep -q "$NAME"; then + echo "STATE=EXITED iter=$i" + echo "--- last log ---"; docker logs "$NAME" 2>&1 | tail -25 + exit 0 + fi + AVAIL=$(free -m | awk '/^Mem:/{print $7}') + if [ "$AVAIL" -lt "$FLOOR_MB" ]; then + echo "STATE=OOM-GUARD iter=$i avail_mb=$AVAIL -> killing $NAME" + docker kill "$NAME" >/dev/null 2>&1 + echo "--- last log ---"; docker logs "$NAME" 2>&1 | tail -25 + exit 0 + fi + LOG=$(docker logs "$NAME" 2>&1) + if echo "$LOG" | grep -qiE "server is fired up|Application startup complete|The server is ready"; then + echo "STATE=READY iter=$i avail_mb=$AVAIL" + echo "--- tail ---"; echo "$LOG" | tail -20 + exit 0 + fi + if echo "$LOG" | grep -qE "Traceback \(most recent call last\)|CUDA out of memory|RuntimeError|AssertionError|ValueError|raise NotImplementedError"; then + echo "STATE=ERROR iter=$i avail_mb=$AVAIL" + echo "--- tail ---"; echo "$LOG" | tail -40 + exit 0 + fi + LAST=$(echo "$LOG" | tail -1 | cut -c1-110) + echo "iter=$i avail_mb=$AVAIL :: $LAST" + sleep 20 +done +echo "STATE=TIMEOUT after $MAX_ITERS iters" +docker logs "$NAME" 2>&1 | tail -20 diff --git a/scripts/run_bank.sh b/scripts/run_bank.sh new file mode 100755 index 0000000..00eccfd --- /dev/null +++ b/scripts/run_bank.sh @@ -0,0 +1,16 @@ +#!/bin/bash +# Run the 4-workload acceptance/throughput bank on whichever server is up. +# $1 = container name (for the authoritative server-log accept_len). Uses the SAME +# prompts as the DFlash run so DFlash vs MTP is apples-to-apples. +CN="${1:-vllm-mtp}" +BU=http://127.0.0.1:8000; M=qwen +acc_tail(){ docker logs "$CN" 2>&1 | grep "Mean acceptance length" | tail -1 \ + | sed -E "s/.*Mean acceptance length: ([0-9.]+).*Per-position acceptance rate: ([0-9., ]+), Avg.*/accept_len=\1 per-pos=[\2]/"; } +run(){ L="$1"; P="$2"; echo "########## $L ##########" + python3 ~/bench_decode.py --base-url $BU --model $M --max-tokens 256 --runs 4 --warmup 1 --label "$L" --prompt "$P" 2>&1 | grep -E "decode tok/s" + echo " server $(acc_tail)"; } +run "PROSE" "You are a careful writer. Write a long flowing essay about the history and philosophy of science, with no lists. Begin: " +run "CODE" "Write a complete Python implementation of a binary search tree class with insert, search, delete, and inorder traversal. Include docstrings and type hints. Begin:\n\nclass BSTNode:" +run "COUNTING" "Write all the whole numbers from 1 to 400, separated by commas, with no other text. Begin: 1, 2, 3, " +echo "########## HERMES (real agent turns) ##########" +python3 ~/hermes_bench.py --base-url $BU --model $M --label "HERMES" --n-samples 10 --max-tokens 200 2>&1 | grep -E "decode tok/s|accept length" diff --git a/tools/build-hybrid-checkpoint.py b/tools/build-hybrid-checkpoint.py new file mode 100644 index 0000000..92efe22 --- /dev/null +++ b/tools/build-hybrid-checkpoint.py @@ -0,0 +1,463 @@ +#!/usr/bin/env python3 +""" +Build a hybrid GPTQ-INT4 + FP8 checkpoint for Qwen3.5-122B-A10B. + +Takes MoE expert weights from the GPTQ-INT4 checkpoint (0.5 bytes/param), +and dense layers (attention, shared experts, embeddings) from the official +FP8 checkpoint (1 byte/param + calibrated block scales). + +Result: a checkpoint that is ~9 GB smaller than the GPTQ-INT4 original +while using properly calibrated FP8 scales (not naive cast), yielding +better decode throughput on bandwidth-limited hardware. + +NOTE: Requires the hybrid FP8 dispatch patch from https://github.com/rmstxrx/vllm/tree/v0.17.1-hybrid-fp8 + +Usage: + python build-hybrid-checkpoint.py \ + --gptq-dir ~/inference/models/hf/qwen3.5-122b-a10b-gptq-int4 \ + --fp8-repo Qwen/Qwen3.5-122B-A10B-FP8 \ + --output ~/inference/models/hf/qwen3.5-122b-a10b-fp8hybrid +""" + +import argparse +import json +import logging +import shutil +from pathlib import Path + +import torch +from huggingface_hub import hf_hub_download +from safetensors import safe_open +from safetensors.torch import load_file, save_file + + +logger = logging.getLogger(__name__) + + +def get_fp8_non_expert_manifest(fp8_repo: str) -> dict[str, str]: + """Get the weight map for non-expert tensors from the FP8 checkpoint. + + Args: + fp8_repo: Hugging Face repo ID for the FP8 checkpoint. + + Returns: + Mapping of tensor names to shard filenames for non-expert tensors. + """ + idx_path = hf_hub_download(fp8_repo, "model.safetensors.index.json") + with open(idx_path, encoding="utf-8") as f: + idx = json.load(f) + + wm = idx["weight_map"] + return {k: v for k, v in wm.items() if ".experts." not in k} + + +def download_fp8_shards(fp8_repo: str, shards: set[str], cache_dir: Path) -> dict[str, Path]: + """Download only the needed FP8 shards. + + Args: + fp8_repo: Hugging Face repo ID for the FP8 checkpoint. + shards: Shard filenames that contain non-expert tensors. + cache_dir: Local cache directory for downloaded shards. + + Returns: + Mapping of shard filename to downloaded local path. + """ + shard_paths: dict[str, Path] = {} + for shard in sorted(shards): + logger.info(" Downloading %s...", shard) + path = hf_hub_download(fp8_repo, shard, local_dir=str(cache_dir)) + shard_paths[shard] = Path(path) + logger.info(" -> %s", path) + return shard_paths + + +def extract_fp8_tensors(shard_paths: dict[str, Path], wanted: dict[str, str]) -> dict[str, torch.Tensor]: + """Extract the requested FP8 tensors from downloaded shards. + + Args: + shard_paths: Mapping of shard filename to local path. + wanted: Mapping of tensor name to shard filename. + + Returns: + Mapping of tensor name to loaded FP8 tensor. + """ + tensors: dict[str, torch.Tensor] = {} + for shard_name, shard_path in shard_paths.items(): + keys_in_shard = [k for k, v in wanted.items() if v == shard_name] + if not keys_in_shard: + continue + + logger.info(" Extracting %d tensors from %s...", len(keys_in_shard), shard_name) + with safe_open(str(shard_path), framework="pt") as f: + for k in keys_in_shard: + tensors[k] = f.get_tensor(k) + + return tensors + + +def find_model_safetensors_files(directory: Path) -> list[Path]: + """Find model safetensors files in a GPTQ checkpoint directory. + + Args: + directory: Directory containing GPTQ checkpoint files. + + Returns: + Sorted list of shard paths, or a single-element list for single-file checkpoints. + + Raises: + FileNotFoundError: If no supported model safetensors files are found. + """ + gptq_shards = sorted(directory.glob("model.safetensors-*")) + if gptq_shards: + return gptq_shards + + # AutoRound / HF naming: model-NNNNN-of-NNNNN.safetensors + gptq_shards = sorted(directory.glob("model-*-of-*.safetensors")) + if gptq_shards: + return gptq_shards + + single_file = directory / "model.safetensors" + if single_file.is_file(): + return [single_file] + + raise FileNotFoundError( + f"No model.safetensors files found in {directory}. Expected a sharded " + "checkpoint (model.safetensors-NNNNN-of-NNNNN)." + ) + + +def validate_gptq_input(gptq_dir: Path) -> None: + """Validate the GPTQ checkpoint directory before any downloads. + + Args: + gptq_dir: Path to the local GPTQ checkpoint directory. + + Raises: + FileNotFoundError: If the directory or required files do not exist. + NotADirectoryError: If the path exists but is not a directory. + """ + if not gptq_dir.exists(): + raise FileNotFoundError(f"GPTQ directory does not exist: {gptq_dir}") + if not gptq_dir.is_dir(): + raise NotADirectoryError(f"GPTQ path is not a directory: {gptq_dir}") + if not any(path.is_file() for path in gptq_dir.glob("*.safetensors*")): + raise FileNotFoundError(f"No .safetensors files found in {gptq_dir}") + if not (gptq_dir / "config.json").is_file(): + raise FileNotFoundError(f"Missing config.json in {gptq_dir}") + + +def validate_output_dir(output_dir: Path, force: bool) -> None: + """Validate the output directory before building. + + Args: + output_dir: Path where the hybrid checkpoint will be written. + force: Whether destructive cleanup is allowed. + + Raises: + FileExistsError: If the output directory is non-empty without `force`. + NotADirectoryError: If the output path exists but is not a directory. + """ + if output_dir.exists() and not output_dir.is_dir(): + raise NotADirectoryError(f"Output path is not a directory: {output_dir}") + + if output_dir.exists() and any(output_dir.iterdir()) and not force: + raise FileExistsError( + f"Output directory {output_dir} exists and is not empty. Use --force " + "to remove existing model.safetensors* and config.json files before building." + ) + + +def prepare_output_dir(output_dir: Path, force: bool) -> None: + """Create or clean the output directory before writing files. + + Args: + output_dir: Path where the hybrid checkpoint will be written. + force: Whether destructive cleanup is allowed. + """ + if not output_dir.exists(): + output_dir.mkdir(parents=True, exist_ok=True) + return + + if not force: + return + + delete_targets = sorted( + { + path + for pattern in ("model.safetensors*", "config.json") + for path in output_dir.glob(pattern) + if path.is_file() + } + ) + for path in delete_targets: + path.unlink() + logger.info(" Deleted %s", path) + + +def build_hybrid_checkpoint( + gptq_dir: Path, + fp8_tensors: dict[str, torch.Tensor], + output_dir: Path, + force: bool, +) -> tuple[int, int, int]: + """Build the hybrid checkpoint from GPTQ and FP8 tensors. + + Args: + gptq_dir: Path to the local GPTQ checkpoint directory. + fp8_tensors: FP8 tensors keyed by tensor name. + output_dir: Output directory for rewritten checkpoint shards. + force: Whether to continue if many unexpected FP8 tensors are unmatched. + + Returns: + Tuple of replaced tensor count, added scale tensor count, and bytes saved. + + Raises: + FileNotFoundError: If no supported GPTQ model files are found. + RuntimeError: If too many FP8 tensors cannot be matched and `force` is not set. + ValueError: If a matched FP8 tensor has an incompatible shape. + """ + # 1. Copy non-safetensors files + for f in gptq_dir.iterdir(): + if f.suffix != ".safetensors" and not f.name.startswith("."): + if f.is_file(): + shutil.copy2(f, output_dir / f.name) + + # 2. Process each GPTQ shard + gptq_shards = find_model_safetensors_files(gptq_dir) + total_replaced = 0 + total_added = 0 + total_saved_bytes = 0 + + # Track which FP8 tensors have been placed + placed_fp8: set[str] = set() + + for i, shard_path in enumerate(gptq_shards): + shard_name = shard_path.name + logger.info(" [%d/%d] %s...", i + 1, len(gptq_shards), shard_name) + + gptq_tensors = load_file(str(shard_path)) + output_tensors: dict[str, torch.Tensor] = {} + replaced = 0 + + for name, tensor in gptq_tensors.items(): + if name in fp8_tensors: + # Replace BF16 tensor with FP8 version + fp8_tensor = fp8_tensors[name] + if tensor.shape != fp8_tensor.shape: + raise ValueError( + f"Shape mismatch for {name}: GPTQ={tensor.shape}, FP8={fp8_tensor.shape}" + ) + old_bytes = tensor.numel() * tensor.element_size() + new_bytes = fp8_tensor.numel() * fp8_tensor.element_size() + total_saved_bytes += old_bytes - new_bytes + output_tensors[name] = fp8_tensor + placed_fp8.add(name) + replaced += 1 + + # Also add the scale tensor if it exists + scale_name = name.replace(".weight", ".weight_scale_inv") + if scale_name in fp8_tensors and scale_name != name: + output_tensors[scale_name] = fp8_tensors[scale_name] + placed_fp8.add(scale_name) + total_added += 1 + else: + output_tensors[name] = tensor + + total_replaced += replaced + + save_file(output_tensors, str(output_dir / shard_name)) + logger.info(" replaced=%d, tensors=%d", replaced, len(output_tensors)) + + # 3. Check for any FP8 tensors not yet placed (e.g. scale tensors + # for weights that exist in shards we already processed) + unplaced = set(fp8_tensors.keys()) - placed_fp8 + if unplaced: + expected_unplaced = { + name + for name in unplaced + if name.endswith(".weight_scale_inv") + and f"{name.removesuffix('.weight_scale_inv')}.weight" in placed_fp8 + } + unexpected_unplaced = sorted(unplaced - expected_unplaced) + + for name in unexpected_unplaced: + logger.warning("WARNING: unexpected unmatched FP8 tensor %s", name) + + if len(unexpected_unplaced) > 5 and not force: + message = ( + f"{len(unexpected_unplaced)} FP8 tensors could not be matched to GPTQ tensors. " + "This likely indicates a naming mismatch between the GPTQ and FP8 checkpoints. " + "Use --force to proceed anyway." + ) + logger.error("ERROR: %s", message) + raise RuntimeError(message) + if len(unexpected_unplaced) > 5 and force: + logger.warning( + "WARNING: proceeding despite %d unexpected unmatched FP8 tensors because --force was provided", + len(unexpected_unplaced), + ) + + if expected_unplaced: + logger.info(" Adding %d expected unplaced FP8 scale tensors to final shard...", len(expected_unplaced)) + + if expected_unplaced: + # Load last shard, add expected scale tensors, re-save + last_shard = output_dir / gptq_shards[-1].name + existing = load_file(str(last_shard)) + for name in sorted(expected_unplaced): + existing[name] = fp8_tensors[name] + total_added += 1 + save_file(existing, str(last_shard)) + + return total_replaced, total_added, total_saved_bytes + + +def update_safetensors_index(output_dir: Path) -> None: + """Rebuild `model.safetensors.index.json` from actual shard contents. + + Args: + output_dir: Directory containing rewritten model shards. + """ + weight_map: dict[str, str] = {} + total_size = 0 + + for shard_path in find_model_safetensors_files(output_dir): + with safe_open(str(shard_path), framework="pt") as f: + for key in f.keys(): + weight_map[key] = shard_path.name + tensor = f.get_tensor(key) + total_size += tensor.numel() * tensor.element_size() + + index = { + "metadata": {"total_size": total_size}, + "weight_map": weight_map + } + + with open(output_dir / "model.safetensors.index.json", "w", encoding="utf-8") as f: + json.dump(index, f, indent=2, sort_keys=True) + + logger.info(" Index rebuilt: %d tensors, %.2f GB", len(weight_map), total_size / 1e9) + + +def update_config(output_dir: Path) -> None: + """Update `config.json` with hybrid quantization metadata. + + Args: + output_dir: Directory containing the hybrid checkpoint. + """ + config_path = output_dir / "config.json" + with open(config_path, encoding="utf-8") as f: + config = json.load(f) + + config["_hybrid_quant_info"] = { + "description": "Hybrid GPTQ-INT4 + FP8 checkpoint for single-GPU deployment", + "moe_experts": "GPTQ INT4 (group_size=128, sym=True, Marlin kernels)", + "dense_layers": "FP8 E4M3 block-128 (from official Qwen/Qwen3.5-122B-A10B-FP8, calibrated scales)", + "norms_gates_embeddings": "Preserved at source dtype (BF16 for norms/gates, FP8 for others)", + "source_gptq": "Qwen/Qwen3.5-122B-A10B-GPTQ-Int4", + "source_fp8": "Qwen/Qwen3.5-122B-A10B-FP8", + "vllm_patch": "https://github.com/rmstxrx/vllm/tree/v0.17.1-hybrid-fp8", + "target_hardware": "NVIDIA DGX Spark (GB10, 128GB unified, 273 GB/s)", + "converter": "build-hybrid-checkpoint.py" + } + + with open(config_path, "w", encoding="utf-8") as f: + json.dump(config, f, indent=2) + + +def main() -> int: + """Build a hybrid GPTQ-INT4 + FP8 checkpoint. + + Returns: + Process exit code. + """ + logging.basicConfig(level=logging.INFO, format="%(message)s") + parser = argparse.ArgumentParser(description="Build hybrid GPTQ-INT4 + FP8 checkpoint") + parser.add_argument("--gptq-dir", required=True, help="Path to GPTQ-INT4 model") + parser.add_argument("--fp8-repo", default="Qwen/Qwen3.5-122B-A10B-FP8", help="HF repo for FP8 model") + parser.add_argument("--output", required=True, help="Output directory") + parser.add_argument("--dry-run", action="store_true") + parser.add_argument( + "--force", + action="store_true", + help="Allow output cleanup and continue despite many unmatched FP8 tensors", + ) + args = parser.parse_args() + + gptq_dir = Path(args.gptq_dir) + output_dir = Path(args.output) + + validate_gptq_input(gptq_dir) + find_model_safetensors_files(gptq_dir) + validate_output_dir(output_dir, args.force) + + logger.info("=== Hybrid GPTQ-INT4 + FP8 Checkpoint Builder ===") + logger.info(" GPTQ source: %s", gptq_dir) + logger.info(" FP8 source: %s", args.fp8_repo) + logger.info(" Output: %s", output_dir) + logger.info("") + + # Step 1: Get manifest of non-expert tensors from FP8 checkpoint + logger.info("[1/5] Fetching FP8 tensor manifest...") + fp8_manifest = get_fp8_non_expert_manifest(args.fp8_repo) + shards_needed = set(fp8_manifest.values()) + logger.info(" Non-expert tensors: %d", len(fp8_manifest)) + logger.info(" Shards to download: %s", sorted(shards_needed)) + + if args.dry_run: + logger.info("") + logger.info("[DRY RUN] Would download shards and build hybrid. Exiting.") + return 0 + + prepare_output_dir(output_dir, args.force) + cache_dir = output_dir / ".fp8_cache" + cache_dir.mkdir(exist_ok=True) + + # Step 2: Download needed FP8 shards + logger.info("") + logger.info("[2/5] Downloading %d FP8 shards...", len(shards_needed)) + shard_paths = download_fp8_shards(args.fp8_repo, shards_needed, cache_dir) + + # Step 3: Extract non-expert FP8 tensors + logger.info("") + logger.info("[3/5] Extracting FP8 tensors...") + fp8_tensors = extract_fp8_tensors(shard_paths, fp8_manifest) + logger.info(" Extracted: %d tensors", len(fp8_tensors)) + + # Show dtype breakdown + dtypes: dict[str, int] = {} + for name, t in fp8_tensors.items(): + d = str(t.dtype) + dtypes[d] = dtypes.get(d, 0) + 1 + logger.info(" Dtypes: %s", dtypes) + + # Step 4: Build hybrid checkpoint + logger.info("") + logger.info("[4/5] Building hybrid checkpoint...") + replaced, added, saved = build_hybrid_checkpoint( + gptq_dir, + fp8_tensors, + output_dir, + args.force, + ) + + # Step 5: Update index and config + logger.info("") + logger.info("[5/5] Updating index and config...") + update_safetensors_index(output_dir) + update_config(output_dir) + + # Cleanup downloaded FP8 shards + shutil.rmtree(cache_dir, ignore_errors=True) + + logger.info("") + logger.info("=== Complete ===") + logger.info(" Tensors replaced (BF16→FP8): %d", replaced) + logger.info(" Scale tensors added: %d", added) + logger.info(" Bytes saved: %.2f GB", saved / 1e9) + logger.info(" Output: %s", output_dir) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/inspect_ckpt.py b/tools/inspect_ckpt.py new file mode 100644 index 0000000..7d2877b --- /dev/null +++ b/tools/inspect_ckpt.py @@ -0,0 +1,115 @@ +#!/usr/bin/env python3 +"""Inspect the Intel INT4 AutoRound checkpoint to decide if albond's hybrid +INT4+FP8 build will work: it swaps BF16 *dense* (non-expert) tensors for FP8 by +name. If AutoRound quantized the dense linears to INT4 (.qweight), the swap is a +near no-op. We need attention/shared_expert dense weights stored as BF16 (.weight). + +Also pulls the FP8 repo's index to confirm it exists and that names line up. +""" +import json +import sys +from collections import Counter +from pathlib import Path + +INT4_DIR = Path(sys.argv[1]) if len(sys.argv) > 1 else Path( + "/home/ent/.cache/huggingface/hub/models--Intel--Qwen3.5-122B-A10B-int4-AutoRound/" + "snapshots/3045d02bb737effc4581da91bddbad3be02934e4") +FP8_REPO = sys.argv[2] if len(sys.argv) > 2 else "Qwen/Qwen3.5-122B-A10B-FP8" + +idx = json.loads((INT4_DIR / "model.safetensors.index.json").read_text()) +wm = idx["weight_map"] +names = list(wm.keys()) + +# Non-expert = no ".experts." (routed experts stay INT4 in the hybrid) +nonexp = [n for n in names if ".experts." not in n] + +def bucket(n): + if ".experts." in n: + return "ROUTED-EXPERT" + if "shared_expert" in n: + return "shared_expert" + if "self_attn" in n or "linear_attn" in n or ".attn" in n: + return "attention" + if "embed_tokens" in n or "lm_head" in n: + return "embed/head" + if "mtp" in n: + return "mtp" + return "other" + +print(f"=== Intel INT4 checkpoint: {INT4_DIR.name} ===") +print(f"total tensors: {len(names)} non-expert: {len(nonexp)}") + +# Suffix histogram tells us quant scheme: .qweight/.scales/.qzeros => INT4; .weight => dense +suffix = Counter(n.rsplit(".", 1)[-1] for n in nonexp) +print("\nnon-expert tensor SUFFIX histogram (qweight/scales/qzeros = INT4-packed; weight = dense):") +for s, c in suffix.most_common(): + print(f" .{s:20s} {c}") + +# For each functional group, does it have .weight (BF16 dense) or .qweight (INT4)? +print("\nper-group quant scheme (sample names):") +groups = {} +for n in nonexp: + g = bucket(n) + groups.setdefault(g, {"weight": 0, "qweight": 0, "scale": 0, "other": 0, "ex": None}) + suf = n.rsplit(".", 1)[-1] + if suf == "weight": + groups[g]["weight"] += 1 + elif suf == "qweight": + groups[g]["qweight"] += 1 + elif "scale" in suf or suf in ("qzeros",): + groups[g]["scale"] += 1 + else: + groups[g]["other"] += 1 + if groups[g]["ex"] is None and suf in ("weight", "qweight"): + groups[g]["ex"] = n +for g, d in sorted(groups.items()): + scheme = "INT4(.qweight)" if d["qweight"] else ("DENSE(.weight)" if d["weight"] else "?") + print(f" {g:16s} weight={d['weight']:4d} qweight={d['qweight']:4d} scale={d['scale']:4d} -> {scheme}") + print(f" e.g. {d['ex']}") + +# Dtypes of a few non-expert .weight tensors (open the shard header only). +print("\ndtypes of sample non-expert '.weight' tensors (FP8 swap needs BF16 here):") +from safetensors import safe_open # noqa: E402 +sample = [n for n in nonexp if n.endswith(".weight") + and ("self_attn" in n or "shared_expert" in n or "embed" in n or "lm_head" in n)] +seen_shards = {} +shown = 0 +for n in sample: + shard = wm[n] + f = seen_shards.get(shard) + if f is None: + f = safe_open(str(INT4_DIR / shard), framework="pt") + seen_shards[shard] = f + try: + t = f.get_slice(n) + print(f" {n:60s} {t.get_dtype()} {tuple(t.get_shape())}") + except Exception as e: + print(f" {n:60s} ") + shown += 1 + if shown >= 12: + break + +# FP8 repo: confirm exists + list its non-expert names/dtypes for name-match sanity. +print(f"\n=== FP8 repo manifest: {FP8_REPO} ===") +try: + from huggingface_hub import hf_hub_download + p = hf_hub_download(FP8_REPO, "model.safetensors.index.json") + fidx = json.loads(Path(p).read_text()) + fwm = fidx["weight_map"] + fnon = [n for n in fwm if ".experts." not in n] + print(f"FP8 total tensors: {len(fwm)} non-expert: {len(fnon)}") + fsuf = Counter(n.rsplit('.', 1)[-1] for n in fnon) + print("FP8 non-expert suffix histogram:") + for s, c in fsuf.most_common(10): + print(f" .{s:20s} {c}") + # how many FP8 non-expert .weight names also exist in INT4 as .weight? + int4_weight = {n for n in nonexp if n.endswith('.weight')} + fp8_weight = {n for n in fnon if n.endswith('.weight')} + match = int4_weight & fp8_weight + print(f"\nname overlap (FP8 '.weight' that also exist as '.weight' in INT4): {len(match)} / {len(fp8_weight)} FP8 weights") + only_fp8 = sorted(fp8_weight - int4_weight)[:8] + print(f"FP8 '.weight' NOT present as '.weight' in INT4 (would not swap): {len(fp8_weight - int4_weight)}") + for n in only_fp8: + print(f" {n} (INT4 has: {'qweight' if n[:-7]+'.qweight' in set(names) else 'MISSING'})") +except Exception as e: + print(f"") diff --git a/tools/test_int8_lmhead2.py b/tools/test_int8_lmhead2.py new file mode 100644 index 0000000..7c9915b --- /dev/null +++ b/tools/test_int8_lmhead2.py @@ -0,0 +1,96 @@ +#!/usr/bin/env python3 +"""v2 harness: proper bf16 GEMV baseline + a single BATCHED int8 w8a16 GEMM kernel +(one launch for any B, dot-based, pads B to >=16) vs albond's per-row loop, at the +real lm-head shape and B in {1,5,13} (base + DFlash verify-batch sizes). +""" +import time + +import torch +import triton +import triton.language as tl + +V, H = 248320, 3072 +DEV = "cuda" +torch.manual_seed(0) + +W_bf16 = (torch.randn(V, H, device=DEV, dtype=torch.float32) * 0.02).to(torch.bfloat16) +scales = (W_bf16.float().abs().amax(dim=1) / 127.0).clamp(min=1e-12) +W_int8 = (W_bf16.float() / scales.unsqueeze(1)).round().clamp(-127, 127).to(torch.int8) +scales_f16 = scales.to(torch.float16) +INT8_BYTES = V * H +BF16_BYTES = V * H * 2 + + +def bench(fn, iters=30, warmup=5): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + t0 = time.perf_counter() + for _ in range(iters): + fn() + torch.cuda.synchronize() + return (time.perf_counter() - t0) / iters + + +# ---- batched int8 w8a16 GEMM: out[B,N] = (x[B,K] @ (W_int8[N,K]*s[N]).T) ---- +@triton.jit +def _k_batched(x_ptr, w_ptr, s_ptr, o_ptr, B, N, K, + sxb, sxk, swn, swk, sob, son, + BLOCK_B: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr): + pid_n = tl.program_id(0) + offs_b = tl.arange(0, BLOCK_B) + offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + offs_k = tl.arange(0, BLOCK_K) + x_ptrs = x_ptr + offs_b[:, None] * sxb + offs_k[None, :] * sxk + w_ptrs = w_ptr + offs_n[:, None] * swn + offs_k[None, :] * swk + acc = tl.zeros((BLOCK_B, BLOCK_N), dtype=tl.float32) + for k in range(0, K, BLOCK_K): + km = (offs_k[None, :] + k) < K + x = tl.load(x_ptrs, mask=(offs_b[:, None] < B) & km, other=0.0).to(tl.float16) + w = tl.load(w_ptrs, mask=(offs_n[:, None] < N) & km, other=0).to(tl.float16) + acc += tl.dot(x, w.T) # [BB,BK] @ [BK,BN] -> [BB,BN] + x_ptrs += BLOCK_K * sxk + w_ptrs += BLOCK_K * swk + s = tl.load(s_ptr + offs_n, mask=offs_n < N, other=0.0).to(tl.float32) + acc = acc * s[None, :] + o_ptrs = o_ptr + offs_b[:, None] * sob + offs_n[None, :] * son + tl.store(o_ptrs, acc.to(tl.float16), mask=(offs_b[:, None] < B) & (offs_n[None, :] < N)) + + +def run_batched(x, BLOCK_N=128, BLOCK_K=64, num_warps=4, num_stages=3): + B = x.shape[0] + BLOCK_B = max(16, triton.next_power_of_2(B)) + out = torch.empty(B, V, dtype=torch.float16, device=DEV) + xf = x.to(torch.float16) + grid = ((V + BLOCK_N - 1) // BLOCK_N,) + _k_batched[grid](xf, W_int8, scales_f16, out, B, V, H, + xf.stride(0), xf.stride(1), W_int8.stride(0), W_int8.stride(1), + out.stride(0), out.stride(1), + BLOCK_B=BLOCK_B, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, + num_warps=num_warps, num_stages=num_stages) + return out + + +for B in (1, 5, 13): + print(f"\n=== B={B} ===") + x = torch.randn(B, H, device=DEV, dtype=torch.bfloat16) * 0.1 + ref_bf16 = (x.float() @ W_bf16.float().T) + ref_deq = (x.float() @ (W_int8.float() * scales.unsqueeze(1)).T) + am_floor = (ref_bf16.argmax(-1) == ref_deq.argmax(-1)).float().mean().item() + print(f" quant floor: argmax_match={am_floor*100:.1f}% maxerr={(ref_deq-ref_bf16).abs().max():.4f}") + + # real bf16 GEMV baseline (stays bf16, reads 1.5GB) + dt = bench(lambda: torch.matmul(x, W_bf16.t())) + print(f" {'bf16 GEMV (real baseline)':30s} {'':33s}{dt*1e3:7.3f} ms {BF16_BYTES/dt/1e9:6.1f} GB/s") + + # batched int8 kernel, a few configs + for (bn, bk, nw, ns) in [(128, 64, 4, 3), (256, 64, 8, 3), (128, 128, 4, 3), (64, 128, 4, 3)]: + try: + out = run_batched(x, bn, bk, nw, ns) + err = (out.float() - ref_deq.float()).abs().max().item() + am = (out.float().argmax(-1) == ref_bf16.argmax(-1)).float().mean().item() + dt = bench(lambda: run_batched(x, bn, bk, nw, ns)) + print(f" batched N{bn}/K{bk}/w{nw}/s{ns:<2d} maxerr={err:8.4f} argmax={am*100:5.1f}% " + f"{dt*1e3:7.3f} ms {INT8_BYTES/dt/1e9:6.1f} GB/s") + except Exception as e: + print(f" batched N{bn}/K{bk}/w{nw}/s{ns}: ERR {str(e)[:70]}") diff --git a/tools/validate_inc_patch.py b/tools/validate_inc_patch.py new file mode 100644 index 0000000..20819c5 --- /dev/null +++ b/tools/validate_inc_patch.py @@ -0,0 +1,27 @@ +#!/usr/bin/env python3 +"""Validate patch_inc_hybrid.py landed on AEON 0.23's inc.py. Run AFTER the patch +in the same container: python3 /host/patch_inc_hybrid.py && python3 /host/validate_inc_patch.py +""" +import inspect + +import vllm.model_executor.layers.quantization.inc as m + +C = m.INCConfig +src = inspect.getsource(m) +print("import OK") +print("sentinel in source :", "spark-dflash-hybrid-fp8" in src) +print("maybe_update_config OWN :", "maybe_update_config" in C.__dict__) +print("_is_layer_fp8 OWN :", "_is_layer_fp8" in C.__dict__) +print("maybe_update_config hasattr:", hasattr(C, "maybe_update_config")) +print("_is_layer_fp8 hasattr :", hasattr(C, "_is_layer_fp8")) +# signature must accept hf_config kw (config/vllm.py calls it that way) +try: + sig = inspect.signature(C.maybe_update_config) + print("maybe_update_config sig :", str(sig)) + print("accepts hf_config kw :", "hf_config" in sig.parameters) +except Exception as e: + print("sig err:", e) +# count FP8 dispatch sites +print("Fp8LinearMethod dispatch ct:", src.count("return Fp8LinearMethod(self.fp8_config)")) +# byte-compile sanity already implied by import; show line count +print("inc.py lines :", len(src.splitlines())) diff --git a/tools/validate_int8_v3.py b/tools/validate_int8_v3.py new file mode 100644 index 0000000..7d95caf --- /dev/null +++ b/tools/validate_int8_v3.py @@ -0,0 +1,45 @@ +#!/usr/bin/env python3 +"""Validate patch_int8_lmhead_v3 landed + the helper runs on GPU. Run AFTER the +patch in the same container.""" +import torch + +import vllm.model_executor.layers.logits_processor as lp + +print("import OK | helpers:", + hasattr(lp, "_spark_int8_gemm"), + hasattr(lp, "_spark_int8_lmhead_apply"), + hasattr(lp, "_spark_k_int8")) +src = open(lp.__file__).read() +print("sentinel in _get_logits:", "DGX_SPARK_INT8_LMHEAD_V3: int8 w8a16" in src) + +V, H = 4096, 512 +torch.manual_seed(0) +W = torch.randn(V, H, device="cuda") * 0.02 +s = (W.abs().amax(1) / 127).clamp(min=1e-12) +wi = (W / s.unsqueeze(1)).round().clamp(-127, 127).to(torch.int8).contiguous() +sf = s.to(torch.float16) +for B in (1, 5, 13): + x = torch.randn(B, H, device="cuda", dtype=torch.bfloat16) * 0.1 + out = lp._spark_int8_gemm(x, wi, sf) + ref = x.float() @ (wi.float() * s.unsqueeze(1)).T + am = (out.argmax(-1) == ref.argmax(-1)).float().mean().item() + print(f" B={B}: out{tuple(out.shape)} {out.dtype} argmax={am*100:.0f}% " + f"maxerr={(out - ref).abs().max().item():.4f}") + +# Exercise the FULL apply path (print + quantize-once + gemm) with a mock lm_head +# (vocab > 100k to trigger the int8 path). self is unused -> None. +print("--- full _spark_int8_lmhead_apply path (mock lm_head, V=131072) ---") +Vbig = 131072 + +class _MockLMHead: + pass + +mh = _MockLMHead() +mh.weight = (torch.randn(Vbig, H, device="cuda", dtype=torch.float32) * 0.02).to(torch.bfloat16) +hs = torch.randn(3, H, device="cuda", dtype=torch.bfloat16) * 0.1 +o1 = lp._spark_int8_lmhead_apply(None, mh, hs, None) # first call: quantizes + prints +o2 = lp._spark_int8_lmhead_apply(None, mh, hs, None) # second call: reuses int8 +refb = hs.float() @ mh.weight.float().T +am = (o2.argmax(-1) == refb.argmax(-1)).float().mean().item() +print(f" apply: out{tuple(o2.shape)} {o2.dtype} argmax_vs_bf16={am*100:.0f}% " + f"int8_ready={getattr(mh, '_spark_int8_ready', None)} weight_kept={mh.weight.numel() > 0}")