diff --git a/src/speculators/convert/eagle/__init__.py b/src/speculators/convert/eagle/__init__.py index 94d6e15db..31cf9e3a0 100644 --- a/src/speculators/convert/eagle/__init__.py +++ b/src/speculators/convert/eagle/__init__.py @@ -3,6 +3,5 @@ """ from speculators.convert.eagle.eagle3_converter import Eagle3Converter -from speculators.convert.eagle.eagle_converter import EagleConverter -__all__ = ["Eagle3Converter", "EagleConverter"] +__all__ = ["Eagle3Converter"] diff --git a/src/speculators/convert/eagle/eagle_converter.py b/src/speculators/convert/eagle/eagle_converter.py deleted file mode 100644 index 2b2e110fa..000000000 --- a/src/speculators/convert/eagle/eagle_converter.py +++ /dev/null @@ -1,347 +0,0 @@ -""" -Eagle checkpoint converter with loguru logging. -""" - -from pathlib import Path - -import torch -from loguru import logger -from transformers import LlamaConfig, PretrainedConfig - -from speculators.config import SpeculatorsConfig, VerifierConfig -from speculators.convert.eagle.eagle_legacy_model import ( - EagleSpeculator, - EagleSpeculatorConfig, -) -from speculators.convert.eagle.utils import ( - build_llama_config_rope_kwargs, - detect_fusion_bias_and_layernorms, -) -from speculators.convert.utils import ( - ensure_checkpoint_is_local, - load_checkpoint_config, - load_checkpoint_weights, -) -from speculators.proposals.greedy import GreedyTokenProposalConfig - - -class EagleConverter: - """ - Converter for Eagle/HASS checkpoints to speculators format. - - This converter handles the transformation of Eagle-style checkpoints - (including HASS variants) into the standardized speculators format. - It supports automatic feature detection, weight remapping, and - optional validation. - - :Example: - - >>> converter = EagleConverter() - >>> converter.convert( - ... "yuhuili/EAGLE-LLaMA3.1-Instruct-8B", - ... "./output", - ... "meta-llama/Meta-Llama-3.1-8B-Instruct" - ... ) - """ - - EAGLE_TO_SPECULATORS_LAYERNORM_MAPPINGS = { - "embed_layernorm.weight": "embedding_layernorm.weight", - "lm_head_layernorm.weight": "pre_lm_head_layernorm.weight", - } - - def convert( - self, - input_path: str | Path, - output_path: str | Path, - base_model: str, - fusion_bias: bool = False, - layernorms: bool = False, - validate: bool = True, - cache_dir: str | Path | None = None, - ) -> None: - """ - Convert an Eagle checkpoint to speculators format. - - This method orchestrates the complete conversion process: - - 1. Ensures the checkpoint is available locally - 2. Loads the original config and weights - 3. Auto-detects features if not explicitly specified (layernorms, fusion bias) - 4. Builds the speculators configuration - 5. Processes and remaps the weights - 6. Saves the converted checkpoint - 7. Optionally validates the result by running a forward pass - - :param input_path: Path to Eagle checkpoint (local or HuggingFace ID) - :param output_path: Where to save converted checkpoint - :param base_model: Base model name (e.g., meta-llama/Llama-3.1-8B-Instruct) - :param fusion_bias: Enable fusion bias (auto-detected if not specified) - :param layernorms: Enable extra layernorms (auto-detected if not specified) - :param validate: Whether to validate the converted checkpoint - :param cache_dir: Optional cache directory for downloads - - :Example: - - >>> # Convert standard Eagle checkpoint - >>> converter = EagleConverter() - >>> converter.convert( - ... "yuhuili/EAGLE-LLaMA3.1-Instruct-8B", - ... "./eagle-converted", - ... "meta-llama/Meta-Llama-3.1-8B-Instruct", - ... validate=True - ... ) - - >>> # Convert HASS checkpoint with layernorms - >>> converter.convert( - ... "nm-testing/Eagle_Speculator_Llama_3_1_8B_TTT", - ... "./hass-converted", - ... "meta-llama/Meta-Llama-3.1-8B-Instruct", - ... layernorms=True - ... ) - """ - logger.info(f"Converting Eagle checkpoint: {input_path}") - - local_checkpoint_path = ensure_checkpoint_is_local(input_path, cache_dir) - - eagle_config = load_checkpoint_config(local_checkpoint_path) - weights = load_checkpoint_weights(local_checkpoint_path) - logger.info(f"Loaded {len(weights)} weights") - - detected_fusion_bias, detected_layernorms = detect_fusion_bias_and_layernorms( - weights - ) - fusion_bias = fusion_bias or detected_fusion_bias - layernorms = layernorms or detected_layernorms - - speculator_config = self._build_eagle_speculator_config( - eagle_config, base_model, fusion_bias, layernorms - ) - - processed_weights = self._process_checkpoint_weights(weights, layernorms) - - # Save the converted checkpoint using the model's save_pretrained - saved_path = self._save_converted_checkpoint( - config=speculator_config, weights=processed_weights, output_dir=output_path - ) - - logger.success(f"Saved to: {saved_path}") - - if validate: - self._validate_converted_checkpoint(saved_path, verifier_model=base_model) - - def _create_verifier_config(self, base_model: str) -> VerifierConfig: - config_dict, _ = PretrainedConfig.get_config_dict(base_model) - return VerifierConfig( - name_or_path=base_model, - architectures=config_dict.get("architectures", ["LlamaForCausalLM"]), - ) - - def _create_transformer_config_from_eagle(self, eagle_config: dict) -> LlamaConfig: - """ - Create a transformer config for the Eagle model's single decoder layer. - - :param eagle_config: Original Eagle checkpoint config - :return: LlamaConfig for the transformer layer - """ - return LlamaConfig( - vocab_size=eagle_config.get("vocab_size", 32000), - hidden_size=eagle_config.get("hidden_size", 4096), - intermediate_size=eagle_config.get("intermediate_size", 11008), - num_hidden_layers=1, # Eagle always uses a single decoder layer - num_attention_heads=eagle_config.get("num_attention_heads", 32), - num_key_value_heads=eagle_config.get("num_key_value_heads"), - hidden_act=eagle_config.get("hidden_act", "silu"), - max_position_embeddings=eagle_config.get("max_position_embeddings", 4096), - initializer_range=eagle_config.get("initializer_range", 0.02), - rms_norm_eps=eagle_config.get("rms_norm_eps", 1e-6), - use_cache=eagle_config.get("use_cache", True), - pad_token_id=eagle_config.get("pad_token_id"), - bos_token_id=eagle_config.get("bos_token_id", 1), - eos_token_id=eagle_config.get("eos_token_id", 2), - tie_word_embeddings=False, # Eagle uses separate embed_tokens from verifier - attention_bias=eagle_config.get("attention_bias", False), - attention_dropout=eagle_config.get("attention_dropout", 0.0), - mlp_bias=eagle_config.get("mlp_bias", False), - **build_llama_config_rope_kwargs( - rope_theta=eagle_config.get("rope_theta", 10000.0), - rope_scaling=eagle_config.get("rope_scaling"), - ), - ) - - def _build_eagle_speculator_config( - self, - eagle_config: dict, - base_model: str, - fusion_bias: bool, - layernorms: bool, - ) -> EagleSpeculatorConfig: - """ - Build a complete EagleSpeculatorConfig from Eagle checkpoint config. - - :param eagle_config: Original checkpoint config dictionary - :param base_model: Base model name for the verifier - :param fusion_bias: Whether to enable fusion bias - :param layernorms: Whether to enable extra layernorms - :return: Complete Eagle speculator configuration - """ - logger.debug( - f"Building config with fusion_bias={fusion_bias}, layernorms={layernorms}" - ) - - transformer_config = self._create_transformer_config_from_eagle(eagle_config) - verifier_config = self._create_verifier_config(base_model) - - greedy_proposal = GreedyTokenProposalConfig( - proposal_type="greedy", - speculative_tokens=3, - ) - - speculators_config = SpeculatorsConfig( - algorithm="eagle", - proposal_methods=[greedy_proposal], - default_proposal_method="greedy", - verifier=verifier_config, - ) - - return EagleSpeculatorConfig( - transformer_layer_config=transformer_config, - speculators_config=speculators_config, - layernorms=layernorms, - fusion_bias=fusion_bias, - ) - - def _should_skip_weight(self, weight_name: str, has_layernorms: bool) -> bool: - """ - Determine if a weight should be skipped during conversion. - - :param weight_name: Original weight name - :param has_layernorms: Whether layernorms are enabled - :return: True if the weight should be excluded from the output - """ - # Skip embed_tokens - Eagle gets these from the verifier model - if weight_name == "embed_tokens.weight": - logger.debug("Skipping embed_tokens.weight (tied to lm_head)") - return True - - # Skip hidden_layernorm when layernorms are disabled - return weight_name == "hidden_layernorm.weight" and not has_layernorms - - def _remap_weight_name(self, weight_name: str, has_layernorms: bool) -> str: - """ - Remap an Eagle weight name to speculators format. - - :param weight_name: Original weight name - :param has_layernorms: Whether layernorms are enabled - :return: Remapped weight name - """ - # hidden_layernorm maps to the decoder's input_layernorm when layernorms enabled - if weight_name == "hidden_layernorm.weight" and has_layernorms: - return "transformer.input_layernorm.weight" - - if ( - has_layernorms - and weight_name in self.EAGLE_TO_SPECULATORS_LAYERNORM_MAPPINGS - ): - return self.EAGLE_TO_SPECULATORS_LAYERNORM_MAPPINGS[weight_name] - - if weight_name.startswith("fc."): - return weight_name.replace("fc.", "fusion_fc.") - - if weight_name.startswith("layers.0."): - return weight_name.replace("layers.0.", "transformer.") - - return weight_name - - def _process_checkpoint_weights( - self, - weights: dict[str, torch.Tensor], - has_layernorms: bool, - ) -> dict[str, torch.Tensor]: - """ - Process and remap all weights from Eagle to speculators format. - - :param weights: Original checkpoint weights - :param has_layernorms: Whether layernorms are enabled - :return: Processed weights with remapped names - """ - logger.debug(f"Processing {len(weights)} weights") - - processed_weights = {} - skipped_weights = [] - remapped_weights = [] - - for original_name, tensor in weights.items(): - if self._should_skip_weight(original_name, has_layernorms): - skipped_weights.append(original_name) - continue - - new_name = self._remap_weight_name(original_name, has_layernorms) - processed_weights[new_name] = tensor - - if new_name != original_name: - remapped_weights.append(f"{original_name} -> {new_name}") - - if skipped_weights: - logger.debug(f"Skipped weights: {skipped_weights}") - if remapped_weights: - logger.debug(f"Remapped weights: {remapped_weights}") - - return processed_weights - - def _save_converted_checkpoint( - self, - config: EagleSpeculatorConfig, - weights: dict[str, torch.Tensor], - output_dir: str | Path, - ) -> Path: - """ - Save the converted checkpoint using the model's save_pretrained method. - - This method initializes an EagleSpeculator model with detached verifier mode - to prevent automatic verifier loading, loads the converted weights, and uses - the model's save_pretrained to ensure proper HuggingFace Hub compatibility. - - The saved checkpoint will include: - - config.json: Model configuration - - model.safetensors: Model weights (excluding verifier-shared components) - - eagle.py: Auto-generated model code for Hub integration - - :param config: The Eagle speculator config - :param weights: The processed weights dictionary - :param output_dir: Directory to save the checkpoint - :return: Path to the saved checkpoint - :raises RuntimeError: If checkpoint saving fails - """ - model = EagleSpeculator( - config=config, verifier=None, verifier_attachment_mode="detached" - ) - # Load the converted weights into the model - model.load_state_dict(weights, strict=False) # type: ignore[attr-defined] - logger.debug(f"Saving model to: {output_dir}") - model.save_pretrained(str(output_dir)) # type: ignore[attr-defined] - return Path(output_dir) - - def _validate_converted_checkpoint( - self, checkpoint_path: Path, verifier_model: str - ) -> None: - """ - Validate that a converted checkpoint can be loaded using from_pretrained. - - :param checkpoint_path: Path to the converted checkpoint - :param verifier_model: verifier model id or local path to attach - :raises Exception: If validation fails - """ - logger.info("Validating converted checkpoint...") - - try: - logger.debug("Loading model with EagleSpeculator.from_pretrained") - EagleSpeculator.from_pretrained( - checkpoint_path, - verifier=verifier_model, - verifier_attachment_mode="detached", - ) - logger.success("Model loaded successfully") - - except Exception as exception: - logger.error(f"Validation failed: {exception}") - raise exception diff --git a/src/speculators/convert/eagle/eagle_legacy_model.py b/src/speculators/convert/eagle/eagle_legacy_model.py deleted file mode 100644 index 2b4a57d10..000000000 --- a/src/speculators/convert/eagle/eagle_legacy_model.py +++ /dev/null @@ -1,726 +0,0 @@ -""" -Speculators implementations providing a unified implementation -for EAGLE v1, EAGLE v2, and HASS variants for spec decoding: - - Eagle / Eagle v1: https://arxiv.org/abs/2401.15077 - - Eagle v2: https://arxiv.org/abs/2406.16858 - - HASS: https://arxiv.org/abs/2408.15766 - -Classes: - EagleSpeculatorConfig: Configuration class for EAGLE/HASS model variants - EagleSpeculator: Main model implementation for EAGLE/HASS speculators -""" - -import importlib -import inspect -import os -import re -import warnings -from typing import Any, ClassVar, Literal, cast - -import torch -from pydantic import Field, field_serializer, field_validator, model_validator -from torch import nn -from transformers import ( - AutoConfig, - AutoModelForCausalLM, - PretrainedConfig, - PreTrainedModel, -) -from transformers.modeling_attn_mask_utils import _prepare_4d_causal_attention_mask -from transformers.modeling_outputs import CausalLMOutputWithPast -from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING -from transformers.models.llama.configuration_llama import LlamaConfig -from typing_extensions import Self - -from speculators import SpeculatorModel, SpeculatorModelConfig -from speculators.config import SpeculatorsConfig, VerifierConfig -from speculators.proposals.greedy import GreedyTokenProposalConfig - -__all__ = [ - "EagleSpeculator", - "EagleSpeculatorConfig", -] - - -@SpeculatorModelConfig.register("eagle") -class EagleSpeculatorConfig(SpeculatorModelConfig): - """ - A SpeculatorModelConfig implementation to be used with the EagleSpeculator - for EAGLE and HASS variants for spec decoding: - - Eagle / Eagle v1: https://arxiv.org/abs/2401.15077 - - Eagle v2: https://arxiv.org/abs/2406.16858 - - HASS: https://arxiv.org/abs/2408.15766 - - Model Configurations: - - EAGLE1: layernorms=False, fusion_bias=False - - EAGLE2: layernorms=False, fusion_bias=False - - HASS: layernorms=False, fusion_bias=True - """ - - speculators_model_type: Literal["eagle"] = "eagle" - architectures: list[str] = Field( - default_factory=lambda: ["EagleSpeculator"], - description=( - "List of model architectures that can be used with the model " - "pretrained weights. Automatically includes the transformer layer " - "architecture to ensure compatibility during model loading and " - "validation." - ), - ) - - transformer_layer_architecture: str = Field( - default="auto", - description=( - "The architecture class name of the transformer layer to use for " - "the speculator's decoder layer. Must correspond to a valid " - "transformer decoder layer class (e.g., 'LlamaDecoderLayer')." - ), - ) - transformer_layer_config: PretrainedConfig = Field( - default_factory=LlamaConfig, - description=( - "Configuration object for the transformer layer architecture. " - "Must be a PretrainedConfig instance that matches the requirements " - "of the transformer_layer_architecture. Contains parameters such as " - "hidden_size, num_attention_heads, intermediate_size, vocab_size, " - "and other architecture-specific settings." - ), - ) - layernorms: bool = Field( - default=False, - description=( - "Whether to include additional layer normalization layers in the " - "model architecture. When True, adds RMSNorm layers after the " - "verifier's hidden state (embedding_layernorm), after the fusion " - "layer output, and before the language model head (pre_lm_head_layernorm). " - "When False, these layers are not included and the output layernorm " - "within the transformer architecture is removed as well. " - "Standard EAGLE1, EAGLE2, and HASS implementations use False." - ), - ) - fusion_bias: bool = Field( - default=False, - description=( - "Whether to add a learnable bias term to the fusion (fully connected) " - "layer that combines input embeddings with verifier hidden states. " - "The fusion layer concatenates input embeddings and hidden states, " - "then projects to hidden_size dimensions. Standard EAGLE1 and EAGLE2 " - "use False, while HASS uses True." - ), - ) - - @model_validator(mode="after") - def check_add_architectures(self) -> Self: - """ - Automatically adds the transformer layer architecture to the - architectures list if it's not already present. - - :return: The validated configuration instance with updated architectures - """ - if ( - self.transformer_layer_architecture != "auto" - and self.transformer_layer_architecture not in self.architectures - ): - self.architectures.append(self.transformer_layer_architecture) - - return self - - @field_serializer("transformer_layer_config") - def serialize_transformer_layer_config(self, value: PretrainedConfig) -> dict: - """ - Serialize the transformer_layer_config to a dictionary for JSON storage. - - Converts the PretrainedConfig object to its dictionary representation - using to_diff_dict() to only include non-default values. - - :param value: The PretrainedConfig instance to serialize - :return: Dictionary representation of the transformer layer configuration - """ - return value.to_diff_dict() - - @field_validator("transformer_layer_config", mode="before") - @classmethod - def validate_transformer_layer_config(cls, value: Any) -> PretrainedConfig: - """ - Validate and convert transformer_layer_config to a PretrainedConfig instance. - - Accepts either a dictionary that can be converted to a PretrainedConfig - or an existing PretrainedConfig instance. - - :param value: The value to validate (dict or PretrainedConfig) - :return: A validated PretrainedConfig instance - :raises ValueError: If the value cannot be converted to a PretrainedConfig - """ - if isinstance(value, dict): - return AutoConfig.for_model(**value) - if isinstance(value, PretrainedConfig): - return value - - raise ValueError( - "transformer_layer_config must be a PretrainedConfig instance or a " - "dictionary that can be converted to a PretrainedConfig." - ) - - -@SpeculatorModel.register("eagle") -class EagleSpeculator(SpeculatorModel): - """ - A SpeculatorModel implementation for EAGLE and HASS variants for spec decoding: - - Eagle / Eagle v1: https://arxiv.org/abs/2401.15077 - - Eagle v2: https://arxiv.org/abs/2406.16858 - - HASS: https://arxiv.org/abs/2408.15766 - - Architecture Overview: - The EAGLE speculator consists of: - 1. Input embedding layer (shared with verifier) - 2. Optional embedding layer normalization - 3. Fusion layer: Concatenates and projects input embeddings + verifier hidden - states to a latent space of hidden_size - 4. Single transformer decoder layer for candidate token generation - 5. Optional pre-LM head layer normalization - 6. Language model head (shared with verifier) - - Speculative Decoding Process: - 1. Verifier model processes input and generates hidden states - 2. EAGLE speculator uses these hidden states + input embeddings to predict - next tokens - 3. Multiple candidate tokens generated in parallel using token proposal methods - 4. Verifier validates candidates and accepts/rejects based on probability - thresholds - 5. Process continues iteratively for multi-token speculation - """ - - # PreTrainedModel settings - config_class: ClassVar[type[EagleSpeculatorConfig]] = EagleSpeculatorConfig # type: ignore[misc] - _keys_to_ignore_on_load_missing: ClassVar[list[str]] = [ # type: ignore[assignment,misc] - "verifier*", - "embed_tokens*", - "lm_head*", - ] - - _keys_to_ignore_on_save: ClassVar[list[str]] = [ # type: ignore[assignment,misc] - "embed_tokens.weight", - "lm_head.weight", - "lm_head.bias", - "verifier*", - ] - - @classmethod - def from_training_args( - cls, - verifier_config: PretrainedConfig, - **kwargs, - ) -> "EagleSpeculator": - """Create EAGLE model from training arguments. - - Args: - verifier_config: Verifier model configuration - **kwargs: Training arguments with EAGLE-specific params - - layernorms: Whether to include layer normalization layers - - fusion_bias: Whether to add bias to fusion layer - - transformer_layer_architecture: Name of transformer decoder layer - class - - verifier_name_or_path: Path to verifier model - - Returns: - Initialized EagleSpeculator - """ - config = EagleSpeculatorConfig( - transformer_layer_config=verifier_config, - layernorms=kwargs.get("layernorms", False), - fusion_bias=kwargs.get("fusion_bias", False), - transformer_layer_architecture=kwargs.get( - "transformer_layer_architecture", "auto" - ), - speculators_config=SpeculatorsConfig( - algorithm="eagle", - proposal_methods=[GreedyTokenProposalConfig()], - default_proposal_method="greedy", - verifier=VerifierConfig.from_config( - verifier_config, name_or_path=kwargs["verifier_name_or_path"] - ), - ), - ) - - return cls(config=config) - - @staticmethod - def get_trainer_kwargs(**kwargs) -> tuple[dict, dict]: # noqa: ARG004 - """Get training and validation kwargs for EAGLE. - - EAGLE doesn't require any special forward pass arguments during training, - so this returns empty dictionaries. - - Args: - **kwargs: Training arguments (unused) - - Returns: - Tuple of (train_call_kwargs, val_call_kwargs), both empty dicts - """ - return {}, {} - - def __init__( - self, - config: EagleSpeculatorConfig, - verifier: str | os.PathLike | PreTrainedModel | None = None, - verifier_attachment_mode: Literal["detached", "full", "train_only"] - | None = None, - ): - """ - Initializes an EAGLE speculator architecture with configurable components based - on the provided configuration. The model starts with verifier-dependent layers - (embed_tokens, rotary_emb, lm_head) set to None until a verifier is attached. - - :param config: Configuration object specifying model architecture, layer - settings, and speculative decoding parameters. Must be an instance of - EagleSpeculatorConfig containing transformer layer configuration and - EAGLE-specific settings. - :param verifier: Optional verifier model to attach for speculative decoding. - Can be a path to a model directory, Hugging Face model identifier, or - PreTrainedModel instance. If None, must be attached later via - attach_verifier() before using the model. - :param verifier_attachment_mode: Mode for verifier attachment. "detached" - prevents attachment even if verifier is provided. "full" enables - complete integration for both training and generation. "train_only" - attaches only components needed for training, optimizing memory usage. - """ - if not isinstance(config, EagleSpeculatorConfig): - raise ValueError( - "config must be an instance of EagleSpeculatorConfig, " - f"got {type(config)} instead." - ) - - # Initialize model parameters from config - self.vocab_size = config.transformer_layer_config.vocab_size - self.hidden_size = config.transformer_layer_config.hidden_size - self.padding_idx = config.transformer_layer_config.pad_token_id - - # Set layers pulled from the verifier to None until attach is called - self.embed_tokens: nn.Embedding | None = None - self.rotary_emb: nn.Module | None = None - self.lm_head: nn.Linear | None = None - - super().__init__(config=config) - self.verifier: PreTrainedModel | None = None - self.verifier_attachment_mode: Literal["detached", "full", "train_only"] = ( - "detached" - ) - - verifier = verifier or config.speculators_config.verifier.name_or_path - if verifier is not None and verifier_attachment_mode != "detached": - self.attach_verifier(verifier, mode=verifier_attachment_mode) - - self._decoder_class, self._layernorm_class = self._import_model_classes() - # Initialize layers based on the configuration - self.embedding_layernorm: nn.Module | None = self._create_layernorm() - self.fusion_fc: nn.Linear = nn.Linear( - 2 * self.hidden_size, - self.hidden_size, - bias=config.fusion_bias, - ) - self.transformer: nn.Module = self._create_transformer_layer() - self.pre_lm_head_layernorm: nn.Module | None = self._create_layernorm() - - self.post_init() # type: ignore[attr-defined] - - def resolve_verifier( - self, verifier: str | os.PathLike | PreTrainedModel - ) -> PreTrainedModel: - """ - Resolves the verifier model from a given path or identifier. - - This method loads the verifier model from a specified path or identifier, - ensuring it is compatible with the speculator's configuration. If the - verifier is already attached, it returns the existing verifier instance. - - :param verifier: The verifier model to resolve. Can be a path to a local - model directory, a Hugging Face model identifier, or an instance of - PreTrainedModel. - :return: The resolved PreTrainedModel instance for the verifier. - """ - if not verifier: - raise ValueError( - "Verifier must be provided as a path, identifier, or PreTrainedModel. " - ) - - if not isinstance(verifier, (str, os.PathLike, PreTrainedModel)): - raise TypeError( - f"Expected verifier to be a PreTrainedModel, a string path, " - f"or an os.PathLike object, got {type(verifier)} {verifier}." - ) - - if isinstance(verifier, PreTrainedModel): - return verifier - - return AutoModelForCausalLM.from_pretrained(verifier) - - def state_dict( - self, - *, - destination: dict[str, Any] = None, # type: ignore[assignment] - prefix: str = "", - keep_vars: bool = False, - ): - """ - Overrides the state_dict method from PyTorch to ensure that save pathways - within Transformers PreTrainedModel do not include the verifier model's - parameters. This is important to ensure that the speculator model - can be saved and loaded without including the verifier's state, which - is expected to be managed separately. - - :param destination: Optional dictionary to store the state. - :param prefix: Optional prefix for parameter names. - :param keep_vars: Whether to keep Variables in the state_dict. - :return: A dictionary containing the state of the speculator model, - excluding the verifier model's parameters. This dictionary can be used - to save the model's state to disk or for further processing. - """ - tmp_verifier = self.verifier - self.verifier = None - state = super().state_dict( # type: ignore[misc] - destination=destination, prefix=prefix, keep_vars=keep_vars - ) - self.verifier = tmp_verifier - - return state - - def attach_verifier( - self, - verifier: str | os.PathLike | PreTrainedModel, - mode: Literal["full", "train_only"] | None = None, - ): - """ - Attach a verifier model to the EagleSpeculator for speculative decoding. - Utilizes the verifier's embed_tokens, rotary_emb, and lm_head layers - for the speculator's forward pass and generation methods. - Additionally, for `generate`, it uses the verifier's hidden states - to generate speculative token predictions. - - If mode is "full", the verifier is fully integrated for use with - both `generate` and `forward` methods. - - If mode is "train_only", only the verifier's layers required for a forward pass - are attached, allowing for better resource utilization during training. - `generate` will not be available until a full verifier is attached. - - Example: - ```python - # Load and attach a verifier - verifier = EagleSpeculator(...) - - # For generation - speculator.attach_verifier(verifier) - outputs = speculator.generate(input_ids) - speculator.detach_verifier() - - # For training - speculator.attach_verifier(verifier, mode="train_only") - outputs = speculator(input_ids, hidden_states) - speculator.detach_verifier() - ``` - - :param verifier: The verifier model to attach. This can be a path to a local - model directory, a Hugging Face model identifier, or an instance of - PreTrainedModel. If a path or identifier is provided, the model will be - loaded automatically. If an instance is provided, it will be used directly. - :param mode: The mode for attaching the verifier. Can be "full" or "train_only". - If None, defaults to "full". In "train_only" mode, only the layers - required for a forward pass are attached, and the speculator cannot - perform generation until a full verifier is attached. - :return: The PreTrainedModel instance for the verifier that was attached. - """ - if self.verifier_attachment_mode != "detached": - raise RuntimeError( - "Cannot attach a verifier when the speculator is not in detached mode. " - "Detach the current verifier first using `detach_verifier()`." - ) - - if mode not in {"full", "train_only", None}: - raise ValueError( - f"Invalid verifier_attachment_mode: {mode}. " - "Must be one of 'full', 'train_only', or None." - ) - - self.verifier_attachment_mode = mode or "full" - self.verifier = ( - self.resolve_verifier(verifier) - if self.verifier_attachment_mode == "full" - else None - ) # Expect subclasses to handle references if train_only - - if self.verifier_attachment_mode == "train_only": - verifier_model = self.resolve_verifier(verifier) - elif self.verifier_attachment_mode == "full": - verifier_model = cast("PreTrainedModel", self.verifier) - else: - return - - if hasattr(verifier_model, "model"): - self.embed_tokens = verifier_model.model.embed_tokens # type: ignore[assignment,union-attr] - self.rotary_emb = verifier_model.model.rotary_emb # type: ignore[assignment,union-attr] - else: - # Bare model structure - self.embed_tokens = verifier_model.embed_tokens # type: ignore[assignment,attr-defined] - self.rotary_emb = verifier_model.rotary_emb # type: ignore[assignment,attr-defined] - - # lm_head is always at the top level of the verifier - self.lm_head = verifier_model.lm_head # type: ignore[assignment,attr-defined] - - def detach_verifier(self): - """ - Removes the reference to the attached verifier model and frees up the - associated memory. After calling this method, the speculator will not - be able to perform forward passes or generation until a new verifier - is attached. - """ - if self.verifier_attachment_mode == "detached": - raise RuntimeError( - "Verifier is already detached, cannot be called again until " - "a new verifier is attached." - ) - - if self.verifier is not None: - del self.verifier - - self.verifier = None - self.verifier_attachment_mode = "detached" - - del self.embed_tokens - self.embed_tokens = None - del self.rotary_emb - self.rotary_emb = None - del self.lm_head - self.lm_head = None - - def forward( - self, - input_ids: torch.LongTensor, - hidden_states: torch.FloatTensor, - attention_mask: torch.Tensor | None = None, - position_ids: torch.LongTensor | None = None, - past_key_values: tuple[tuple[torch.FloatTensor]] | None = None, - use_cache: bool | None = None, - output_attentions: bool | None = None, - output_hidden_states: bool | None = None, # noqa: ARG002 - return_dict: bool | None = None, - ) -> torch.FloatTensor | CausalLMOutputWithPast: - """ - Execute the forward pass for speculative token generation. - - Processes input tokens and verifier hidden states through the EAGLE architecture - to generate candidate tokens for speculative decoding. The method combines input - embeddings with verifier hidden states via a fusion layer, processes them - through a transformer decoder layer, and produces logits for next token - prediction. - - :param input_ids: Token IDs for the current input sequence. Shape: (batch_size, - sequence_length). These represent the tokens that will be converted to - embeddings and combined with verifier hidden states. - :param hidden_states: Hidden state representations from the verifier model - corresponding to the input sequence. Shape: (batch_size, sequence_length, - hidden_size). These capture the verifier's understanding of the context. - :param attention_mask: Optional attention mask to avoid attending to padding - tokens. Shape: (batch_size, sequence_length) for 2D or (batch_size, 1, - sequence_length, sequence_length) for 4D causal mask. - :param position_ids: Optional position indices for tokens in the sequence. - Shape: (batch_size, sequence_length). If None, auto-generated based on - sequence length and past key values. - :param past_key_values: Optional cached key-value states from previous forward - passes for efficient generation. Tuple of layer key-value pairs. - :param use_cache: Whether to return key-value states for caching in subsequent - forward passes. Useful for autoregressive generation efficiency. - :param output_attentions: Whether to return attention weights from the - transformer layer. Used for analysis and visualization. - :param output_hidden_states: Whether to return hidden states from the - transformer layer. Currently not implemented in this model. - :param return_dict: Whether to return structured CausalLMOutputWithPast instead - of raw logits. If None, uses config.use_return_dict default. - :return: Either raw logits tensor (batch_size, sequence_length, vocab_size) if - return_dict=False, or CausalLMOutputWithPast containing logits, past key - values, and optional attention weights. - :raises ValueError: If verifier components (embed_tokens, rotary_emb, lm_head) - are not attached. Call attach_verifier() before using forward(). - """ - if self.embed_tokens is None or self.rotary_emb is None or self.lm_head is None: - raise ValueError( - "Verifier model layers not initialized. " - "Call `attach_verifier` to set up the model before using forward." - ) - - return_dict = ( - return_dict if return_dict is not None else self.config.use_return_dict - ) - - inputs_embeds = self.embed_tokens(input_ids) - if self.embedding_layernorm is not None: - inputs_embeds = self.embedding_layernorm(inputs_embeds) - - hidden_states = self.fusion_fc( - torch.cat([inputs_embeds, hidden_states], dim=-1) - ) - hidden_states, attention_mask, position_ids = self._prepare_decoder_inputs( - hidden_states, attention_mask, position_ids, past_key_values - ) - - cos, sin = self.rotary_emb(hidden_states, position_ids) - layer_outputs = self.transformer( - hidden_states, - attention_mask=attention_mask, - position_ids=position_ids, - past_key_value=past_key_values[0] if past_key_values else None, - output_attentions=output_attentions, - use_cache=use_cache, - position_embeddings=(cos, sin), - ) - hidden_states = layer_outputs[0] - - if self.pre_lm_head_layernorm is not None: - hidden_states = self.pre_lm_head_layernorm(hidden_states) - - logits = self.lm_head(hidden_states) - - if not return_dict: - return logits - - return CausalLMOutputWithPast( - logits=logits, - past_key_values=layer_outputs[1] if use_cache else None, - hidden_states=None, - attentions=None, - ) - - def _prepare_decoder_inputs( - self, - hidden_states: torch.FloatTensor, - attention_mask: torch.Tensor | None, - position_ids: torch.LongTensor | None, - past_key_values: tuple[tuple[torch.FloatTensor]] | None, - ) -> tuple[torch.FloatTensor, torch.Tensor | None, torch.LongTensor | None]: - batch_size, seq_length = hidden_states.shape[:2] - - if position_ids is None: - device = hidden_states.device - position_ids = ( - torch.arange(seq_length, dtype=torch.long, device=device) # type: ignore[assignment] - .unsqueeze(0) - .expand(batch_size, -1) - ) - - if attention_mask is not None and attention_mask.dim() == 2: # noqa: PLR2004 - past_key_values_length = ( - past_key_values[0][0].shape[2] if past_key_values else 0 - ) - attention_mask = _prepare_4d_causal_attention_mask( - attention_mask, - (batch_size, seq_length), - hidden_states, - past_key_values_length, - sliding_window=getattr(self.config, "sliding_window", None), - ) - - return hidden_states, attention_mask, position_ids - - def _create_layernorm(self) -> nn.Module | None: - if not self.config.layernorms: - return None - - return self._layernorm_class( - self.hidden_size, eps=self.config.transformer_layer_config.rms_norm_eps - ) - - def _create_transformer_layer(self) -> nn.Module: - layer = self._decoder_class( - self.config.transformer_layer_config, - layer_idx=0, - ) - - if not self.config.layernorms: - # Replace input_layernorm with Identity if layernorms are not used - layer.input_layernorm = nn.Identity() - - return layer - - def _update_config_with_decoder_class(self, decoder_class: type[nn.Module]): - if self.config.transformer_layer_architecture == "auto": - decoder_name = decoder_class.__name__ - self.config.transformer_layer_architecture = decoder_name - if ( - self.config.architectures - and decoder_name not in self.config.architectures - ): - self.config.architectures.append(decoder_name) - - def _import_model_classes(self) -> tuple[type[nn.Module], type[nn.Module]]: - """ - Get the decoder layer class and layer normalization class for a given config - and decoder layer name. Uses the config qualified name to find the corresponding - modeling module to import the decoder layer and normalization classes from. - """ - # e.g. LlamaConfig + LlamaDecoderLayer - # -> get config class: transformers.models.llama.configuration_llama.LlamaConfig - # -> find corresponding causal language model class: - # transformers.models.llama.modeling_llama.LlamaForCausalLM - # -> import module: transformers.models.llama.modeling_llama - # -> get decoder layer class: - # transformers.models.llama.modeling_llama.LlamaDecoderLayer - # -> get layer normalization class: - # transformers.models.llama.modeling_llama.LlamaRMSNorm - - config_class = type(self.config.transformer_layer_config) - if config_class not in MODEL_FOR_CAUSAL_LM_MAPPING: - raise TypeError( - f"Config class {config_class} is not a valid causal language model " - f"config class. Please use a valid config, e.g., LlamaConfig." - ) - - causal_lm_model_class = MODEL_FOR_CAUSAL_LM_MAPPING[config_class] - module_path = causal_lm_model_class.__module__ - - # Import the modeling module - modeling_module = importlib.import_module(module_path) - - ### Get decoder layer class - transformer_arch_name = self.config.transformer_layer_architecture - if transformer_arch_name == "auto": - transformer_arch_name = ( - config_class.__name__.removesuffix("Config") + "DecoderLayer" - ) - - try: - decoder_class = getattr(modeling_module, transformer_arch_name) - except AttributeError as e: - if self.config.transformer_layer_architecture == "auto": - msg = ( - "Unable to automatically determine transformer layer architecture. " - "Please set `transformer_layer_architecture` to a valid " - "architecture, e.g., 'LlamaDecoderLayer' for Llama models." - ) - else: - msg = ( - f"Transformer layer architecture {transformer_arch_name} not found " - f"in module {module_path}. Please set " - "`transformer_layer_architecture` to a valid architecture, " - "e.g., 'LlamaDecoderLayer' for Llama models." - ) - raise ValueError(msg) from e - - self._update_config_with_decoder_class(decoder_class) - - ### Get layer normalization class - if not self.config.layernorms: - # No layernorms being used, so no need to import a layer normalization class - return decoder_class, nn.Identity - - classes = dict(inspect.getmembers(modeling_module, inspect.isclass)) - for pat in [r".*RMSNorm$", r".*Norm$"]: - for name, cls in classes.items(): - if re.match(pat, name): - return decoder_class, cls - - warnings.warn( - "Unable to automatically determine layer normalization class. " - "Falling back to torch.nn.LayerNorm.", - stacklevel=2, - ) - - return decoder_class, torch.nn.LayerNorm diff --git a/src/speculators/convert/entrypoints.py b/src/speculators/convert/entrypoints.py index 3cefcf7ed..ebb5daec1 100644 --- a/src/speculators/convert/entrypoints.py +++ b/src/speculators/convert/entrypoints.py @@ -4,10 +4,7 @@ It supports the following algorithms and conversion from their associated research repositories: -- EAGLE -- EAGLE2 - EAGLE3 -- HASS - MTP - DFlash @@ -24,7 +21,6 @@ from speculators.convert.dflash.converter import DFlashConverter from speculators.convert.eagle.eagle3_converter import Eagle3Converter -from speculators.convert.eagle.eagle_converter import EagleConverter from speculators.convert.mtp.converter import MTPConverter __all__ = ["convert_model", "maybe_convert_external_checkpoint"] @@ -33,7 +29,7 @@ def convert_model( model: str, verifier: str, - algorithm: Literal["eagle", "eagle3", "mtp", "dflash"], + algorithm: Literal["eagle3", "mtp", "dflash"], output_path: str = "converted", validate_device: str | None = None, **kwargs, @@ -42,25 +38,6 @@ def convert_model( Convert a non speculator's model checkpoint to a speculator's model checkpoint for use within the Speculators library, Hugging Face Hub, or vLLM. - algorithm=="eagle": - Eagle v1, v2: https://github.com/SafeAILab/EAGLE - HASS: https://github.com/HArmonizedSS/HASS - :: - # general - convert_model( - model="yuhuili/EAGLE-LLaMA3.1-Instruct-8B", - verifier="meta-llama/Llama-3.1-8B-Instruct", - algorithm="eagle", - ) - # with layernorms and fusion bias enabled - convert_model( - model="./eagle/checkpoint", - verifier="meta-llama/Llama-3.1-8B-Instruct", - algorithm="eagle", - layernorms=True, - fusion_bias=True, - ) - algorithm=="eagle3": Eagle v3: https://github.com/SafeAILab/EAGLE :: @@ -102,25 +79,16 @@ def convert_model( :param verifier: Verifier model checkpoint or Hugging Face model ID to attach as the verification/base model for speculative decoding :param algorithm: The conversion algorithm to use: - "eagle", "eagle3", "mtp", or "dflash". + "eagle3", "mtp", or "dflash". :param output_path: Directory path where the converted model will be saved. :param kwargs: Additional keyword arguments for the conversion algorithm. - Options for Eagle: {"layernorms": true, "fusion_bias": true}. Options for Eagle3: {"norm_before_residual": true, "eagle_aux_hidden_state_layer_ids": [1,23,44]}. Options for MTP: {"num_speculative_steps": 3}. Options for DFlash: {"aux_hidden_state_layer_ids": [2,10,18,26,34]}. """ - if algorithm == "eagle": - EagleConverter().convert( - model, - output_path, - verifier, - validate=validate_device is not None, - **kwargs, - ) - elif algorithm == "eagle3": + if algorithm == "eagle3": Eagle3Converter().convert( model, output_path, diff --git a/src/speculators/model.py b/src/speculators/model.py index fb837b1c0..3c8c3b063 100644 --- a/src/speculators/model.py +++ b/src/speculators/model.py @@ -244,6 +244,7 @@ def from_pretrained( weights_only: bool = True, t2d: torch.Tensor | None = None, d2t: torch.Tensor | None = None, + verifier: str | None = None, **kwargs, ) -> "SpeculatorModel": """ @@ -319,7 +320,7 @@ def from_pretrained( pretrained_model_name_or_path = maybe_convert_external_checkpoint( pretrained_model_name_or_path, - verifier=kwargs.get("verifier"), + verifier=verifier, cache_dir=cache_dir, config_dict=config_dict, ) @@ -362,6 +363,7 @@ def from_pretrained( weights_only=weights_only, t2d=t2d, d2t=d2t, + verifier=verifier, **kwargs, ) diff --git a/tests/integration/convert/test_eagle.py b/tests/integration/convert/test_eagle.py deleted file mode 100644 index 31301f856..000000000 --- a/tests/integration/convert/test_eagle.py +++ /dev/null @@ -1,360 +0,0 @@ -""" -End-to-end tests for Eagle checkpoint conversion. - -Verifies the complete conversion workflow for Eagle and HASS checkpoints: -1. Converting checkpoints to speculators format -2. Loading converted models using from_pretrained -3. Executing forward passes -4. Saving models using save_pretrained -5. Validating saved directories and configs -""" - -import gc -import json -from pathlib import Path - -import pytest -import torch -from loguru import logger - -from speculators.convert.eagle import EagleConverter -from speculators.convert.eagle.eagle_legacy_model import ( - EagleSpeculator, - EagleSpeculatorConfig, -) - - -class TestEagleConversion: - """End-to-end tests for Eagle checkpoint conversion.""" - - def setup_method(self): - """Clear any cached models or state before each test.""" - # Clear transformers model cache to ensure clean state - - gc.collect() - if torch.cuda.is_available(): - torch.cuda.empty_cache() - - @pytest.fixture - def temp_cache_dir(self, tmp_path, monkeypatch): - """Create a temporary cache directory for model downloads.""" - cache_dir = tmp_path / "hf_cache" - cache_dir.mkdir(exist_ok=True) - - # Also set environment variables to ensure HF uses our cache - monkeypatch.setenv("HF_HOME", str(cache_dir)) - monkeypatch.setenv("TRANSFORMERS_CACHE", str(cache_dir)) - monkeypatch.setenv("HUGGINGFACE_HUB_CACHE", str(cache_dir)) - - return cache_dir - - @pytest.fixture - def converter(self): - """Create an Eagle converter instance.""" - return EagleConverter() - - @pytest.fixture - def base_model(self): - """Base model name for conversions.""" - return "meta-llama/Llama-3.1-8B-Instruct" - - @pytest.fixture - def temp_dir(self, tmp_path): - """Create a temporary directory for test outputs.""" - return tmp_path / "e2e_test" - - def verify_config( - self, config_path: Path, expected_type: str, expected_features: dict - ): - """ - Verify the saved config file contains expected values. - - :param config_path: Path to config.json - :param expected_type: Expected speculators_model_type - :param expected_features: Expected feature flags (layernorms, fusion_bias) - """ - assert config_path.exists(), f"Config file not found: {config_path}" - - with config_path.open() as f: - config_dict = json.load(f) - - # Verify model type - assert config_dict.get("speculators_model_type") == expected_type - - # Verify features - for feature, expected_value in expected_features.items(): - assert config_dict.get(feature) == expected_value, ( - f"Expected {feature}={expected_value}, got {config_dict.get(feature)}" - ) - - # Verify essential fields - assert "transformer_layer_config" in config_dict - assert "speculators_config" in config_dict - assert config_dict["speculators_config"]["algorithm"] == "eagle" - assert ( - config_dict["speculators_config"]["verifier"]["name_or_path"] - == "meta-llama/Llama-3.1-8B-Instruct" - ) - - def verify_checkpoint_structure(self, checkpoint_dir: Path): - """ - Verify checkpoint directory structure after conversion. - - After conversion, checkpoints are always stored in safetensors format. - - :param checkpoint_dir: Path to checkpoint directory - """ - assert checkpoint_dir.exists(), ( - f"Checkpoint directory not found: {checkpoint_dir}" - ) - assert (checkpoint_dir / "config.json").exists(), "Missing config.json" - - # Check for weights in safetensors format only - single_safetensors = checkpoint_dir / "model.safetensors" - sharded_safetensors_index = checkpoint_dir / "model.safetensors.index.json" - - has_weights = single_safetensors.exists() or sharded_safetensors_index.exists() - - assert has_weights, "Missing model weights in safetensors format" - - # For sharded models, check that at least one shard exists - if sharded_safetensors_index.exists(): - shard_files = list(checkpoint_dir.glob("model-*.safetensors")) - assert len(shard_files) > 0, "Index file exists but no shard files found" - - def execute_forward_pass(self, model: EagleSpeculator) -> torch.Tensor | None: - """ - Execute a forward pass with the model. - - :param model: EagleSpeculator model instance - :return: Output logits or None if model is on meta device - """ - - # Check if model is on meta device - device = next(model.parameters()).device # type: ignore[attr-defined] - if device.type == "meta": - logger.info("Model is on meta device, skipping forward pass test") - return None - - batch_size = 2 - seq_length = 10 - hidden_size = model.config.transformer_layer_config.hidden_size - vocab_size = model.config.transformer_layer_config.vocab_size - - # Create dummy inputs on the same device as the model - input_ids = torch.randint( - 0, min(1000, vocab_size), (batch_size, seq_length) - ).to(device) - hidden_states = torch.randn(batch_size, seq_length, hidden_size).to(device) - - # Execute forward pass - with torch.no_grad(): - output = model(input_ids=input_ids, hidden_states=hidden_states) # type: ignore[operator] - - # Verify output shape - assert hasattr(output, "logits"), "Output missing logits attribute" - assert output.logits.shape == (batch_size, seq_length, vocab_size), ( - f"Unexpected output shape: {output.logits.shape}" - ) - - # Check for NaN/Inf - assert not torch.isnan(output.logits).any(), "Output contains NaN values" - assert not torch.isinf(output.logits).any(), "Output contains Inf values" - - return output.logits - - @pytest.mark.smoke - @pytest.mark.skip("Missing Llama HF Token") - @pytest.mark.parametrize( - "checkpoint_info", - [ - { - "name": "Eagle Standard", - "input_path": "yuhuili/EAGLE-LLaMA3.1-Instruct-8B", - "expected_features": {"layernorms": False, "fusion_bias": False}, - }, - { - "name": "HASS with Layernorms", - "input_path": "nm-testing/Eagle_Speculator_Llama_3_1_8B_TTT", - "expected_features": {"layernorms": True, "fusion_bias": False}, - }, - ], - ) - def test_eagle_checkpoint_conversion( - self, checkpoint_info, converter, base_model, temp_dir, temp_cache_dir - ): - """ - Test end-to-end conversion workflow for Eagle checkpoints. - - This test: - 1. Converts the checkpoint to speculators format - 2. Loads the converted model - 3. Executes a forward pass - 4. Saves the model again - 5. Validates the saved checkpoint - """ - name = checkpoint_info["name"] - input_path = checkpoint_info["input_path"] - expected_features = checkpoint_info["expected_features"] - - # Create test directories - converted_dir = temp_dir / f"{name.lower().replace(' ', '_')}_converted" - resaved_dir = temp_dir / f"{name.lower().replace(' ', '_')}_resaved" - - logger.info(f"Testing: {name}") - logger.info(f"Input: {input_path}") - logger.info(f"Expected features: {expected_features}") - - # Step 1: Convert checkpoint - logger.info("Converting checkpoint...") - converter.convert( - input_path=input_path, - output_path=converted_dir, - base_model=base_model, - validate=True, # This already tests loading and forward pass - cache_dir=temp_cache_dir, - ) - - # Verify converted checkpoint structure - assert converted_dir.exists(), f"Converted directory not found: {converted_dir}" - assert (converted_dir / "config.json").exists(), "Missing config.json" - assert (converted_dir / "model.safetensors").exists(), ( - "Missing model.safetensors" - ) - - # Verify config - self.verify_config( - converted_dir / "config.json", - expected_type="eagle", - expected_features=expected_features, - ) - logger.success("Conversion successful") - - # Step 2: Load converted model - logger.info("Loading converted model...") - model = EagleSpeculator.from_pretrained(converted_dir) - assert isinstance(model, EagleSpeculator), "Wrong model type loaded" - assert isinstance(model.config, EagleSpeculatorConfig), "Wrong config type" - - # Verify config attributes - assert model.config.layernorms == expected_features["layernorms"] - assert model.config.fusion_bias == expected_features["fusion_bias"] - logger.success("Model loaded successfully") - - # Step 3: Execute forward pass - logger.info("Executing forward pass...") - logits = self.execute_forward_pass(model) - if logits is not None: - logger.success(f"Forward pass successful, output shape: {logits.shape}") - else: - logger.info("Forward pass skipped (model on meta device)") - - # Step 4: Save model using save_pretrained - logger.info("Saving model using save_pretrained...") - model.save_pretrained(resaved_dir) # type: ignore[attr-defined] - logger.success(f"Model saved to: {resaved_dir}") - - # Step 5: Validate saved checkpoint - logger.info("Validating saved checkpoint...") - self.verify_checkpoint_structure(resaved_dir) - self.verify_config( - resaved_dir / "config.json", - expected_type="eagle", - expected_features=expected_features, - ) - - # Load the resaved model to ensure it works - logger.info("Loading resaved model...") - model2 = EagleSpeculator.from_pretrained(resaved_dir) - assert isinstance(model2, EagleSpeculator) - assert isinstance(model2.config, EagleSpeculatorConfig) - - # Verify configs match - assert model2.config.layernorms == model.config.layernorms - assert model2.config.fusion_bias == model.config.fusion_bias - assert ( - model2.config.transformer_layer_config.vocab_size - == model.config.transformer_layer_config.vocab_size - ) - - # Execute forward pass on resaved model - self.execute_forward_pass(model2) - logger.success("Resaved model forward pass successful") - - logger.success(f"{name} - All tests passed!") - - @pytest.mark.smoke - @pytest.mark.skip("Missing Llama HF Token") - def test_conversion_with_explicit_features( - self, converter, base_model, temp_dir, temp_cache_dir - ): - """ - Test conversion with explicitly set features overriding auto-detection. - """ - # Use the standard Eagle checkpoint but force fusion_bias=True - input_path = "yuhuili/EAGLE-LLaMA3.1-Instruct-8B" - output_dir = temp_dir / "eagle_forced_fusion_bias" - - logger.info("Testing explicit feature override") - - # Convert with forced fusion_bias - converter.convert( - input_path=input_path, - output_path=output_dir, - base_model=base_model, - fusion_bias=True, # Force this even though checkpoint doesn't have fc.bias - layernorms=False, - validate=True, - cache_dir=temp_cache_dir, - ) - - # Load and verify - model = EagleSpeculator.from_pretrained(output_dir) - assert model.config.fusion_bias is True, "fusion_bias should be True" - assert model.config.layernorms is False, "layernorms should be False" - - # Check that fc layer has bias - assert model.fusion_fc.bias is not None, ( # type: ignore[union-attr,attr-defined] - "fusion_fc layer should have bias parameter" - ) - - logger.success("Explicit feature override successful") - - @pytest.mark.smoke - @pytest.mark.skip("Missing Llama HF Token") - @pytest.mark.parametrize("validate", [True, False]) - def test_validation_flag( - self, converter, base_model, temp_dir, temp_cache_dir, validate - ): - """ - Test that the validate flag works correctly. - """ - input_path = "yuhuili/EAGLE-LLaMA3.1-Instruct-8B" - output_dir = temp_dir / f"eagle_validate_{validate}" - - logger.info(f"Testing validation flag: validate={validate}") - - # Convert with specified validation setting - converter.convert( - input_path=input_path, - output_path=output_dir, - base_model=base_model, - validate=validate, - cache_dir=temp_cache_dir, - ) - - # Conversion should succeed regardless of validation - assert output_dir.exists() - assert (output_dir / "config.json").exists() - assert (output_dir / "model.safetensors").exists() - - # Try loading the model - should work even if validation was skipped - model = EagleSpeculator.from_pretrained(output_dir) - self.execute_forward_pass(model) # type: ignore[arg-type] - - logger.success(f"Conversion with validate={validate} successful") - - -if __name__ == "__main__": - # Run tests with pytest - pytest.main([__file__, "-v", "-s"]) diff --git a/tests/unit/models/test_eagle_config.py b/tests/unit/models/test_eagle_config.py deleted file mode 100644 index 1a5b8d901..000000000 --- a/tests/unit/models/test_eagle_config.py +++ /dev/null @@ -1,685 +0,0 @@ -""" -Unit tests for the eagle model module in the Speculators library. -""" - -import tempfile -from pathlib import Path - -import pytest -from pydantic import BaseModel, ValidationError -from transformers import PretrainedConfig -from transformers.models.deepseek_v3.configuration_deepseek_v3 import DeepseekV3Config -from transformers.models.gemma.configuration_gemma import GemmaConfig -from transformers.models.granite.configuration_granite import GraniteConfig -from transformers.models.llama.configuration_llama import LlamaConfig -from transformers.models.mistral.configuration_mistral import MistralConfig -from transformers.models.mixtral.configuration_mixtral import MixtralConfig -from transformers.models.qwen3.configuration_qwen3 import Qwen3Config - -from speculators import ( - SpeculatorModelConfig, - SpeculatorsConfig, - VerifierConfig, -) -from speculators.convert.eagle.eagle_legacy_model import EagleSpeculatorConfig -from speculators.proposals import GreedyTokenProposalConfig - -# ===== Fixtures ===== - - -@pytest.fixture -def sample_verifier_config(): - return VerifierConfig( - name_or_path="test/verifier", - architectures=["LlamaForCausalLM"], - ) - - -@pytest.fixture -def sample_token_proposal_config(): - return GreedyTokenProposalConfig( - speculative_tokens=5, - verifier_accept_k=1, - accept_tolerance=0.0, - ) - - -@pytest.fixture -def sample_speculators_config(sample_token_proposal_config, sample_verifier_config): - return SpeculatorsConfig( - algorithm="eagle", - proposal_methods=[sample_token_proposal_config], - default_proposal_method="greedy", - verifier=sample_verifier_config, - ) - - -@pytest.fixture -def sample_llama_config(): - return LlamaConfig( - vocab_size=32000, - hidden_size=768, - intermediate_size=3072, - num_hidden_layers=12, - num_attention_heads=12, - max_position_embeddings=2048, - ) - - -@pytest.fixture -def eagle12_config_dict(): - return { - "speculators_model_type": "eagle", - "architectures": ["EagleSpeculator", "LlamaDecoderLayer"], - "transformer_layer_architecture": "LlamaDecoderLayer", - "transformer_layer_config": { - "model_type": "llama", - "vocab_size": 32000, - "hidden_size": 768, - "intermediate_size": 3072, - "num_hidden_layers": 12, - "num_attention_heads": 12, - "max_position_embeddings": 2048, - }, - "layernorms": False, - "fusion_bias": False, - "speculators_config": { - "algorithm": "eagle", - "proposal_methods": [ - { - "proposal_type": "greedy", - "speculative_tokens": 5, - "verifier_accept_k": 1, - "accept_tolerance": 0.0, - } - ], - "default_proposal_method": "greedy", - "verifier": { - "name_or_path": "test/verifier", - "architectures": ["LlamaForCausalLM"], - "hidden_size": 768, - "intermediate_size": 3072, - "vocab_size": 32000, - "max_position_embeddings": 2048, - "bos_token_id": 1, - "eos_token_id": 2, - }, - }, - } - - -@pytest.fixture -def hass_config_dict(eagle12_config_dict): - config_dict = eagle12_config_dict.copy() - config_dict["fusion_bias"] = True # Key difference for HASS - - return config_dict - - -# ===== Config Classes ===== - -LAYER_TYPES: list[tuple[str, type[PretrainedConfig]]] = [ - ("LlamaDecoderLayer", LlamaConfig), - ("MistralDecoderLayer", MistralConfig), - ("Qwen3DecoderLayer", Qwen3Config), - ("GemmaDecoderLayer", GemmaConfig), - ("MixtralDecoderLayer", MixtralConfig), - ("DeepseekV3DecoderLayer", DeepseekV3Config), - ("GraniteDecoderLayer", GraniteConfig), -] - - -def create_layer_config(config_class: type[PretrainedConfig]) -> PretrainedConfig: - """Create a config instance for the given config class with standard parameters.""" - base_params = { - "vocab_size": 32000, - "hidden_size": 768, - "intermediate_size": 3072, - "num_hidden_layers": 12, - "num_attention_heads": 12, - "max_position_embeddings": 2048, - } - - # Add extra parameters for specific config types - if config_class in (MixtralConfig, DeepseekV3Config, GraniteConfig): - base_params["num_key_value_heads"] = 12 - - if config_class == MixtralConfig: - base_params.update( - { - "num_local_experts": 8, - "num_experts_per_tok": 2, - } - ) - - return config_class(**base_params) # type: ignore[arg-type] - - -# ===== EagleSpeculatorConfig Tests ===== - - -@pytest.mark.smoke -def test_eagle_speculator_config_initialization(): - """Test default initialization of EagleSpeculatorConfig.""" - config = EagleSpeculatorConfig() - - # Verify Eagle-specific defaults - assert config.speculators_model_type == "eagle" - assert config.architectures == ["EagleSpeculator"] - assert config.transformer_layer_architecture == "auto" - assert isinstance(config.transformer_layer_config, LlamaConfig) - assert config.layernorms is False - assert config.fusion_bias is False - - # Verify base class defaults - assert config.model_type == "speculator_model" - assert config.speculators_config is None - - -@pytest.mark.smoke -def test_eagle_speculator_config_custom_initialization( - sample_speculators_config, sample_llama_config -): - """Test custom initialization of EagleSpeculatorConfig.""" - config = EagleSpeculatorConfig( - architectures=["CustomEagleSpeculator"], - transformer_layer_architecture="CustomDecoderLayer", - transformer_layer_config=sample_llama_config, - layernorms=True, - fusion_bias=True, - speculators_config=sample_speculators_config, - ) - - # Verify custom values - assert config.speculators_model_type == "eagle" - assert "CustomEagleSpeculator" in config.architectures - assert "CustomDecoderLayer" in config.architectures - assert config.transformer_layer_architecture == "CustomDecoderLayer" - assert config.transformer_layer_config == sample_llama_config - assert config.layernorms is True - assert config.fusion_bias is True - assert config.speculators_config == sample_speculators_config - - -@pytest.mark.smoke -@pytest.mark.parametrize(("layer_architecture", "config_class"), LAYER_TYPES) -def test_eagle_speculator_config_with_different_configs( - layer_architecture, config_class, sample_speculators_config -): - """Test EagleSpeculatorConfig with different transformer layer configurations.""" - layer_config = create_layer_config(config_class) - - config = EagleSpeculatorConfig( - transformer_layer_architecture=layer_architecture, - transformer_layer_config=layer_config, - speculators_config=sample_speculators_config, - ) - - # Verify the configuration - assert config.transformer_layer_architecture == layer_architecture - assert isinstance(config.transformer_layer_config, config_class) - assert config.transformer_layer_config.vocab_size == 32000 - assert config.transformer_layer_config.hidden_size == 768 - assert layer_architecture in config.architectures - assert "EagleSpeculator" in config.architectures - assert config.transformer_layer_config == layer_config - - -@pytest.mark.smoke -def test_eagle_speculator_config_base_initialization(sample_speculators_config): - # Create EagleSpeculatorConfig with custom values - original_config = EagleSpeculatorConfig( - transformer_layer_architecture="TestDecoderLayer", - layernorms=True, - fusion_bias=True, - speculators_config=sample_speculators_config, - ) - - # Convert to dict and validate through base class - config_dict = original_config.model_dump() - recreated_config = SpeculatorModelConfig.model_validate(config_dict) - - # Verify type and values preservation - assert isinstance(recreated_config, EagleSpeculatorConfig) - assert recreated_config.speculators_model_type == "eagle" - assert "TestDecoderLayer" in recreated_config.architectures - assert recreated_config.transformer_layer_architecture == "TestDecoderLayer" - assert recreated_config.layernorms is True - assert recreated_config.fusion_bias is True - assert recreated_config.speculators_config == sample_speculators_config - - -@pytest.mark.regression -def test_eagle_speculator_config_nested_initialization(): - class ParentModel(BaseModel): - single_config: EagleSpeculatorConfig - config_list: list[EagleSpeculatorConfig] - config_dict: dict[str, EagleSpeculatorConfig] - - parent = ParentModel( - single_config=EagleSpeculatorConfig(fusion_bias=True), - config_list=[ - EagleSpeculatorConfig(layernorms=True), - EagleSpeculatorConfig(fusion_bias=True), - ], - config_dict={ - "eagle1": EagleSpeculatorConfig(layernorms=False), - "hass": EagleSpeculatorConfig(fusion_bias=True), - }, - ) - - # Verify single config - assert isinstance(parent.single_config, EagleSpeculatorConfig) - assert parent.single_config.fusion_bias is True - - # Verify config list - assert len(parent.config_list) == 2 - assert all(isinstance(c, EagleSpeculatorConfig) for c in parent.config_list) - assert parent.config_list[0].layernorms is True - assert parent.config_list[1].fusion_bias is True - - # Verify config dict - assert len(parent.config_dict) == 2 - assert all( - isinstance(c, EagleSpeculatorConfig) for c in parent.config_dict.values() - ) - assert parent.config_dict["eagle1"].layernorms is False - assert parent.config_dict["hass"].fusion_bias is True - - -@pytest.mark.smoke -def test_eagle_speculator_config_invalid_initialization(): - # Test invalid speculators_model_type - with pytest.raises(ValidationError) as exc_info: - EagleSpeculatorConfig(speculators_model_type="invalid") # type: ignore[arg-type] - assert "speculators_model_type" in str(exc_info.value) - - # Test invalid architectures type - with pytest.raises(ValidationError) as exc_info: - EagleSpeculatorConfig(architectures="not_a_list") # type: ignore[arg-type] - assert "architectures" in str(exc_info.value) - - # Test invalid transformer_layer_architecture type - with pytest.raises(ValidationError) as exc_info: - EagleSpeculatorConfig(transformer_layer_architecture=123) # type: ignore[arg-type] - assert "transformer_layer_architecture" in str(exc_info.value) - - # Test invalid layernorms type - with pytest.raises(ValidationError) as exc_info: - EagleSpeculatorConfig(layernorms="not_a_bool") # type: ignore[arg-type] - assert "layernorms" in str(exc_info.value) - - # Test invalid fusion_bias type - with pytest.raises(ValidationError) as exc_info: - EagleSpeculatorConfig(fusion_bias="not_a_bool") # type: ignore[arg-type] - assert "fusion_bias" in str(exc_info.value) - - -@pytest.mark.smoke -def test_eagle_speculator_config_auto_registry(): - registered_classes = SpeculatorModelConfig.registered_classes() - class_names = [cls.__name__ for cls in registered_classes] - - # Verify EagleSpeculatorConfig is registered - assert "EagleSpeculatorConfig" in class_names - - # Verify registry key mapping - assert SpeculatorModelConfig.registry is not None - assert "eagle" in SpeculatorModelConfig.registry - assert SpeculatorModelConfig.registry["eagle"] == EagleSpeculatorConfig - - -@pytest.mark.smoke -def test_eagle_speculator_config_marshalling(sample_speculators_config): - original_config = EagleSpeculatorConfig( - transformer_layer_architecture="TestDecoderLayer", - layernorms=True, - fusion_bias=True, - speculators_config=sample_speculators_config, - ) - - # Test model_dump() - config_dict = original_config.model_dump() - assert isinstance(config_dict, dict) - assert config_dict["speculators_model_type"] == "eagle" - assert "TestDecoderLayer" in config_dict["architectures"] - assert config_dict["layernorms"] is True - assert config_dict["fusion_bias"] is True - - # Test model_validate() on base class - recreated_base = SpeculatorModelConfig.model_validate(config_dict) - assert isinstance(recreated_base, EagleSpeculatorConfig) - assert recreated_base.transformer_layer_architecture == "TestDecoderLayer" - assert recreated_base.layernorms is True - assert recreated_base.fusion_bias is True - - # Test model_validate() on derived class - recreated_derived = EagleSpeculatorConfig.model_validate(config_dict) - assert isinstance(recreated_derived, EagleSpeculatorConfig) - assert recreated_derived.transformer_layer_architecture == "TestDecoderLayer" - assert recreated_derived.layernorms is True - assert recreated_derived.fusion_bias is True - - -@pytest.mark.smoke -@pytest.mark.parametrize(("layer_architecture", "config_class"), LAYER_TYPES) -def test_eagle_speculator_config_marshalling_different_layers( - layer_architecture, config_class, sample_speculators_config -): - """Test marshalling with different layer architectures.""" - layer_config = create_layer_config(config_class) - - original_config = EagleSpeculatorConfig( - transformer_layer_architecture=layer_architecture, - transformer_layer_config=layer_config, - layernorms=True, - fusion_bias=True, - speculators_config=sample_speculators_config, - ) - - # Test model_dump() - config_dict = original_config.model_dump() - assert isinstance(config_dict, dict) - assert config_dict["speculators_model_type"] == "eagle" - assert config_dict["transformer_layer_architecture"] == layer_architecture - assert layer_architecture in config_dict["architectures"] - - # Test model_validate() on base class - recreated_base = SpeculatorModelConfig.model_validate(config_dict) - assert isinstance(recreated_base, EagleSpeculatorConfig) - assert recreated_base.transformer_layer_architecture == layer_architecture - assert recreated_base.layernorms is True - assert recreated_base.fusion_bias is True - - # Test model_validate() roundtrip - recreated_config = EagleSpeculatorConfig.model_validate(config_dict) - assert isinstance(recreated_config, EagleSpeculatorConfig) - assert recreated_config.transformer_layer_architecture == layer_architecture - assert recreated_config.layernorms is True - assert recreated_config.fusion_bias is True - - -@pytest.mark.smoke -def test_eagle_speculator_config_model_validator(): - config1 = EagleSpeculatorConfig(transformer_layer_architecture="CustomDecoderLayer") - assert "CustomDecoderLayer" in config1.architectures - assert "EagleSpeculator" in config1.architectures - - # Test with custom architectures already containing the layer - config2 = EagleSpeculatorConfig( - architectures=["CustomSpeculator", "CustomDecoderLayer"], - transformer_layer_architecture="CustomDecoderLayer", - ) - # Should not duplicate - architecture_count = config2.architectures.count("CustomDecoderLayer") - assert architecture_count == 1 - assert "CustomSpeculator" in config2.architectures - - # Test with custom architectures not containing the layer - config3 = EagleSpeculatorConfig( - architectures=["CustomSpeculator"], - transformer_layer_architecture="NewDecoderLayer", - ) - assert "CustomSpeculator" in config3.architectures - assert "NewDecoderLayer" in config3.architectures - - -# # ====== EagleSpeculatorConfig Eagle 1 / Eagle 2 Tests ====== - - -@pytest.mark.smoke -def test_eagle_speculator_config_eagle12_backwards_compatibility(eagle12_config_dict): - config_derived = EagleSpeculatorConfig.model_validate(eagle12_config_dict) - assert isinstance(config_derived, EagleSpeculatorConfig) - assert config_derived.speculators_model_type == "eagle" - assert "LlamaDecoderLayer" in config_derived.architectures - assert config_derived.transformer_layer_architecture == "LlamaDecoderLayer" - assert config_derived.layernorms is False - assert config_derived.fusion_bias is False - assert config_derived.speculators_config.algorithm == "eagle" - - # Test loading with base SpeculatorModelConfig.model_validate - config_base = SpeculatorModelConfig.model_validate(eagle12_config_dict) - assert isinstance(config_base, EagleSpeculatorConfig) - assert config_base.speculators_model_type == "eagle" - assert config_base.transformer_layer_architecture == "LlamaDecoderLayer" - assert config_base.layernorms is False - assert config_base.fusion_bias is False - assert config_base.speculators_config.algorithm == "eagle" - - -@pytest.mark.smoke -@pytest.mark.parametrize(("layer_architecture", "config_class"), LAYER_TYPES) -def test_eagle_speculator_config_backwards_compatibility_different_layers( - layer_architecture, config_class, sample_speculators_config -): - """Test backwards compatibility with different layer architectures.""" - model_type = layer_architecture.lower().replace("decoderlayer", "") - if model_type == "deepseekv3": - model_type = "deepseek_v3" - - config_dict = { - "speculators_model_type": "eagle", - "architectures": ["EagleSpeculator", layer_architecture], - "transformer_layer_architecture": layer_architecture, - "transformer_layer_config": { - "model_type": model_type, - "vocab_size": 32000, - "hidden_size": 768, - "intermediate_size": 3072, - "num_hidden_layers": 12, - "num_attention_heads": 12, - "max_position_embeddings": 2048, - }, - "layernorms": False, - "fusion_bias": False, - "speculators_config": sample_speculators_config.model_dump(), - } - - config = EagleSpeculatorConfig.model_validate(config_dict) - assert isinstance(config, EagleSpeculatorConfig) - assert config.speculators_model_type == "eagle" - assert config.transformer_layer_architecture == layer_architecture - assert layer_architecture in config.architectures - assert isinstance(config.transformer_layer_config, config_class) - - -@pytest.mark.smoke -def test_eagle_speculator_config_eagle12_dict_marshalling(eagle12_config_dict): - original_config = EagleSpeculatorConfig.model_validate(eagle12_config_dict) - - # Convert to dict with model_dump - config_dict = original_config.model_dump() - assert isinstance(config_dict, dict) - assert config_dict["speculators_model_type"] == "eagle" - assert config_dict["fusion_bias"] is False - - # Load with from_dict on base class - recreated_base = SpeculatorModelConfig.from_dict(config_dict) - assert isinstance(recreated_base, EagleSpeculatorConfig) - assert recreated_base.fusion_bias is False - assert recreated_base.layernorms is False - assert recreated_base.transformer_layer_architecture == "LlamaDecoderLayer" - - # Load with from_dict on derived class (should work through inheritance) - recreated_derived = EagleSpeculatorConfig.model_validate(config_dict) - assert isinstance(recreated_derived, EagleSpeculatorConfig) - assert recreated_derived.fusion_bias is False - assert recreated_derived.layernorms is False - assert recreated_derived.transformer_layer_architecture == "LlamaDecoderLayer" - - -@pytest.mark.smoke -def test_eagle_speculator_config_eagle12_from_pretrained_local_marshalling( - eagle12_config_dict, -): - original_config = EagleSpeculatorConfig.model_validate(eagle12_config_dict) - - with tempfile.TemporaryDirectory() as temp_dir: - temp_path = Path(temp_dir) - - # Save with save_pretrained - original_config.save_pretrained(temp_path) - - # Verify config.json was created - config_file = temp_path / "config.json" - assert config_file.exists() - - # Load with from_pretrained on base class - loaded_base = SpeculatorModelConfig.from_pretrained(temp_path) - assert isinstance(loaded_base, EagleSpeculatorConfig) - assert loaded_base.speculators_model_type == "eagle" - assert loaded_base.fusion_bias is False - assert loaded_base.layernorms is False - assert loaded_base.transformer_layer_architecture == "LlamaDecoderLayer" - - # Load with from_pretrained on derived class - loaded_derived = EagleSpeculatorConfig.from_pretrained(temp_path) - assert isinstance(loaded_derived, EagleSpeculatorConfig) - assert loaded_derived.speculators_model_type == "eagle" - assert loaded_derived.fusion_bias is False - assert loaded_derived.layernorms is False - assert loaded_derived.transformer_layer_architecture == "LlamaDecoderLayer" - - -@pytest.mark.smoke -@pytest.mark.parametrize(("layer_architecture", "config_class"), LAYER_TYPES) -def test_eagle_speculator_config_from_pretrained_different_layers( - layer_architecture, config_class, sample_speculators_config -): - """Test from_pretrained with different layer architectures.""" - layer_config = create_layer_config(config_class) - - original_config = EagleSpeculatorConfig( - transformer_layer_architecture=layer_architecture, - transformer_layer_config=layer_config, - layernorms=False, - fusion_bias=False, - speculators_config=sample_speculators_config, - ) - - with tempfile.TemporaryDirectory() as temp_dir: - temp_path = Path(temp_dir) - - # Save with save_pretrained - original_config.save_pretrained(temp_path) - - # Verify config.json was created - config_file = temp_path / "config.json" - assert config_file.exists() - - # Load with from_pretrained - loaded_config = EagleSpeculatorConfig.from_pretrained(temp_path) - assert isinstance(loaded_config, EagleSpeculatorConfig) - assert loaded_config.speculators_model_type == "eagle" - assert loaded_config.transformer_layer_architecture == layer_architecture - assert layer_architecture in loaded_config.architectures - - -# ====== EagleSpeculatorConfig HASS Tests ====== - - -@pytest.mark.smoke -def test_eagle_speculator_config_hass_backwards_compatibility(hass_config_dict): - config_derived = EagleSpeculatorConfig.model_validate(hass_config_dict) - assert isinstance(config_derived, EagleSpeculatorConfig) - assert config_derived.speculators_model_type == "eagle" - assert "LlamaDecoderLayer" in config_derived.architectures - assert config_derived.transformer_layer_architecture == "LlamaDecoderLayer" - assert config_derived.layernorms is False - assert config_derived.fusion_bias is True # Key difference for HASS - assert config_derived.speculators_config.algorithm == "eagle" - - # Test loading with base SpeculatorModelConfig.model_validate - config_base = SpeculatorModelConfig.model_validate(hass_config_dict) - assert isinstance(config_base, EagleSpeculatorConfig) - assert config_base.speculators_model_type == "eagle" - assert config_base.transformer_layer_architecture == "LlamaDecoderLayer" - assert config_base.layernorms is False - assert config_base.fusion_bias is True # Key difference for HASS - assert config_base.speculators_config.algorithm == "eagle" - - -@pytest.mark.smoke -@pytest.mark.parametrize(("layer_architecture", "config_class"), LAYER_TYPES) -def test_eagle_speculator_config_hass_different_layers( - layer_architecture, config_class -): - """Test HASS configuration with different layer architectures.""" - layer_config = create_layer_config(config_class) - - config = EagleSpeculatorConfig( - transformer_layer_architecture=layer_architecture, - transformer_layer_config=layer_config, - layernorms=False, - fusion_bias=True, # Key difference for HASS - ) - - assert isinstance(config, EagleSpeculatorConfig) - assert config.speculators_model_type == "eagle" - assert config.transformer_layer_architecture == layer_architecture - assert layer_architecture in config.architectures - assert config.fusion_bias is True # Key difference for HASS - assert config.layernorms is False - - -@pytest.mark.smoke -def test_eagle_speculator_config_hass_dict_marshalling(hass_config_dict): - original_config = EagleSpeculatorConfig.model_validate(hass_config_dict) - - # Convert to dict with model_dump - config_dict = original_config.model_dump() - assert isinstance(config_dict, dict) - assert config_dict["speculators_model_type"] == "eagle" - assert config_dict["fusion_bias"] is True # Key difference for HASS - - # Load with from_dict on base class - recreated_base = SpeculatorModelConfig.from_dict(config_dict) - assert isinstance(recreated_base, EagleSpeculatorConfig) - assert recreated_base.fusion_bias is True # Key difference for HASS - assert recreated_base.layernorms is False - assert recreated_base.transformer_layer_architecture == "LlamaDecoderLayer" - assert isinstance(recreated_base.transformer_layer_config, LlamaConfig) - - # Load with from_dict on derived class (should work through inheritance) - recreated_derived = EagleSpeculatorConfig.model_validate(config_dict) - assert isinstance(recreated_derived, EagleSpeculatorConfig) - assert recreated_derived.fusion_bias is True # Key difference for HASS - assert recreated_derived.layernorms is False - assert recreated_derived.transformer_layer_architecture == "LlamaDecoderLayer" - assert isinstance(recreated_derived.transformer_layer_config, LlamaConfig) - - -@pytest.mark.smoke -def test_eagle_speculator_config_hass_from_pretrained_local_marshalling( - hass_config_dict, -): - original_config = EagleSpeculatorConfig.model_validate(hass_config_dict) - - with tempfile.TemporaryDirectory() as temp_dir: - temp_path = Path(temp_dir) - - # Save with save_pretrained - original_config.save_pretrained(temp_path) - - # Verify config.json was created - config_file = temp_path / "config.json" - assert config_file.exists() - - # Load with from_pretrained on base class - loaded_base = SpeculatorModelConfig.from_pretrained(temp_path) - assert isinstance(loaded_base, EagleSpeculatorConfig) - assert loaded_base.speculators_model_type == "eagle" - assert loaded_base.fusion_bias is True # Key difference for HASS - assert loaded_base.layernorms is False - assert loaded_base.transformer_layer_architecture == "LlamaDecoderLayer" - assert isinstance(loaded_base.transformer_layer_config, LlamaConfig) - - # Load with from_pretrained on derived class - loaded_derived = EagleSpeculatorConfig.from_pretrained(temp_path) - assert isinstance(loaded_derived, EagleSpeculatorConfig) - assert loaded_derived.speculators_model_type == "eagle" - assert loaded_derived.fusion_bias is True # Key difference for HASS - assert loaded_derived.layernorms is False - assert loaded_derived.transformer_layer_architecture == "LlamaDecoderLayer" - assert isinstance(loaded_derived.transformer_layer_config, LlamaConfig) diff --git a/tests/unit/models/test_eagle_model.py b/tests/unit/models/test_eagle_model.py deleted file mode 100644 index 6c9502ab9..000000000 --- a/tests/unit/models/test_eagle_model.py +++ /dev/null @@ -1,803 +0,0 @@ -""" -Unit tests for the EagleSpeculator model in the Speculators library. -""" - -import copy -import tempfile -from unittest.mock import patch - -import pytest -import torch -from torch import nn -from transformers import PreTrainedModel -from transformers.configuration_utils import PretrainedConfig -from transformers.models.deepseek_v3.configuration_deepseek_v3 import DeepseekV3Config -from transformers.models.deepseek_v3.modeling_deepseek_v3 import ( - DeepseekV3DecoderLayer, - DeepseekV3RMSNorm, -) -from transformers.models.gemma.configuration_gemma import GemmaConfig -from transformers.models.gemma.modeling_gemma import GemmaDecoderLayer, GemmaRMSNorm -from transformers.models.granite.configuration_granite import GraniteConfig -from transformers.models.granite.modeling_granite import ( - GraniteDecoderLayer, - GraniteRMSNorm, -) -from transformers.models.llama.configuration_llama import LlamaConfig -from transformers.models.llama.modeling_llama import ( - LlamaDecoderLayer, - LlamaRMSNorm, - LlamaRotaryEmbedding, -) -from transformers.models.mistral.configuration_mistral import MistralConfig -from transformers.models.mistral.modeling_mistral import ( - MistralDecoderLayer, - MistralRMSNorm, -) -from transformers.models.mixtral.configuration_mixtral import MixtralConfig -from transformers.models.mixtral.modeling_mixtral import ( - MixtralDecoderLayer, - MixtralRMSNorm, -) -from transformers.models.qwen3.configuration_qwen3 import Qwen3Config -from transformers.models.qwen3.modeling_qwen3 import Qwen3DecoderLayer, Qwen3RMSNorm - -from speculators import ( - SpeculatorModel, - SpeculatorsConfig, - VerifierConfig, -) -from speculators.convert.eagle.eagle_legacy_model import ( - EagleSpeculator, - EagleSpeculatorConfig, -) -from speculators.proposals import GreedyTokenProposalConfig - -# ===== Layer Types Constants ===== - -LAYER_TYPES: dict[str, tuple[type, type, type]] = { - # Format: "LayerName": (LayerClass, NormClass, ConfigClass) - "LlamaDecoderLayer": (LlamaDecoderLayer, LlamaRMSNorm, LlamaConfig), - "MistralDecoderLayer": (MistralDecoderLayer, MistralRMSNorm, MistralConfig), - "Qwen3DecoderLayer": (Qwen3DecoderLayer, Qwen3RMSNorm, Qwen3Config), - "GemmaDecoderLayer": (GemmaDecoderLayer, GemmaRMSNorm, GemmaConfig), - "MixtralDecoderLayer": (MixtralDecoderLayer, MixtralRMSNorm, MixtralConfig), - "DeepseekV3DecoderLayer": ( - DeepseekV3DecoderLayer, - DeepseekV3RMSNorm, - DeepseekV3Config, - ), - "GraniteDecoderLayer": (GraniteDecoderLayer, GraniteRMSNorm, GraniteConfig), -} - -LAYER_ARCHITECTURES = list(LAYER_TYPES.keys()) - -# ===== Test Helper Classes ===== - - -class MockVerifier(PreTrainedModel): - def __init__(self, config): - super().__init__(config) - self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) - self.rotary_emb = LlamaRotaryEmbedding(config) - self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) - - def forward(self, input_ids, **kwargs): - embeddings = self.embed_tokens(input_ids) - return type( - "MockOutput", - (), - {"last_hidden_state": embeddings, "hidden_states": (embeddings,)}, - )() - - -# ===== Fixtures ===== - - -@pytest.fixture -def sample_llama_config(): - return LlamaConfig( - attention_bias=False, - attention_dropout=0.0, - bos_token_id=128000, - eos_token_id=128001, - head_dim=128, - hidden_act="silu", - hidden_size=4096, - initializer_range=0.02, - intermediate_size=14336, - max_position_embeddings=131072, - mlp_bias=False, - num_attention_heads=32, - num_hidden_layers=32, - num_key_value_heads=8, - pretraining_tp=1, - rms_norm_eps=1e-5, # type: ignore[arg-type] # (bad transformer's type hint, int instead of float) - tie_word_embeddings=False, - transformers_version="4.46.0", - use_cache=True, - vocab_size=128256, - ) - - -# ===== Config Helper Function ===== - - -def create_layer_config_for_architecture(layer_architecture: str): - config_class = LAYER_TYPES[layer_architecture][ - 2 - ] # Third element is the config class - base_params = { - "vocab_size": 32000, - "hidden_size": 768, - "intermediate_size": 3072, - "num_hidden_layers": 12, - "num_attention_heads": 12, - "max_position_embeddings": 2048, - } - - # Add extra parameters for specific config types - if config_class in (MixtralConfig, DeepseekV3Config, GraniteConfig): - base_params["num_key_value_heads"] = 12 - - if config_class == MixtralConfig: - base_params.update( - { - "num_local_experts": 8, - "num_experts_per_tok": 2, - } - ) - - return config_class(**base_params) - - -@pytest.fixture -def sample_verifier_config(): - return VerifierConfig( - name_or_path="test/verifier", - architectures=["LlamaForCausalLM"], - ) - - -@pytest.fixture -def sample_speculators_config(sample_verifier_config): - return SpeculatorsConfig( - algorithm="eagle_v1", - proposal_methods=[GreedyTokenProposalConfig()], - default_proposal_method="greedy", - verifier=sample_verifier_config, - ) - - -@pytest.fixture -def eagle_speculator_config(sample_llama_config, sample_speculators_config): - return EagleSpeculatorConfig( - transformer_layer_config=sample_llama_config, - speculators_config=sample_speculators_config, - ) - - -@pytest.fixture -def eagle_speculator_config_layernorms(sample_llama_config, sample_speculators_config): - return EagleSpeculatorConfig( - transformer_layer_config=sample_llama_config, - speculators_config=sample_speculators_config, - layernorms=True, - fusion_bias=True, - ) - - -@pytest.fixture -def mock_verifier(sample_llama_config): - return MockVerifier(sample_llama_config) - - -# ===== EagleSpeculator Class Attributes Tests ===== - - -@pytest.mark.smoke -def test_eagle_speculator_class_attributes(): - assert EagleSpeculator.auto_package == "speculators.models" - assert EagleSpeculator.registry_auto_discovery is True - assert EagleSpeculator.config_class == EagleSpeculatorConfig - assert EagleSpeculator.base_model_prefix == "model" - assert EagleSpeculator.main_input_name == "input_ids" - - -# ===== EagleSpeculator Registry Tests ===== - - -@pytest.mark.smoke -def test_eagle_speculator_registry(): - assert SpeculatorModel.registry is not None - assert "eagle" in SpeculatorModel.registry - assert SpeculatorModel.registry["eagle"] == EagleSpeculator - - -@pytest.mark.smoke -def test_eagle_speculator_registered_model_class_from_config(eagle_speculator_config): - model_class = SpeculatorModel.registered_model_class_from_config( - eagle_speculator_config - ) - assert model_class == EagleSpeculator - - -# ===== EagleSpeculator Initialization Tests ===== - - -@pytest.mark.smoke -def test_eagle_speculator_initialization_without_verifier(eagle_speculator_config): - eagle_speculator_config = copy.deepcopy(eagle_speculator_config) - eagle_speculator_config.speculators_config.verifier.name_or_path = None - model = EagleSpeculator(eagle_speculator_config) - - assert model.config == eagle_speculator_config - assert model.verifier is None - assert model.verifier_attachment_mode == "detached" - - # Verifier-dependent layers should be None - assert model.embed_tokens is None - assert model.rotary_emb is None - assert model.lm_head is None - - # Model-specific layers should be initialized - assert model.fusion_fc is not None - assert model.transformer is not None - assert isinstance(model.fusion_fc, nn.Linear) - assert isinstance(model.transformer, LlamaDecoderLayer) - - -@pytest.mark.smoke -def test_eagle_speculator_initialization_with_verifier( - eagle_speculator_config, mock_verifier -): - model = EagleSpeculator(eagle_speculator_config, verifier=mock_verifier) - - assert model.config == eagle_speculator_config - assert model.verifier == mock_verifier - assert model.verifier_attachment_mode == "full" - - # Verifier-dependent layers should be attached - assert model.embed_tokens is not None - assert model.rotary_emb is not None - assert model.lm_head is not None - assert model.embed_tokens == mock_verifier.embed_tokens - assert model.rotary_emb == mock_verifier.rotary_emb - assert model.lm_head == mock_verifier.lm_head - - -@pytest.mark.smoke -def test_eagle_speculator_initialization_with_verifier_path( - eagle_speculator_config, mock_verifier -): - with patch( - "transformers.AutoModelForCausalLM.from_pretrained", return_value=mock_verifier - ): - verifier_path = "path/to/verifier/model" - model = EagleSpeculator( - eagle_speculator_config, - verifier=verifier_path, - verifier_attachment_mode=None, - ) - - assert model.config == eagle_speculator_config - assert model.verifier == mock_verifier - assert model.verifier_attachment_mode == "full" - assert model.embed_tokens is not None - assert model.rotary_emb is not None - assert model.lm_head is not None - assert model.embed_tokens == mock_verifier.embed_tokens - assert model.rotary_emb == mock_verifier.rotary_emb - assert model.lm_head == mock_verifier.lm_head - - -@pytest.mark.smoke -def test_eagle_speculator_initialization_with_verifier_train_only( - eagle_speculator_config, mock_verifier -): - model = EagleSpeculator( - eagle_speculator_config, - verifier=mock_verifier, - verifier_attachment_mode="train_only", - ) - - assert model.config == eagle_speculator_config - assert model.verifier is None - assert model.verifier_attachment_mode == "train_only" - assert model.embed_tokens is not None - assert model.rotary_emb is not None - assert model.lm_head is not None - assert model.embed_tokens == mock_verifier.embed_tokens - assert model.rotary_emb == mock_verifier.rotary_emb - assert model.lm_head == mock_verifier.lm_head - - -@pytest.mark.smoke -def test_eagle_speculator_initialization_with_verifier_detached( - eagle_speculator_config, mock_verifier -): - model = EagleSpeculator( - eagle_speculator_config, - verifier=mock_verifier, - verifier_attachment_mode="detached", - ) - - assert model.config == eagle_speculator_config - assert model.verifier is None - assert model.verifier_attachment_mode == "detached" - assert model.embed_tokens is None - assert model.rotary_emb is None - assert model.lm_head is None - - -# ===== EagleSpeculator from_pretrained Tests ===== - - -@pytest.mark.smoke -def test_eagle_speculator_from_pretrained_config( - eagle_speculator_config, mock_verifier -): - eagle_speculator_config = copy.deepcopy(eagle_speculator_config) - state_dict = EagleSpeculator( - eagle_speculator_config, verifier_attachment_mode="detached" - ).state_dict() - model = SpeculatorModel.from_pretrained( - None, - config=eagle_speculator_config, - verifier=mock_verifier, - state_dict=state_dict, - ) - - eagle_speculator_config.torch_dtype = torch.float32 - assert isinstance(model, EagleSpeculator) - assert model.config == eagle_speculator_config - assert model.verifier is not None - assert model.verifier_attachment_mode == "full" - assert model.embed_tokens == mock_verifier.embed_tokens - assert model.rotary_emb == mock_verifier.rotary_emb - assert model.lm_head == mock_verifier.lm_head - - -@pytest.mark.smoke -def test_eagle_speculator_from_pretrained_local_marshalling( - eagle_speculator_config, mock_verifier -): - eagle_speculator_config = copy.deepcopy(eagle_speculator_config) - state_dict = EagleSpeculator( - eagle_speculator_config, verifier_attachment_mode="detached" - ).state_dict() - - with tempfile.TemporaryDirectory() as tmpdir: - model = SpeculatorModel.from_pretrained( - None, - config=eagle_speculator_config, - verifier=mock_verifier, - state_dict=state_dict, - ) - model.save_pretrained(tmpdir) # type: ignore[attr-defined] - - loaded_model = SpeculatorModel.from_pretrained(tmpdir, verifier=mock_verifier) - eagle_speculator_config.torch_dtype = torch.float32 - - assert isinstance(loaded_model, EagleSpeculator) - assert isinstance(loaded_model.config, EagleSpeculatorConfig) - assert ( - loaded_model.config.transformer_layer_architecture - == eagle_speculator_config.transformer_layer_architecture - ) - assert loaded_model.config.layernorms == eagle_speculator_config.layernorms - assert loaded_model.config.fusion_bias == eagle_speculator_config.fusion_bias - assert ( - loaded_model.config.speculators_config - == eagle_speculator_config.speculators_config - ) - assert loaded_model.verifier == mock_verifier - assert loaded_model.verifier_attachment_mode == "full" - assert loaded_model.embed_tokens == mock_verifier.embed_tokens - assert loaded_model.rotary_emb == mock_verifier.rotary_emb - assert loaded_model.lm_head == mock_verifier.lm_head - - -# ===== EagleSpeculator Architecture Tests ===== - - -@pytest.mark.smoke -def test_eagle_speculator_architecture_eagle(eagle_speculator_config, mock_verifier): - model = EagleSpeculator( - eagle_speculator_config, verifier=mock_verifier, verifier_attachment_mode="full" - ) - llama_config: LlamaConfig = eagle_speculator_config.transformer_layer_config - - assert isinstance(model, EagleSpeculator) - assert isinstance(model.config, EagleSpeculatorConfig) - assert model.embed_tokens is not None - assert isinstance(model.embed_tokens, nn.Embedding) - assert model.embed_tokens.weight.shape == ( - llama_config.vocab_size, - llama_config.hidden_size, - ) - assert model.rotary_emb is not None - assert isinstance(model.rotary_emb, LlamaRotaryEmbedding) - assert model.lm_head is not None - assert isinstance(model.lm_head, nn.Linear) - assert model.lm_head.weight.shape == ( - llama_config.vocab_size, - llama_config.hidden_size, - ) - assert model.lm_head.bias is None - assert model.embedding_layernorm is None - assert model.fusion_fc is not None - assert isinstance(model.fusion_fc, nn.Linear) - assert model.fusion_fc.weight.shape == ( - llama_config.hidden_size, - 2 * llama_config.hidden_size, - ) - assert model.fusion_fc.bias is None - assert model.transformer is not None - assert isinstance(model.transformer, LlamaDecoderLayer) - assert model.transformer.self_attn.config.hidden_size == llama_config.hidden_size - assert isinstance(model.transformer.input_layernorm, nn.Identity) - assert model.pre_lm_head_layernorm is None - - -@pytest.mark.smoke -def test_eagle_speculator_architecture_hass( - eagle_speculator_config_layernorms, mock_verifier -): - model = EagleSpeculator( - eagle_speculator_config_layernorms, - verifier=mock_verifier, - verifier_attachment_mode="full", - ) - llama_config: LlamaConfig = ( - eagle_speculator_config_layernorms.transformer_layer_config - ) - - assert isinstance(model, EagleSpeculator) - assert isinstance(model.config, EagleSpeculatorConfig) - assert model.embed_tokens is not None - assert isinstance(model.embed_tokens, nn.Embedding) - assert model.embed_tokens.weight.shape == ( - llama_config.vocab_size, - llama_config.hidden_size, - ) - assert model.rotary_emb is not None - assert isinstance(model.rotary_emb, LlamaRotaryEmbedding) - assert model.lm_head is not None - assert isinstance(model.lm_head, nn.Linear) - assert model.lm_head.weight.shape == ( - llama_config.vocab_size, - llama_config.hidden_size, - ) - assert model.embedding_layernorm is not None - assert isinstance(model.embedding_layernorm, LlamaRMSNorm) - assert model.embedding_layernorm.weight.shape == (llama_config.hidden_size,) - assert model.fusion_fc is not None - assert isinstance(model.fusion_fc, nn.Linear) - assert llama_config.hidden_size is not None # typing - assert model.fusion_fc.weight.shape == ( - llama_config.hidden_size, - 2 * llama_config.hidden_size, - ) - assert model.fusion_fc.bias is not None - assert model.transformer is not None - assert isinstance(model.transformer, LlamaDecoderLayer) - assert model.transformer.self_attn.config.hidden_size == llama_config.hidden_size - assert isinstance(model.transformer.input_layernorm, LlamaRMSNorm) - assert model.pre_lm_head_layernorm is not None - assert isinstance(model.pre_lm_head_layernorm, LlamaRMSNorm) - - -# ===== EagleSpeculator Architecture Tests with Different Layer Types ===== - - -@pytest.mark.smoke -@pytest.mark.parametrize("layer_architecture", LAYER_ARCHITECTURES) -def test_eagle_speculator_initialization_different_layers( - layer_architecture, sample_speculators_config -): - """Test EagleSpeculator initialization with different layer architectures.""" - layer_config = create_layer_config_for_architecture(layer_architecture) - - # Create EagleSpeculatorConfig with the specific layer architecture - eagle_config = EagleSpeculatorConfig( - transformer_layer_architecture=layer_architecture, - transformer_layer_config=layer_config, - speculators_config=sample_speculators_config, - ) - - # Create mock verifier for this config - mock_verifier = MockVerifier(layer_config) - - model = EagleSpeculator(eagle_config, verifier=mock_verifier) - - assert model.config == eagle_config - assert model.verifier == mock_verifier - assert model.verifier_attachment_mode == "full" - - # Verify basic architecture - assert model.embed_tokens is not None - assert model.embed_tokens == mock_verifier.embed_tokens - assert model.rotary_emb is not None - assert model.rotary_emb == mock_verifier.rotary_emb - assert model.lm_head is not None - assert model.lm_head == mock_verifier.lm_head - assert model.fusion_fc is not None - assert model.transformer is not None - - -@pytest.mark.smoke -@pytest.mark.parametrize("layer_architecture", LAYER_ARCHITECTURES) -def test_eagle_speculator_architecture_different_layers( - layer_architecture, sample_speculators_config -): - """Test EagleSpeculator architecture with different layer types.""" - layer_config = create_layer_config_for_architecture(layer_architecture) - - # Create EagleSpeculatorConfig with the specific layer architecture - eagle_config = EagleSpeculatorConfig( - transformer_layer_architecture=layer_architecture, - transformer_layer_config=layer_config, - speculators_config=sample_speculators_config, - ) - - # Create mock verifier for this config - mock_verifier = MockVerifier(layer_config) - - model = EagleSpeculator( - eagle_config, verifier=mock_verifier, verifier_attachment_mode="full" - ) - - assert isinstance(model, EagleSpeculator) - assert isinstance(model.config, EagleSpeculatorConfig) - - # Verify embedding layer - assert model.embed_tokens is not None - assert isinstance(model.embed_tokens, nn.Embedding) - assert model.embed_tokens.weight.shape == ( - layer_config.vocab_size, - layer_config.hidden_size, - ) - - # Verify rotary embedding - assert model.rotary_emb is not None - assert isinstance(model.rotary_emb, LlamaRotaryEmbedding) - - # Verify language model head - assert model.lm_head is not None - assert isinstance(model.lm_head, nn.Linear) - assert model.lm_head.weight.shape == ( - layer_config.vocab_size, - layer_config.hidden_size, - ) - assert model.lm_head.bias is None - - # Verify fusion layer - assert model.fusion_fc is not None - assert isinstance(model.fusion_fc, nn.Linear) - assert model.fusion_fc.weight.shape == ( - layer_config.hidden_size, - 2 * layer_config.hidden_size, - ) - assert model.fusion_fc.bias is None - - # Verify transformer layer - assert model.transformer is not None - assert isinstance(model.transformer, LAYER_TYPES[layer_architecture][0]) - assert model.transformer.self_attn.config.hidden_size == layer_config.hidden_size - - # Verify no layernorms by default - assert model.embedding_layernorm is None - assert model.pre_lm_head_layernorm is None - - -@pytest.mark.smoke -@pytest.mark.parametrize("layer_architecture", LAYER_ARCHITECTURES) -def test_eagle_speculator_architecture_different_layers_with_layernorms( - layer_architecture, sample_speculators_config -): - """Test EagleSpeculator with layernorms enabled for different decoder layers.""" - layer_config = create_layer_config_for_architecture(layer_architecture) - - # Create EagleSpeculatorConfig with layernorms enabled - eagle_config = EagleSpeculatorConfig( - transformer_layer_architecture=layer_architecture, - transformer_layer_config=layer_config, - speculators_config=sample_speculators_config, - layernorms=True, - fusion_bias=True, - ) - - # Create mock verifier for this config - mock_verifier = MockVerifier(layer_config) - - model = EagleSpeculator( - eagle_config, verifier=mock_verifier, verifier_attachment_mode="full" - ) - - assert isinstance(model, EagleSpeculator) - assert isinstance(model.config, EagleSpeculatorConfig) - - # Verify embedding layer - assert model.embed_tokens is not None - assert isinstance(model.embed_tokens, nn.Embedding) - assert model.embed_tokens.weight.shape == ( - layer_config.vocab_size, - layer_config.hidden_size, - ) - - # Verify embedding layernorm - assert model.embedding_layernorm is not None - assert isinstance(model.embedding_layernorm, LAYER_TYPES[layer_architecture][1]) - assert model.embedding_layernorm.weight.shape == (layer_config.hidden_size,) # type: ignore[attr-defined] - - # Verify fusion layer with bias - assert model.fusion_fc is not None - assert isinstance(model.fusion_fc, nn.Linear) - assert model.fusion_fc.weight.shape == ( - layer_config.hidden_size, - 2 * layer_config.hidden_size, - ) - assert model.fusion_fc.bias is not None - - # Verify pre-lm-head layernorm - assert model.pre_lm_head_layernorm is not None - assert isinstance(model.pre_lm_head_layernorm, LAYER_TYPES[layer_architecture][1]) - - -@pytest.mark.smoke -@pytest.mark.parametrize("layer_architecture", LAYER_ARCHITECTURES) -def test_eagle_speculator_from_pretrained_different_layers( - layer_architecture, sample_speculators_config -): - """Test EagleSpeculator from_pretrained with different layer architectures.""" - layer_config = create_layer_config_for_architecture(layer_architecture) - - # Create EagleSpeculatorConfig with the specific layer architecture - eagle_config = EagleSpeculatorConfig( - transformer_layer_architecture=layer_architecture, - transformer_layer_config=layer_config, - speculators_config=sample_speculators_config, - ) - - # Create state dict from a detached model - state_dict = EagleSpeculator( - eagle_config, verifier_attachment_mode="detached" - ).state_dict() - - # Create mock verifier for this config - mock_verifier = MockVerifier(layer_config) - - # Load model using from_pretrained - model = SpeculatorModel.from_pretrained( - None, - config=eagle_config, - verifier=mock_verifier, - state_dict=state_dict, - ) - - assert isinstance(model, EagleSpeculator) - assert model.config.transformer_layer_architecture == layer_architecture - assert model.verifier == mock_verifier - assert model.verifier_attachment_mode == "full" - assert model.embed_tokens == mock_verifier.embed_tokens - assert model.rotary_emb == mock_verifier.rotary_emb - assert model.lm_head == mock_verifier.lm_head - - -@pytest.mark.smoke -@pytest.mark.parametrize("layer_architecture", LAYER_ARCHITECTURES) -def test_eagle_speculator_local_marshalling_different_layers( - layer_architecture, sample_speculators_config -): - """Test EagleSpeculator local marshalling with different layer architectures.""" - layer_config = create_layer_config_for_architecture(layer_architecture) - - # Create EagleSpeculatorConfig with the specific layer architecture - eagle_config = EagleSpeculatorConfig( - transformer_layer_architecture=layer_architecture, - transformer_layer_config=layer_config, - speculators_config=sample_speculators_config, - ) - - # Create state dict from a detached model - state_dict = EagleSpeculator( - eagle_config, verifier_attachment_mode="detached" - ).state_dict() - - # Create mock verifier for this config - mock_verifier = MockVerifier(layer_config) - - with tempfile.TemporaryDirectory() as tmpdir: - # Create and save model - model = SpeculatorModel.from_pretrained( - None, - config=eagle_config, - verifier=mock_verifier, - state_dict=state_dict, - ) - model.save_pretrained(tmpdir) # type: ignore[attr-defined] - - # Load model from saved directory - loaded_model = SpeculatorModel.from_pretrained(tmpdir, verifier=mock_verifier) - - assert isinstance(loaded_model, EagleSpeculator) - assert isinstance(loaded_model.config, EagleSpeculatorConfig) - assert loaded_model.config.transformer_layer_architecture == layer_architecture - assert loaded_model.verifier == mock_verifier - assert loaded_model.verifier_attachment_mode == "full" - assert loaded_model.embed_tokens == mock_verifier.embed_tokens - assert loaded_model.rotary_emb == mock_verifier.rotary_emb - assert loaded_model.lm_head == mock_verifier.lm_head - - -# ===== EagleSpeculator Architecture Auto-Detection Tests ===== - - -@pytest.mark.smoke -def test_eagle_speculator_auto_architecture_derivation(sample_speculators_config): - layer_config = LlamaConfig() - - # Create config with auto architecture - eagle_config = EagleSpeculatorConfig( - transformer_layer_architecture="auto", - transformer_layer_config=layer_config, - speculators_config=sample_speculators_config, - ) - - # Create mock verifier - mock_verifier = MockVerifier(layer_config) - - # Create model - this should work with auto architecture - model = EagleSpeculator(eagle_config, verifier=mock_verifier) - - # Verify the model was created successfully - assert isinstance(model, EagleSpeculator) - # This value is set during initialization when a decoder layer class is found - assert model.config.transformer_layer_architecture == "LlamaDecoderLayer" - assert model.config.architectures is not None - assert "LlamaDecoderLayer" in model.config.architectures - assert model.verifier == mock_verifier - assert model.verifier_attachment_mode == "full" - - # Verify the transformer layer is the correct type - assert isinstance(model.transformer, LlamaDecoderLayer) - - -@pytest.mark.smoke -def test_eagle_speculator_auto_architecture_error_handling(): - # Create a custom config class that doesn't have a corresponding decoder layer - class CustomConfig(PretrainedConfig): - model_type = "custom" - - def __init__(self, **kwargs): - super().__init__(**kwargs) - self.vocab_size = 32000 - self.hidden_size = 768 - self.intermediate_size = 3072 - self.num_hidden_layers = 12 - self.num_attention_heads = 12 - self.max_position_embeddings = 2048 - self.pad_token_id = 0 - - # This config is not in MODEL_FOR_CAUSAL_LM_MAPPING, so it should fail - custom_config = CustomConfig() - - eagle_config = EagleSpeculatorConfig( - transformer_layer_architecture="auto", - transformer_layer_config=custom_config, - speculators_config=SpeculatorsConfig( - algorithm="eagle", - proposal_methods=[GreedyTokenProposalConfig()], - default_proposal_method="greedy", - verifier=VerifierConfig( - name_or_path="test/verifier", - architectures=["CustomForCausalLM"], - ), - ), - ) - - with pytest.raises( - TypeError, match="is not a valid causal language model config class" - ): - EagleSpeculator(eagle_config, verifier_attachment_mode="detached")