Download source/src/speculators/model.py from khazic/spec-b300: direct link, hf CLI and curl.
- Browser
- Download file 25.8 kB
-
https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/model.py
- Command line
-
hf download hf://khazic/spec-b300/source/src/speculators/model.py
-
curl -L -o model.py https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/model.py
25.8 kB
| """ | |
| 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] | |
| 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 | |
| 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())}." | |
| ) | |
| 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" | |
| ) | |
| 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." | |
| ) | |
| 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." | |
| ) | |