Skip to content
Open
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 .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -38,4 +38,4 @@ repos:
language: system
types: [python]
pass_filenames: false
args: [src/]
args: [model2vec/]
5 changes: 3 additions & 2 deletions model2vec/inference/evaluation.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,17 +56,18 @@ def _precision_recall_f1_support(

def evaluate_single_or_multi_label(
predictions: np.ndarray,
y: list[int] | list[str] | list[list[int]] | list[list[str]],
y: Sequence[Any],
) -> dict[str, dict[str, float]]:
"""Evaluate the classifier on a given dataset using a classification report.

This function computes per-class precision, recall and f1-score (via one-vs-rest / multi-hot encoding), plus
overall accuracy, macro average, and weighted average.

:param predictions: The predictions.
:param y: The ground truth labels.
:param y: The ground truth labels as a sequence.
:return: A classification report, as a dictionary.
"""
y = [label.tolist() if hasattr(label, "tolist") else label for label in y]
if _is_multi_label_shaped(y):
y = cast(list[list[str]] | list[list[int]], y)
predictions = cast(np.ndarray, predictions)
Expand Down
2 changes: 1 addition & 1 deletion model2vec/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -459,7 +459,7 @@ def _encode_batch(self, sentences: Sequence[str], normalize: bool) -> np.ndarray
ids = self.tokenize(sentences=sentences)
dtype = self.embedding.dtype
if dtype == np.int8:
dtype = np.float32
dtype = np.dtype(np.float32)
out = np.zeros((len(ids), self.dim), dtype=dtype)

if self.token_mapping is None and self.weights is None:
Expand Down
37 changes: 28 additions & 9 deletions model2vec/train/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ distilled_model = distill("baai/bge-base-en-v1.5")
classifier = StaticModelForClassification.from_static_model(model=distilled_model)

# From a pre-trained model: potion is the default
classifier = StaticModelForClassification.from_pretrained(model_name="minishlab/potion-base-32m")
classifier = StaticModelForClassification.from_pretrained(path="minishlab/potion-base-32m")
```

This creates a very simple classifier: a StaticModel with a single 512-unit hidden layer on top. You can adjust the number of hidden layers and the number units through some parameters on both functions. Note that the default for `from_pretrained` is [potion-base-32m](https://huggingface.co/minishlab/potion-base-32M), our best model to date. This is our recommended path if you're working with general English data.
Expand All @@ -35,6 +35,7 @@ Now that you have created the classifier, let's just train a model. The example
```python
import numpy as np
from datasets import load_dataset
from time import perf_counter

# Load the subj dataset
ds = load_dataset("setfit/subj")
Expand All @@ -45,13 +46,13 @@ s = perf_counter()
classifier = classifier.fit(train["text"], train["label"])

print(f"Training took {int(perf_counter() - s)} seconds.")
# Training took 81 seconds
# Training took 31 seconds
classification_report = classifier.evaluate(ds["test"]["text"], ds["test"]["label"])
print(classification_report)
# Achieved 91.0 test accuracy
# Achieved 92.0 test accuracy
```

As you can see, we got a pretty nice 91% accuracy, with only 81 seconds of training.
As you can see, we got a pretty nice 92% accuracy, with only 31 seconds of training.

The training loop is a plain PyTorch loop (see [`model2vec/train/trainer.py`](trainer.py)). By default the training loop splits the data into a train and validation split, with 90% of the data being used for training and 10% for validation. By default, it runs with early stopping on the validation set accuracy, with a patience of 5.

Expand All @@ -63,7 +64,7 @@ from time import perf_counter
s = perf_counter()
classifier.predict(test["text"])
print(f"Took {int((perf_counter() - s) * 1000)} milliseconds for {len(test)} instances on CPU.")
# Took 67 milliseconds for 2000 instances on CPU.
# Took 66 milliseconds for 2000 instances on CPU.
```

## Multi-label classification
Expand All @@ -75,7 +76,7 @@ from datasets import load_dataset
from model2vec.train import StaticModelForClassification

# Initialize a classifier from a pre-trained model
classifier = StaticModelForClassification.from_pretrained(model_name="minishlab/potion-base-32M")
classifier = StaticModelForClassification.from_pretrained(path="minishlab/potion-base-32M")

# Load a multi-label dataset
ds = load_dataset("google-research-datasets/go_emotions")
Expand Down Expand Up @@ -113,6 +114,25 @@ Because the other pairs in a batch serve as negatives, the training and validati

The InfoNCE temperature can be set with `temperature` (default `0.05`). It must be positive.

## Large datasets

Training data is read and tokenized per batch, so `fit` also accepts the columns of a Hugging Face dataset. These are not loaded into memory:

```python
from datasets import load_dataset

dataset = load_dataset("sentence-transformers/gooaq", split="train")
model.fit(text_a=dataset["question"], text_b=dataset["answer"])
```

Without an explicit validation set, `fit` holds out `test_size` of the data for validation, capped at 10,000 rows; pass an int to hold out an exact number of rows. For single-label classification, the split is stratified by class.

Because batches are shuffled, training reads the rows of a dataset in random order. If the dataset is stored on disk and does not fit in memory, this can be slow, especially on a network file system.

Columns of a dataset with a transform, set with `with_transform`, are not accepted. Apply the transform first with `dataset.map(transform, batched=True)`. For a dataset loaded from disk or the Hub, this writes the result to the cache on disk, so it is still not loaded into memory.

Columns of an iterable dataset, such as one loaded with `streaming=True`, are not accepted.

# Persistence

You can turn a classifier into a lightweight inference pipeline, as follows:
Expand Down Expand Up @@ -148,12 +168,11 @@ Our training architecture is set up to be extensible, with each task having a sp
The core functionality of the `StaticModelForClassification` is contained in a couple of functions:

* `construct_head`: This function constructs the classifier on top of the staticmodel. For example, if you want to create a model that has LayerNorm, just subclass, and replace this function. This should be the main function to update if you want to change model behavior.
* `train_test_split`: governs the train test split before classification.
* `prepare_dataset`: Selects the `torch.Dataset` that will be used in the `Dataloader` during training.
* `_create_datasets`: splits off the validation data, and creates the `torch.Dataset`s that will be used in the `Dataloader` during training.
* `_encode`: The encoding function used in the model.
* `fit`: contains all the fitting logic.

The training loop itself lives in `model2vec.train.trainer.run_training_loop`, a plain torch loop that is fairly basic and easy to modify. Each task passes in its own loss function (and, for classification, a small function that computes extra validation metrics like accuracy).
The training loop itself is defined in `model2vec.train.trainer.run_training_loop`, a plain torch loop that is fairly basic and easy to modify. Each task passes in its own loss function (and, for classification, a small function that computes extra validation metrics like accuracy).

# Results

Expand Down
200 changes: 120 additions & 80 deletions model2vec/train/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,24 +2,27 @@

import copy
import logging
from collections.abc import Sequence
from collections.abc import Sequence, Sized
from typing import Any, TypeVar

import numpy as np
import torch
from datasets import Column, DatasetDict, IterableColumn, IterableDataset, IterableDatasetDict
from datasets import Dataset as HFDataset
from tokenizers import Encoding, Tokenizer
from torch import nn
from torch.nn.utils.rnn import pad_sequence
from tqdm import trange

from model2vec.inference import StaticModelPipeline
from model2vec.model import DEFAULT_MAX_LENGTH, PathLike, StaticModel, _disable_padding, _get_unk_token_id
from model2vec.train.dataset import PairDataset, TextDataset
from model2vec.train.dataset import ColumnRows, PairDataset, TextDataset, has_only_strings, has_transform
from model2vec.train.trainer import MetricsFn, default_metrics, resolve_device, run_training_loop
from model2vec.train.utils import (
MAX_VALIDATION_SIZE,
get_probable_pad_token_id,
split_indices,
to_pipeline,
train_test_split,
)

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -97,11 +100,24 @@ def __init__(
_disable_padding(self.tokenizer)
self.unk_token_id = _get_unk_token_id(self.tokenizer)

def _remove_unk(self, token_ids: list[int]) -> list[int]:
"""Drop unknown tokens, mirroring `StaticModel.tokenize`."""
if self.unk_token_id is None:
return token_ids
return [token_id for token_id in token_ids if token_id != self.unk_token_id]
def _tokenize_ids(self, texts: Sequence[str]) -> list[list[int]]:
"""Tokenize texts into lists of token ids, dropping unknown tokens and truncating to `max_length` tokens.

:param texts: The texts to tokenize.
:return: The token ids of each text.
"""
if self.max_length is not None:
truncate_length = self.max_length * 10
texts = [text[:truncate_length] for text in texts]
encoded: list[Encoding] = self.tokenizer.encode_batch_fast(texts, add_special_tokens=False)
ids = [encoding.ids for encoding in encoded]
if self.unk_token_id is not None:
ids = [[token_id for token_id in token_ids if token_id != self.unk_token_id] for token_ids in ids]
return [token_ids[: self.max_length] for token_ids in ids]

def _to_targets(self, labels: Any) -> torch.Tensor:
"""Turn a batch of labels, such as vectors, into a float tensor of targets."""
return torch.as_tensor(labels, dtype=torch.float32)

def construct_weights(self) -> nn.Parameter:
"""Construct the weights for the model."""
Expand Down Expand Up @@ -275,11 +291,7 @@ def tokenize(self, texts: list[str]) -> torch.Tensor:
:param texts: The texts to tokenize.
:return: A 2D padded tensor
"""
max_length = self.max_length
encoded: list[Encoding] = self.tokenizer.encode_batch_fast(texts, add_special_tokens=False)
encoded_ids: list[torch.Tensor] = [
torch.Tensor(self._remove_unk(encoding.ids)[:max_length]).long() for encoding in encoded
]
encoded_ids: list[torch.Tensor] = [torch.LongTensor(token_ids) for token_ids in self._tokenize_ids(texts)]
return pad_sequence(encoded_ids, batch_first=True, padding_value=self.pad_id)

@property
Expand Down Expand Up @@ -327,31 +339,6 @@ def _determine_batch_size(self, batch_size: int | None, train_length: int) -> in

return batch_size

def _check_val_split(
self,
X: list[str],
y: list,
X_val: list[str] | None,
y_val: list | None,
test_size: float,
) -> tuple[list[str], list[str], Sequence, Sequence]:
if (X_val is not None) != (y_val is not None):
raise ValueError("Both X_val and y_val must be provided together, or neither.")

if X_val is not None and y_val is not None:
# Additional check to ensure y_val is of the same type as y
if type(y_val[0]) != type(y[0]):
raise ValueError("X_val and y_val must be of the same type as X and y.")

train_texts = X
train_labels = y
validation_texts = X_val
validation_labels = y_val
else:
train_texts, validation_texts, train_labels, validation_labels = train_test_split(X, y, test_size=test_size)

return train_texts, validation_texts, train_labels, validation_labels

def _train(
self,
loss_function: nn.Module,
Expand Down Expand Up @@ -420,59 +407,112 @@ def _determine_val_check_interval(

return val_check_interval, check_val_every_epoch

def _tokenize_texts(self, X: list[str], max_length: int | None) -> list[list[int]]:
"""Tokenize a list of texts into lists of token ids.
@staticmethod
def _check_inputs(**arguments: object) -> None:
"""Check that every argument of `fit` is a sequence, an array, a tensor, or a column of a Hugging Face dataset.

:param **arguments: The arguments, by name. None is skipped.
:raises ValueError: If an argument is a Hugging Face `Dataset` or `DatasetDict`, an iterable dataset or one of
its columns, a column of a dataset with a transform, a single string, or any other object that is not a
sequence, an array, or a tensor.
"""
for name, value in arguments.items():
if isinstance(value, (IterableDataset, IterableDatasetDict, IterableColumn)):
raise ValueError(
f"{name} comes from an iterable Hugging Face dataset, which has no length. Pass a column of a "
"regular dataset instead, such as one loaded without `streaming=True`."
)
if isinstance(value, (HFDataset, DatasetDict)):
raise ValueError(
f"{name} is a Hugging Face dataset. Pass one of its columns instead, such as `dataset['text']`."
)
if isinstance(value, Column) and has_transform(value):
raise ValueError(
f"{name} is a column of a Hugging Face dataset with a transform. Apply the transform first with "
"`dataset.map(transform, batched=True)`, or pass a list."
)
if isinstance(value, str) or (
value is not None and not isinstance(value, (Sequence, np.ndarray, torch.Tensor))
):
raise ValueError(
f"{name} must be a list, a tuple, an array, a tensor, or a column of a Hugging Face dataset, got "
f"{type(value).__name__}."
)

@staticmethod
def _check_aligned(**columns: Sized) -> None:
"""Check that the columns of the training or validation data have the same length.

:param X: The texts to tokenize.
:param max_length: The maximum length of the input in tokens. If this is None, no truncation is done.
:return: The tokenized texts.
:param **columns: The columns, by name.
:raises ValueError: If the columns don't all have the same length.
"""
batch_size = 1024
tokenized: list[list[int]] = []
for batch_idx in trange(0, len(X), batch_size, desc="Tokenizing data"):
batch = X[batch_idx : batch_idx + batch_size]
if max_length is not None:
truncate_length = max_length * 10
batch = [x[:truncate_length] for x in batch]
encoded = self.tokenizer.encode_batch_fast(batch, add_special_tokens=False)
tokenized.extend([self._remove_unk(encoding.ids)[:max_length] for encoding in encoded])

return tokenized

def _prepare_dataset(self, X: list[str], y: torch.Tensor, max_length: int | None) -> TextDataset:
"""Prepare a dataset.

:param X: The texts.
:param y: The labels.
:param max_length: The maximum length of the input in tokens. If this is None, no truncation is done.
:return: A TextDataset.
lengths = {name: len(column) for name, column in columns.items()}
if len(set(lengths.values())) > 1:
raise ValueError(f"{' and '.join(lengths)} must have the same length, got {lengths}.")

@staticmethod
def _check_texts(**texts: Sequence[str] | None) -> None:
"""Check that all texts are strings.

:param **texts: The texts, by name. None is skipped.
:raises ValueError: If a text is missing or is not a string.
"""
return TextDataset(self._tokenize_texts(X, max_length), y, pad_id=self.pad_id)
for name, values in texts.items():
if values is None:
continue
is_valid = (
has_only_strings(values)
if isinstance(values, Column)
else all(isinstance(text, str) for text in values)
)
if not is_valid:
raise ValueError(f"All texts in {name} must be strings.")

def _labels_to_tensor(self, labels: Any) -> torch.Tensor:
"""Turn the labels into a tensor."""
return labels
def _text_dataset(self, rows: ColumnRows, indices: np.ndarray | None = None) -> TextDataset:
"""Create a dataset of labeled texts that are tokenized per batch.

:param rows: The labeled texts, in a `text` and a `label` column.
:param indices: The indices of the rows that belong to the dataset. If None, all rows belong to it.
:return: The dataset.
"""
return TextDataset(rows, self._tokenize_ids, self._to_targets, indices, pad_id=self.pad_id)

def _create_datasets(
self,
X: list[str],
X: Sequence[str],
y: Any,
X_val: list[str] | None,
X_val: Sequence[str] | None,
y_val: Any | None,
test_size: float,
test_size: float | int,
stratify_by: Sequence[Any] | None = None,
) -> tuple[TextDataset, TextDataset]:
train_texts, validation_texts, train_labels, validation_labels = self._check_val_split(
X, y, X_val, y_val, test_size
)
y_tensor = self._labels_to_tensor(train_labels)
y_val_tensor = self._labels_to_tensor(validation_labels)

logger.info("Preparing train dataset.")
train_dataset = self._prepare_dataset(train_texts, y_tensor, self.max_length)
logger.info("Preparing validation dataset.")
val_dataset = self._prepare_dataset(validation_texts, y_val_tensor, self.max_length)
"""Create the training and validation datasets.

:param X: The training texts.
:param y: The training labels.
:param X_val: The validation texts. If None, the validation data is split off from `X` and `y`.
:param y_val: The validation labels.
:param test_size: The size of the validation split if `X_val` is None: a fraction of the data, capped at
`MAX_VALIDATION_SIZE` rows, or a number of rows if it is an int.
:param stratify_by: Validated single labels to stratify the validation split by. If None, the split is not
stratified.
:return: The train and validation datasets.
:raises ValueError: If only one of `X_val` and `y_val` is given, or if the texts and labels have different
lengths.
"""
if (X_val is None) != (y_val is None):
raise ValueError("Both X_val and y_val must be provided together, or neither.")
self._check_aligned(X=X, y=y)
self._check_texts(X=X, X_val=X_val)
rows = ColumnRows(text=X, label=y)
if X_val is not None and y_val is not None:
self._check_aligned(X_val=X_val, y_val=y_val)
return self._text_dataset(rows), self._text_dataset(ColumnRows(text=X_val, label=y_val))

return train_dataset, val_dataset
train_indices, val_indices = split_indices(
len(rows), test_size, max_test_size=MAX_VALIDATION_SIZE, stratify_by=stratify_by
)
return self._text_dataset(rows, train_indices), self._text_dataset(rows, val_indices)


T = TypeVar("T", bound=BaseFinetuneable)
Loading
Loading