""" Base model classes for the Speculators library. This module contains the base model classes for speculative decoding implementations in the Speculators library. """ import os from abc import abstractmethod from typing import ClassVar import torch from torch import nn from transformers import PretrainedConfig, PreTrainedModel from speculators.config import SpeculatorModelConfig from speculators.utils import ClassRegistryMixin class DraftVocabMixin(nn.Module): """ Mixin for speculator models that use draft vocabulary mapping. Initializes vocab mapping buffers, token embeddings, and LM heads for models that implement draft-to-target vocabulary speculation. Requires the config to have ``transformer_layer_config`` and ``draft_vocab_size`` fields. """ t2d: torch.Tensor | None d2t: torch.Tensor | None embed_tokens: nn.Embedding lm_head: nn.Linear verifier_lm_head: nn.Linear def _init_vocab(self, config): """Initialize vocab mappings, token embeddings, and LM heads. Must be called after ``super().__init__(config=config)`` in the concrete model's ``__init__``. """ # VOCAB MAPPINGS tl_config = config.transformer_layer_config self.draft_vocab_size = config.draft_vocab_size self.verifier_vocab_size = tl_config.vocab_size self.hidden_size = tl_config.hidden_size self.use_draft_vocab = self.draft_vocab_size != self.verifier_vocab_size t2d: torch.Tensor | None = None d2t: torch.Tensor | None = None if self.use_draft_vocab: t2d = torch.zeros((self.verifier_vocab_size,), dtype=torch.bool) d2t = torch.zeros((self.draft_vocab_size,), dtype=torch.long) self.register_buffer("t2d", t2d) self.register_buffer("d2t", d2t) # TOKEN EMBEDDINGS self.embed_tokens = nn.Embedding( self.verifier_vocab_size, self.hidden_size, padding_idx=getattr(tl_config, "pad_token_id", None), ) self.embed_tokens.weight.requires_grad_(False) # LM HEADS self.lm_head = nn.Linear(self.hidden_size, self.draft_vocab_size, bias=False) self.verifier_lm_head = nn.Linear( self.hidden_size, self.draft_vocab_size, bias=False ) self.verifier_lm_head.weight.requires_grad = False self.lm_head.weight.requires_grad = False # Initialize weights to nan so it's easy to detect if they're never loaded torch.nn.init.constant_(self.lm_head.weight, torch.nan) torch.nn.init.constant_(self.embed_tokens.weight, torch.nan) torch.nn.init.constant_(self.verifier_lm_head.weight, torch.nan) self.lm_head._is_hf_initialized = True # type: ignore[assignment] # noqa: SLF001 self.embed_tokens._is_hf_initialized = True # type: ignore[assignment] # noqa: SLF001 self.verifier_lm_head._is_hf_initialized = True # type: ignore[assignment] # noqa: SLF001 def load_vocab_mappings(self, t2d: torch.Tensor | None, d2t: torch.Tensor | None): """Load target-to-draft and draft-to-target vocabulary mapping tensors. Args: t2d: Target-to-draft vocabulary mapping tensor. d2t: Draft-to-target vocabulary mapping tensor. """ if t2d is None and d2t is None: return elif t2d is None or d2t is None: raise ValueError( "Both t2d and d2t must be provided together, or both must be None. " f"Got t2d={'provided' if t2d is not None else 'None'}, " f"d2t={'provided' if d2t is not None else 'None'}" ) if not self.use_draft_vocab: # draft_vocab_size == verifier_vocab_size, so the mappings would be # identity no-ops; accept and ignore them rather than erroring. Real # data directories may carry full-vocab t2d/d2t files alongside a # checkpoint that does not prune the vocabulary. return if t2d.shape[0] != self.verifier_vocab_size: raise ValueError( f"t2d.shape[0] ({t2d.shape[0]}) must match" f" verifier_vocab_size ({self.verifier_vocab_size})." ) if int(t2d.sum(dtype=torch.long).item()) != self.draft_vocab_size: raise ValueError( f"t2d has {int(t2d.sum(dtype=torch.long).item())} non-zero values, " f"expected {self.draft_vocab_size}." ) if d2t.shape[0] != self.draft_vocab_size: raise ValueError( f"d2t.shape[0] ({d2t.shape[0]}) must match" f" draft_vocab_size ({self.draft_vocab_size})." ) self.load_state_dict({"t2d": t2d, "d2t": d2t}, strict=False) def load_verifier_weights(self) -> None: """Load verifier-owned weights without replacing checkpoint values.""" self._load_verifier_weights() def _load_verifier_weights( # noqa: C901 self, *, overwrite_embed_tokens: bool = False, overwrite_lm_head: bool = False, ) -> None: """Load verifier model weights (embeddings, lm_head, etc.). Loads embed_tokens, lm_head, and verifier_lm_head weights from the verifier model. Handles draft vocab masking via t2d when use_draft_vocab is True. Subclasses can override to load additional weights (e.g. norms, tokenizer) by calling super().load_verifier_weights() first. Checkpoint weights take precedence unless an internal caller explicitly requests overwrite. """ import warnings # noqa: PLC0415 from speculators.utils.loading import load_model_layers # noqa: PLC0415 speculators_config = getattr( getattr(self, "config", None), "speculators_config", None ) if speculators_config is None: return verifier_config = speculators_config.verifier if verifier_config.name_or_path is None: return # Determine which weights to load based on model attributes weights_to_load = ["embed_tokens.weight", "lm_head.weight"] if hasattr(self, "verifier_norm"): weights_to_load.append("model.norm.weight") verifier_weights = load_model_layers( weights_to_load, verifier_config.name_or_path, ) embed_tokens_weight = verifier_weights["embed_tokens.weight"] lm_head_weight = verifier_weights.get("lm_head.weight", embed_tokens_weight) # Load embed_tokens if not already loaded (NaN means uninitialized) if overwrite_embed_tokens or self.embed_tokens.weight.isnan().any(): self.embed_tokens.load_state_dict({"weight": embed_tokens_weight}) if self.use_draft_vocab: if self.t2d is None or not torch.any(self.t2d).item(): # type: ignore[arg-type] raise ValueError( "t2d tensor hasn't been set. Please call " "`.load_vocab_mappings(t2d, d2t)` before `.load_verifier_weights()`" ) lm_head_weight = lm_head_weight[ self.t2d.to(device=lm_head_weight.device, dtype=torch.bool), : # type: ignore[union-attr,index] ] if overwrite_lm_head or self.lm_head.weight.isnan().any(): self.lm_head.load_state_dict( {"weight": lm_head_weight.detach().clone()}, strict=False ) self.verifier_lm_head.load_state_dict( {"weight": lm_head_weight.detach().clone()}, strict=False ) # Load verifier norm weights if the model has verifier_norm if hasattr(self, "verifier_norm"): if "model.norm.weight" not in verifier_weights: warnings.warn( f"Could not find final norm weights in " f"{verifier_config.name_or_path}. " "Using default initialization (weight=1.0).", UserWarning, stacklevel=2, ) else: verifier_norm_sd = {"weight": verifier_weights["model.norm.weight"]} self.verifier_norm.load_state_dict(verifier_norm_sd) # type: ignore[union-attr] # HF's from_pretrained resets requires_grad=True on all parameters. # Re-freeze verifier weights that should never be trained. self.embed_tokens.weight.requires_grad_(False) self.lm_head.weight.requires_grad_(False) self.verifier_lm_head.weight.requires_grad_(False) if hasattr(self, "verifier_norm"): self.verifier_norm.weight.requires_grad_(False) # type: ignore[union-attr] class SpeculatorModel(ClassRegistryMixin, PreTrainedModel): # type: ignore[misc] """ Abstract base class for all speculator models in the Speculators library. This class provides the foundation for implementing speculative decoding models that can generate candidate tokens to be verified by a base verifier model. It combines the functionality of Hugging Face's PreTrainedModel and GenerationMixin with automatic model registration and discovery capabilities. All concrete speculator model implementations must inherit from this class, register with `SpeculatorModel.register(NAME)`, and implement the abstract forward method. Example: ```python # Load a speculator model with automatic class resolution model = SpeculatorModel.from_pretrained("path/to/speculator") ``` """ # Registry configuration auto_package: ClassVar[str] = "speculators.models" registry_auto_discovery: ClassVar[bool] = True # PreTrainedModel settings 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] @classmethod def from_pretrained( cls: type["SpeculatorModel"], pretrained_model_name_or_path: str | os.PathLike | None, *model_args, config: PretrainedConfig | str | os.PathLike | None = None, cache_dir: str | os.PathLike | None = None, ignore_mismatched_sizes: bool = False, force_download: bool = False, local_files_only: bool = False, token: str | bool | None = None, revision: str = "main", use_safetensors: bool | None = None, weights_only: bool = True, t2d: torch.Tensor | None = None, d2t: torch.Tensor | None = None, verifier: str | None = None, **kwargs, ) -> "SpeculatorModel": """ Load a pretrained speculator model from the Hugging Face Hub or local directory. This method automatically resolves the correct speculator model class based on the configuration type and loads the model with the appropriate weights. If called on the base SpeculatorModel class, it will automatically determine and instantiate the correct subclass based on the model configuration. Example: ```python # Load with automatic class resolution model = SpeculatorModel.from_pretrained("RedHatAI/speculator-llama-7b") # Load from local directory model = SpeculatorModel.from_pretrained("./my_speculator") # Load with custom config config = SpeculatorModelConfig.from_pretrained("RedHatAI/eagle-llama-7b") model = SpeculatorModel.from_pretrained( None, config=config, state_dict=state_dict ) ``` :param pretrained_model_name_or_path: The model identifier on Hugging Face Hub, or path to a local directory containing the model files. Can be None if config is provided as a path. :param model_args: Additional positional arguments passed to the model constructor. :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. :param cache_dir: Directory to cache downloaded files. If None, uses default transformers cache directory. :param ignore_mismatched_sizes: Whether to ignore size mismatches when loading pretrained weights. Useful for loading models with different architectures. :param force_download: Whether to force re-download of model files even if they exist in cache. :param local_files_only: Whether to avoid downloading files and only use local cached files. Raises an error if files are not found locally. :param token: Optional authentication token for accessing private models on Hugging Face Hub. Can be a string token or True to use saved token. :param revision: The specific model revision to load (branch name, tag, or commit hash). Defaults to "main". :param use_safetensors: Whether to use safetensors format for loading weights. If None, automatically detects the available format. :param weights_only: Whether to only load model weights without optimizer states or other training artifacts. :param verifier: Verifier model id/path used to auto-convert an external (non-speculators) checkpoint; ignored for speculators checkpoints. :param kwargs: Additional keyword arguments passed to the model constructor and loading process. :return: A SpeculatorModel instance of the appropriate subclass, loaded with the pretrained weights and configuration. """ if not config: if not pretrained_model_name_or_path: raise ValueError( "Either `config` or `pretrained_model_name_or_path` must be " "provided to load a SpeculatorModel." ) # Auto-convert external (non-speculators) checkpoints so one # `from_pretrained` pathway finetunes both formats. Detect format # once here and only invoke the converter when needed. config_dict, _ = PretrainedConfig.get_config_dict( pretrained_model_name_or_path, cache_dir=cache_dir ) if "speculators_model_type" not in config_dict: from speculators.convert.entrypoints import ( # noqa: PLC0415 maybe_convert_external_checkpoint, ) pretrained_model_name_or_path = maybe_convert_external_checkpoint( pretrained_model_name_or_path, verifier=verifier, cache_dir=cache_dir, config_dict=config_dict, ) config = cls.config_class.from_pretrained( pretrained_model_name_or_path, cache_dir=cache_dir, force_download=force_download, local_files_only=local_files_only, token=token, revision=revision, ) if not isinstance(config, SpeculatorModelConfig): raise TypeError( f"Expected config to be an instance of SpeculatorModelConfig, " f"got {type(config)}." ) if not pretrained_model_name_or_path and not kwargs.get("state_dict"): raise ValueError( "Either `pretrained_model_name_or_path` or `state_dict` must be " "provided to load a SpeculatorModel." ) if cls is SpeculatorModel: # generic call to from_pretrained on this class, need to resolve the # specific model class to use for loading based on the config and registry model_class = cls.registered_model_class_from_config(config) return model_class.from_pretrained( pretrained_model_name_or_path, *model_args, config=config, cache_dir=cache_dir, ignore_mismatched_sizes=ignore_mismatched_sizes, force_download=force_download, local_files_only=local_files_only, token=token, revision=revision, use_safetensors=use_safetensors, weights_only=weights_only, t2d=t2d, d2t=d2t, verifier=verifier, **kwargs, ) model: SpeculatorModel = super().from_pretrained( # type: ignore[misc] pretrained_model_name_or_path, *model_args, config=config, cache_dir=cache_dir, ignore_mismatched_sizes=ignore_mismatched_sizes, force_download=force_download, local_files_only=local_files_only, token=token, revision=revision, use_safetensors=use_safetensors, weights_only=weights_only, **kwargs, ) if hasattr(model, "load_vocab_mappings"): model.load_vocab_mappings(t2d, d2t) # type: ignore[operator,attr-defined] if hasattr(model, "load_verifier_weights"): model.load_verifier_weights() # type: ignore[operator,attr-defined] return model @classmethod def registered_model_class_from_config( cls, config: SpeculatorModelConfig ) -> type["SpeculatorModel"]: """ Looks up the appropriate speculator model class from the registry based on the configuration type. It matches the config class to the corresponding model class that was registered during auto-discovery or manual registration. :param config: The configuration for which to find the registered model class. Must be an instance of a SpeculatorModelConfig subclass. :return: The registered model class that matches the configuration type. """ if not isinstance(config, SpeculatorModelConfig): raise TypeError( f"Expected config to be an instance of SpeculatorModelConfig, " f"got {type(config)} {config}." ) if type(config) is SpeculatorModelConfig: raise TypeError( "Received a SpeculatorModelConfig instance but expected a subclass. " "Use the specific subclass of SpeculatorModelConfig instead. " f"Received: {type(config)} {config}" ) if not cls.registry: raise ValueError( "No registered model classes found. " "Ensure that models are registered with " "`SpeculatorModel.register(NAME)` or that auto-discovery is enabled." ) for _, model_class in cls.registry.items(): model_config_class: type[SpeculatorModelConfig] = model_class.config_class if type(config) is model_config_class: return model_class raise ValueError( f"No registered model class found for config type {type(config)}. " f"Available registered model classes: {list(cls.registry.keys())}." ) @classmethod def verify_training_compatible(cls, model: "SpeculatorModel") -> None: """Verify that a model instance is compatible with training infrastructure. This method validates that the given model is: 1. An instance of SpeculatorModel 2. Registered in the SpeculatorModel registry 3. Has a `layers` attribute (required for FSDP wrapping) Args: model: The model instance to verify Raises: TypeError: If model is not a SpeculatorModel instance ValueError: If model's class is not in the registry AttributeError: If model doesn't have a `layers` attribute """ if not isinstance(model, SpeculatorModel): raise TypeError( f"Model must be a SpeculatorModel, got {type(model).__name__}" ) model_class = type(model) registry = cls.registry if registry is None or model_class not in registry.values(): raise ValueError( f"Model {model_class.__name__} is not registered in " f"SpeculatorModel.registry. " f"Available models: {list(registry.keys()) if registry else []}" ) if not hasattr(model, "layers"): raise AttributeError( f"Model {model_class.__name__} must have a 'layers' attribute " f"containing decoder layers for FSDP wrapping" ) @classmethod @abstractmethod def from_training_args( cls, verifier_config: PretrainedConfig, **kwargs ) -> "SpeculatorModel": """Create model instance from training arguments. This factory method is used by the training script to instantiate models from command-line arguments. Each algorithm must implement this to support the training infrastructure. Args: verifier_config: Configuration from the verifier/base model. **kwargs: Training arguments as keyword arguments. Each algorithm extracts the parameters it needs. Returns: Initialized model instance ready for training. Example: ```python @classmethod def from_training_args(cls, verifier_config, **kwargs): config = MySpeculatorConfig( transformer_layer_config=verifier_config, num_layers=kwargs['num_layers'], ... ) return cls(config=config, t2d=kwargs.get('t2d'), d2t=kwargs.get('d2t')) ``` """ raise NotImplementedError( f"{cls.__name__} must implement from_training_args() classmethod " "to support training infrastructure." ) @staticmethod @abstractmethod def get_trainer_kwargs(**kwargs) -> tuple[dict, dict]: """Get algorithm-specific kwargs for training and validation. This method extracts algorithm-specific parameters from the training arguments and returns separate kwargs dictionaries for training and validation forward passes. Args: **kwargs: Training arguments containing algorithm-specific parameters. Returns: Tuple of (train_kwargs, val_kwargs) where: - train_kwargs: Dict passed to model.forward() during training - val_kwargs: Dict passed to model.forward() during validation Example: ```python @staticmethod def get_trainer_kwargs(**kwargs): train_kwargs = { "num_steps": kwargs["num_steps"], "use_special_mode": True, } val_kwargs = { "num_steps": kwargs["num_steps"], "use_special_mode": False, } return train_kwargs, val_kwargs ``` """ raise NotImplementedError( "Model must implement get_trainer_kwargs() staticmethod " "to support training infrastructure." ) def on_training_step( self, global_step: int, total_steps: int | None = None, ) -> None: """Hook called by the trainer before each training forward pass. The per-forward kwargs from :meth:`get_trainer_kwargs` are fixed for the whole run, so objectives whose weighting depends on training progress (Domino's decaying base-loss weight, for example) read it here instead. The trainer restores ``global_step`` when resuming, so a schedule driven from this hook resumes at the right point. Implementations should write into a tensor buffer rather than a Python attribute: a changing Python scalar read inside a ``torch.compile``d forward is specialized on by value and forces a recompile every step. Args: global_step: Optimizer steps completed so far in the run. total_steps: The run's step horizon, or None when it is unknown. """ def __init__(self, config: SpeculatorModelConfig, **kwargs): """ Initialize a SpeculatorModel instance. :param config: The configuration for the speculator model. Must be a SpeculatorModelConfig instance containing model hyperparameters and speculative decoding settings. :param kwargs: Additional keyword arguments passed to the parent PreTrainedModel constructor. """ if not config: raise ValueError( "Config must be provided to initialize a SpeculatorModel. " "Use SpeculatorModelConfig to create a valid configuration." ) if not isinstance(config, SpeculatorModelConfig): raise TypeError( f"Expected config to be an instance of SpeculatorModelConfig, " f"got {type(config)} {config}." ) config.tie_word_embeddings = False super().__init__(config, **kwargs) self.config: SpeculatorModelConfig = config def forward(self, *args, **kwargs): raise NotImplementedError( "The forward method is only supported on concrete " "speculator model subclasses." )