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..bc47e27 100644 --- a/model2vec/train/base.py +++ b/model2vec/train/base.py @@ -173,14 +173,47 @@ def from_pretrained( 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 = 0, + hidden_dim: int = 256, + out_dim: int = 2, + freeze: bool = False, + normalize: bool = True, + freeze_weights: bool = False, **kwargs: Any, ) -> 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. + :param **kwargs: Additional keyword arguments passed to the constructor. + :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, + **kwargs, + ) @classmethod def from_static_model( @@ -189,6 +222,12 @@ def from_static_model( model: StaticModel, 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, **kwargs: Any, ) -> T: """Load the model from a static model. @@ -197,29 +236,23 @@ 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. + :param **kwargs: Additional keyword arguments passed to the constructor. :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, + **_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, **kwargs, ) @@ -516,3 +549,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: + 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..3c163ae 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,95 @@ 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, + **kwargs: Any, + ) -> 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. + :param **kwargs: Additional keyword arguments passed to the constructor. + :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, + **kwargs, + ) + + @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, + **kwargs: Any, + ) -> 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. + :param **kwargs: Additional keyword arguments passed to the constructor. + :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, + **kwargs, + ) + @property def classes(self) -> np.ndarray: """Return all clasess in the correct order.""" @@ -359,3 +448,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..a9e792c 100644 --- a/model2vec/train/pairs.py +++ b/model2vec/train/pairs.py @@ -2,15 +2,15 @@ import logging from collections.abc import Sequence -from typing import TypeVar +from typing import Any, TypeVar import numpy as np import torch 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,97 @@ 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, + **kwargs: Any, + ) -> 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. + :param **kwargs: Additional keyword arguments passed to the constructor. + :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, + **kwargs, + ) + + @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, + **kwargs: Any, + ) -> 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. + :param **kwargs: Additional keyword arguments passed to the constructor. + :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, + **kwargs, + ) + 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..4928b07 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,95 @@ 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, + **kwargs: Any, + ) -> 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. + :param **kwargs: Additional keyword arguments passed to the constructor. + :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, + **kwargs, + ) + + @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, + **kwargs: Any, + ) -> 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. + :param **kwargs: Additional keyword arguments passed to the constructor. + :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, + **kwargs, + ) + def fit( self: T, X: Sequence[str], diff --git a/tests/test_trainable.py b/tests/test_trainable.py index 3e78d59..ddd99a6 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,143 @@ 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", + "kwargs", +} + + +@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 + + 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 + + +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)))