diff --git a/docs/cli/train.md b/docs/cli/train.md index b02294adf..03c9b5a35 100644 --- a/docs/cli/train.md +++ b/docs/cli/train.md @@ -136,7 +136,9 @@ torchrun --standalone --nproc_per_node=4 scripts/train.py \ - **`--embed-requires-grad` / `--no-embed-requires-grad`** (flag, default: `False`) Whether to train embedding layer weights. -- **`--norm-before-fc`** (flag) Use RMSNorm before FC layer in draft path (e.g., for gpt-oss models). +- **`--norm-before-fc`** (flag, default: `False`) Use RMSNorm before FC layer in draft path (e.g., for Eagle 3.1 / gpt-oss models). + +- **`--norm-output`** (flag, default: `False`) Feed post-norm hidden states back across TTT steps to stabilize magnitude drift across speculation depths (Eagle 3.1). - **`--ttt-steps`** (int, default: `3`) Number of test-time training steps diff --git a/scripts/train.py b/scripts/train.py index c6dfbf1b1..545e5b52a 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -993,8 +993,16 @@ def parse_args(): parser.add_argument( "--norm-before-fc", action="store_true", - help="Use RMSNorm before fc in Eagle3 draft path " - "(e.g. for gpt-oss). Omit for other models.", + default=False, + help="Use RMSNorm before FC layer in draft path " + "(e.g., for Eagle 3.1 / gpt-oss models).", + ) + parser.add_argument( + "--norm-output", + action="store_true", + default=False, + help="Feed post-norm hidden states back across TTT steps to stabilize " + "magnitude drift across speculation depths (Eagle 3.1).", ) # D-Flash specific parameters parser.add_argument( diff --git a/src/speculators/convert/eagle/eagle3_converter.py b/src/speculators/convert/eagle/eagle3_converter.py index f66b6e3f6..98f417658 100644 --- a/src/speculators/convert/eagle/eagle3_converter.py +++ b/src/speculators/convert/eagle/eagle3_converter.py @@ -39,6 +39,8 @@ def convert( base_model: str, validate: bool = True, norm_before_residual: bool = False, + norm_before_fc: bool = False, + norm_output: bool = False, eagle_aux_hidden_state_layer_ids: list[int] | None = None, cache_dir: str | Path | None = None, ) -> None: @@ -77,6 +79,8 @@ def convert( eagle_config, base_model, norm_before_residual, + norm_before_fc, + norm_output, eagle_aux_hidden_state_layer_ids, ) @@ -106,6 +110,8 @@ def _build_eagle3_speculator_config( eagle_config: dict, base_model: str, norm_before_residual: bool = False, + norm_before_fc: bool = False, + norm_output: bool = False, eagle_aux_hidden_state_layer_ids: list[int] | None = None, ) -> Eagle3SpeculatorConfig: transformer_config = self._create_transformer_config_from_eagle( @@ -130,6 +136,8 @@ def _build_eagle3_speculator_config( speculators_config=speculators_config, draft_vocab_size=eagle_config.get("draft_vocab_size", 32000), norm_before_residual=norm_before_residual, + norm_before_fc=norm_before_fc or eagle_config.get("norm_before_fc", False), + norm_output=norm_output or eagle_config.get("norm_output", False), target_hidden_size=eagle_config.get("target_hidden_size"), eagle_aux_hidden_state_layer_ids=eagle_aux_hidden_state_layer_ids, ) diff --git a/src/speculators/models/eagle3/config.py b/src/speculators/models/eagle3/config.py index dbcdb7bdb..c394c2fd2 100644 --- a/src/speculators/models/eagle3/config.py +++ b/src/speculators/models/eagle3/config.py @@ -58,9 +58,16 @@ class Eagle3SpeculatorConfig(SpeculatorModelConfig): norm_before_fc: bool = Field( default=False, description=( - "If True, vLLM will add and apply RMSNorm before the fc layer when loading " - "this draft model (e.g. for gpt-oss draft checkpoints). Set in config when " - "converting or saving gpt-oss draft models." + "Use RMSNorm before FC layer in draft path " + "(e.g., for Eagle 3.1 / gpt-oss models)." + ), + ) + + norm_output: bool = Field( + default=False, + description=( + "Feed post-norm hidden states back across TTT steps to stabilize " + "magnitude drift across speculation depths (Eagle 3.1)." ), ) diff --git a/src/speculators/models/eagle3/core.py b/src/speculators/models/eagle3/core.py index b69cd03a0..fe51e7ff0 100644 --- a/src/speculators/models/eagle3/core.py +++ b/src/speculators/models/eagle3/core.py @@ -93,7 +93,6 @@ def __init__(self, config: Eagle3SpeculatorConfig): self.verifier_norm = norm_class(self.hidden_size, eps=tl_config.rms_norm_eps) self.verifier_norm.weight.requires_grad = False - # Normalize draft path input (gpt-oss only) if config.norm_before_fc: self.input_norm = self._model_definitions.norm_class( 3 * self.hidden_size, @@ -233,7 +232,11 @@ def forward( # noqa: C901 **kwargs, ) - logits = self.lm_head(self.norm(hidden_states)) + if self.config.norm_output: + hidden_states = self.norm(hidden_states) + logits = self.lm_head(hidden_states) + else: + logits = self.lm_head(self.norm(hidden_states)) # shape: [1, total_seq_len, draft_vocab_size] if return_loss: @@ -327,6 +330,7 @@ def from_training_args( draft_vocab_size=kwargs["draft_vocab_size"], norm_before_residual=kwargs["norm_before_residual"], norm_before_fc=kwargs.get("norm_before_fc", False), + norm_output=kwargs.get("norm_output", False), embed_requires_grad=kwargs.get("embed_requires_grad", False), eagle_aux_hidden_state_layer_ids=target_layer_ids, speculators_config=SpeculatorsConfig( diff --git a/src/speculators/models/peagle/core.py b/src/speculators/models/peagle/core.py index ebd633e18..c83f59ba9 100644 --- a/src/speculators/models/peagle/core.py +++ b/src/speculators/models/peagle/core.py @@ -115,6 +115,8 @@ def forward( ).unsqueeze(0) # [1, total_sampled, 3*hidden_size] # Project concatenated hidden states (3*hidden_size) -> hidden_size + if self.input_norm is not None: + sampled_hidden = self.input_norm(sampled_hidden) sampled_hidden = self.fc(sampled_hidden) # [1, total_sampled, hidden_size] layer_input = torch.cat( @@ -213,6 +215,8 @@ def from_training_args( transformer_layer_config=verifier_config, draft_vocab_size=kwargs["draft_vocab_size"], norm_before_residual=kwargs.get("norm_before_residual", False), + norm_before_fc=kwargs.get("norm_before_fc", False), + norm_output=kwargs.get("norm_output", False), eagle_aux_hidden_state_layer_ids=target_layer_ids, num_depths=kwargs.get("num_depths", 8), down_sample_ratio=kwargs.get("down_sample_ratio", 0.7), diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index acdab3e07..934d92e99 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -97,6 +97,8 @@ def make_eagle3_model( *, draft_vocab_size: int = 64, norm_before_residual: bool = False, + norm_before_fc: bool = False, + norm_output: bool = False, draft_attn_impl: str | None = None, device: str = "cuda:0", dtype: torch.dtype = torch.bfloat16, @@ -108,6 +110,8 @@ def make_eagle3_model( transformer_layer_config=transformer_config, draft_vocab_size=draft_vocab_size, norm_before_residual=norm_before_residual, + norm_before_fc=norm_before_fc, + norm_output=norm_output, embed_requires_grad=False, speculators_config=SpeculatorsConfig( algorithm="eagle3", @@ -166,6 +170,8 @@ def make_peagle_model( draft_vocab_size: int = 64, num_depths: int = 4, down_sample_ratio: float = 0.7, + norm_before_fc: bool = False, + norm_output: bool = False, draft_attn_impl: str | None = None, device: str = "cuda:0", dtype: torch.dtype = torch.bfloat16, @@ -177,6 +183,8 @@ def make_peagle_model( transformer_layer_config=transformer_config, draft_vocab_size=draft_vocab_size, norm_before_residual=False, + norm_before_fc=norm_before_fc, + norm_output=norm_output, embed_requires_grad=True, num_depths=num_depths, down_sample_ratio=down_sample_ratio, diff --git a/tests/integration/models/test_model_forward.py b/tests/integration/models/test_model_forward.py index 139d438dd..7647b53ce 100644 --- a/tests/integration/models/test_model_forward.py +++ b/tests/integration/models/test_model_forward.py @@ -356,6 +356,46 @@ def test_attention_backends_match(self, seq_lengths): ) +@requires_cuda +class TestNormOutputParams: + """Tests for Eagle 3.1: norm_before_fc + norm_output.""" + + def test_norm_output(self): + model = make_eagle3_model(norm_before_fc=True, norm_output=True) + assert model.input_norm is not None + samples = _make_samples([128]) + batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) + draft_tokens, loss, _metrics = model(**batch, ttt_steps=3) + + assert len(draft_tokens) == 3 + assert loss.isfinite() + loss.backward() + + def test_norm_output_without_norm_before_fc(self): + model = make_eagle3_model(norm_output=True) + assert model.input_norm is None + samples = _make_samples([128]) + batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) + draft_tokens, loss, _metrics = model(**batch, ttt_steps=3) + + assert len(draft_tokens) == 3 + assert loss.isfinite() + loss.backward() + + def test_peagle_norm_before_fc(self): + model = make_peagle_model() + assert model.input_norm is None + + model = make_peagle_model(norm_before_fc=True) + assert model.input_norm is not None + samples = _make_samples([128]) + batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) + _draft_tokens, loss, _metrics = model(**batch) + + assert loss.isfinite() + loss.backward() + + @requires_cuda class TestPEagleParams: @pytest.mark.parametrize("num_depths", [2, 4, 8]) diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py index 12102db22..3506da76b 100644 --- a/tests/unit/test_config.py +++ b/tests/unit/test_config.py @@ -2,6 +2,7 @@ Unit tests for the config module in the Speculators library. """ +import copy import json import tempfile from pathlib import Path @@ -19,6 +20,8 @@ VerifierConfig, reload_schemas, ) +from speculators.models.eagle3 import Eagle3SpeculatorConfig +from speculators.proposals.greedy import GreedyTokenProposalConfig # ===== TokenProposalConfig Tests ===== @@ -449,3 +452,89 @@ def test_speculator_model_config_from_pretrained_conversion(sample_speculators_c assert "Loading a non-speculator model config is not supported yet" in str( exc_info.value ) + + +# ===== Eagle3SpeculatorConfig Tests ===== + +TINY_LLAMA_CONFIG = LlamaConfig( + vocab_size=64, + hidden_size=32, + intermediate_size=128, + num_hidden_layers=1, + num_attention_heads=4, + num_key_value_heads=4, + head_dim=8, + max_position_embeddings=32, + rms_norm_eps=1e-6, + tie_word_embeddings=False, +) + + +def _make_eagle3_speculators_config(): + return SpeculatorsConfig( + algorithm="eagle3", + proposal_methods=[GreedyTokenProposalConfig(speculative_tokens=1)], + default_proposal_method="greedy", + verifier=VerifierConfig( + name_or_path=None, + architectures=["LlamaForCausalLM"], + ), + ) + + +@pytest.mark.sanity +def test_eagle3_config_norm_output_roundtrip(): + original = Eagle3SpeculatorConfig( + transformer_layer_config=copy.deepcopy(TINY_LLAMA_CONFIG), + draft_vocab_size=32000, + norm_before_residual=False, + norm_output=True, + speculators_config=_make_eagle3_speculators_config(), + ) + + config_dict = original.model_dump() + assert config_dict["norm_output"] is True + + recreated = Eagle3SpeculatorConfig.model_validate(config_dict) + assert recreated.norm_output is True + + +@pytest.mark.sanity +def test_eagle3_config_norm_output_dict_roundtrip(): + original = Eagle3SpeculatorConfig( + transformer_layer_config=copy.deepcopy(TINY_LLAMA_CONFIG), + draft_vocab_size=32000, + norm_output=True, + speculators_config=_make_eagle3_speculators_config(), + ) + + config_dict = original.to_dict() + assert config_dict["norm_output"] is True + + reloaded = SpeculatorModelConfig.from_dict(config_dict) + assert reloaded.norm_output is True + + +@pytest.mark.sanity +def test_eagle3_config_norm_output_pretrained_roundtrip(): + original = Eagle3SpeculatorConfig( + transformer_layer_config=copy.deepcopy(TINY_LLAMA_CONFIG), + draft_vocab_size=32000, + norm_output=True, + speculators_config=_make_eagle3_speculators_config(), + ) + + with tempfile.TemporaryDirectory() as tmp_dir: + original.save_pretrained(tmp_dir) + reloaded = SpeculatorModelConfig.from_pretrained(tmp_dir) + + assert reloaded.norm_output is True + + +@pytest.mark.sanity +def test_eagle3_config_norm_output_defaults(): + config = Eagle3SpeculatorConfig( + transformer_layer_config=copy.deepcopy(TINY_LLAMA_CONFIG), + speculators_config=_make_eagle3_speculators_config(), + ) + assert config.norm_output is False