--- license: apache-2.0 base_model: Qwen/Qwen3-4B tags: - block-attention - long-context - rag - kv-cache - efficient-inference --- # Qwen3-4B Block-Distilled (+ decode-side per-step router) A block-attention–distilled **Qwen3-4B** for efficient long-context / RAG serving, plus a light **decode-side router** that re-selects the active blocks **every generated token**. Distilled on one H200 with LoRA (merged into the released weights). `infer.py` is a runnable reference and the spec for a custom serving kernel (e.g. vLLM). ## Block attention Input layout `[system] [block_1] … [block_n] [query]`. Each context **block** attends only to itself and the system prefix (block-diagonal), **not** to other blocks — so a block's KV is context-independent and can be **computed once and reused across requests** (the main RAG serving win: no per-request document re-encoding). Each block is prefixed with **4 sink tokens** (`\n`) and carries a **summary** = its **last 8 content tokens**, which stay resident. The model is distilled so block-attention tracks a frozen full-attention teacher (KL + damage-weighted CE) — block ≈ full in quality. ## Decode paths (`--path`) All three share one block-diagonal prefill (query sees all blocks, KV cached once); they differ only at decode: | path | behavior | note | |---|---|---| | `dense` | decode attends to all blocks | quality reference | | `reuse-shift` | **select-once**: router picks top-k at step 0, RoPE-shifts them + query to compact positions, decodes reusing that compact KV | == static KV **pruning** | | `reroute` | **per-step**: all block KV stays resident; every step re-selects top-k from the current token's summary-attention and re-RoPE-shifts to compact positions | follows the block as it **moves** during generation | RoPE key-shift is exact (relative RoPE: `R(Δ)·R(p)k = R(p+Δ)k`), so cached post-RoPE keys are retargeted to compact positions without recompute (verified: offset-invariance 6/6 in fp32). ## What we validated (honest) Measured on this 4B model, real LongBench + a controlled probe: - **Machinery is exact.** With `k = all blocks`, `dense == reuse-shift == reroute` token-for-token (0 mismatch). - **On short-answer QA, static (pruning) already ≈ dense**, and per-step reroute ties it — the answer comes from a few blocks that don't move, so there's little for per-step to gain. (Real LongBench single-hop & multi-hop, n=100–150: reroute/static within ~1–2 examples of dense; EM identical.) - **Per-step reroute becomes *necessary* when the output spans many blocks.** In a controlled multi-target recall (12 blocks, recite a fact from 8 of them in order, k=3): static is structurally capped at **0.24** recall (it holds only k=3 blocks), while an **oracle per-step selector recovers 0.43 ≈ dense 0.46** — a ~2× gap that pruning cannot close because the relevant block moves across the generation. **This is the regime the per-step router targets** (long-form / multi-doc generation), and the contribution over static KV pruning. ## The router (`router.pt`) — work in progress Trained by **oracle-distillation**: it predicts, from the serve-time per-step summary-attention feature, the blocks the full-attention model actually attends to at that step. On the multi-target probe this lifts learned reroute from static's 0.24 toward the 0.43 oracle ceiling (first cut ~0.32; the static→oracle gap is the active research target). Small MLP over per-(layer,head) summary attention; `summary_tokens=8`, backbone frozen. `infer.py:load_router` reads `{in_dim, arch, state_dict, summary_tokens}`. ## Usage `router.weights` is the current **oracle-distilled per-step router** (load with `torch.load`; same format as a `.pt`, non-LFS name for reliable upload). The legacy `router.pt` is the older single-landmark linear router — prefer `router.weights`. ```bash python infer.py --model --router router.weights --k 3 --path reroute python infer.py --model --router router.weights --k 3 --path reuse-shift # static pruning baseline python infer.py --model --path dense # reference ``` ## Serving spec (vLLM etc.) `infer.py` is the reference: block-diagonal context KV store (reusable per doc chunk across requests), resident per-block summaries, per-step top-k block gather + RoPE compact-shift at decode. Savings are **KV-read bandwidth** (~k/n of blocks, plus resident summaries), realized in the **long-context × batched** regime where decode is KV-bound; at batch-1 short-context the win is small (decode is weight-bound). The bigger serving win is **prefill KV reuse** of document chunks across requests. ## Caveats - Block layout / 4× `\n` sink format must match training. See `infer.py`. - The per-step router is WIP: it beats static on span-output but has headroom to the oracle; short-answer QA does not need per-step at all (static suffices there). - Distilled with LoRA on segmented LongBench + SemanticSeg. Sibling: `hxia7/qwen3-14b-blockdist`.