Skip to content
8 changes: 0 additions & 8 deletions src/speculators/models/dflash/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)",
Expand Down
21 changes: 12 additions & 9 deletions src/speculators/models/dflash/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -299,15 +300,15 @@ 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(
total_seq_len, dtype=torch.long, device=device
).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
Expand Down Expand Up @@ -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(
Expand All @@ -400,6 +402,7 @@ def forward(
verifier_last_hidden_states,
document_ids,
position_ids,
max_anchors=max_anchors,
**kwargs,
)
loss, metrics = compute_metrics(
Expand Down
6 changes: 5 additions & 1 deletion src/speculators/models/dspark/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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,
):
Expand All @@ -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.
Expand Down
24 changes: 0 additions & 24 deletions src/speculators/models/peagle/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
"""

Expand All @@ -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",
Expand Down
31 changes: 17 additions & 14 deletions src/speculators/models/peagle/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
):
"""
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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,
)

Expand All @@ -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
Expand All @@ -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",
Expand Down Expand Up @@ -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)
6 changes: 0 additions & 6 deletions tests/integration/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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(
Expand All @@ -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,
Expand All @@ -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",
Expand Down
6 changes: 3 additions & 3 deletions tests/integration/models/test_dflash_vllm_parity.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,18 +67,18 @@ 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()
return result

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,
Expand Down
Loading
Loading