diff --git a/scripts/train.py b/scripts/train.py index e5eaca887..a9b99fae6 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -177,8 +177,6 @@ def main(args: argparse.Namespace): ) # Get model class from registry and create model using its factory method - if SpeculatorModel.registry_auto_discovery: - SpeculatorModel.auto_populate_registry() if args.speculator_type not in SpeculatorModel.registry: raise ValueError( diff --git a/src/speculators/__init__.py b/src/speculators/__init__.py index a5f8bf282..0d34de881 100644 --- a/src/speculators/__init__.py +++ b/src/speculators/__init__.py @@ -23,22 +23,23 @@ from .config import ( SpeculatorModelConfig, SpeculatorsConfig, - TokenProposalConfig, VerifierConfig, - reload_and_populate_configs, + reload_schemas, ) -from .model import SpeculatorModel, reload_and_populate_models +from .model import SpeculatorModel +from .models import Eagle3DraftModel, Eagle3SpeculatorConfig +from .proposals import TokenProposalConfig __all__ = [ + "Eagle3DraftModel", + "Eagle3SpeculatorConfig", "SpeculatorModel", "SpeculatorModelConfig", "SpeculatorsConfig", "TokenProposalConfig", "VerifierConfig", - "reload_and_populate_configs", - "reload_and_populate_models", + "reload_schemas", ] # base imports complete, run auto loading for base classes -reload_and_populate_configs() -reload_and_populate_models() +reload_schemas() diff --git a/src/speculators/config.py b/src/speculators/config.py index 1efc5ce6b..1df6b4ccb 100644 --- a/src/speculators/config.py +++ b/src/speculators/config.py @@ -25,49 +25,17 @@ from pydantic import BaseModel, ConfigDict, Field from transformers import PretrainedConfig +from speculators.proposals import TokenProposalConfig from speculators.utils import PydanticClassRegistryMixin, ReloadableBaseModel __all__ = [ "SpeculatorModelConfig", "SpeculatorsConfig", - "TokenProposalConfig", "VerifierConfig", - "reload_and_populate_configs", + "reload_schemas", ] -class TokenProposalConfig(PydanticClassRegistryMixin): - """ - The base config for a token proposal method which defines how tokens are generated - by the speculator, how they are passed to the verifier, and how they are scored - for acceptance or rejection. All implementations of token proposal methods - must inherit from this class, set the proposal_type to a unique value, and - add any additional parameters needed to instantiate and implement the method. - - It uses pydantic to validate the parameters, provide default values, and - enable automatic serialization and deserialization of the correct class - types based on the proposal_type field. - """ - - @classmethod - def __pydantic_schema_base_type__(cls) -> type["TokenProposalConfig"]: - if cls.__name__ == "TokenProposalConfig": - return cls - - return TokenProposalConfig - - auto_package: ClassVar[str] = "speculators.proposals" - registry_auto_discovery: ClassVar[bool] = True - schema_discriminator: ClassVar[str] = "proposal_type" - - proposal_type: str = Field( - description=( - "The type of token proposal the config is for. " - "Must be a supported proposal type from the Speculators repo." - ), - ) - - class VerifierConfig(BaseModel): """ The base config for a verifier model which defines the parameters that are required @@ -330,12 +298,12 @@ def to_diff_dict(self) -> dict[str, Any]: return super().to_diff_dict() -def reload_and_populate_configs(): +def reload_schemas(): """ Automatically populates the registry for all PydanticClassRegistryMixin subclasses and reloads schemas for all Config classes to ensure their schemas are up-to-date with the current registry state. """ - TokenProposalConfig.auto_populate_registry() + TokenProposalConfig.reload_schema() SpeculatorsConfig.reload_schema() - SpeculatorModelConfig.auto_populate_registry() + SpeculatorModelConfig.reload_schema() diff --git a/src/speculators/convert/eagle/__init__.py b/src/speculators/convert/eagle/__init__.py index 64777b87f..94d6e15db 100644 --- a/src/speculators/convert/eagle/__init__.py +++ b/src/speculators/convert/eagle/__init__.py @@ -2,6 +2,7 @@ Eagle checkpoint conversion utilities. """ +from speculators.convert.eagle.eagle3_converter import Eagle3Converter from speculators.convert.eagle.eagle_converter import EagleConverter -__all__ = ["EagleConverter"] +__all__ = ["Eagle3Converter", "EagleConverter"] diff --git a/src/speculators/convert/eagle/eagle3_legacy_model.py b/src/speculators/convert/eagle/eagle3_legacy_model.py index 4ba8dc186..49f458fa1 100644 --- a/src/speculators/convert/eagle/eagle3_legacy_model.py +++ b/src/speculators/convert/eagle/eagle3_legacy_model.py @@ -242,7 +242,7 @@ class Eagle3Speculator(SpeculatorModel): """ config_class: ClassVar[type[Eagle3SpeculatorConfig]] = Eagle3SpeculatorConfig # type: ignore[misc] - _keys_to_ignore_on_load_missing: ClassVar[list[str]] = [ # type: ignore[misc] + _keys_to_ignore_on_load_missing: ClassVar[list[str]] = [ # type: ignore[assignment,misc] "verifier*", ] _keys_to_ignore_on_save: ClassVar[list[str]] = [] # type: ignore[misc,assignment] diff --git a/src/speculators/convert/eagle/eagle_converter.py b/src/speculators/convert/eagle/eagle_converter.py index c792e30ed..df5985488 100644 --- a/src/speculators/convert/eagle/eagle_converter.py +++ b/src/speculators/convert/eagle/eagle_converter.py @@ -9,13 +9,16 @@ 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 ( detect_fusion_bias_and_layernorms, ensure_checkpoint_is_local, load_checkpoint_config, load_checkpoint_weights, ) -from speculators.models.eagle import EagleSpeculator, EagleSpeculatorConfig from speculators.proposals.greedy import GreedyTokenProposalConfig diff --git a/src/speculators/models/eagle.py b/src/speculators/convert/eagle/eagle_legacy_model.py similarity index 86% rename from src/speculators/models/eagle.py rename to src/speculators/convert/eagle/eagle_legacy_model.py index da0841e2d..2b4a57d10 100644 --- a/src/speculators/models/eagle.py +++ b/src/speculators/convert/eagle/eagle_legacy_model.py @@ -20,7 +20,12 @@ import torch from pydantic import Field, field_serializer, field_validator, model_validator from torch import nn -from transformers import AutoConfig, PretrainedConfig, PreTrainedModel +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 @@ -50,28 +55,6 @@ class EagleSpeculatorConfig(SpeculatorModelConfig): - EAGLE1: layernorms=False, fusion_bias=False - EAGLE2: layernorms=False, fusion_bias=False - HASS: layernorms=False, fusion_bias=True - - Example: - ```python - from speculators import SpeculatorsConfig, VerifierConfig - from speculators.models import EagleSpeculatorConfig - from speculators.proposals import GreedyTokenProposalConfig - from transformers import AutoConfig - - config = EagleSpeculatorConfig( - transformer_layer_config=AutoConfig.from_pretrained("meta-llama/Llama-3.1-8B-Instruct"), - speculators_config=SpeculatorsConfig( - algorithm="eagle", - proposal_methods=[ - GreedyTokenProposalConfig(), - ], - default_proposal_method="greedy", - verifier=VerifierConfig( - name_or_path="meta-llama/Llama-3.1-8B-Instruct", - architectures=["LlamaForCausalLM"], - ) - ) - ``` """ speculators_model_type: Literal["eagle"] = "eagle" @@ -205,44 +188,21 @@ class EagleSpeculator(SpeculatorModel): 4. Verifier validates candidates and accepts/rejects based on probability thresholds 5. Process continues iteratively for multi-token speculation - - Example: - ```python - from speculators import SpeculatorsConfig, VerifierConfig - from speculators.models import EagleSpeculator, EagleSpeculatorConfig - from speculators.proposals import GreedyTokenProposalConfig - from transformers import AutoConfig, AutoTokenizer - - config = EagleSpeculatorConfig( - transformer_layer_config=AutoConfig.from_pretrained("meta-llama/Llama-3.1-8B-Instruct"), - speculators_config=SpeculatorsConfig( - algorithm="eagle", - proposal_methods=[ - GreedyTokenProposalConfig(), - ], - default_proposal_method="greedy", - verifier=VerifierConfig( - name_or_path="meta-llama/Llama-3.1-8B-Instruct", - architectures=["LlamaForCausalLM"], - ) - ) - speculator = EagleSpeculator( - config, verifier=verifier, verifier_attachment_mode="full" - ) - ``` """ # PreTrainedModel settings config_class: ClassVar[type[EagleSpeculatorConfig]] = EagleSpeculatorConfig # type: ignore[misc] - _keys_to_ignore_on_load_missing: ClassVar[list[str]] = [ # 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 @@ -340,13 +300,16 @@ def __init__( self.rotary_emb: nn.Module | None = None self.lm_head: nn.Linear | None = None - # Delayed initialization to ensure everything needed for attach_verifier is set - super().__init__( - config=config, - verifier=verifier, - verifier_attachment_mode=verifier_attachment_mode, + 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() @@ -360,6 +323,67 @@ def __init__( 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, @@ -405,7 +429,24 @@ def attach_verifier( perform generation until a full verifier is attached. :return: The PreTrainedModel instance for the verifier that was attached. """ - super().attach_verifier(verifier=verifier, mode=mode) + 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) @@ -432,7 +473,17 @@ def detach_verifier(self): be able to perform forward passes or generation until a new verifier is attached. """ - super().detach_verifier() + 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 diff --git a/src/speculators/model.py b/src/speculators/model.py index 00d830720..8fd80fb76 100644 --- a/src/speculators/model.py +++ b/src/speculators/model.py @@ -2,47 +2,20 @@ Base model classes for the Speculators library. This module contains the base model classes for speculative decoding implementations -in the Speculators library. These classes provide the foundation for creating -speculator models that can perform speculative token generation with verifier -models for accelerated inference. - -The models extend Hugging Face's PreTrainedModel and GenerationMixin to maintain -full compatibility with the transformers ecosystem while adding speculative -decoding capabilities. They support automatic model registration and discovery, -dynamic model loading based on configuration, and flexible verifier attachment. - -Classes: - SpeculatorModel: Abstract base class for all speculator models with transformers - compatibility, automatic registry support, and speculative generation methods - -Functions: - reload_and_populate_models: Automatically populates the model registry for - discovery and instantiation of registered speculator models +in the Speculators library. """ import os from abc import abstractmethod -from collections.abc import Callable -from typing import Any, ClassVar, Literal, Optional - -import torch -from transformers import ( - AutoModelForCausalLM, - GenerationConfig, - GenerationMixin, - PretrainedConfig, - PreTrainedModel, -) -from transformers.generation.logits_process import LogitsProcessorList -from transformers.generation.stopping_criteria import StoppingCriteriaList -from transformers.generation.streamers import BaseStreamer -from transformers.generation.utils import GenerateOutput +from typing import ClassVar + +from transformers import PretrainedConfig, PreTrainedModel from speculators.config import SpeculatorModelConfig from speculators.utils import ClassRegistryMixin -class SpeculatorModel(ClassRegistryMixin, PreTrainedModel, GenerationMixin): # type: ignore[misc] +class SpeculatorModel(ClassRegistryMixin, PreTrainedModel): # type: ignore[misc] """ Abstract base class for all speculator models in the Speculators library. @@ -58,13 +31,6 @@ class SpeculatorModel(ClassRegistryMixin, PreTrainedModel, GenerationMixin): # ```python # Load a speculator model with automatic class resolution model = SpeculatorModel.from_pretrained("path/to/speculator") - - # Optionally attach a new verifier model - verifier = AutoModel.from_pretrained("path/to/verifier") - model.attach_verifier(verifier) - - # Generate with speculative decoding - outputs = model.generate(input_ids, max_length=100) ``` """ @@ -76,18 +42,13 @@ class SpeculatorModel(ClassRegistryMixin, PreTrainedModel, GenerationMixin): # config_class: ClassVar[type[SpeculatorModelConfig]] = SpeculatorModelConfig # type: ignore[assignment,misc] base_model_prefix: ClassVar[str] = "model" # type: ignore[misc] main_input_name: ClassVar[str] = "input_ids" # type: ignore[misc] - _keys_to_ignore_on_load_missing: ClassVar[list[str]] = [ # type: ignore[assignment,misc] - "verifier*", - ] + _keys_to_ignore_on_load_missing: ClassVar[list[str]] = [] # type: ignore[assignment,misc] @classmethod def from_pretrained( cls: type["SpeculatorModel"], pretrained_model_name_or_path: str | os.PathLike | None, *model_args, - verifier: str | os.PathLike | PreTrainedModel | None = None, - verifier_attachment_mode: Literal["detached", "full", "train_only"] - | None = None, config: PretrainedConfig | str | os.PathLike | None = None, cache_dir: str | os.PathLike | None = None, ignore_mismatched_sizes: bool = False, @@ -127,22 +88,6 @@ def from_pretrained( config is provided as a path. :param model_args: Additional positional arguments passed to the model constructor. - :param verifier: Optional verifier model to attach the speculator to. - Can be a path to a local model directory, a Hugging Face model identifier, - or an instance of PreTrainedModel. If provided, the speculator will use this - verifier for speculative decoding. If None, the speculator will load the - verifier from the config if specified, or it must be attached later - using the `attach_verifier` method. - :param verifier_attachment_mode: Optional mode for how the verifier is - attached to the speculator. If "detached", any verifier passed in or - resolved from the config will not be ignored. - If "full", the verifier is fully integrated into the - speculator's forward pass and generation methods. - If "train_only", only the portions of the verifier needed for training - are attached, allowing for better resource utilization during training. - If None and a verifier is provided, it defaults to "full". - If a verifier is not provided and None is found in the config, - this parameter is ignored. :param config: Optional configuration for the model. Can be a SpeculatorModelConfig instance, a path to a config file, or None to load from model directory. @@ -203,8 +148,6 @@ def from_pretrained( return model_class.from_pretrained( pretrained_model_name_or_path, *model_args, - verifier=verifier, - verifier_attachment_mode=verifier_attachment_mode, config=config, cache_dir=cache_dir, ignore_mismatched_sizes=ignore_mismatched_sizes, @@ -220,8 +163,6 @@ def from_pretrained( return super().from_pretrained( # type: ignore[misc] pretrained_model_name_or_path, *model_args, - verifier=verifier, - verifier_attachment_mode=verifier_attachment_mode, config=config, cache_dir=cache_dir, ignore_mismatched_sizes=ignore_mismatched_sizes, @@ -389,42 +330,13 @@ def get_trainer_kwargs(**kwargs): "to support training infrastructure." ) - def __init__( - self, - config: SpeculatorModelConfig, - verifier: str | os.PathLike | PreTrainedModel | None, - verifier_attachment_mode: Literal["detached", "full", "train_only"] | None, - **kwargs, - ): + def __init__(self, config: SpeculatorModelConfig, **kwargs): """ Initialize a SpeculatorModel instance. - Sets up the basic structure for a speculator model, including configuration - storage and optional verifier model attachment. The verifier model is used - during speculative decoding to validate the tokens proposed by the speculator. - - If no verifier is provided during initialization, it must be attached later - using the attach_verifier method before calling generate. - :param config: The configuration for the speculator model. Must be a SpeculatorModelConfig instance containing model hyperparameters and speculative decoding settings. - :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 provided, the speculator will use this verifier for - speculative decoding. If None, the speculator will load the verifier from - the config if specified, or it must be attached later using the - `attach_verifier` method. - :param verifier_attachment_mode: Optional mode for how the verifier is - attached to the speculator. If "detach", any verifier passed in or - resolved from the config will not be attached. - If "full", the verifier is fully integrated into the - speculator's forward pass and generation methods. - If "train_only", only the portions of the verifier needed for training - are attached, allowing for better resource utilization during training. - If None and a verifier is provided, it defaults to "full". - If a verifier is not provided and None is found in the config, - this parameter is ignored. :param kwargs: Additional keyword arguments passed to the parent PreTrainedModel constructor. """ @@ -442,240 +354,9 @@ def __init__( super().__init__(config, **kwargs) self.config: SpeculatorModelConfig = 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) - - 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 attach_verifier( - self, - verifier: str | os.PathLike | PreTrainedModel, - mode: Literal["full", "train_only"] | None = None, - ): - """ - Attach a verifier model for the speculator that is used to attach to - for running inference/training with the speculator and validates the - candidate tokens generated by the speculator during the - speculative decoding process. It should be compatible - with the speculator's configuration in terms of vocabulary, architecture, - and tokenization. - - Example: - ```python - # Load and attach a verifier - verifier = AutoModel.from_pretrained("meta-llama/Llama-2-7b-hf") - speculator.attach_verifier(verifier) - - # Now ready for generation - outputs = speculator.generate(input_ids) - ``` - - :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: Optional mode for how the verifier is attached to the speculator. - If "full", the verifier is fully integrated into the speculator's forward - pass and generation methods. If "train_only", only the portions of the - verifier needed for training are attached, allowing for better resource - utilization during training. If None, defaults to "full". - :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 - - 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" - - 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 forward(self, *args, **kwargs): - """ - Defines the forward pass computation for the speculator model. - - This method must be implemented by all concrete speculator model - subclasses. It defines how the model processes inputs to generate candidate - tokens or logits specifically for training pipelines. - - Use `model.generate` for generation tasks, which will handle - speculative decoding with the attached verifier. - - :param args: Positional arguments for the forward pass, typically including - input_ids and potentially attention_mask, position_ids, etc. - :param kwargs: Keyword arguments for the forward pass, which may include - various model-specific parameters and options. - :return: Model outputs, typically including logits or candidate token - sequences, depending on the specific speculator implementation. - """ raise NotImplementedError( "The forward method is only supported on concrete " "speculator model subclasses." ) - - @torch.no_grad() - def generate( - self, - inputs: torch.Tensor | None = None, # noqa: ARG002 - generation_config: GenerationConfig | None = None, # noqa: ARG002 - logits_processor: LogitsProcessorList | None = None, # noqa: ARG002 - stopping_criteria: StoppingCriteriaList | None = None, # noqa: ARG002 - prefix_allowed_tokens_fn: Callable[[int, torch.Tensor], list[int]] # noqa: ARG002 - | None = None, - synced_gpus: bool | None = None, # noqa: ARG002 - assistant_model: Optional["PreTrainedModel"] = None, # type: ignore[override] # noqa: ARG002 - streamer: Optional["BaseStreamer"] = None, # noqa: ARG002 - negative_prompt_ids: torch.Tensor | None = None, # noqa: ARG002 - negative_prompt_attention_mask: torch.Tensor | None = None, # noqa: ARG002 - use_model_defaults: bool | None = None, # noqa: ARG002 - custom_generate: str | Callable[..., Any] | None = None, # noqa: ARG002 - **kwargs, # noqa: ARG002 - ) -> GenerateOutput | torch.LongTensor: - """ - Generate text using speculative decoding with the attached verifier model. - The method follows the standard transformers generation interface, making it - compatible with existing generation workflows while adding speculative - decoding capabilities allowing for faster generation. - - :param inputs: The input token IDs to generate from. Can be None if input_ids - are provided in kwargs. - :param generation_config: Configuration for generation parameters like - max_length, temperature, top_p, etc. If None, uses model defaults. - :param logits_processor: List of logits processors to apply during generation - for tasks like repetition penalty, top-k filtering, etc. - :param stopping_criteria: List of stopping criteria to determine when to - stop generation (e.g., max length, end-of-sequence tokens). - :param prefix_allowed_tokens_fn: Function to constrain generation to allowed - tokens based on the current prefix. Useful for structured generation. - :param synced_gpus: Whether to synchronize GPUs during distributed generation. - Relevant for multi-GPU setups. - :param assistant_model: An assistant model to use for generation. This - parameter maintains compatibility with transformers but may not be - used in speculative decoding. - :param streamer: A streamer to output tokens as they are generated, enabling - real-time streaming of the generation process. - :param negative_prompt_ids: Token IDs for negative prompting to steer - generation away from certain content. - :param negative_prompt_attention_mask: Attention mask for negative prompt - tokens to properly handle padding. - :param use_model_defaults: Whether to use model-specific default generation - parameters instead of transformers defaults. - :param kwargs: Additional keyword arguments for generation, including - input_ids, attention_mask, max_length, etc. - :return: Generated token sequences as either a GenerateOutput object - (containing additional metadata) or a LongTensor of token IDs. - """ - if self.verifier is None: - raise ValueError( - "Verifier model is not attached. Please attach a verifier model " - "before calling generate." - ) - - raise NotImplementedError( - "The generate method for speculator models is not implemented yet." - ) - - -def reload_and_populate_models(): - """ - Triggers the automatic discovery and registration of all - SpeculatorModel subclasses found in the speculators.models package - that have been registered with `SpeculatorModel.register(NAME)`. This - enables dynamic model loading and instantiation based on configuration - types without requiring explicit imports. - """ - SpeculatorModel.auto_populate_registry() diff --git a/src/speculators/models/__init__.py b/src/speculators/models/__init__.py index 670d40e69..e22068c94 100644 --- a/src/speculators/models/__init__.py +++ b/src/speculators/models/__init__.py @@ -1,13 +1,3 @@ -from .eagle import EagleSpeculator, EagleSpeculatorConfig from .eagle3 import Eagle3DraftModel, Eagle3SpeculatorConfig -from .independent import IndependentSpeculatorConfig -from .mlp import MLPSpeculatorConfig -__all__ = [ - "Eagle3DraftModel", - "Eagle3SpeculatorConfig", - "EagleSpeculator", - "EagleSpeculatorConfig", - "IndependentSpeculatorConfig", - "MLPSpeculatorConfig", -] +__all__ = ["Eagle3DraftModel", "Eagle3SpeculatorConfig"] diff --git a/src/speculators/models/eagle3/core.py b/src/speculators/models/eagle3/core.py index 2ad3a8dbd..603dac1ac 100644 --- a/src/speculators/models/eagle3/core.py +++ b/src/speculators/models/eagle3/core.py @@ -173,11 +173,7 @@ def __init__( t2d: torch.Tensor | None, d2t: torch.Tensor | None, ): - super().__init__( - config=config, - verifier=None, - verifier_attachment_mode="train_only", - ) + super().__init__(config=config) self.hidden_size = config.transformer_layer_config.hidden_size self.draft_vocab_size = config.draft_vocab_size diff --git a/src/speculators/models/independent.py b/src/speculators/models/independent.py deleted file mode 100644 index ca35b8796..000000000 --- a/src/speculators/models/independent.py +++ /dev/null @@ -1,31 +0,0 @@ -from transformers import PretrainedConfig - -from speculators import SpeculatorModelConfig, SpeculatorsConfig - -__all__ = ["IndependentSpeculatorConfig"] - - -@SpeculatorModelConfig.register("independent") -class IndependentSpeculatorConfig(SpeculatorModelConfig): - @classmethod - def from_pretrained_config( - cls, pretrained_config: PretrainedConfig, speculators_config: SpeculatorsConfig - ) -> "IndependentSpeculatorConfig": - pretrained_dict = pretrained_config.to_dict() - pretrained_dict["model_type"] = pretrained_config.model_type - - return cls(**pretrained_dict, speculators_config=speculators_config) - - speculators_model_type: str = "independent" - - def __init__(self, **kwargs): - super().__init__(**kwargs) - - # ensure we set the model_type to the one from the original config - self._model_type = kwargs.get("model_type") - - def to_dict(self): - config_dict = super().to_dict() - config_dict["model_type"] = self._model_type - del config_dict["_model_type"] - return config_dict diff --git a/src/speculators/models/mlp.py b/src/speculators/models/mlp.py deleted file mode 100644 index 5b1d8b2e9..000000000 --- a/src/speculators/models/mlp.py +++ /dev/null @@ -1,72 +0,0 @@ -from pydantic import Field - -from speculators.config import SpeculatorModelConfig - -__all__ = ["MLPSpeculatorConfig"] - - -@SpeculatorModelConfig.register("mlp") -class MLPSpeculatorConfig(SpeculatorModelConfig): - """ - TODO - """ - - architectures: list[str] = Field( - default_factory=lambda: ["MLPSpeculator"], - description=("The architectures this speculator uses."), - ) - torch_dtype: str = Field( - default="bfloat16", - description=( - "The torch dtype this speculator uses. " - "This is used to set the dtype of the model." - ), - ) - inputs: list[str] = Field( - default_factory=lambda: ["input_embeddings", "hidden_states[-1]"], - description=( - "The inputs from the verifier that this speculator uses to generate " - "proposal tokens for verification." - ), - ) - inputs_hidden_states_normalized: bool = Field( - default=False, - description=( - "Whether to use the hidden states of the verifier after the layer norm is " - "applied. If False, the hidden states are used before the " - "layer norm is applied." - ), - ) - hidden_size: int = Field( - default=4096, - description=("The hidden size from the verifier that this speculator targets."), - ) - intermediate_size: int = Field( - default=4096, - description=( - "The intermediate size the MLP speculator uses for predicting tokens from." - ), - ) - vocab_size: int = Field( - default=128256, - description=( - "The size of the vocabulary the MLP speculator supports for " - "predicting tokens from." - ), - ) - num_layers: int = Field( - default=5, - description=( - "The number of layers in the MLP speculator which ties directly to the " - "maximum number of tokens the speculator can predict." - ), - ) - tie_weights: bool = Field( - default=True, - description=( - "Whether to tie the weights across all of the MLP layers together so they " - "are shared. Reduces the overall number of parameters in the model. " - "If False, each layer will have its own set of embeddings, linear weights, " - "and head weights." - ), - ) diff --git a/src/speculators/proposals/__init__.py b/src/speculators/proposals/__init__.py index 5c0a0687e..cba9bcf35 100644 --- a/src/speculators/proposals/__init__.py +++ b/src/speculators/proposals/__init__.py @@ -1,5 +1,7 @@ +from .base import TokenProposalConfig from .greedy import GreedyTokenProposalConfig __all__ = [ "GreedyTokenProposalConfig", + "TokenProposalConfig", ] diff --git a/src/speculators/proposals/base.py b/src/speculators/proposals/base.py new file mode 100644 index 000000000..f9e0b9cd0 --- /dev/null +++ b/src/speculators/proposals/base.py @@ -0,0 +1,39 @@ +from typing import ClassVar + +from pydantic import Field + +from speculators.utils import PydanticClassRegistryMixin + +__all__ = ["TokenProposalConfig"] + + +class TokenProposalConfig(PydanticClassRegistryMixin): + """ + The base config for a token proposal method which defines how tokens are generated + by the speculator, how they are passed to the verifier, and how they are scored + for acceptance or rejection. All implementations of token proposal methods + must inherit from this class, set the proposal_type to a unique value, and + add any additional parameters needed to instantiate and implement the method. + + It uses pydantic to validate the parameters, provide default values, and + enable automatic serialization and deserialization of the correct class + types based on the proposal_type field. + """ + + @classmethod + def __pydantic_schema_base_type__(cls) -> type["TokenProposalConfig"]: + if cls.__name__ == "TokenProposalConfig": + return cls + + return TokenProposalConfig + + auto_package: ClassVar[str] = "speculators.proposals" + registry_auto_discovery: ClassVar[bool] = True + schema_discriminator: ClassVar[str] = "proposal_type" + + proposal_type: str = Field( + description=( + "The type of token proposal the config is for. " + "Must be a supported proposal type from the Speculators repo." + ), + ) diff --git a/src/speculators/proposals/greedy.py b/src/speculators/proposals/greedy.py index 6be1a75f9..e8b7a8b76 100644 --- a/src/speculators/proposals/greedy.py +++ b/src/speculators/proposals/greedy.py @@ -16,7 +16,7 @@ from pydantic import Field -from speculators.config import TokenProposalConfig +from speculators.proposals.base import TokenProposalConfig __all__ = ["GreedyTokenProposalConfig"] diff --git a/src/speculators/utils/__init__.py b/src/speculators/utils/__init__.py index ebe8d1406..b75b7aa5f 100644 --- a/src/speculators/utils/__init__.py +++ b/src/speculators/utils/__init__.py @@ -1,4 +1,3 @@ -from .auto_importer import AutoImporterMixin from .pydantic_utils import PydanticClassRegistryMixin, ReloadableBaseModel from .registry import ClassRegistryMixin diff --git a/src/speculators/utils/auto_importer.py b/src/speculators/utils/auto_importer.py deleted file mode 100644 index 5b8fccd08..000000000 --- a/src/speculators/utils/auto_importer.py +++ /dev/null @@ -1,100 +0,0 @@ -""" -Automatic module importing utilities for dynamic class discovery. - -This module provides a mixin class for automatic module importing within a package, -enabling dynamic discovery of classes and implementations without explicit imports. -It is particularly useful for auto-registering classes in a registry pattern where -subclasses need to be discoverable at runtime. - -The AutoImporterMixin can be combined with registration mechanisms to create -extensible systems where new implementations are automatically discovered and -registered when they are placed in the correct package structure. - -Classes: - - AutoImporterMixin: A mixin class that provides functionality to automatically - import all modules within a specified package or list of packa -""" - -import importlib -import pkgutil -import sys -from typing import ClassVar - -__all__ = ["AutoImporterMixin"] - - -class AutoImporterMixin: - """ - A mixin class that provides functionality to automatically import all modules - within a specified package or list of packages. - - This mixin is designed to be used with class registration mechanisms to enable - automatic discovery and registration of classes without explicit imports. When - a class inherits from AutoImporterMixin, it can define the package(s) to scan - for modules by setting the `auto_package` class variable. - - Usage Example: - ```python - from speculators.utils import AutoImporterMixin - class MyRegistry(AutoImporterMixin): - auto_package = "my_package.implementations" - - MyRegistry.auto_import_package_modules() - ``` - - :cvar auto_package: The package name or tuple of names to import modules from. - :cvar auto_ignore_modules: Optional tuple of module names to ignore during import. - :cvar auto_imported_modules: List tracking which modules have been imported. - """ - - auto_package: ClassVar[str | tuple[str, ...] | None] = None - auto_ignore_modules: ClassVar[tuple[str, ...] | None] = None - auto_imported_modules: ClassVar[list | None] = None - - @classmethod - def auto_import_package_modules(cls): - """ - Automatically imports all modules within the specified package(s). - - This method scans the package(s) defined in the `auto_package` class variable - and imports all modules found, tracking them in `auto_imported_modules`. It - skips packages (directories) and any modules listed in `auto_ignore_modules`. - - :raises ValueError: If the `auto_package` class variable is not set - """ - if cls.auto_package is None: - raise ValueError( - "The class variable 'auto_package' must be set to the package name to " - "import modules from." - ) - - cls.auto_imported_modules = [] - packages = ( - cls.auto_package - if isinstance(cls.auto_package, tuple) - else (cls.auto_package,) - ) - - for package_name in packages: - package = importlib.import_module(package_name) - - for _, module_name, is_pkg in pkgutil.walk_packages( - package.__path__, package.__name__ + "." - ): - if ( - is_pkg - or ( - cls.auto_ignore_modules is not None - and module_name in cls.auto_ignore_modules - ) - or module_name in cls.auto_imported_modules - ): - # Skip packages and ignored modules - continue - - if module_name in sys.modules: - # Avoid circular imports - cls.auto_imported_modules.append(module_name) - else: - importlib.import_module(module_name) - cls.auto_imported_modules.append(module_name) diff --git a/src/speculators/utils/pydantic_utils.py b/src/speculators/utils/pydantic_utils.py index 408a6cb6b..77876e017 100644 --- a/src/speculators/utils/pydantic_utils.py +++ b/src/speculators/utils/pydantic_utils.py @@ -189,23 +189,3 @@ def __pydantic_generate_base_schema__( :return: A CoreSchema object representing the base schema """ return core_schema.any_schema() - - @classmethod - def auto_populate_registry(cls) -> bool: - """ - Ensures that all registered classes in the registry are properly initialized. - - This method is called automatically by Pydantic when the model is instantiated - or validated. It ensures that all classes in the registry are loaded and ready - for use. - - This is particularly useful for ensuring that all subclasses are registered - before any validation occurs. - - :return: True if the registry was populated, False if it was already populated - :raises ValueError: If called when registry_auto_discovery is False - """ - populated = super().auto_populate_registry() - cls.reload_schema() - - return populated diff --git a/src/speculators/utils/registry.py b/src/speculators/utils/registry.py index e0c276de9..73ebacf3f 100644 --- a/src/speculators/utils/registry.py +++ b/src/speculators/utils/registry.py @@ -19,12 +19,10 @@ from collections.abc import Callable from typing import Any, ClassVar -from speculators.utils.auto_importer import AutoImporterMixin - __all__ = ["ClassRegistryMixin"] -class ClassRegistryMixin(AutoImporterMixin): +class ClassRegistryMixin: """ A mixin class that provides a registration system for tracking class implementations with optional auto-discovery capabilities. @@ -164,37 +162,6 @@ class ExampleClass: return clazz - @classmethod - def auto_populate_registry(cls) -> bool: - """ - Ensures that all modules in the specified auto_package are imported. - - This method is called automatically by registered_classes when - registry_auto_discovery==True to ensure that all available implementations are - discovered and registered before returning the list of registered classes. - - To enable auto-discovery: - 1. Set registry_auto_discovery = True on the class - 2. Define an auto_package class variable with the package path to import - - :return: True if the registry was populated, False if it was already populated. - :raises ValueError: If called when registry_auto_discovery is False - """ - if not cls.registry_auto_discovery: - raise ValueError( - "ClassRegistryMixin.auto_populate_registry() cannot be called " - "because registry_auto_discovery is set to False. " - "Set registry_auto_discovery to True to enable auto-discovery." - ) - - if cls.registry_populated: - return False - - cls.auto_import_package_modules() - cls.registry_populated = True - - return True - @classmethod def registered_classes(cls) -> tuple[type[Any], ...]: """ @@ -209,8 +176,6 @@ def registered_classes(cls) -> tuple[type[Any], ...]: those discovered through auto-importing when registry_auto_discovery==True. :raises ValueError: If called before any classes have been registered. """ - if cls.registry_auto_discovery: - cls.auto_populate_registry() if cls.registry is None: raise ValueError( diff --git a/tests/integration/convert/test_eagle.py b/tests/integration/convert/test_eagle.py index c23bfae87..31301f856 100644 --- a/tests/integration/convert/test_eagle.py +++ b/tests/integration/convert/test_eagle.py @@ -18,7 +18,10 @@ from loguru import logger from speculators.convert.eagle import EagleConverter -from speculators.models.eagle import EagleSpeculator, EagleSpeculatorConfig +from speculators.convert.eagle.eagle_legacy_model import ( + EagleSpeculator, + EagleSpeculatorConfig, +) class TestEagleConversion: diff --git a/tests/unit/models/test_eagle_config.py b/tests/unit/models/test_eagle_config.py index f89fd7ece..1a5b8d901 100644 --- a/tests/unit/models/test_eagle_config.py +++ b/tests/unit/models/test_eagle_config.py @@ -21,7 +21,7 @@ SpeculatorsConfig, VerifierConfig, ) -from speculators.models import EagleSpeculatorConfig +from speculators.convert.eagle.eagle_legacy_model import EagleSpeculatorConfig from speculators.proposals import GreedyTokenProposalConfig # ===== Fixtures ===== diff --git a/tests/unit/models/test_eagle_model.py b/tests/unit/models/test_eagle_model.py index 4a1f73d0f..3622b1740 100644 --- a/tests/unit/models/test_eagle_model.py +++ b/tests/unit/models/test_eagle_model.py @@ -47,7 +47,10 @@ SpeculatorsConfig, VerifierConfig, ) -from speculators.models import EagleSpeculator, EagleSpeculatorConfig +from speculators.convert.eagle.eagle_legacy_model import ( + EagleSpeculator, + EagleSpeculatorConfig, +) from speculators.proposals import GreedyTokenProposalConfig # ===== Layer Types Constants ===== diff --git a/tests/unit/proposals/test_greedy.py b/tests/unit/proposals/test_greedy.py index 40e2a4367..a024397e0 100644 --- a/tests/unit/proposals/test_greedy.py +++ b/tests/unit/proposals/test_greedy.py @@ -5,8 +5,7 @@ import pytest from pydantic import BaseModel, ValidationError -from speculators.config import TokenProposalConfig -from speculators.proposals import GreedyTokenProposalConfig +from speculators.proposals import GreedyTokenProposalConfig, TokenProposalConfig # ===== GreedyTokenProposalConfig Tests ===== diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py index a0378a73a..abf59ce47 100644 --- a/tests/unit/test_config.py +++ b/tests/unit/test_config.py @@ -17,7 +17,7 @@ SpeculatorsConfig, TokenProposalConfig, VerifierConfig, - reload_and_populate_configs, + reload_schemas, ) # ===== TokenProposalConfig Tests ===== @@ -30,7 +30,7 @@ class TokenProposalConfigTest(TokenProposalConfig): # Ensure the schemas are reloaded to include the test proposal type -reload_and_populate_configs() +reload_schemas() @pytest.mark.smoke @@ -237,7 +237,7 @@ class SpeculatorModelConfigTest(SpeculatorModelConfig): # Ensure the schemas are reloaded to include the test proposal type -reload_and_populate_configs() +reload_schemas() @pytest.fixture @@ -274,9 +274,7 @@ def test_speculator_model_config_auto_registry(): classes = SpeculatorModelConfig.registered_classes() class_names = [cls.__name__ for cls in classes] assert len(class_names) > 0 - assert "EagleSpeculatorConfig" in class_names - assert "IndependentSpeculatorConfig" in class_names - assert "MLPSpeculatorConfig" in class_names + assert "Eagle3SpeculatorConfig" in class_names @pytest.mark.smoke diff --git a/tests/unit/test_model.py b/tests/unit/test_model.py index b17a7df99..b00ab4732 100644 --- a/tests/unit/test_model.py +++ b/tests/unit/test_model.py @@ -2,23 +2,19 @@ Unit tests for the model module in the Speculators library. """ -import os import tempfile from typing import Literal -from unittest.mock import MagicMock import pytest import torch from torch import nn -from transformers import PreTrainedModel from speculators import ( SpeculatorModel, SpeculatorModelConfig, SpeculatorsConfig, VerifierConfig, - reload_and_populate_configs, - reload_and_populate_models, + reload_schemas, ) from speculators.proposals import GreedyTokenProposalConfig @@ -35,20 +31,8 @@ class SpeculatorModelTestConfig(SpeculatorModelConfig): class SpeculatorTestModel(SpeculatorModel): config_class = SpeculatorModelTestConfig # type: ignore[misc] - def __init__( - self, - config: SpeculatorModelTestConfig, - verifier: str | os.PathLike | PreTrainedModel | None = None, - verifier_attachment_mode: Literal["detached", "full", "train_only"] - | None = None, - **kwargs, - ): - super().__init__( - config, - verifier=verifier, - verifier_attachment_mode=verifier_attachment_mode, - **kwargs, - ) + def __init__(self, config: SpeculatorModelTestConfig, **kwargs): + super().__init__(config, **kwargs) self.test_module = nn.Linear(10, 10) self.post_init() # type: ignore[attr-defined] @@ -78,8 +62,7 @@ def get_trainer_kwargs(**kwargs): # Reload registries to include test classes -reload_and_populate_configs() -reload_and_populate_models() +reload_schemas() @pytest.fixture @@ -182,80 +165,18 @@ class UnregisteredConfig(SpeculatorModelConfig): @pytest.mark.smoke -def test_speculator_model_initialization_without_verifier(speculator_model_test_config): +def test_speculator_model_initialization(speculator_model_test_config): model = SpeculatorTestModel(speculator_model_test_config) assert model.config == speculator_model_test_config - assert model.verifier is None - assert model.verifier_attachment_mode == "detached" - - -@pytest.mark.smoke -def test_speculator_model_initialization_with_verifier(speculator_model_test_config): - verifier = MagicMock(spec=PreTrainedModel) - model = SpeculatorTestModel(speculator_model_test_config, verifier=verifier) - assert model.config == speculator_model_test_config - assert model.verifier == verifier - assert model.verifier_attachment_mode == "full" - - -@pytest.mark.smoke -def test_speculator_model_initialization_with_verifier_path( - speculator_model_test_config, monkeypatch -): - mock_model = MagicMock(spec=PreTrainedModel) - mock_from_pretrained = MagicMock(return_value=mock_model) - monkeypatch.setattr( - "transformers.AutoModelForCausalLM.from_pretrained", mock_from_pretrained - ) - - verifier_path = "path/to/verifier/model" - model = SpeculatorTestModel(speculator_model_test_config, verifier=verifier_path) - - mock_from_pretrained.assert_called_once_with(verifier_path) - assert model.config == speculator_model_test_config - assert model.verifier is mock_model - assert model.verifier_attachment_mode == "full" - - -@pytest.mark.smoke -def test_speculator_model_initialization_with_verifier_train_only( - speculator_model_test_config, -): - verifier = MagicMock(spec=PreTrainedModel) - model = SpeculatorTestModel( - speculator_model_test_config, - verifier=verifier, - verifier_attachment_mode="train_only", - ) - assert model.config == speculator_model_test_config - assert model.verifier is None - assert model.verifier_attachment_mode == "train_only" - - -@pytest.mark.smoke -def test_speculator_model_initialization_with_verifier_detached( - speculator_model_test_config, -): - verifier = MagicMock(spec=PreTrainedModel) - model = SpeculatorTestModel( - speculator_model_test_config, - verifier=verifier, - verifier_attachment_mode="detached", - ) - assert model.config == speculator_model_test_config - assert model.verifier is None - assert model.verifier_attachment_mode == "detached" @pytest.mark.sanity -def test_speculator_model_initialization_invalid( - speculator_model_test_config, -): +def test_speculator_model_initialization_invalid(): # No config with pytest.raises( ValueError, match="Config must be provided to initialize a SpeculatorModel" ): - SpeculatorModel(config=None, verifier=None, verifier_attachment_mode=None) # type: ignore[abstract, arg-type] + SpeculatorModel(config=None) # type: ignore[abstract, arg-type] # Invalid config type with pytest.raises( @@ -263,26 +184,6 @@ def test_speculator_model_initialization_invalid( ): SpeculatorModel( # type: ignore[abstract] config="invalid_config", # type: ignore[arg-type] - verifier=None, - verifier_attachment_mode=None, # type: ignore[arg-type] - ) - - # Invalid verifier type - with pytest.raises( - TypeError, match="Expected verifier to be a PreTrainedModel, a string path," - ): - SpeculatorModel( # type: ignore[abstract] - config=speculator_model_test_config, - verifier=123, # type: ignore[arg-type] - verifier_attachment_mode=None, - ) - - # Invalid verifier attachment mode - with pytest.raises(ValueError, match="Invalid verifier_attachment_mode: "): - SpeculatorModel( # type: ignore[abstract] - config=speculator_model_test_config, - verifier=MagicMock(spec=PreTrainedModel), - verifier_attachment_mode="invalid_mode", # type: ignore[arg-type] ) @@ -333,62 +234,14 @@ def test_speculator_model_from_pretrained_local_marshalling( @pytest.mark.smoke -def test_speculator_model_from_pretrained_verifier( +def test_speculator_model_from_pretrained( speculator_model_test_config, ): state_dict = SpeculatorTestModel(speculator_model_test_config).state_dict() # type: ignore[attr-defined] - verifier = MagicMock(spec=PreTrainedModel) model = SpeculatorModel.from_pretrained( - None, - config=speculator_model_test_config, - verifier=verifier, - state_dict=state_dict, - ) - assert isinstance(model, SpeculatorTestModel) - assert model.verifier == verifier - assert model.verifier_attachment_mode == "full" - assert isinstance(model.config, SpeculatorModelTestConfig) - assert model.config.speculators_model_type == "test_speculator_model" - assert model.config.test_param == 456 - - -@pytest.mark.smoke -def test_speculator_model_from_pretrained_verifier_train_only( - speculator_model_test_config, -): - state_dict = SpeculatorTestModel(speculator_model_test_config).state_dict() # type: ignore[attr-defined] - verifier = MagicMock(spec=PreTrainedModel) - model = SpeculatorModel.from_pretrained( - None, - config=speculator_model_test_config, - verifier=verifier, - verifier_attachment_mode="train_only", - state_dict=state_dict, - ) - assert isinstance(model, SpeculatorTestModel) - assert model.verifier is None - assert model.verifier_attachment_mode == "train_only" - assert isinstance(model.config, SpeculatorModelTestConfig) - assert model.config.speculators_model_type == "test_speculator_model" - assert model.config.test_param == 456 - - -@pytest.mark.smoke -def test_speculator_model_from_pretrained_verifier_detached( - speculator_model_test_config, -): - state_dict = SpeculatorTestModel(speculator_model_test_config).state_dict() # type: ignore[attr-defined] - verifier = MagicMock(spec=PreTrainedModel) - model = SpeculatorModel.from_pretrained( - None, - config=speculator_model_test_config, - verifier=verifier, - verifier_attachment_mode="detached", - state_dict=state_dict, + None, config=speculator_model_test_config, state_dict=state_dict ) assert isinstance(model, SpeculatorTestModel) - assert model.verifier is None - assert model.verifier_attachment_mode == "detached" assert isinstance(model.config, SpeculatorModelTestConfig) assert model.config.speculators_model_type == "test_speculator_model" assert model.config.test_param == 456 @@ -445,84 +298,3 @@ def test_speculator_model_forward_abstract(speculator_model_test_config): NotImplementedError, match="The forward method is only supported on concrete" ): model.forward() - - -# ===== SpeculatorModel Verifier Management Tests ===== - - -@pytest.mark.smoke -def test_speculator_model_attachment_lifecycle(speculator_model_test_config): - model = SpeculatorTestModel(config=speculator_model_test_config) - assert model.verifier is None - assert model.verifier_attachment_mode == "detached" - - # Attach a verifier - verifier = MagicMock(spec=PreTrainedModel) - model.attach_verifier(verifier) - assert model.verifier == verifier - assert model.verifier_attachment_mode == "full" - - # Ensure attachment before detaching raises an error - with pytest.raises( - RuntimeError, - match="Cannot attach a verifier when the speculator is not in detached mode.", - ): - model.attach_verifier(verifier) - - # Detach the verifier - model.detach_verifier() - assert model.verifier is None - assert model.verifier_attachment_mode == "detached" - - # Ensure detaching again raises an error - with pytest.raises( - RuntimeError, - match="Verifier is already detached, cannot be called again until", - ): - model.detach_verifier() - - # Attach train_only verifier - model.attach_verifier(verifier, mode="train_only") - assert model.verifier is None - assert model.verifier_attachment_mode == "train_only" - - # Detach again - model.detach_verifier() - assert model.verifier is None - assert model.verifier_attachment_mode == "detached" - - # Attach different verifier - new_verifier = MagicMock(spec=PreTrainedModel) - model.attach_verifier(new_verifier, mode="full") - assert model.verifier == new_verifier - assert model.verifier_attachment_mode == "full" - assert model.verifier != verifier - - -@pytest.mark.sanity -def test_speculator_model_attach_verifier_invalid( - speculator_model_test_config, -): - model = SpeculatorTestModel(config=speculator_model_test_config) - - # Invalid verifier type - with pytest.raises( - TypeError, match="Expected verifier to be a PreTrainedModel, a string path," - ): - model.attach_verifier(123) # type: ignore[arg-type] - - model = SpeculatorTestModel(config=speculator_model_test_config) - # Invalid attachment mode - with pytest.raises( - ValueError, match="Invalid verifier_attachment_mode: invalid_mode" - ): - model.attach_verifier(verifier=None, mode="invalid_mode") # type: ignore[arg-type] - - # Attaching when not in detached mode - model = SpeculatorTestModel(config=speculator_model_test_config) - model.verifier_attachment_mode = "full" - with pytest.raises( - RuntimeError, - match="Cannot attach a verifier when the speculator is not in detached mode.", - ): - model.attach_verifier(MagicMock(spec=PreTrainedModel)) diff --git a/tests/unit/utils/test_auto_importer.py b/tests/unit/utils/test_auto_importer.py deleted file mode 100644 index 77640a8d5..000000000 --- a/tests/unit/utils/test_auto_importer.py +++ /dev/null @@ -1,196 +0,0 @@ -""" -Unit tests for the auto_importer module in the Speculators library. -""" - -from unittest import mock - -import pytest - -from speculators.utils.auto_importer import AutoImporterMixin - -# ===== Basic Functionality Tests ===== - - -@pytest.mark.smoke -def test_auto_importer_initialization(): - class TestAutoImporterClass(AutoImporterMixin): - auto_package = "test_package.modules" - - assert AutoImporterMixin.auto_package is None - assert AutoImporterMixin.auto_ignore_modules is None - assert AutoImporterMixin.auto_imported_modules is None - - -@pytest.mark.smoke -def test_auto_importer_subclass_attributes(): - class TestAutoImporterClass(AutoImporterMixin): - auto_package = "test_package.modules" - - assert TestAutoImporterClass.auto_package == "test_package.modules" - assert TestAutoImporterClass.auto_ignore_modules is None - assert TestAutoImporterClass.auto_imported_modules is None - - -@pytest.mark.smoke -def test_no_package_raises_error(): - class TestAutoImporterClass(AutoImporterMixin): ... - - with pytest.raises(ValueError) as exc_info: - TestAutoImporterClass.auto_import_package_modules() - - assert "auto_package" in str(exc_info.value) - assert "must be set" in str(exc_info.value) - - -# ===== Module Import Tests ===== - - -@pytest.mark.smoke -def test_single_package_import(): - class TestAutoImporterClass(AutoImporterMixin): - auto_package = "test_package.modules" - - with ( - mock.patch("pkgutil.walk_packages") as mock_walk, - mock.patch("importlib.import_module") as mock_import, - ): - # Create a mock package with the necessary attributes - mock_package = mock.MagicMock() - mock_package.__path__ = ["test_package/modules"] - mock_package.__name__ = "test_package.modules" - - def import_module(name: str): - if name == "test_package.modules": - return mock_package - elif name == "test_package.modules.module1": - module = mock.MagicMock() - module.__name__ = "test_package.modules.module1" - return module - elif name == "test_package.modules.module2": - module = mock.MagicMock() - module.__name__ = "test_package.modules.module2" - return module - else: - raise ImportError(f"No module named {name}") - - def walk_packages(package_path, package_name): - if package_name == "test_package.modules.": - return [ - (None, "test_package.modules.module1", False), - (None, "test_package.modules.module2", False), - ] - else: - raise ValueError(f"Unknown package: {package_name}") - - mock_walk.side_effect = walk_packages - mock_import.side_effect = import_module - TestAutoImporterClass.auto_import_package_modules() - - mock_import.assert_any_call("test_package.modules") - assert TestAutoImporterClass.auto_imported_modules == [ - "test_package.modules.module1", - "test_package.modules.module2", - ] - - -@pytest.mark.sanity -def test_multiple_package_import(): - class TestAutoImporterClass(AutoImporterMixin): - auto_package = ("test_package.modules1", "test_package.modules2") - - with ( - mock.patch("pkgutil.walk_packages") as mock_walk, - mock.patch("importlib.import_module") as mock_import, - ): - # Create a mock package with the necessary attributes - mock_package1 = mock.MagicMock() - mock_package1.__path__ = ["test_package/modules1"] - mock_package1.__name__ = "test_package.modules1" - - mock_package2 = mock.MagicMock() - mock_package2.__path__ = ["test_package/modules2"] - mock_package2.__name__ = "test_package.modules2" - - def import_module(name: str): - if name == "test_package.modules1": - return mock_package1 - elif name == "test_package.modules2": - return mock_package2 - elif name == "test_package.modules1.moduleA": - module = mock.MagicMock() - module.__name__ = "test_package.modules1.moduleA" - return module - elif name == "test_package.modules2.moduleB": - module = mock.MagicMock() - module.__name__ = "test_package.modules2.moduleB" - return module - else: - raise ImportError(f"No module named {name}") - - def walk_packages(package_path, package_name): - if package_name == "test_package.modules1.": - return [ - (None, "test_package.modules1.moduleA", False), - ] - elif package_name == "test_package.modules2.": - return [ - (None, "test_package.modules2.moduleB", False), - ] - else: - raise ValueError(f"Unknown package: {package_name}") - - mock_walk.side_effect = walk_packages - mock_import.side_effect = import_module - TestAutoImporterClass.auto_import_package_modules() - - assert TestAutoImporterClass.auto_imported_modules == [ - "test_package.modules1.moduleA", - "test_package.modules2.moduleB", - ] - - -@pytest.mark.sanity -def test_ignore_modules(): - class TestAutoImporterClass(AutoImporterMixin): - auto_package = "test_package.modules" - auto_ignore_modules = ("test_package.modules.module1",) - - with ( - mock.patch("pkgutil.walk_packages") as mock_walk, - mock.patch("importlib.import_module") as mock_import, - ): - # Create a mock package with the necessary attributes - mock_package = mock.MagicMock() - mock_package.__path__ = ["test_package/modules"] - mock_package.__name__ = "test_package.modules" - - def import_module(name: str): - if name == "test_package.modules": - return mock_package - elif name == "test_package.modules.module1": - module = mock.MagicMock() - module.__name__ = "test_package.modules.module1" - return module - elif name == "test_package.modules.module2": - module = mock.MagicMock() - module.__name__ = "test_package.modules.module2" - return module - else: - raise ImportError(f"No module named {name}") - - def walk_packages(package_path, package_name): - if package_name == "test_package.modules.": - return [ - (None, "test_package.modules.module1", False), - (None, "test_package.modules.module2", False), - ] - else: - raise ValueError(f"Unknown package: {package_name}") - - mock_walk.side_effect = walk_packages - mock_import.side_effect = import_module - TestAutoImporterClass.auto_import_package_modules() - - assert TestAutoImporterClass.auto_imported_modules == [ - "test_package.modules.module2", - ] diff --git a/tests/unit/utils/test_registry.py b/tests/unit/utils/test_registry.py index e653c83e7..14e7f0ca6 100644 --- a/tests/unit/utils/test_registry.py +++ b/tests/unit/utils/test_registry.py @@ -2,8 +2,6 @@ Unit tests for the registry module in the Speculators library. """ -from unittest import mock - import pytest from speculators.utils.registry import ClassRegistryMixin @@ -184,127 +182,3 @@ class TestAutoRegistry(ClassRegistryMixin): assert TestAutoRegistry.registry_populated is False assert TestAutoRegistry.auto_package == "test_package.modules" assert TestAutoRegistry.registry_auto_discovery is True - - -@pytest.mark.smoke -def test_auto_populate_registry(): - class TestAutoRegistry(ClassRegistryMixin): - registry_auto_discovery = True - auto_package = "test_package.modules" - - with mock.patch.object( - TestAutoRegistry, "auto_import_package_modules" - ) as mock_import: - TestAutoRegistry.auto_populate_registry() - mock_import.assert_called_once() - assert TestAutoRegistry.registry_populated is True - - # Second call should not trigger another import since already populated - TestAutoRegistry.auto_populate_registry() - mock_import.assert_called_once() - - -@pytest.mark.sanity -def test_auto_populate_registry_disabled(): - class TestDisabledAutoRegistry(ClassRegistryMixin): - # registry_auto_discovery is False by default - auto_package = "test_package.modules" - - with pytest.raises(ValueError) as exc_info: - TestDisabledAutoRegistry.auto_populate_registry() - - assert "registry_auto_discovery is set to False" in str(exc_info.value) - - -@pytest.mark.sanity -def test_auto_registered_classes(): - class TestAutoRegistry(ClassRegistryMixin): - registry_auto_discovery = True - auto_package = "test_package.modules" - - with mock.patch.object(TestAutoRegistry, "auto_populate_registry") as mock_populate: - # Mock the registry content - TestAutoRegistry.registry = {"Class1": "class1", "Class2": "class2"} # type: ignore[dict-item] - classes = TestAutoRegistry.registered_classes() - mock_populate.assert_called_once() - assert classes == ("class1", "class2") - - -@pytest.mark.regression -def test_auto_registry_integration(): - class TestAutoRegistry(ClassRegistryMixin): - registry_auto_discovery = True - auto_package = "test_package.modules" - - with ( - mock.patch("pkgutil.walk_packages") as mock_walk, - mock.patch("importlib.import_module") as mock_import, - ): - # Create a mock package with the necessary attributes - mock_package = mock.MagicMock() - mock_package.__path__ = ["test_package/modules"] - mock_package.__name__ = "test_package.modules" - - def import_module(name: str): - if name == "test_package.modules": - return mock_package - elif name == "test_package.modules.module1": - module = mock.MagicMock() - module.__name__ = "test_package.modules.module1" - - class Module1Class: - pass - - TestAutoRegistry.register_decorator(Module1Class, "Module1Class") - return module - else: - raise ImportError(f"No module named {name}") - - def walk_packages(package_path, package_name): - if package_name == "test_package.modules.": - return [(None, "test_package.modules.module1", False)] - else: - raise ValueError(f"Unknown package: {package_name}") - - mock_walk.side_effect = walk_packages - mock_import.side_effect = import_module - - classes = TestAutoRegistry.registered_classes() - assert len(classes) == 1 - assert TestAutoRegistry.registry_populated is True - assert TestAutoRegistry.registry is not None - assert "Module1Class" in TestAutoRegistry.registry - - -@pytest.mark.regression -def test_auto_registry_with_multiple_packages(): - class TestMultiPackageRegistry(ClassRegistryMixin): - registry_auto_discovery = True - auto_package = ("package1", "package2") - - with mock.patch.object( - TestMultiPackageRegistry, "auto_import_package_modules" - ) as mock_import: - # Mock the registry to avoid ValueError when getting registered_classes - TestMultiPackageRegistry.registry = {} - TestMultiPackageRegistry.registered_classes() - mock_import.assert_called_once() - assert TestMultiPackageRegistry.registry_populated is True - - -@pytest.mark.regression -def test_auto_registry_no_package(): - class TestNoPackageRegistry(ClassRegistryMixin): - registry_auto_discovery = True - # No auto_package defined - - with mock.patch.object( - TestNoPackageRegistry, - "auto_import_package_modules", - side_effect=ValueError("auto_package must be set"), - ) as mock_import: - with pytest.raises(ValueError) as exc_info: - TestNoPackageRegistry.auto_populate_registry() - - mock_import.assert_called_once() - assert "auto_package must be set" in str(exc_info.value)