mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[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:
@@ -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).
|
||||
|
||||
Reference in New Issue
Block a user