diff --git a/src/speculators/models/dflash/config.py b/src/speculators/models/dflash/config.py index 7521c1595..0e504e3ab 100644 --- a/src/speculators/models/dflash/config.py +++ b/src/speculators/models/dflash/config.py @@ -48,14 +48,6 @@ class DFlashSpeculatorConfig(SpeculatorModelConfig): ), ) - max_anchors: int = Field( - default=256, - description=( - "Maximum number of anchor positions to sample during training " - "(controls memory usage and training efficiency)" - ), - ) - target_hidden_size: int | None = Field( default=None, description="Hidden size of the target model (if different from draft model)", diff --git a/src/speculators/models/dflash/core.py b/src/speculators/models/dflash/core.py index fe70569ae..af04c5222 100644 --- a/src/speculators/models/dflash/core.py +++ b/src/speculators/models/dflash/core.py @@ -131,7 +131,6 @@ def from_training_args( **kwargs: Training arguments with DFlash-specific params - draft_vocab_size: Size of draft vocabulary - block_size: Block size for draft predictions (default: 8) - - max_anchors: Max anchor positions during training (default: 256) - verifier_name_or_path: Path to verifier model Returns: @@ -179,9 +178,6 @@ def _build_base_config_kwargs( "transformer_layer_config": verifier_config, "draft_vocab_size": kwargs["draft_vocab_size"], "block_size": block_size, - "max_anchors": ( - 3072 if kwargs.get("max_anchors") is None else kwargs["max_anchors"] - ), "aux_hidden_state_layer_ids": target_layer_ids, "mask_token_id": kwargs.get("mask_token_id"), "sliding_window_non_causal": kwargs.get("sliding_window_non_causal", False), @@ -210,7 +206,12 @@ def get_trainer_kwargs(**kwargs) -> tuple[dict, dict]: """ loss_config = resolve_loss_config(kwargs["loss_fn"]) gamma = kwargs.get("dflash_decay_gamma", 4.0) - shared = {"loss_config": loss_config, "gamma": gamma} + max_anchors = kwargs.get("max_anchors", 3072) + shared = { + "loss_config": loss_config, + "gamma": gamma, + "max_anchors": max_anchors, + } return dict(shared), dict(shared) @property @@ -251,11 +252,11 @@ def _create_attention_mask( ) @torch.compiler.disable - def _build_attention_mask(self, loss_mask, document_ids, device): + def _build_attention_mask(self, loss_mask, max_anchors, document_ids, device): total_seq_len = loss_mask.shape[1] anchor_positions, anchor_valid = select_anchors( - loss_mask, self.config.max_anchors, self.block_size + loss_mask, max_anchors, self.block_size ) full_attn_mask = None @@ -299,7 +300,7 @@ def _backbone_forward( """ device = hidden_states.device total_seq_len = hidden_states.shape[1] - num_anchors = self.config.max_anchors + num_anchors = kwargs.pop("max_anchors", 3072) if position_ids is None: position_ids = torch.arange( @@ -307,7 +308,7 @@ def _backbone_forward( ).unsqueeze(0) full_attn_mask, sliding_window_attn_mask, anchor_positions, anchor_valid = ( - self._build_attention_mask(loss_mask, document_ids, device) + self._build_attention_mask(loss_mask, num_anchors, document_ids, device) ) mask_tokens_size = num_anchors * self.block_size @@ -391,6 +392,7 @@ def forward( position_ids: torch.Tensor | None = None, # shape: [1, total_seq_len] loss_config: LossConfig | None = None, gamma: float = 4.0, + max_anchors: int = 3072, **kwargs, ): _, logits, targets, aligned_loss_mask, _ = self._backbone_forward( @@ -400,6 +402,7 @@ def forward( verifier_last_hidden_states, document_ids, position_ids, + max_anchors=max_anchors, **kwargs, ) loss, metrics = compute_metrics( diff --git a/src/speculators/models/dspark/core.py b/src/speculators/models/dspark/core.py index 01df7e5f8..6c4726a49 100644 --- a/src/speculators/models/dspark/core.py +++ b/src/speculators/models/dspark/core.py @@ -82,10 +82,12 @@ def get_trainer_kwargs(**kwargs) -> tuple[dict, dict]: """Resolve DSpark's compound loss from ``--loss-fn``.""" loss_config = resolve_loss_config(kwargs["loss_fn"]) gamma = kwargs.get("dflash_decay_gamma", 4.0) + max_anchors = kwargs.get("max_anchors", 3072) confidence_head_alpha = kwargs.get("confidence_head_alpha", 1.0) shared = { "loss_config": loss_config, "gamma": gamma, + "max_anchors": max_anchors, "confidence_head_alpha": confidence_head_alpha, } return dict(shared), dict(shared) @@ -101,6 +103,7 @@ def forward( position_ids: torch.Tensor | None = None, # [1, total_seq_len] loss_config: LossConfig | None = None, gamma: float = 4.0, + max_anchors: int = 3072, confidence_head_alpha: float = 1.0, **kwargs, ): @@ -112,12 +115,13 @@ def forward( verifier_last_hidden_states, document_ids, position_ids, + max_anchors=max_anchors, **kwargs, ) ) # DSpark: add the Markov logit bias and predict per-position confidence. - num_blocks = self.config.max_anchors + num_blocks = max_anchors block = self.block_size mask_tokens_size = num_blocks * block # Ground-truth block tokens (verifier vocab); position 0 is the anchor. diff --git a/src/speculators/models/peagle/config.py b/src/speculators/models/peagle/config.py index 6c860720a..15e65bc90 100644 --- a/src/speculators/models/peagle/config.py +++ b/src/speculators/models/peagle/config.py @@ -18,9 +18,6 @@ class PEagleSpeculatorConfig(Eagle3SpeculatorConfig): P-EAGLE extends EAGLE-3 with parallel multi-token prediction using Conditional Drop Token (COD) sampling for memory-efficient training. - :param num_depths: Number of parallel prediction groups (typically 8) - :param down_sample_ratio: Geometric decay ratio for COD sampling (r in [0,1]) - :param down_sample_ratio_min: Minimum retention ratio floor :param mask_token_id: Token ID used for masking """ @@ -30,27 +27,6 @@ class PEagleSpeculatorConfig(Eagle3SpeculatorConfig): description="Model architectures that can load these weights", ) - num_depths: int = Field( - default=8, - description="Number of parallel prediction groups (num_depths)", - ge=1, - le=16, - ) - - down_sample_ratio: float = Field( - default=0.7, - description="Geometric decay ratio for COD sampling (retention rate r)", - gt=0.0, - le=1.0, - ) - - down_sample_ratio_min: float = Field( - default=0.2, - description="Minimum retention ratio floor to prevent over-sampling", - gt=0.0, - le=1.0, - ) - mask_token_id: int | None = Field( default=None, description="Token ID used for padding unused positions in parallel groups", diff --git a/src/speculators/models/peagle/core.py b/src/speculators/models/peagle/core.py index 88c84a04a..2d819b9be 100644 --- a/src/speculators/models/peagle/core.py +++ b/src/speculators/models/peagle/core.py @@ -38,9 +38,6 @@ def __init__( ): super().__init__(config=config) - self.num_depths = config.num_depths - self.down_sample_ratio = config.down_sample_ratio - self.down_sample_ratio_min = config.down_sample_ratio_min self.mask_token_id = config.mask_token_id # Learnable mask_hidden parameter for padding unsampled positions @@ -64,6 +61,9 @@ def forward( verifier_last_hidden_states: torch.Tensor | None = None, loss_config: LossConfig | None = None, max_anchors: int | None = None, + num_depths: int = 8, + down_sample_ratio: float = 0.7, + down_sample_ratio_min: float = 0.2, **kwargs, ): """ @@ -95,9 +95,9 @@ def forward( anchor_pos, depth = generate_cod_sample_indices( seq_length=seq_length, loss_mask=loss_mask, - num_depths=self.num_depths, - down_sample_ratio=self.down_sample_ratio, - down_sample_ratio_min=self.down_sample_ratio_min, + num_depths=num_depths, + down_sample_ratio=down_sample_ratio, + down_sample_ratio_min=down_sample_ratio_min, max_anchors=max_anchors, ) total_sampled = anchor_pos.shape[0] @@ -184,7 +184,7 @@ def forward( loss_mask=loss_mask, anchor_pos=anchor_pos, depth=depth, - num_depths=self.num_depths, + num_depths=num_depths, loss_config=loss_config, ) @@ -206,9 +206,6 @@ def from_training_args( **kwargs: Training arguments with P-EAGLE-specific params - draft_vocab_size: Size of draft vocabulary - norm_before_residual: Whether to normalize before residual - - num_depths: Number of parallel groups (default 8) - - down_sample_ratio: COD sampling ratio (default 0.7) - - down_sample_ratio_min: Minimum sampling ratio (default 0.2) - mask_token_id: Mask token ID - t2d: Target-to-draft vocabulary mapping - d2t: Draft-to-target vocabulary mapping @@ -234,9 +231,6 @@ def from_training_args( fc_norm=kwargs.get("fc_norm", 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), - down_sample_ratio_min=kwargs.get("down_sample_ratio_min", 0.2), mask_token_id=kwargs.get("mask_token_id"), speculators_config=SpeculatorsConfig( algorithm="peagle", @@ -270,5 +264,14 @@ def get_trainer_kwargs(**kwargs) -> tuple[dict, dict]: """ loss_config = resolve_loss_config(kwargs["loss_fn"]) max_anchors = kwargs.get("max_anchors") - shared = {"loss_config": loss_config, "max_anchors": max_anchors} + num_depths = kwargs.get("num_depths", 8) + down_sample_ratio = kwargs.get("down_sample_ratio", 0.7) + down_sample_ratio_min = kwargs.get("down_sample_ratio_min", 0.2) + shared = { + "loss_config": loss_config, + "max_anchors": max_anchors, + "num_depths": num_depths, + "down_sample_ratio": down_sample_ratio, + "down_sample_ratio_min": down_sample_ratio_min, + } return dict(shared), dict(shared) diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 55df572eb..0a1e0459e 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -134,7 +134,6 @@ def make_dflash_model( *, draft_vocab_size: int = 64, block_size: int = 4, - max_anchors: int = 8, draft_attn_impl: str | None = None, device: str = "cuda:0", dtype: torch.dtype = torch.bfloat16, @@ -147,7 +146,6 @@ def make_dflash_model( transformer_layer_config=transformer_config, draft_vocab_size=draft_vocab_size, block_size=block_size, - max_anchors=max_anchors, aux_hidden_state_layer_ids=[0, 1, 2], mask_token_id=0, speculators_config=SpeculatorsConfig( @@ -171,7 +169,6 @@ def make_peagle_model( *, draft_vocab_size: int = 64, num_depths: int = 4, - down_sample_ratio: float = 0.7, norm_before_fc: bool = False, fc_norm: bool = False, norm_output: bool = False, @@ -190,9 +187,6 @@ def make_peagle_model( fc_norm=fc_norm, norm_output=norm_output, embed_requires_grad=True, - num_depths=num_depths, - down_sample_ratio=down_sample_ratio, - down_sample_ratio_min=0.2, mask_token_id=0, speculators_config=SpeculatorsConfig( algorithm="peagle", diff --git a/tests/integration/models/test_dflash_vllm_parity.py b/tests/integration/models/test_dflash_vllm_parity.py index 8387f931d..164de04eb 100644 --- a/tests/integration/models/test_dflash_vllm_parity.py +++ b/tests/integration/models/test_dflash_vllm_parity.py @@ -67,10 +67,10 @@ class TestDFlashAttentionParity: def _run_backend(self, backend, batch, state): torch.manual_seed(42) - m = make_dflash_model(block_size=4, max_anchors=4, draft_attn_impl=backend) + m = make_dflash_model(block_size=4, draft_attn_impl=backend) m.load_state_dict(state) with torch.no_grad(): - _, loss, _ = m(**batch) + _, loss, _ = m(**batch, max_anchors=4) result = loss.item() del m torch.cuda.empty_cache() @@ -78,7 +78,7 @@ def _run_backend(self, backend, batch, state): def _make_shared_inputs(self): torch.manual_seed(42) - ref = make_dflash_model(block_size=4, max_anchors=4, draft_attn_impl="eager") + ref = make_dflash_model(block_size=4, draft_attn_impl="eager") samples = [ make_sample( seq_len=64, diff --git a/tests/integration/models/test_model_forward.py b/tests/integration/models/test_model_forward.py index 876221617..228bdbf03 100644 --- a/tests/integration/models/test_model_forward.py +++ b/tests/integration/models/test_model_forward.py @@ -66,11 +66,21 @@ class ModelSpec: batch_factory: Callable[..., Any] = make_batch -DFLASH_SPEC = ModelSpec(name="dflash", factory=make_dflash_model) +DFLASH_SPEC = ModelSpec( + name="dflash", factory=make_dflash_model, forward_kwargs={"max_anchors": 8} +) EAGLE3_SPEC = ModelSpec( name="eagle3", factory=make_eagle3_model, forward_kwargs={"ttt_steps": 2} ) -PEAGLE_SPEC = ModelSpec(name="peagle", factory=make_peagle_model) +PEAGLE_SPEC = ModelSpec( + name="peagle", + factory=make_peagle_model, + forward_kwargs={ + "num_depths": 4, + "down_sample_ratio": 0.7, + "down_sample_ratio_min": 0.2, + }, +) MTP_SPEC = ModelSpec( name="mtp", factory=make_mtp_model, @@ -240,20 +250,20 @@ def test_boundary_tokens(self, draft_vocab_model): class TestDFlashParams: @pytest.mark.parametrize("block_size", [2, 4, 8]) def test_varying_block_size(self, block_size): - model = make_dflash_model(block_size=block_size, max_anchors=4) + model = make_dflash_model(block_size=block_size) samples = _make_samples([128]) batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) - draft_tokens, loss, metrics = model(**batch) + draft_tokens, loss, metrics = model(**batch, max_anchors=4) assert loss.isfinite() loss.backward() @pytest.mark.parametrize("max_anchors", [2, 8, 16]) def test_varying_max_anchors(self, max_anchors): - model = make_dflash_model(max_anchors=max_anchors) + model = make_dflash_model() samples = _make_samples([128]) batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) - draft_tokens, loss, metrics = model(**batch) + draft_tokens, loss, metrics = model(**batch, max_anchors=max_anchors) assert loss.isfinite() loss.backward() @@ -264,7 +274,7 @@ def test_attention_backend(self, draft_attn_impl, seq_lengths): model = make_dflash_model(draft_attn_impl=draft_attn_impl) samples = _make_samples(seq_lengths) batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) - draft_tokens, loss, metrics = model(**batch) + draft_tokens, loss, metrics = model(**batch, max_anchors=8) assert loss.isfinite() loss.backward() @@ -283,7 +293,7 @@ def test_attention_backends_match(self, seq_lengths): batch = make_batch( max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE ) - _, loss, _ = model(**batch) + _, loss, _ = model(**batch, max_anchors=8) results[backend] = loss.detach().cpu() del model torch.cuda.empty_cache() @@ -401,7 +411,7 @@ def test_peagle_fc_norm(self): assert len(model.fc_norm) == 3 samples = _make_samples([128]) batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) - _draft_tokens, loss, _metrics = model(**batch) + _draft_tokens, loss, _metrics = model(**batch, num_depths=4) assert loss.isfinite() loss.backward() @@ -414,7 +424,7 @@ def test_peagle_norm_before_fc(self): 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) + _draft_tokens, loss, _metrics = model(**batch, num_depths=4) assert loss.isfinite() loss.backward() @@ -427,17 +437,19 @@ def test_varying_num_depths(self, num_depths): model = make_peagle_model(num_depths=num_depths) samples = _make_samples([128]) batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) - draft_tokens, loss, metrics = model(**batch) + draft_tokens, loss, metrics = model(**batch, num_depths=num_depths) assert loss.isfinite() loss.backward() @pytest.mark.parametrize("down_sample_ratio", [0.3, 0.7, 1.0]) def test_varying_down_sample_ratio(self, down_sample_ratio): - model = make_peagle_model(down_sample_ratio=down_sample_ratio) + model = make_peagle_model() samples = _make_samples([128]) batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) - draft_tokens, loss, metrics = model(**batch) + draft_tokens, loss, metrics = model( + **batch, num_depths=4, down_sample_ratio=down_sample_ratio + ) assert loss.isfinite() loss.backward() @@ -448,7 +460,7 @@ def test_attention_backend(self, draft_attn_impl, seq_lengths): model = make_peagle_model(draft_attn_impl=draft_attn_impl) samples = _make_samples(seq_lengths) batch = make_batch(max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE) - draft_tokens, loss, metrics = model(**batch) + draft_tokens, loss, metrics = model(**batch, num_depths=4) assert loss.isfinite() loss.backward() @@ -467,7 +479,7 @@ def test_attention_backends_match(self, seq_lengths): batch = make_batch( max_len=MAX_LEN, samples=samples, hidden_size=HIDDEN_SIZE ) - _, loss, _ = model(**batch) + _, loss, _ = model(**batch, num_depths=4) results[backend] = loss.detach().cpu() del model torch.cuda.empty_cache()