From ec3552e5f01d940c5290884688c92af1e75d0d56 Mon Sep 17 00:00:00 2001 From: stephantul Date: Mon, 5 Oct 2026 20:42:48 +0200 Subject: [PATCH 1/2] feat: add explicit kwargs to all trainable --- model2vec/train/README.md | 2 +- model2vec/train/base.py | 125 ++++++++++++++++++++++++++-------- model2vec/train/classifier.py | 92 ++++++++++++++++++++++++- model2vec/train/pairs.py | 89 +++++++++++++++++++++++- model2vec/train/similarity.py | 87 ++++++++++++++++++++++- tests/test_trainable.py | 97 +++++++++++++++++++++++++- 6 files changed, 453 insertions(+), 39 deletions(-) diff --git a/model2vec/train/README.md b/model2vec/train/README.md index b45839a..0dea77b 100644 --- a/model2vec/train/README.md +++ b/model2vec/train/README.md @@ -106,7 +106,7 @@ The scores are competitive with the popular [roberta-base-go_emotions](https://h ```python from model2vec.train import StaticModelForPairSimilarity -model = StaticModelForPairSimilarity.from_pretrained(model_name="minishlab/potion-base-32M") +model = StaticModelForPairSimilarity.from_pretrained(path="minishlab/potion-base-32M") model.fit(text_a=queries, text_b=documents) ``` diff --git a/model2vec/train/base.py b/model2vec/train/base.py index 4b84317..3f89d81 100644 --- a/model2vec/train/base.py +++ b/model2vec/train/base.py @@ -173,14 +173,44 @@ def from_pretrained( path: PathLike = "minishlab/potion-base-32m", *, token: str | None = None, - **kwargs: Any, + model_name: PathLike | None = None, + pad_token: str | None = None, + max_length: int | None = None, + n_layers: int = 0, + hidden_dim: int = 256, + out_dim: int = 2, + freeze: bool = False, + normalize: bool = True, + freeze_weights: bool = False, ) -> T: - """Load the model from a pretrained model2vec model.""" - if model_name := kwargs.pop("model_name", None): - logger.warning("The 'model_name' argument is deprecated. Use 'path' instead.") - path = model_name - model = StaticModel.from_pretrained(path, token=token) - return cls.from_static_model(model=model, **kwargs) + """Load the model from a pretrained model2vec model. + + :param path: The path to the folder containing the model, or a repository on the Hugging Face Hub. + :param token: The token to use to download the model from the hub. + :param model_name: Deprecated alias for `path`. + :param pad_token: The token to use for padding. If None, it is inferred from the tokenizer. + :param max_length: The default maximum sequence length to use for tokenization. If None, the + static model's `max_length` is used. + :param n_layers: The number of layers in the head. + :param hidden_dim: The hidden dimension of the head. + :param out_dim: The output dimension of the head. + :param freeze: Whether to freeze the embeddings. + :param normalize: Whether to normalize the embeddings. + :param freeze_weights: Whether to freeze the learned token weights. + :return: The initialized model. + """ + model = _load_static_model(path, token=token, model_name=model_name) + return cls.from_static_model( + model=model, + pad_token=pad_token, + max_length=max_length, + n_layers=n_layers, + hidden_dim=hidden_dim, + out_dim=out_dim, + freeze=freeze, + normalize=normalize, + freeze_weights=freeze_weights, + ) @classmethod def from_static_model( @@ -189,7 +219,12 @@ def from_static_model( model: StaticModel, pad_token: str | None = None, max_length: int | None = None, - **kwargs: Any, + n_layers: int = 0, + hidden_dim: int = 256, + out_dim: int = 2, + freeze: bool = False, + normalize: bool = True, + freeze_weights: bool = False, ) -> T: """Load the model from a static model. @@ -197,30 +232,22 @@ def from_static_model( :param pad_token: The token to use for padding. If None, it is inferred from the tokenizer. :param max_length: The default maximum sequence length to use for tokenization. If None, the static model's `max_length` is used. - :param **kwargs: Any additional keyword arguments to pass to the constructor. + :param n_layers: The number of layers in the head. + :param hidden_dim: The hidden dimension of the head. + :param out_dim: The output dimension of the head. + :param freeze: Whether to freeze the embeddings. + :param normalize: Whether to normalize the embeddings. + :param freeze_weights: Whether to freeze the learned token weights. :return: The initialized model. """ - model.embedding = np.nan_to_num(model.embedding) - weights = torch.from_numpy(model.weights) if model.weights is not None else None - embeddings_converted = torch.from_numpy(model.embedding) - if model.token_mapping is not None: - token_mapping = model.token_mapping.tolist() - else: - token_mapping = None - if pad_token is not None: - pad_id = model.tokenizer.get_vocab()[pad_token] - else: - pad_id = get_probable_pad_token_id(model.tokenizer) - if max_length is None: - max_length = model.max_length return cls( - vectors=embeddings_converted, - pad_id=pad_id, - tokenizer=model.tokenizer, - token_mapping=token_mapping, - weights=weights, - max_length=max_length, - **kwargs, + **_static_model_arguments(model, pad_token=pad_token, max_length=max_length), + n_layers=n_layers, + hidden_dim=hidden_dim, + out_dim=out_dim, + freeze=freeze, + normalize=normalize, + freeze_weights=freeze_weights, ) def _apply_token_dropout(self, keep_mask: torch.Tensor) -> torch.Tensor: @@ -516,3 +543,43 @@ def _create_datasets( T = TypeVar("T", bound=BaseFinetuneable) + + +def _load_static_model(path: PathLike, *, token: str | None, model_name: PathLike | None) -> StaticModel: + """Load a static model, resolving the deprecated `model_name` argument. + + :param path: The path to the folder containing the model, or a repository on the Hugging Face Hub. + :param token: The token to use to download the model from the hub. + :param model_name: Deprecated alias for `path`. If given, it overrides `path`. + :return: The loaded static model. + """ + if model_name is not None: + logger.warning("The 'model_name' argument is deprecated. Use 'path' instead.") + path = model_name + return StaticModel.from_pretrained(path, token=token) + + +def _static_model_arguments(model: StaticModel, *, pad_token: str | None, max_length: int | None) -> dict[str, Any]: + """Derive the constructor arguments of a finetuneable model from a static model. + + :param model: The static model to derive the arguments from. + :param pad_token: The token to use for padding. If None, it is inferred from the tokenizer. + :param max_length: The default maximum sequence length to use for tokenization. If None, the + static model's `max_length` is used. + :return: The constructor arguments. + """ + model.embedding = np.nan_to_num(model.embedding) + weights = torch.from_numpy(model.weights) if model.weights is not None else None + token_mapping = model.token_mapping.tolist() if model.token_mapping is not None else None + if pad_token is not None: + pad_id = model.tokenizer.get_vocab()[pad_token] + else: + pad_id = get_probable_pad_token_id(model.tokenizer) + return { + "vectors": torch.from_numpy(model.embedding), + "pad_id": pad_id, + "tokenizer": model.tokenizer, + "token_mapping": token_mapping, + "weights": weights, + "max_length": model.max_length if max_length is None else max_length, + } diff --git a/model2vec/train/classifier.py b/model2vec/train/classifier.py index 5ca5da7..7508c29 100644 --- a/model2vec/train/classifier.py +++ b/model2vec/train/classifier.py @@ -4,7 +4,7 @@ from collections import Counter from collections.abc import Mapping, Sequence from itertools import chain -from typing import Any, Literal, cast +from typing import Any, Literal, TypeVar, cast import numpy as np import torch @@ -14,8 +14,8 @@ from tqdm import trange from model2vec.inference import evaluate_single_or_multi_label -from model2vec.model import DEFAULT_MAX_LENGTH -from model2vec.train.base import BaseFinetuneable +from model2vec.model import DEFAULT_MAX_LENGTH, PathLike, StaticModel +from model2vec.train.base import BaseFinetuneable, _load_static_model, _static_model_arguments from model2vec.train.dataset import read_label_column from model2vec.train.utils import DEFAULT_RANDOM_SEED, seed_everything @@ -111,6 +111,89 @@ def __init__( max_length=max_length, ) + @classmethod + def from_pretrained( + cls: type[T], + path: PathLike = "minishlab/potion-base-32m", + *, + token: str | None = None, + model_name: PathLike | None = None, + pad_token: str | None = None, + max_length: int | None = None, + n_layers: int = 1, + hidden_dim: int = 512, + out_dim: int = 2, + freeze: bool = False, + normalize: bool = True, + freeze_weights: bool = False, + ) -> T: + """Load a classifier from a pretrained model2vec model. + + :param path: The path to the folder containing the model, or a repository on the Hugging Face Hub. + :param token: The token to use to download the model from the hub. + :param model_name: Deprecated alias for `path`. + :param pad_token: The token to use for padding. If None, it is inferred from the tokenizer. + :param max_length: The default maximum sequence length to use for tokenization. If None, the + static model's `max_length` is used. + :param n_layers: The number of hidden layers in the head. + :param hidden_dim: The hidden dimension of the head. + :param out_dim: The number of classes. This is reset when calling `fit`. + :param freeze: Whether to freeze the embeddings. + :param normalize: Whether to normalize the embeddings. + :param freeze_weights: Whether to freeze the learned token weights. + :return: The initialized classifier. + """ + model = _load_static_model(path, token=token, model_name=model_name) + return cls.from_static_model( + model=model, + pad_token=pad_token, + max_length=max_length, + n_layers=n_layers, + hidden_dim=hidden_dim, + out_dim=out_dim, + freeze=freeze, + normalize=normalize, + freeze_weights=freeze_weights, + ) + + @classmethod + def from_static_model( + cls: type[T], + *, + model: StaticModel, + pad_token: str | None = None, + max_length: int | None = None, + n_layers: int = 1, + hidden_dim: int = 512, + out_dim: int = 2, + freeze: bool = False, + normalize: bool = True, + freeze_weights: bool = False, + ) -> T: + """Load a classifier from a static model. + + :param model: The static model to load from. + :param pad_token: The token to use for padding. If None, it is inferred from the tokenizer. + :param max_length: The default maximum sequence length to use for tokenization. If None, the + static model's `max_length` is used. + :param n_layers: The number of hidden layers in the head. + :param hidden_dim: The hidden dimension of the head. + :param out_dim: The number of classes. This is reset when calling `fit`. + :param freeze: Whether to freeze the embeddings. + :param normalize: Whether to normalize the embeddings. + :param freeze_weights: Whether to freeze the learned token weights. + :return: The initialized classifier. + """ + return cls( + **_static_model_arguments(model, pad_token=pad_token, max_length=max_length), + n_layers=n_layers, + hidden_dim=hidden_dim, + out_dim=out_dim, + freeze=freeze, + normalize=normalize, + freeze_weights=freeze_weights, + ) + @property def classes(self) -> np.ndarray: """Return all clasess in the correct order.""" @@ -359,3 +442,6 @@ def _to_targets(self, labels: Any) -> torch.Tensor: return targets except KeyError as error: raise ValueError(f"Label {error.args[0]!r} is not one of the classes {self.classes_}.") from None + + +T = TypeVar("T", bound=StaticModelForClassification) diff --git a/model2vec/train/pairs.py b/model2vec/train/pairs.py index cec1ac5..00f1e6b 100644 --- a/model2vec/train/pairs.py +++ b/model2vec/train/pairs.py @@ -9,8 +9,8 @@ from tokenizers import Tokenizer from torch import nn -from model2vec.model import DEFAULT_MAX_LENGTH -from model2vec.train.base import BaseFinetuneable +from model2vec.model import DEFAULT_MAX_LENGTH, PathLike, StaticModel +from model2vec.train.base import BaseFinetuneable, _load_static_model, _static_model_arguments from model2vec.train.dataset import ColumnRows, PairDataset from model2vec.train.utils import DEFAULT_RANDOM_SEED, MAX_VALIDATION_SIZE, seed_everything, split_indices @@ -104,6 +104,91 @@ def __init__( max_length=max_length, ) + @classmethod + def from_pretrained( + cls: type[T], + path: PathLike = "minishlab/potion-base-32m", + *, + token: str | None = None, + model_name: PathLike | None = None, + pad_token: str | None = None, + max_length: int | None = None, + n_layers: int = 1, + hidden_dim: int = 512, + out_dim: int | None = None, + freeze: bool = False, + normalize: bool = True, + freeze_weights: bool = False, + ) -> T: + """Load the model from a pretrained model2vec model. + + :param path: The path to the folder containing the model, or a repository on the Hugging Face Hub. + :param token: The token to use to download the model from the hub. + :param model_name: Deprecated alias for `path`. + :param pad_token: The token to use for padding. If None, it is inferred from the tokenizer. + :param max_length: The default maximum sequence length to use for tokenization. If None, the + static model's `max_length` is used. + :param n_layers: The number of layers in the head. If this is 0 and `out_dim` equals the embedding + dimension, the model has no head, and the embeddings are used as is. + :param hidden_dim: The hidden dimension of the head. + :param out_dim: The output embedding dimension. If None, defaults to the input embedding dimension. + :param freeze: Whether to freeze the embeddings. + :param normalize: Whether to normalize the embeddings. + :param freeze_weights: Whether to freeze the learned token weights. + :return: The initialized model. + """ + model = _load_static_model(path, token=token, model_name=model_name) + return cls.from_static_model( + model=model, + pad_token=pad_token, + max_length=max_length, + n_layers=n_layers, + hidden_dim=hidden_dim, + out_dim=out_dim, + freeze=freeze, + normalize=normalize, + freeze_weights=freeze_weights, + ) + + @classmethod + def from_static_model( + cls: type[T], + *, + model: StaticModel, + pad_token: str | None = None, + max_length: int | None = None, + n_layers: int = 1, + hidden_dim: int = 512, + out_dim: int | None = None, + freeze: bool = False, + normalize: bool = True, + freeze_weights: bool = False, + ) -> T: + """Load the model from a static model. + + :param model: The static model to load from. + :param pad_token: The token to use for padding. If None, it is inferred from the tokenizer. + :param max_length: The default maximum sequence length to use for tokenization. If None, the + static model's `max_length` is used. + :param n_layers: The number of layers in the head. If this is 0 and `out_dim` equals the embedding + dimension, the model has no head, and the embeddings are used as is. + :param hidden_dim: The hidden dimension of the head. + :param out_dim: The output embedding dimension. If None, defaults to the input embedding dimension. + :param freeze: Whether to freeze the embeddings. + :param normalize: Whether to normalize the embeddings. + :param freeze_weights: Whether to freeze the learned token weights. + :return: The initialized model. + """ + return cls( + **_static_model_arguments(model, pad_token=pad_token, max_length=max_length), + n_layers=n_layers, + hidden_dim=hidden_dim, + out_dim=out_dim, + freeze=freeze, + normalize=normalize, + freeze_weights=freeze_weights, + ) + def forward( # type: ignore[override] self, input_ids: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: diff --git a/model2vec/train/similarity.py b/model2vec/train/similarity.py index 74ff0a0..ddf1878 100644 --- a/model2vec/train/similarity.py +++ b/model2vec/train/similarity.py @@ -10,8 +10,8 @@ from tokenizers import Tokenizer from torch import nn -from model2vec.model import DEFAULT_MAX_LENGTH -from model2vec.train.base import BaseFinetuneable +from model2vec.model import DEFAULT_MAX_LENGTH, PathLike, StaticModel +from model2vec.train.base import BaseFinetuneable, _load_static_model, _static_model_arguments from model2vec.train.dataset import get_vector_dims_from_column from model2vec.train.utils import DEFAULT_RANDOM_SEED, seed_everything @@ -97,6 +97,89 @@ def __init__( max_length=max_length, ) + @classmethod + def from_pretrained( + cls: type[T], + path: PathLike = "minishlab/potion-base-32m", + *, + token: str | None = None, + model_name: PathLike | None = None, + pad_token: str | None = None, + max_length: int | None = None, + n_layers: int = 1, + hidden_dim: int = 512, + out_dim: int = 2, + freeze: bool = False, + normalize: bool = True, + freeze_weights: bool = False, + ) -> T: + """Load the model from a pretrained model2vec model. + + :param path: The path to the folder containing the model, or a repository on the Hugging Face Hub. + :param token: The token to use to download the model from the hub. + :param model_name: Deprecated alias for `path`. + :param pad_token: The token to use for padding. If None, it is inferred from the tokenizer. + :param max_length: The default maximum sequence length to use for tokenization. If None, the + static model's `max_length` is used. + :param n_layers: The number of hidden layers in the head. + :param hidden_dim: The hidden dimension of the head. + :param out_dim: The output dimension of the head. This is reset when calling `fit`. + :param freeze: Whether to freeze the embeddings. + :param normalize: Whether to normalize the embeddings. + :param freeze_weights: Whether to freeze the learned token weights. + :return: The initialized model. + """ + model = _load_static_model(path, token=token, model_name=model_name) + return cls.from_static_model( + model=model, + pad_token=pad_token, + max_length=max_length, + n_layers=n_layers, + hidden_dim=hidden_dim, + out_dim=out_dim, + freeze=freeze, + normalize=normalize, + freeze_weights=freeze_weights, + ) + + @classmethod + def from_static_model( + cls: type[T], + *, + model: StaticModel, + pad_token: str | None = None, + max_length: int | None = None, + n_layers: int = 1, + hidden_dim: int = 512, + out_dim: int = 2, + freeze: bool = False, + normalize: bool = True, + freeze_weights: bool = False, + ) -> T: + """Load the model from a static model. + + :param model: The static model to load from. + :param pad_token: The token to use for padding. If None, it is inferred from the tokenizer. + :param max_length: The default maximum sequence length to use for tokenization. If None, the + static model's `max_length` is used. + :param n_layers: The number of hidden layers in the head. + :param hidden_dim: The hidden dimension of the head. + :param out_dim: The output dimension of the head. This is reset when calling `fit`. + :param freeze: Whether to freeze the embeddings. + :param normalize: Whether to normalize the embeddings. + :param freeze_weights: Whether to freeze the learned token weights. + :return: The initialized model. + """ + return cls( + **_static_model_arguments(model, pad_token=pad_token, max_length=max_length), + n_layers=n_layers, + hidden_dim=hidden_dim, + out_dim=out_dim, + freeze=freeze, + normalize=normalize, + freeze_weights=freeze_weights, + ) + def fit( self: T, X: Sequence[str], diff --git a/tests/test_trainable.py b/tests/test_trainable.py index 3e78d59..322349b 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -1,3 +1,4 @@ +import inspect import logging from collections import Counter, UserList from tempfile import TemporaryDirectory @@ -98,7 +99,7 @@ def test_init_base_from_model(mock_vectors: np.ndarray, mock_tokenizer: Tokenize with TemporaryDirectory() as temp_dir: model.save_pretrained(temp_dir) - s = BaseFinetuneable.from_pretrained(model_name=temp_dir) + s = BaseFinetuneable.from_pretrained(path=temp_dir) assert s.vectors.shape == mock_vectors.shape assert s.w.shape[0] == mock_vectors.shape[0] @@ -112,11 +113,103 @@ def test_init_classifier_from_model(mock_vectors: np.ndarray, mock_tokenizer: To with TemporaryDirectory() as temp_dir: model.save_pretrained(temp_dir) - s = StaticModelForClassification.from_pretrained(model_name=temp_dir) + s = StaticModelForClassification.from_pretrained(path=temp_dir) assert s.vectors.shape == mock_vectors.shape assert s.w.shape[0] == mock_vectors.shape[0] +FINETUNEABLE_CLASSES = [ + BaseFinetuneable, + StaticModelForClassification, + StaticModelForSimilarity, + StaticModelForRegression, + StaticModelForPairSimilarity, +] +NON_HEAD_PARAMETERS = { + "self", + "cls", + "path", + "token", + "model_name", + "model", + "pad_token", + "vectors", + "tokenizer", + "pad_id", + "token_mapping", + "weights", + "max_length", +} + + +@pytest.mark.parametrize("cls", FINETUNEABLE_CLASSES) +@pytest.mark.parametrize("loader_name", ["from_pretrained", "from_static_model"]) +def test_loader_signature_matches_init(cls: type[BaseFinetuneable], loader_name: str) -> None: + """Test that the loaders expose every constructor argument explicitly, with the constructor's defaults.""" + loader_parameters = inspect.signature(getattr(cls, loader_name)).parameters + assert not any(p.kind is inspect.Parameter.VAR_KEYWORD for p in loader_parameters.values()) + + init_parameters = inspect.signature(cls.__init__).parameters + expected = {name: p.default for name, p in init_parameters.items() if name not in NON_HEAD_PARAMETERS} + actual = {name: p.default for name, p in loader_parameters.items() if name not in NON_HEAD_PARAMETERS} + assert actual == expected + + +@pytest.mark.parametrize("cls", FINETUNEABLE_CLASSES) +def test_from_pretrained_explicit_arguments( + cls: type[BaseFinetuneable], mock_vectors: np.ndarray, mock_tokenizer: Tokenizer +) -> None: + """Test that from_pretrained passes each argument on to the model.""" + model = StaticModel(vectors=mock_vectors, tokenizer=mock_tokenizer) + with TemporaryDirectory() as temp_dir: + model.save_pretrained(temp_dir) + s = cls.from_pretrained( + temp_dir, + max_length=7, + n_layers=2, + hidden_dim=3, + out_dim=4, + freeze=True, + normalize=False, + freeze_weights=True, + ) + assert s.max_length == 7 + assert s.n_layers == 2 + assert s.hidden_dim == 3 + assert s.out_dim == 4 + assert s.freeze + assert not s.embeddings.weight.requires_grad + assert not s.normalize + assert s.freeze_weights + assert not s.w.requires_grad + + +@pytest.mark.parametrize("cls", FINETUNEABLE_CLASSES) +def test_from_pretrained_forwards_token_and_path( + cls: type[BaseFinetuneable], mock_vectors: np.ndarray, mock_tokenizer: Tokenizer +) -> None: + """Test that from_pretrained loads the static model from the given path with the given token.""" + model = StaticModel(vectors=mock_vectors, tokenizer=mock_tokenizer) + with patch("model2vec.train.base.StaticModel.from_pretrained", return_value=model) as mock_from_pretrained: + cls.from_pretrained("fake/repo-id", token="secret") + mock_from_pretrained.assert_called_once_with("fake/repo-id", token="secret") + + +@pytest.mark.parametrize("cls", FINETUNEABLE_CLASSES) +def test_from_pretrained_model_name_deprecated( + cls: type[BaseFinetuneable], mock_vectors: np.ndarray, mock_tokenizer: Tokenizer, caplog: pytest.LogCaptureFixture +) -> None: + """Test that the deprecated model_name argument overrides path and warns.""" + model = StaticModel(vectors=mock_vectors, tokenizer=mock_tokenizer) + with ( + patch("model2vec.train.base.StaticModel.from_pretrained", return_value=model) as mock_from_pretrained, + caplog.at_level(logging.WARNING, logger="model2vec.train.base"), + ): + cls.from_pretrained(model_name="fake/repo-id") + mock_from_pretrained.assert_called_once_with("fake/repo-id", token=None) + assert "The 'model_name' argument is deprecated" in caplog.text + + def test_init_classifier_from_model_w(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: """Test initializion from a static model.""" model = StaticModel(vectors=mock_vectors, tokenizer=mock_tokenizer, weights=np.ones(len(mock_vectors))) From 9fbef9f2320467906c4b95efc8e6381d7ed40a9a Mon Sep 17 00:00:00 2001 From: stephantul Date: Mon, 5 Oct 2026 20:55:21 +0200 Subject: [PATCH 2/2] fix greptile comments --- model2vec/train/base.py | 8 ++++++- model2vec/train/classifier.py | 6 +++++ model2vec/train/pairs.py | 8 ++++++- model2vec/train/similarity.py | 6 +++++ tests/test_trainable.py | 42 ++++++++++++++++++++++++++++++++++- 5 files changed, 67 insertions(+), 3 deletions(-) diff --git a/model2vec/train/base.py b/model2vec/train/base.py index 3f89d81..bc47e27 100644 --- a/model2vec/train/base.py +++ b/model2vec/train/base.py @@ -182,6 +182,7 @@ def from_pretrained( freeze: bool = False, normalize: bool = True, freeze_weights: bool = False, + **kwargs: Any, ) -> T: """Load the model from a pretrained model2vec model. @@ -197,6 +198,7 @@ def from_pretrained( :param freeze: Whether to freeze the embeddings. :param normalize: Whether to normalize the embeddings. :param freeze_weights: Whether to freeze the learned token weights. + :param **kwargs: Additional keyword arguments passed to the constructor. :return: The initialized model. """ model = _load_static_model(path, token=token, model_name=model_name) @@ -210,6 +212,7 @@ def from_pretrained( freeze=freeze, normalize=normalize, freeze_weights=freeze_weights, + **kwargs, ) @classmethod @@ -225,6 +228,7 @@ def from_static_model( freeze: bool = False, normalize: bool = True, freeze_weights: bool = False, + **kwargs: Any, ) -> T: """Load the model from a static model. @@ -238,6 +242,7 @@ def from_static_model( :param freeze: Whether to freeze the embeddings. :param normalize: Whether to normalize the embeddings. :param freeze_weights: Whether to freeze the learned token weights. + :param **kwargs: Additional keyword arguments passed to the constructor. :return: The initialized model. """ return cls( @@ -248,6 +253,7 @@ def from_static_model( freeze=freeze, normalize=normalize, freeze_weights=freeze_weights, + **kwargs, ) def _apply_token_dropout(self, keep_mask: torch.Tensor) -> torch.Tensor: @@ -553,7 +559,7 @@ def _load_static_model(path: PathLike, *, token: str | None, model_name: PathLik :param model_name: Deprecated alias for `path`. If given, it overrides `path`. :return: The loaded static model. """ - if model_name is not None: + if model_name: logger.warning("The 'model_name' argument is deprecated. Use 'path' instead.") path = model_name return StaticModel.from_pretrained(path, token=token) diff --git a/model2vec/train/classifier.py b/model2vec/train/classifier.py index 7508c29..3c163ae 100644 --- a/model2vec/train/classifier.py +++ b/model2vec/train/classifier.py @@ -126,6 +126,7 @@ def from_pretrained( freeze: bool = False, normalize: bool = True, freeze_weights: bool = False, + **kwargs: Any, ) -> T: """Load a classifier from a pretrained model2vec model. @@ -141,6 +142,7 @@ def from_pretrained( :param freeze: Whether to freeze the embeddings. :param normalize: Whether to normalize the embeddings. :param freeze_weights: Whether to freeze the learned token weights. + :param **kwargs: Additional keyword arguments passed to the constructor. :return: The initialized classifier. """ model = _load_static_model(path, token=token, model_name=model_name) @@ -154,6 +156,7 @@ def from_pretrained( freeze=freeze, normalize=normalize, freeze_weights=freeze_weights, + **kwargs, ) @classmethod @@ -169,6 +172,7 @@ def from_static_model( freeze: bool = False, normalize: bool = True, freeze_weights: bool = False, + **kwargs: Any, ) -> T: """Load a classifier from a static model. @@ -182,6 +186,7 @@ def from_static_model( :param freeze: Whether to freeze the embeddings. :param normalize: Whether to normalize the embeddings. :param freeze_weights: Whether to freeze the learned token weights. + :param **kwargs: Additional keyword arguments passed to the constructor. :return: The initialized classifier. """ return cls( @@ -192,6 +197,7 @@ def from_static_model( freeze=freeze, normalize=normalize, freeze_weights=freeze_weights, + **kwargs, ) @property diff --git a/model2vec/train/pairs.py b/model2vec/train/pairs.py index 00f1e6b..a9e792c 100644 --- a/model2vec/train/pairs.py +++ b/model2vec/train/pairs.py @@ -2,7 +2,7 @@ import logging from collections.abc import Sequence -from typing import TypeVar +from typing import Any, TypeVar import numpy as np import torch @@ -119,6 +119,7 @@ def from_pretrained( freeze: bool = False, normalize: bool = True, freeze_weights: bool = False, + **kwargs: Any, ) -> T: """Load the model from a pretrained model2vec model. @@ -135,6 +136,7 @@ def from_pretrained( :param freeze: Whether to freeze the embeddings. :param normalize: Whether to normalize the embeddings. :param freeze_weights: Whether to freeze the learned token weights. + :param **kwargs: Additional keyword arguments passed to the constructor. :return: The initialized model. """ model = _load_static_model(path, token=token, model_name=model_name) @@ -148,6 +150,7 @@ def from_pretrained( freeze=freeze, normalize=normalize, freeze_weights=freeze_weights, + **kwargs, ) @classmethod @@ -163,6 +166,7 @@ def from_static_model( freeze: bool = False, normalize: bool = True, freeze_weights: bool = False, + **kwargs: Any, ) -> T: """Load the model from a static model. @@ -177,6 +181,7 @@ def from_static_model( :param freeze: Whether to freeze the embeddings. :param normalize: Whether to normalize the embeddings. :param freeze_weights: Whether to freeze the learned token weights. + :param **kwargs: Additional keyword arguments passed to the constructor. :return: The initialized model. """ return cls( @@ -187,6 +192,7 @@ def from_static_model( freeze=freeze, normalize=normalize, freeze_weights=freeze_weights, + **kwargs, ) def forward( # type: ignore[override] diff --git a/model2vec/train/similarity.py b/model2vec/train/similarity.py index ddf1878..4928b07 100644 --- a/model2vec/train/similarity.py +++ b/model2vec/train/similarity.py @@ -112,6 +112,7 @@ def from_pretrained( freeze: bool = False, normalize: bool = True, freeze_weights: bool = False, + **kwargs: Any, ) -> T: """Load the model from a pretrained model2vec model. @@ -127,6 +128,7 @@ def from_pretrained( :param freeze: Whether to freeze the embeddings. :param normalize: Whether to normalize the embeddings. :param freeze_weights: Whether to freeze the learned token weights. + :param **kwargs: Additional keyword arguments passed to the constructor. :return: The initialized model. """ model = _load_static_model(path, token=token, model_name=model_name) @@ -140,6 +142,7 @@ def from_pretrained( freeze=freeze, normalize=normalize, freeze_weights=freeze_weights, + **kwargs, ) @classmethod @@ -155,6 +158,7 @@ def from_static_model( freeze: bool = False, normalize: bool = True, freeze_weights: bool = False, + **kwargs: Any, ) -> T: """Load the model from a static model. @@ -168,6 +172,7 @@ def from_static_model( :param freeze: Whether to freeze the embeddings. :param normalize: Whether to normalize the embeddings. :param freeze_weights: Whether to freeze the learned token weights. + :param **kwargs: Additional keyword arguments passed to the constructor. :return: The initialized model. """ return cls( @@ -178,6 +183,7 @@ def from_static_model( freeze=freeze, normalize=normalize, freeze_weights=freeze_weights, + **kwargs, ) def fit( diff --git a/tests/test_trainable.py b/tests/test_trainable.py index 322349b..ddd99a6 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -139,6 +139,7 @@ def test_init_classifier_from_model(mock_vectors: np.ndarray, mock_tokenizer: To "token_mapping", "weights", "max_length", + "kwargs", } @@ -147,7 +148,6 @@ def test_init_classifier_from_model(mock_vectors: np.ndarray, mock_tokenizer: To def test_loader_signature_matches_init(cls: type[BaseFinetuneable], loader_name: str) -> None: """Test that the loaders expose every constructor argument explicitly, with the constructor's defaults.""" loader_parameters = inspect.signature(getattr(cls, loader_name)).parameters - assert not any(p.kind is inspect.Parameter.VAR_KEYWORD for p in loader_parameters.values()) init_parameters = inspect.signature(cls.__init__).parameters expected = {name: p.default for name, p in init_parameters.items() if name not in NON_HEAD_PARAMETERS} @@ -210,6 +210,46 @@ def test_from_pretrained_model_name_deprecated( assert "The 'model_name' argument is deprecated" in caplog.text +class _CustomClassifier(StaticModelForClassification): + def __init__(self, *, extra: str = "default", **kwargs: Any) -> None: + """Initialize a classifier with an extra constructor argument.""" + self.extra = extra + super().__init__(**kwargs) + + +@pytest.mark.parametrize("loader_name", ["from_pretrained", "from_static_model"]) +def test_loader_forwards_extra_arguments_to_constructor( + loader_name: str, mock_vectors: np.ndarray, mock_tokenizer: Tokenizer +) -> None: + """Test that the loaders pass arguments they do not know on to a subclass constructor.""" + model = StaticModel(vectors=mock_vectors, tokenizer=mock_tokenizer) + with patch("model2vec.train.base.StaticModel.from_pretrained", return_value=model): + if loader_name == "from_pretrained": + s = _CustomClassifier.from_pretrained("fake/repo-id", extra="custom", hidden_dim=3) + else: + s = _CustomClassifier.from_static_model(model=model, extra="custom", hidden_dim=3) + assert s.extra == "custom" + assert s.hidden_dim == 3 + + +@pytest.mark.parametrize("cls", FINETUNEABLE_CLASSES) +def test_loader_rejects_unknown_arguments( + cls: type[BaseFinetuneable], mock_vectors: np.ndarray, mock_tokenizer: Tokenizer +) -> None: + """Test that an argument the constructor does not know is rejected.""" + model = StaticModel(vectors=mock_vectors, tokenizer=mock_tokenizer) + with pytest.raises(TypeError, match="n_hidden"): + cls.from_static_model(model=model, n_hidden=3) + + +def test_from_pretrained_empty_model_name_ignored(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: + """Test that an empty model_name does not override path.""" + model = StaticModel(vectors=mock_vectors, tokenizer=mock_tokenizer) + with patch("model2vec.train.base.StaticModel.from_pretrained", return_value=model) as mock_from_pretrained: + StaticModelForClassification.from_pretrained(model_name="") + mock_from_pretrained.assert_called_once_with("minishlab/potion-base-32m", token=None) + + def test_init_classifier_from_model_w(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: """Test initializion from a static model.""" model = StaticModel(vectors=mock_vectors, tokenizer=mock_tokenizer, weights=np.ones(len(mock_vectors)))