91 lines
5.0 KiB
Markdown
91 lines
5.0 KiB
Markdown
---
|
||
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 <this-repo> --router router.weights --k 3 --path reroute
|
||
python infer.py --model <this-repo> --router router.weights --k 3 --path reuse-shift # static pruning baseline
|
||
python infer.py --model <this-repo> --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`.
|