Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
17 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion docs/cli/train.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
12 changes: 10 additions & 2 deletions scripts/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
8 changes: 8 additions & 0 deletions src/speculators/convert/eagle/eagle3_converter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -77,6 +79,8 @@ def convert(
eagle_config,
base_model,
norm_before_residual,
norm_before_fc,
norm_output,
eagle_aux_hidden_state_layer_ids,
)

Expand Down Expand Up @@ -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(
Expand All @@ -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),
Comment thread
orestis-z marked this conversation as resolved.
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,
)
Expand Down
13 changes: 10 additions & 3 deletions src/speculators/models/eagle3/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)."
),
)

Expand Down
8 changes: 6 additions & 2 deletions src/speculators/models/eagle3/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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(
Expand Down
4 changes: 4 additions & 0 deletions src/speculators/models/peagle/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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),
Expand Down
8 changes: 8 additions & 0 deletions tests/integration/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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",
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down
40 changes: 40 additions & 0 deletions tests/integration/models/test_model_forward.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()


Comment thread
orestis-z marked this conversation as resolved.
@requires_cuda
class TestPEagleParams:
@pytest.mark.parametrize("num_depths", [2, 4, 8])
Expand Down
89 changes: 89 additions & 0 deletions tests/unit/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
Unit tests for the config module in the Speculators library.
"""

import copy
import json
import tempfile
from pathlib import Path
Expand All @@ -19,6 +20,8 @@
VerifierConfig,
reload_schemas,
)
from speculators.models.eagle3 import Eagle3SpeculatorConfig
from speculators.proposals.greedy import GreedyTokenProposalConfig

# ===== TokenProposalConfig Tests =====

Expand Down Expand Up @@ -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
Loading