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
159 changes: 78 additions & 81 deletions model2vec/train/base.py

Large diffs are not rendered by default.

21 changes: 8 additions & 13 deletions model2vec/train/classifier.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,12 +84,11 @@ def __init__(
n_layers: int = 1,
hidden_dim: int = 512,
out_dim: int = 2,
pad_id: int = 0,
token_mapping: list[int] | None = None,
weights: torch.Tensor | None = None,
freeze: bool = False,
normalize: bool = True,
freeze_weights: bool = False,
freeze_weights: bool | None = None,
max_length: int | None = DEFAULT_MAX_LENGTH,
) -> None:
"""Initialize a standard classifier model."""
Expand All @@ -100,7 +99,6 @@ def __init__(
super().__init__(
vectors=vectors,
out_dim=out_dim,
pad_id=pad_id,
tokenizer=tokenizer,
token_mapping=token_mapping,
weights=weights,
Expand All @@ -119,37 +117,35 @@ def from_pretrained(
*,
token: str | None = None,
model_name: PathLike | None = None,
pad_token: str | None = None,
max_length: int | None | _UnsetType = _UNSET,
n_layers: int = 1,
hidden_dim: int = 512,
out_dim: int = 2,
freeze: bool = False,
normalize: bool = True,
freeze_weights: bool = False,
freeze_weights: bool | None = None,
**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 not passed, the
static model's `max_length` is used. Pass None to disable truncation.
: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 freeze_weights: Whether to freeze the token weights. If None, the model's own weights are trained,
and a model without weights gets none. If False, a model without weights learns weights that start at 1.
: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,
Expand All @@ -165,33 +161,32 @@ def from_static_model(
cls: type[T],
*,
model: StaticModel,
pad_token: str | None = None,
max_length: int | None | _UnsetType = _UNSET,
n_layers: int = 1,
hidden_dim: int = 512,
out_dim: int = 2,
freeze: bool = False,
normalize: bool = True,
freeze_weights: bool = False,
freeze_weights: bool | None = None,
**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 not passed, the
static model's `max_length` is used. Pass None to disable truncation.
: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 freeze_weights: Whether to freeze the token weights. If None, the model's own weights are trained,
and a model without weights gets none. If False, a model without weights learns weights that start at 1.
: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),
**_static_model_arguments(model, max_length=max_length),
n_layers=n_layers,
hidden_dim=hidden_dim,
out_dim=out_dim,
Expand Down
83 changes: 57 additions & 26 deletions model2vec/train/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
from abc import ABC, abstractmethod
from collections import Counter, defaultdict
from collections.abc import Callable, Iterator, Mapping, Sequence
from dataclasses import dataclass
from itertools import chain
from typing import Any

import numpy as np
Expand All @@ -12,7 +14,6 @@
import torch
from datasets import Column
from datasets import Dataset as HFDataset
from torch.nn.utils.rnn import pad_sequence
from torch.utils.data import BatchSampler, DataLoader, Dataset, RandomSampler, SequentialSampler

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -215,17 +216,56 @@ def __getitem__(self, indices: list[int]) -> dict[str, Any]:
}


@dataclass(frozen=True)
class TokenBatch:
"""A batch of token id sequences, stored as one flat tensor of ids with the offset and length of each sequence."""

ids: torch.Tensor
offsets: torch.Tensor
lengths: torch.Tensor

@classmethod
def from_lengths(cls, ids: torch.Tensor, lengths: torch.Tensor) -> TokenBatch:
"""Create a batch from a flat tensor of ids and the length of each sequence."""
return cls(ids, lengths.cumsum(0) - lengths, lengths)

@classmethod
def from_token_ids(cls, token_ids: Sequence[Sequence[int]]) -> TokenBatch:
"""Flatten lists of token ids into a batch."""
lengths = np.fromiter(map(len, token_ids), dtype=np.int64, count=len(token_ids))
ids = np.fromiter(chain.from_iterable(token_ids), dtype=np.int64, count=int(lengths.sum()))
return cls.from_lengths(torch.from_numpy(ids), torch.from_numpy(lengths))

def __len__(self) -> int:
"""Return the number of sequences."""
return len(self.lengths)

def to(self, device: torch.device | str) -> TokenBatch:
"""Move the batch to a device."""
return TokenBatch(self.ids.to(device), self.offsets.to(device), self.lengths.to(device))


@dataclass(frozen=True)
class PairBatch:
"""A batch of pairs: the first texts followed by the second texts, and an id per text that is shared by identical texts."""

tokens: TokenBatch
text_ids: torch.Tensor

def to(self, device: torch.device | str) -> PairBatch:
"""Move the batch to a device."""
return PairBatch(self.tokens.to(device), self.text_ids.to(device))


class _Batches(Dataset, ABC):
def __init__(self, rows: ColumnRows, indices: np.ndarray | None, pad_id: int) -> None:
def __init__(self, rows: ColumnRows, indices: np.ndarray | None) -> None:
"""A dataset that fetches rows and turns them into items per batch.

:param rows: The rows to draw items from.
:param indices: The indices of the rows that belong to this dataset. If None, all rows belong to it.
:param pad_id: The id used to pad batches. Must match the `pad_id` of the model being trained.
"""
self.rows = rows
self.indices = np.arange(len(rows)) if indices is None else indices
self.pad_id = pad_id

def __len__(self) -> int:
"""Return the length of the dataset."""
Expand All @@ -244,7 +284,7 @@ def _to_items(self, rows: Mapping[str, Any]) -> list[Any]:
"""Turn a batch of rows into items."""

@abstractmethod
def collate_fn(self, batch: list[Any]) -> tuple[torch.Tensor, torch.Tensor]:
def collate_fn(self, batch: list[Any]) -> tuple[TokenBatch | PairBatch, torch.Tensor]:
"""Collate a batch of items into model inputs and targets."""

def _drop_last(self, batch_size: int) -> bool:
Expand All @@ -268,32 +308,26 @@ def __init__(
tokenize: Callable[[list[str]], list[list[int]]],
to_targets: Callable[[Any], torch.Tensor],
indices: np.ndarray | None = None,
pad_id: int = 0,
) -> None:
"""A dataset of labeled texts, which are tokenized per batch.

:param rows: The labeled texts, in a `text` and a `label` column.
:param tokenize: Turns a batch of texts into lists of token ids.
:param to_targets: Turns a batch of labels into a tensor of targets.
:param indices: The indices of the rows that belong to this dataset. If None, all rows belong to it.
:param pad_id: The id used to pad batches. Must match the `pad_id` of the model being trained.
"""
super().__init__(rows, indices, pad_id)
super().__init__(rows, indices)
self.tokenize = tokenize
self.to_targets = to_targets

def _to_items(self, rows: Mapping[str, Any]) -> list[tuple[list[int], torch.Tensor]]:
"""Tokenize the texts and turn the labels into targets."""
return list(zip(self.tokenize(rows[TEXT_COLUMN]), self.to_targets(rows[LABEL_COLUMN])))

def collate_fn(self, batch: list[tuple[list[int], torch.Tensor]]) -> tuple[torch.Tensor, torch.Tensor]:
def collate_fn(self, batch: list[tuple[list[int], torch.Tensor]]) -> tuple[TokenBatch, torch.Tensor]:
"""Collate function."""
texts, targets = zip(*batch)

tensors: list[torch.Tensor] = [torch.LongTensor(x) for x in texts]
padded = pad_sequence(tensors, batch_first=True, padding_value=self.pad_id)

return padded, torch.stack(targets)
return TokenBatch.from_token_ids(texts), torch.stack(targets)


class PairDataset(_Batches):
Expand All @@ -302,36 +336,33 @@ def __init__(
rows: ColumnRows,
tokenize: Callable[[list[str]], list[list[int]]],
indices: np.ndarray | None = None,
pad_id: int = 0,
) -> None:
"""A dataset of aligned text pairs, which are tokenized per batch.

:param rows: The pairs, in a `text_a` and a `text_b` column.
:param tokenize: Turns a batch of texts into lists of token ids.
:param indices: The indices of the rows that belong to this dataset. If None, all rows belong to it.
:param pad_id: The id used to pad batches. Must match the `pad_id` of the model being trained.
"""
super().__init__(rows, indices, pad_id)
super().__init__(rows, indices)
self.tokenize = tokenize

def _to_items(self, rows: Mapping[str, Any]) -> list[tuple[list[int], list[int]]]:
"""Tokenize both halves of each pair."""
return list(zip(self.tokenize(rows[TEXT_A_COLUMN]), self.tokenize(rows[TEXT_B_COLUMN])))

def collate_fn(self, batch: list[tuple[list[int], list[int]]]) -> tuple[torch.Tensor, torch.Tensor]:
def collate_fn(self, batch: list[tuple[list[int], list[int]]]) -> tuple[PairBatch, torch.Tensor]:
"""Collate function.

Both halves are padded together so they end up with the same sequence length, then
stacked into a single (2, batch_size, seq_len) tensor. The targets are the index of each
pair's second text within the batch.
The first texts and the second texts are put into a single batch, first texts first. The targets
are the index of each pair's second text within the batch.
"""
texts_a, texts_b = zip(*batch)
texts = (*texts_a, *texts_b)
ids_by_text: dict[tuple[int, ...], int] = {}
text_ids = [ids_by_text.setdefault(tuple(text), len(ids_by_text)) for text in texts]

tensors: list[torch.Tensor] = [torch.LongTensor(x) for x in (*texts_a, *texts_b)]
padded = pad_sequence(tensors, batch_first=True, padding_value=self.pad_id)
padded_a, padded_b = padded[: len(texts_a)], padded[len(texts_a) :]

return torch.stack([padded_a, padded_b]), torch.arange(len(texts_a))
pair_batch = PairBatch(TokenBatch.from_token_ids(texts), torch.tensor(text_ids, dtype=torch.int64))
return pair_batch, torch.arange(len(texts_a))

def _drop_last(self, batch_size: int) -> bool:
"""Drop a final batch with a single pair, unless it is the only pair."""
Expand Down
Loading
Loading