[Feat]: Add Final Norm for vLLM Hidden Extractor (#1846)

### What does this PR do?

**Type of change:** Bug fix

vLLM captures the final-layer hidden state *before* the model's final
norm, but the
offline/streaming distillation path fed it straight into `lm_head`, so
the reconstructed
base logits (the KD target) were computed from un-normed hidden states.

This PR re-applies the base model's final norm before `lm_head` when the
producer declares
a pre-norm capture (`base_hidden_prenorm`), for both DFlash and EAGLE:

- Producer sets `base_hidden_prenorm` (streaming: `True`; offline: from
the dump).
- Consumer (`_maybe_apply_base_final_norm`) re-applies the base final
norm, and **fails loud**
if pre-norm is declared but the model's norm type isn't supported (no
silent corruption).
- `FakeBaseModel` now loads the base final norm (+
`rope_theta`/`rms_norm_eps`); norm type is
gated by an explicit `model_type` allowlist (gpt_oss excluded pending a
matching norm class).

### Testing

`tests/unit/torch/speculative/plugins/test_modeling_final_norm.py`;
DFlash/EAGLE streaming
training verified end-to-end.

- Backward compatible?: ✅ (post-norm captures declare
`base_hidden_prenorm=False` → unchanged)
- New tests: ✅


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Expanded speculative decoding/distillation to reconstruct missing
base-model logits by optionally applying a base model’s final
pre–LM-head normalization when pre-norm hidden states are provided.
* Streaming dataset now emits `base_hidden_prenorm` and can configure
RDMA backends from environment settings.
* **Bug Fixes**
  * Rejects mixed `base_hidden_prenorm` values within a batch.
  * Fails fast on streaming token-length mismatches.
* Streaming dataset loading supports directory inputs by expanding
sorted JSONL shards (and errors if none are found).
* **Tests**
* Added/updated unit and dataset tests for final-norm behavior and the
new batch field.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com>
This commit is contained in:
h-guo18
2026-07-07 05:07:33 +00:00
committed by GitHub
parent 32925cdf95
commit bc5bc1ac5f
17 changed files with 461 additions and 46 deletions
@@ -6,8 +6,10 @@
#
# data.mode=streaming sets dflash_offline so the DFlash module consumes streamed
# hidden states instead of running the fake base.
# Capture ids = [2,16,31,45,59,60] (kimi_k25/deepseek_v3, 61 layers): 5 DFlash
# target layers + base 60. n_captured = num_target_layers + 1.
# Capture ids = [2,16,31,45,59,61] (kimi_k25/deepseek_v3, 61 layers): 5 DFlash
# target layers + true final layer 61. n_captured = num_target_layers + 1.
# (61 = true final hidden; requires a vLLM with the aux-capture fix vllm#46788.
# Without it use 60 — the 2nd-to-last layer — which caps acceptance length.)
#
# answer_only_loss=true: Kimi ships only a slow tokenizer, so it can't derive the
# assistant mask the standard way (return_assistant_tokens_mask needs a fast
@@ -72,7 +74,7 @@ pipeline:
environment:
- HF_MODEL_CKPT: <<global_vars.hf_model>>
# No spaces in values: nemo_run emits `export FOO=value` unquoted.
- EAGLE_CAPTURE_IDS: "[2,16,31,45,59,60]"
- EAGLE_CAPTURE_IDS: "[2,16,31,45,59,61]"
- SERVE_TP: "4"
# Per-rank in-flight fetches; keep low so the cold NVFP4-MoE serve isn't flooded past its execute-model timeout (kills EngineCore).
- STREAMING_NUM_WORKERS: "1"
@@ -7,8 +7,10 @@
# on one 4-GPU node, so each serve replica owns a whole node.
#
# Capture ids: build_target_layer_ids(num_orig=61, num_draft=5)=[1,15,30,44,58]
# -> +1 for embedding = [2,16,31,45,59], append base 60 (final layer uncapturable).
# -> +1 for embedding = [2,16,31,45,59], append true final layer 61.
# 6 captured = 5 aux layers, matching the 5-layer DFlash draft block.
# (61 = true final hidden; requires a vLLM with the aux-capture fix vllm#46788.
# Without it use 60 — the 2nd-to-last layer — which caps acceptance length.)
#
# Run ON the cluster login node (paramiko can't reach the cluster through its login proxy):
# export SLURM_HOST=localhost SLURM_ACCOUNT=<your_account> \
@@ -72,7 +74,7 @@ pipeline:
environment:
- HF_MODEL_CKPT: <<global_vars.hf_model>>
# See header for derivation.
- EAGLE_CAPTURE_IDS: "[2,16,31,45,59,60]"
- EAGLE_CAPTURE_IDS: "[2,16,31,45,59,61]"
- SERVE_NODES: "2"
- SERVE_TP: "4"
# Per-rank in-flight fetches; keep low so the cold NVFP4-MoE serve isn't flooded past its execute-model timeout (kills EngineCore).
@@ -3,8 +3,9 @@
# Requires GB200: native NVFP4 + 192 GB/GPU fits the ~551 GB model at TP=4 on one node.
# node 0 = vllm serve (TP=4), node 1 = EAGLE3 trainer (fake base); 4 GPUs each.
#
# Capture ids: deepseek_v3 arch, 61 layers, indexed by layer input (0..60);
# [2,30,58] aux + [60] base (final layer not capturable).
# Capture ids: deepseek_v3 arch, 61 layers; [2,30,58] aux + [61] true final layer.
# (61 = true final hidden; requires a vLLM with the aux-capture fix vllm#46788.
# Without it use 60 — the 2nd-to-last layer — which caps acceptance length.)
#
# Run ON the cluster login node (paramiko can't reach the cluster through its login proxy):
# export SLURM_HOST=localhost SLURM_ACCOUNT=<your_account> \
@@ -58,7 +59,7 @@ pipeline:
environment:
- HF_MODEL_CKPT: <<global_vars.hf_model>>
# No spaces: nemo_run emits `export FOO=value` unquoted.
- EAGLE_CAPTURE_IDS: "[2,30,58,60]"
- EAGLE_CAPTURE_IDS: "[2,30,58,61]"
- SERVE_TP: "4"
# Per-rank in-flight fetches; keep low so the cold NVFP4-MoE serve isn't flooded past its execute-model timeout (kills EngineCore).
- STREAMING_NUM_WORKERS: "1"
@@ -4,8 +4,9 @@
# common/eagle3/train_eagle_streaming.sh header.
#
# Requires GB200: native NVFP4 + 192 GB/GPU fits ~551 GB Kimi at TP=4 on one node.
# Capture ids = [2,30,58] aux + [60] base = 4 (kimi_k25/deepseek_v3, 61 layers;
# layer 60 is the last capturable, used as base).
# Capture ids = [2,30,58] aux + [61] true final = 4 (kimi_k25/deepseek_v3, 61 layers;
# layer 61 is the true final hidden, used as base). 61 requires a vLLM with the
# aux-capture fix vllm#46788; without it use 60 — the 2nd-to-last layer (caps AL).
#
# Run ON the cluster login node (paramiko can't reach the cluster through its login proxy):
# export SLURM_HOST=localhost SLURM_ACCOUNT=<your_account> \
@@ -57,7 +58,7 @@ pipeline:
environment:
- HF_MODEL_CKPT: <<global_vars.hf_model>>
# No spaces in values: nemo_run emits `export FOO=value` unquoted.
- EAGLE_CAPTURE_IDS: "[2,30,58,60]"
- EAGLE_CAPTURE_IDS: "[2,30,58,61]"
- SERVE_NODES: "2"
- SERVE_TP: "4"
# Per-rank in-flight fetches; keep low so the cold NVFP4-MoE serve isn't flooded past its execute-model timeout (kills EngineCore).