Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion model2vec/train/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)
```

Expand Down
125 changes: 99 additions & 26 deletions model2vec/train/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Comment thread
stephantul marked this conversation as resolved.

@classmethod
def from_static_model(
Expand All @@ -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.
Expand All @@ -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,
)

Expand Down Expand Up @@ -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,
}
98 changes: 95 additions & 3 deletions model2vec/train/classifier.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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:
Comment thread
stephantul marked this conversation as resolved.
"""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."""
Expand Down Expand Up @@ -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)
Loading
Loading