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 == reroutetoken-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.
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×
\nsink format must match training. Seeinfer.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.
- Downloads last month
- 93