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
14 changes: 4 additions & 10 deletions model2vec/train/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -100,24 +100,18 @@ The scores are competitive with the popular [roberta-base-go_emotions](https://h

## Pair similarity

`StaticModelForPairSimilarity` trains a model to embed pairs of related texts (e.g. queries and their matching documents) close together, by encoding both sides with the same model and minimizing the cosine distance between them:
`StaticModelForPairSimilarity` trains a model to embed pairs of related texts (e.g. queries and their matching documents) close together, by encoding both sides with the same model. It is trained with an InfoNCE loss with in-batch negatives: each `text_a` is pulled towards its paired `text_b` and pushed away from every other `text_b` in the batch:

```python
from model2vec.train import StaticModelForPairSimilarity

model = StaticModelForPairSimilarity.from_pretrained(model_name="minishlab/potion-base-32M")
model.fit(text_a=["how tall is the eiffel tower?"], text_b=["the eiffel tower is 330 meters tall."])
model.fit(text_a=queries, text_b=documents)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Example uses undefined variables The example calls model.fit with queries and documents but never defines them. Readers cannot run the snippet as shown; sample lists or clearly marked placeholders would make it usable.

Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!

```

Pairs can also be labeled: pairs labeled `1` are pushed together (cosine similarity towards 1), while pairs labeled `0` are pushed towards a cosine similarity of 0. If `labels` is omitted, every pair is treated as positive:
Because the other pairs in a batch serve as negatives, the training and validation sets each need at least two pairs. Pairs with the same `text_a` are treated as alternative positives for that text, so they don't serve as negatives for each other.

```python
model.fit(
text_a=["how tall is the eiffel tower?", "how tall is the eiffel tower?"],
text_b=["the eiffel tower is 330 meters tall.", "paris is the capital of france."],
labels=[1, 0],
)
```
The InfoNCE temperature can be set with `temperature` (default `0.05`). It must be positive.

# Persistence

Expand Down
27 changes: 11 additions & 16 deletions model2vec/train/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,50 +45,45 @@ def __init__(
self,
tokenized_texts_a: list[list[int]],
tokenized_texts_b: list[list[int]],
labels: list[int] | torch.Tensor | None = None,
pad_id: int = 0,
) -> None:
"""A dataset of aligned text pairs.

:param tokenized_texts_a: The tokenized first half of each pair. Each text is a list of token ids.
:param tokenized_texts_b: The tokenized second half of each pair. Each text is a list of token ids.
:param labels: The label for each pair: 1 if the pair should be pushed together, 0 if it should be
pushed towards a cosine similarity of 0. If None, every pair is labeled 1.
:param pad_id: The id used to pad batches. Must match the `pad_id` of the model being trained.
:raises ValueError: If the two halves don't have the same number of texts, or if `labels` doesn't
have one entry per pair.
:raises ValueError: If the two halves don't have the same number of texts.
"""
if len(tokenized_texts_a) != len(tokenized_texts_b):
raise ValueError("The two halves of a pair dataset must have the same number of texts.")
if labels is not None and len(labels) != len(tokenized_texts_a):
raise ValueError("labels must have one entry per pair.")
self.tokenized_texts_a = tokenized_texts_a
self.tokenized_texts_b = tokenized_texts_b
self.labels = torch.ones(len(tokenized_texts_a)) if labels is None else torch.as_tensor(labels).float()
self.pad_id = pad_id

def __len__(self) -> int:
"""Return the length of the dataset."""
return len(self.tokenized_texts_a)

def __getitem__(self, index: int) -> tuple[list[int], list[int], torch.Tensor]:
def __getitem__(self, index: int) -> tuple[list[int], list[int]]:
"""Gets an item."""
return self.tokenized_texts_a[index], self.tokenized_texts_b[index], self.labels[index]
return self.tokenized_texts_a[index], self.tokenized_texts_b[index]

def collate_fn(self, batch: list[tuple[list[int], list[int], torch.Tensor]]) -> tuple[torch.Tensor, torch.Tensor]:
def collate_fn(self, batch: list[tuple[list[int], list[int]]]) -> tuple[torch.Tensor, 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.
stacked into a single (2, batch_size, seq_len) tensor. The targets are the index of each
pair's second text within the batch.
"""
texts_a, texts_b, labels = zip(*batch)
texts_a, texts_b = zip(*batch)

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.stack(labels)
return torch.stack([padded_a, padded_b]), torch.arange(len(texts_a))

def to_dataloader(self, shuffle: bool, batch_size: int = 32) -> DataLoader:
"""Convert the dataset to a DataLoader."""
return DataLoader(self, collate_fn=self.collate_fn, shuffle=shuffle, batch_size=batch_size)
"""Convert the dataset to a DataLoader. A final batch with a single pair is dropped, unless it is the only pair."""
drop_last = len(self) > 1 and len(self) % batch_size == 1

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Sole positive pair gets dropped When a validation set has three pairs labeled [0, 0, 1] and batch_size=2, this drops the only positive pair. The remaining negative-only batch reports zero loss, so checkpoint selection and early stopping never assess a positive match. Shuffling can also drop the only positive training pair for an epoch.

return DataLoader(self, collate_fn=self.collate_fn, shuffle=shuffle, batch_size=batch_size, drop_last=drop_last)
118 changes: 73 additions & 45 deletions model2vec/train/pairs.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,19 +15,40 @@
logger = logging.getLogger(__name__)


class PairCosineLoss(nn.Module):
def __call__(self, head_out: tuple[torch.Tensor, torch.Tensor], y: torch.Tensor) -> torch.Tensor:
"""Returns the cosine loss between the two encoded halves of a pair batch, per pair label.
class PairInfoNCELoss(nn.Module):
def __init__(self, temperature: float = 0.05) -> None:
"""Initialize the InfoNCE loss.

Pairs labeled 1 are pushed towards a cosine similarity of 1, pairs labeled 0 are pushed
towards a cosine similarity of 0.
:param temperature: The temperature by which the cosine similarities are divided. Must be positive.
:raises ValueError: If `temperature` is not positive.
"""
out_a, out_b = head_out
super().__init__()
if temperature <= 0:
raise ValueError(f"temperature must be positive, got {temperature}.")
self.temperature = temperature

def __call__(
self, head_out: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], y: torch.Tensor
) -> torch.Tensor:
"""Returns the InfoNCE loss over a pair batch, using in-batch negatives.

Every first text is an anchor, its paired second text the positive, and all other second texts in
the batch negatives. A second text is not used as a negative for an anchor if it is identical to
the anchor's positive, or if it is paired with a first text identical to the anchor.

:param head_out: The encoded first texts and second texts, followed by an id for each first text and
each second text. Identical texts have the same id.
:param y: For each anchor, the index of its positive among the second texts.
:return: The mean loss over the anchors.
"""
out_a, out_b, ids_a, ids_b = head_out
out_a = torch.nn.functional.normalize(out_a, dim=1)
out_b = torch.nn.functional.normalize(out_b, dim=1)
cosine_sim = torch.sum(out_a * out_b, dim=1)
loss = torch.where(y == 1, 1 - cosine_sim, cosine_sim.abs())
return loss.mean()
logits = (out_a @ out_b.T) / self.temperature
Comment thread
stephantul marked this conversation as resolved.
Comment thread
stephantul marked this conversation as resolved.
false_negatives = (ids_a[:, None] == ids_a[None, :]) | (ids_b[:, None] == ids_b[None, :])
false_negatives[torch.arange(len(y), device=y.device), y] = False
logits = logits.masked_fill(false_negatives, float("-inf"))
return torch.nn.functional.cross_entropy(logits, y)


class StaticModelForPairSimilarity(BaseFinetuneable):
Expand Down Expand Up @@ -81,70 +102,78 @@ def __init__(
max_length=max_length,
)

def forward(self, input_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: # type: ignore[override]
def forward( # type: ignore[override]
self, input_ids: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Encode both halves of a pair batch through the shared embeddings and head.

:param input_ids: A `(2, batch_size, seq_len)` tensor, stacking the two padded text sets.
:return: The head outputs for the first and second set of texts.
:return: The head outputs for the first and second set of texts, followed by an id for each first
text and each second text. Identical texts have the same id.
"""
out_a = self.head(self._encode(input_ids[0]))
out_b = self.head(self._encode(input_ids[1]))
return out_a, out_b
ids_a = torch.unique(input_ids[0], dim=0, return_inverse=True)[1]
ids_b = torch.unique(input_ids[1], dim=0, return_inverse=True)[1]
return out_a, out_b, ids_a, ids_b

def _check_pair_val_split(
self,
text_a: list[str],
text_b: list[str],
labels: list[int],
text_a_val: list[str] | None,
text_b_val: list[str] | None,
labels_val: list[int] | None,
test_size: float,
) -> tuple[list[str], list[str], list[str], list[str], list[int], list[int]]:
) -> tuple[list[str], list[str], list[str], list[str]]:
if len(text_a) != len(text_b):
raise ValueError("text_a and text_b must have the same length.")
if len(labels) != len(text_a):
raise ValueError("labels must have the same length as text_a and text_b.")
if (text_a_val is not None) != (text_b_val is not None):
raise ValueError("Both text_a_val and text_b_val must be provided together, or neither.")

if text_a_val is not None and text_b_val is not None:
if len(text_a_val) != len(text_b_val):
raise ValueError("text_a_val and text_b_val must have the same length.")
labels_val = [1] * len(text_a_val) if labels_val is None else labels_val
if len(labels_val) != len(text_a_val):
raise ValueError("labels_val must have the same length as text_a_val and text_b_val.")
return text_a, text_a_val, text_b, text_b_val, labels, labels_val
return text_a, text_a_val, text_b, text_b_val

pairs = list(zip(text_a, text_b))
train_pairs, val_pairs, train_labels, val_labels = train_test_split(pairs, labels, test_size=test_size)
train_pairs, val_pairs, _, _ = train_test_split(pairs, pairs, test_size=test_size)
train_a, train_b = map(list, zip(*train_pairs)) if train_pairs else ([], [])
val_a, val_b = map(list, zip(*val_pairs)) if val_pairs else ([], [])
return train_a, val_a, train_b, val_b, train_labels, val_labels
return train_a, val_a, train_b, val_b

@staticmethod
def _check_pair_splits(n_train: int, n_val: int) -> None:
"""Check that the training and validation sets each have at least two pairs.

def _prepare_pair_dataset(
self, text_a: list[str], text_b: list[str], labels: list[int], max_length: int | None
) -> PairDataset:
:param n_train: The number of training pairs.
:param n_val: The number of validation pairs.
:raises ValueError: If either set has fewer than two pairs.
"""
for name, n_pairs in (("training", n_train), ("validation", n_val)):
if n_pairs < 2:
raise ValueError(
f"The {name} set needs at least two pairs, got {n_pairs}. Pass more pairs, "
"a different test_size, or an explicit validation set."
)

def _prepare_pair_dataset(self, text_a: list[str], text_b: list[str], max_length: int | None) -> PairDataset:
"""Tokenize both halves of a pair dataset.

:param text_a: The first half of each pair.
:param text_b: The second half of each pair.
:param labels: The label for each pair.
:param max_length: The maximum length of the input in tokens. If this is None, no truncation is done.
:return: A PairDataset.
"""
return PairDataset(
self._tokenize_texts(text_a, max_length),
self._tokenize_texts(text_b, max_length),
labels=labels,
pad_id=self.pad_id,
)

def fit(
self: T,
text_a: list[str],
text_b: list[str],
labels: list[int] | None = None,
learning_rate: float = 1e-3,
batch_size: int | None = None,
min_epochs: int | None = None,
Expand All @@ -154,16 +183,16 @@ def fit(
device: str = "auto",
text_a_val: list[str] | None = None,
text_b_val: list[str] | None = None,
labels_val: list[int] | None = None,
validation_steps: int | None = None,
random_seed: int = DEFAULT_RANDOM_SEED,
temperature: float = 0.05,
) -> T:
"""Fit a model that maximizes the cosine similarity between paired texts.
"""Fit a model that embeds paired texts close together.

This function trains the model with a plain torch training loop. Both `text_a` and `text_b`
are encoded with the same model. Pairs labeled 1 are pushed together, minimizing the cosine
distance between them. Pairs labeled 0 are pushed towards a cosine similarity of 0. We use
early stopping. After training, the weights of the best model are loaded back into the model.
are encoded with the same model, and trained with an InfoNCE loss: each `text_a` is pulled towards
its paired `text_b` and pushed away from all other `text_b` in the batch. Pairs with the same
`text_a` are not used as negatives for each other. We use early stopping. After training, the weights of the best model are loaded back into the model.

This function seeds everything with a seed of 42, so the results are reproducible.
It also splits the data into a train and validation set, again with a random seed.
Expand All @@ -173,8 +202,6 @@ def fit(

:param text_a: The first half of each training pair.
:param text_b: The second half of each training pair.
:param labels: The label for each training pair: 1 if the pair should be pushed together, 0 if
it should be pushed towards a cosine similarity of 0. If None, every pair is labeled 1.
:param learning_rate: The learning rate.
:param batch_size: The batch size. If None, a good batch size is chosen automatically.
:param min_epochs: The minimum number of epochs to train for.
Expand All @@ -186,30 +213,31 @@ def fit(
:param device: The device to train on. If this is "auto", the device is chosen automatically.
:param text_a_val: The first half of each validation pair.
:param text_b_val: The second half of each validation pair.
:param labels_val: The label for each validation pair. If None, every validation pair is labeled 1.
:param validation_steps: The number of steps to run validation for. If None, validation steps are estimated from the data.
:param random_seed: The random seed to use. Defaults to 42.
:param temperature: The temperature of the InfoNCE loss.
:return: The fitted model.
:raises ValueError: If `batch_size` is smaller than 2.
"""
seed_everything(random_seed)
logger.info("Re-initializing model.")
loss_function = PairInfoNCELoss(temperature=temperature)

labels = [1] * len(text_a) if labels is None else labels

train_a, val_a, train_b, val_b, train_labels, val_labels = self._check_pair_val_split(
text_a, text_b, labels, text_a_val, text_b_val, labels_val, test_size
)
train_a, val_a, train_b, val_b = self._check_pair_val_split(text_a, text_b, text_a_val, text_b_val, test_size)
self._check_pair_splits(len(train_a), len(val_a))
self._initialize()

logger.info("Preparing train dataset.")
train_dataset = self._prepare_pair_dataset(train_a, train_b, train_labels, self.max_length)
train_dataset = self._prepare_pair_dataset(train_a, train_b, self.max_length)
logger.info("Preparing validation dataset.")
val_dataset = self._prepare_pair_dataset(val_a, val_b, val_labels, self.max_length)
val_dataset = self._prepare_pair_dataset(val_a, val_b, self.max_length)

batch_size = self._determine_batch_size(batch_size, len(train_dataset))
if batch_size < 2:
raise ValueError(f"batch_size must be at least 2, got {batch_size}.")
Comment on lines +236 to +237

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Invalid batch size resets model When an already trained model receives batch_size=1, this check raises only after _initialize() has replaced its head and embeddings. The call fails, but the caller's trained model has already been reinitialized.


self._train(
loss_function=PairCosineLoss(),
loss_function=loss_function,
learning_rate=learning_rate,
train_dataset=train_dataset,
val_dataset=val_dataset,
Expand Down
6 changes: 3 additions & 3 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,9 +246,9 @@ def mock_trained_pair_similarity_pipeline() -> StaticModelForPairSimilarity:
vectors_torched = torch.randn(len(tokenizer.get_vocab()), 12)
model = StaticModelForPairSimilarity(vectors=vectors_torched, tokenizer=tokenizer, hidden_dim=12).to("cpu")

text_a = ["dog", "cat"]
text_b = ["puppy", "kitten"]
model.fit(text_a, text_b)
text_a = ["dog", "cat", "dog cat", "cat dog"]
text_b = ["puppy", "kitten", "puppy kitten", "kitten puppy"]
model.fit(text_a, text_b, test_size=0.5)

return model

Expand Down
Loading
Loading