From 2021253a9b0fd219fb87f6f44aaaf0b22c69c8b3 Mon Sep 17 00:00:00 2001 From: stephantul Date: Sun, 27 Sep 2026 19:14:33 +0200 Subject: [PATCH 1/3] feat(train): train pair models with InfoNCE Replace the pair trainer's cosine loss with an InfoNCE loss over in-batch negatives: each text_a is pulled towards its own text_b and pushed away from every other text_b in the batch. The temperature can be set with `temperature` (default 0.05). Pairs labeled 0 are no longer pushed towards a cosine similarity of 0. They are not used as anchors, but their text_b still serves as a negative for the other pairs in the batch. Second texts with the same embedding as an anchor's positive, such as duplicates of the positive text, are not used as negatives for that anchor. BREAKING CHANGE: PairCosineLoss is removed, and fit trains with InfoNCE. --- model2vec/train/README.md | 6 +++-- model2vec/train/dataset.py | 4 +-- model2vec/train/pairs.py | 51 ++++++++++++++++++++++++++++---------- tests/test_trainable.py | 42 +++++++++++++++++++++++-------- 4 files changed, 75 insertions(+), 28 deletions(-) diff --git a/model2vec/train/README.md b/model2vec/train/README.md index 98342f4..5128111 100644 --- a/model2vec/train/README.md +++ b/model2vec/train/README.md @@ -100,7 +100,7 @@ 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 @@ -109,7 +109,7 @@ model = StaticModelForPairSimilarity.from_pretrained(model_name="minishlab/potio model.fit(text_a=["how tall is the eiffel tower?"], text_b=["the eiffel tower is 330 meters tall."]) ``` -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: +Pairs can also be labeled. If `labels` is omitted, every pair is treated as positive. Pairs labeled `0` are not used as anchors, but their `text_b` still serves as an additional negative for the other pairs in the batch: ```python model.fit( @@ -119,6 +119,8 @@ model.fit( ) ``` +The InfoNCE temperature can be set with `temperature` (default `0.05`). + # Persistence You can turn a classifier into a lightweight inference pipeline, as follows: diff --git a/model2vec/train/dataset.py b/model2vec/train/dataset.py index 0d7627e..e60311b 100644 --- a/model2vec/train/dataset.py +++ b/model2vec/train/dataset.py @@ -52,8 +52,8 @@ def __init__( :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 labels: The label for each pair: 1 if the pair should be pushed together, 0 otherwise. + 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. diff --git a/model2vec/train/pairs.py b/model2vec/train/pairs.py index 07f9254..e71e1d5 100644 --- a/model2vec/train/pairs.py +++ b/model2vec/train/pairs.py @@ -15,19 +15,41 @@ logger = logging.getLogger(__name__) -class PairCosineLoss(nn.Module): +class PairInfoNCELoss(nn.Module): + def __init__(self, temperature: float = 0.05) -> None: + """Initialize the InfoNCE loss. + + :param temperature: The temperature by which the cosine similarities are divided. + """ + super().__init__() + self.temperature = temperature + 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. + """Returns the InfoNCE loss over a pair batch, using in-batch negatives. + + Every first text labeled 1 is treated as an anchor, its paired second text as the positive, and + all other second texts in the batch as negatives. Pairs labeled 0 are not used as anchors, but + their second texts still serve as negatives for the other anchors. Second texts with the same + embedding as an anchor's positive, such as duplicates of the positive text, are not used as + negatives for that anchor. - Pairs labeled 1 are pushed towards a cosine similarity of 1, pairs labeled 0 are pushed - towards a cosine similarity of 0. + :param head_out: The encoded first texts and second texts. + :param y: The label of each pair. + :return: The mean loss over the anchors. """ out_a, out_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 + targets = torch.arange(len(logits), device=logits.device) + duplicates = (out_b @ out_b.T) > 1 - 1e-6 + duplicates[targets, targets] = False + logits = logits.masked_fill(duplicates, float("-inf")) + loss = torch.nn.functional.cross_entropy(logits, targets, reduction="none") + positive = y == 1 + if not positive.any(): + return (loss * 0).sum() + return loss[positive].mean() class StaticModelForPairSimilarity(BaseFinetuneable): @@ -156,12 +178,14 @@ def fit( 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 + are encoded with the same model, and trained with an InfoNCE loss: each `text_a` labeled 1 is + pulled towards its paired `text_b` and pushed away from all other `text_b` in the batch. Pairs + labeled 0 are not used as anchors, but their `text_b` still serves as an in-batch negative. 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. @@ -172,8 +196,8 @@ 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 labels: The label for each training pair: 1 if the pair should be pushed together, 0 otherwise. + 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. @@ -188,6 +212,7 @@ def fit( :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. """ seed_everything(random_seed) @@ -208,7 +233,7 @@ def fit( batch_size = self._determine_batch_size(batch_size, len(train_dataset)) self._train( - loss_function=PairCosineLoss(), + loss_function=PairInfoNCELoss(temperature=temperature), learning_rate=learning_rate, train_dataset=train_dataset, val_dataset=val_dataset, diff --git a/tests/test_trainable.py b/tests/test_trainable.py index 7d90fb0..3010361 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -16,7 +16,7 @@ from model2vec.train import StaticModelForClassification from model2vec.train.base import BaseFinetuneable from model2vec.train.dataset import PairDataset, TextDataset -from model2vec.train.pairs import PairCosineLoss, StaticModelForPairSimilarity +from model2vec.train.pairs import PairInfoNCELoss, StaticModelForPairSimilarity from model2vec.train.regression import StaticModelForRegression from model2vec.train.similarity import StaticModelForSimilarity from model2vec.train.trainer import _resolve_max_epochs, resolve_device, run_training_loop @@ -420,18 +420,38 @@ def test_pairdataset_labels_mismatched_length() -> None: PairDataset([[1], [2]], [[3], [4]], labels=[1]) -def test_pair_cosine_loss_pushes_towards_label() -> None: - """Label 1 pairs are pushed towards a cosine similarity of 1, label 0 pairs towards 0.""" - loss_fn = PairCosineLoss() - out_a = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) +def test_pair_infonce_loss_is_query_to_document() -> None: + """The loss is the cross-entropy of each first text over all second texts in the batch.""" + torch.manual_seed(0) + out_a, out_b = torch.randn(4, 3), torch.randn(4, 3) + loss_fn = PairInfoNCELoss(temperature=0.1) + logits = torch.nn.functional.normalize(out_a, dim=1) @ torch.nn.functional.normalize(out_b, dim=1).T / 0.1 + expected = torch.nn.functional.cross_entropy(logits, torch.arange(4)) + + assert loss_fn((out_a, out_b), torch.ones(4)).item() == pytest.approx(expected.item(), abs=1e-5) + + +def test_pair_infonce_loss_ignores_negative_anchors() -> None: + """Pairs labeled 0 are not used as anchors, but their second text is still a negative.""" + torch.manual_seed(0) + out_a, out_b = torch.randn(3, 3), torch.randn(3, 3) + loss_fn = PairInfoNCELoss(temperature=0.1) + logits = torch.nn.functional.normalize(out_a, dim=1) @ torch.nn.functional.normalize(out_b, dim=1).T / 0.1 + per_pair = torch.nn.functional.cross_entropy(logits, torch.arange(3), reduction="none") + + loss = loss_fn((out_a, out_b), torch.tensor([1.0, 0.0, 1.0])) + assert loss.item() == pytest.approx(per_pair[[0, 2]].mean().item(), abs=1e-5) + assert loss_fn((out_a, out_b), torch.zeros(3)).item() == 0.0 + - identical = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) - orthogonal = torch.tensor([[0.0, 1.0], [0.0, 1.0]]) +def test_pair_infonce_loss_masks_duplicate_positives() -> None: + """Second texts identical to an anchor's positive are not used as negatives for that anchor.""" + out_a = torch.tensor([[1.0, 0.0], [0.0, 1.0]]) + out_b = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) + loss_fn = PairInfoNCELoss(temperature=0.05) - assert loss_fn((out_a, identical), torch.tensor([1.0, 1.0])).item() == pytest.approx(0.0, abs=1e-6) - assert loss_fn((out_a, orthogonal), torch.tensor([1.0, 1.0])).item() == pytest.approx(1.0) - assert loss_fn((out_a, orthogonal), torch.tensor([0.0, 0.0])).item() == pytest.approx(0.0, abs=1e-6) - assert loss_fn((out_a, identical), torch.tensor([0.0, 0.0])).item() == pytest.approx(1.0) + loss = loss_fn((out_a, out_b), torch.tensor([1.0, 0.0])) + assert loss.item() == pytest.approx(0.0, abs=1e-6) def test_pair_similarity_out_dim_defaults_to_embed_dim(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: From bcb0a43ca78ac5d8d59dedcfb373a2a72dfa03c8 Mon Sep 17 00:00:00 2001 From: stephantul Date: Sun, 27 Sep 2026 20:17:13 +0200 Subject: [PATCH 2/3] address reviewer comments --- model2vec/train/README.md | 14 ++--- model2vec/train/dataset.py | 5 +- model2vec/train/pairs.py | 65 +++++++++++++++++----- tests/conftest.py | 6 +- tests/test_trainable.py | 110 ++++++++++++++++++++++++++++++++++--- 5 files changed, 164 insertions(+), 36 deletions(-) diff --git a/model2vec/train/README.md b/model2vec/train/README.md index 5128111..5c95ae7 100644 --- a/model2vec/train/README.md +++ b/model2vec/train/README.md @@ -106,20 +106,18 @@ The scores are competitive with the popular [roberta-base-go_emotions](https://h 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) ``` -Pairs can also be labeled. If `labels` is omitted, every pair is treated as positive. Pairs labeled `0` are not used as anchors, but their `text_b` still serves as an additional negative for the other pairs in the batch: +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. + +Pairs can also be labeled. If `labels` is omitted, every pair is treated as positive. Pairs labeled `0` are not used as anchors, but their `text_b` still serves as an additional negative for the other pairs in the batch, including pairs with the same `text_a`: ```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], -) +model.fit(text_a=queries, text_b=documents, labels=labels) ``` -The InfoNCE temperature can be set with `temperature` (default `0.05`). +The InfoNCE temperature can be set with `temperature` (default `0.05`). It must be positive. # Persistence diff --git a/model2vec/train/dataset.py b/model2vec/train/dataset.py index e60311b..65a5de3 100644 --- a/model2vec/train/dataset.py +++ b/model2vec/train/dataset.py @@ -90,5 +90,6 @@ def collate_fn(self, batch: list[tuple[list[int], list[int], torch.Tensor]]) -> return torch.stack([padded_a, padded_b]), torch.stack(labels) 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 + return DataLoader(self, collate_fn=self.collate_fn, shuffle=shuffle, batch_size=batch_size, drop_last=drop_last) diff --git a/model2vec/train/pairs.py b/model2vec/train/pairs.py index e71e1d5..e59087d 100644 --- a/model2vec/train/pairs.py +++ b/model2vec/train/pairs.py @@ -19,34 +19,42 @@ class PairInfoNCELoss(nn.Module): def __init__(self, temperature: float = 0.05) -> None: """Initialize the InfoNCE loss. - :param temperature: The temperature by which the cosine similarities are divided. + :param temperature: The temperature by which the cosine similarities are divided. Must be positive. + :raises ValueError: If `temperature` is not positive. """ 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], y: torch.Tensor) -> torch.Tensor: + 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 labeled 1 is treated as an anchor, its paired second text as the positive, and all other second texts in the batch as negatives. Pairs labeled 0 are not used as anchors, but - their second texts still serve as negatives for the other anchors. Second texts with the same - embedding as an anchor's positive, such as duplicates of the positive text, are not used as - negatives for that anchor. + their second texts still serve as negatives for the other anchors. A second text is not used as a + negative for an anchor if it is identical to the anchor's positive, or if it belongs to another + pair labeled 1 with an identical first text. - :param head_out: The encoded first texts and second texts. + :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: The label of each pair. :return: The mean loss over the anchors. """ - out_a, out_b = head_out + out_a, out_b, ids_a, ids_b = head_out + positive = y == 1 out_a = torch.nn.functional.normalize(out_a, dim=1) out_b = torch.nn.functional.normalize(out_b, dim=1) logits = (out_a @ out_b.T) / self.temperature targets = torch.arange(len(logits), device=logits.device) - duplicates = (out_b @ out_b.T) > 1 - 1e-6 - duplicates[targets, targets] = False - logits = logits.masked_fill(duplicates, float("-inf")) + same_positive = ids_b[:, None] == ids_b[None, :] + other_positive = (ids_a[:, None] == ids_a[None, :]) & positive[None, :] + false_negatives = same_positive | other_positive + false_negatives[targets, targets] = False + logits = logits.masked_fill(false_negatives, float("-inf")) loss = torch.nn.functional.cross_entropy(logits, targets, reduction="none") - positive = y == 1 if not positive.any(): return (loss * 0).sum() return loss[positive].mean() @@ -102,15 +110,20 @@ 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, @@ -143,6 +156,23 @@ def _check_pair_val_split( 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 + @staticmethod + def _check_pair_splits(train_labels: list[int], val_labels: list[int]) -> None: + """Check that the training and validation sets each have at least two pairs, and a pair labeled 1. + + :param train_labels: The labels of the training pairs. + :param val_labels: The labels of the validation pairs. + :raises ValueError: If either set has fewer than two pairs, or no pair labeled 1. + """ + for name, split_labels in (("training", train_labels), ("validation", val_labels)): + if len(split_labels) < 2: + raise ValueError( + f"The {name} set needs at least two pairs, got {len(split_labels)}. Pass more pairs, " + "a different test_size, or an explicit validation set." + ) + if not any(label == 1 for label in split_labels): + raise ValueError(f"The {name} set needs at least one pair labeled 1.") + def _prepare_pair_dataset( self, text_a: list[str], text_b: list[str], labels: list[int], max_length: int | None ) -> PairDataset: @@ -214,15 +244,18 @@ def fit( :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 ) + self._check_pair_splits(train_labels, val_labels) self._initialize() logger.info("Preparing train dataset.") @@ -231,9 +264,11 @@ def fit( val_dataset = self._prepare_pair_dataset(val_a, val_b, val_labels, 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}.") self._train( - loss_function=PairInfoNCELoss(temperature=temperature), + loss_function=loss_function, learning_rate=learning_rate, train_dataset=train_dataset, val_dataset=val_dataset, diff --git a/tests/conftest.py b/tests/conftest.py index e1fd172..990c80e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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 diff --git a/tests/test_trainable.py b/tests/test_trainable.py index 3010361..e0efa1e 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -420,6 +420,10 @@ def test_pairdataset_labels_mismatched_length() -> None: PairDataset([[1], [2]], [[3], [4]], labels=[1]) +def _distinct(n: int) -> torch.Tensor: + return torch.arange(n) + + def test_pair_infonce_loss_is_query_to_document() -> None: """The loss is the cross-entropy of each first text over all second texts in the batch.""" torch.manual_seed(0) @@ -428,7 +432,8 @@ def test_pair_infonce_loss_is_query_to_document() -> None: logits = torch.nn.functional.normalize(out_a, dim=1) @ torch.nn.functional.normalize(out_b, dim=1).T / 0.1 expected = torch.nn.functional.cross_entropy(logits, torch.arange(4)) - assert loss_fn((out_a, out_b), torch.ones(4)).item() == pytest.approx(expected.item(), abs=1e-5) + loss = loss_fn((out_a, out_b, _distinct(4), _distinct(4)), torch.ones(4)) + assert loss.item() == pytest.approx(expected.item(), abs=1e-5) def test_pair_infonce_loss_ignores_negative_anchors() -> None: @@ -439,9 +444,9 @@ def test_pair_infonce_loss_ignores_negative_anchors() -> None: logits = torch.nn.functional.normalize(out_a, dim=1) @ torch.nn.functional.normalize(out_b, dim=1).T / 0.1 per_pair = torch.nn.functional.cross_entropy(logits, torch.arange(3), reduction="none") - loss = loss_fn((out_a, out_b), torch.tensor([1.0, 0.0, 1.0])) + loss = loss_fn((out_a, out_b, _distinct(3), _distinct(3)), torch.tensor([1.0, 0.0, 1.0])) assert loss.item() == pytest.approx(per_pair[[0, 2]].mean().item(), abs=1e-5) - assert loss_fn((out_a, out_b), torch.zeros(3)).item() == 0.0 + assert loss_fn((out_a, out_b, _distinct(3), _distinct(3)), torch.zeros(3)).item() == 0.0 def test_pair_infonce_loss_masks_duplicate_positives() -> None: @@ -450,10 +455,47 @@ def test_pair_infonce_loss_masks_duplicate_positives() -> None: out_b = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) loss_fn = PairInfoNCELoss(temperature=0.05) - loss = loss_fn((out_a, out_b), torch.tensor([1.0, 0.0])) + loss = loss_fn((out_a, out_b, _distinct(2), torch.tensor([0, 0])), torch.tensor([1.0, 0.0])) + assert loss.item() == pytest.approx(0.0, abs=1e-6) + + +def test_pair_infonce_loss_masks_other_positives_of_the_same_anchor() -> None: + """Pairs labeled 1 with an identical first text are all positives, so they don't compete.""" + out_a = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) + out_b = torch.tensor([[0.9, 0.1], [0.8, 0.2]]) + loss_fn = PairInfoNCELoss(temperature=0.05) + + loss = loss_fn((out_a, out_b, torch.tensor([0, 0]), _distinct(2)), torch.ones(2)) assert loss.item() == pytest.approx(0.0, abs=1e-6) +def test_pair_infonce_loss_keeps_negatives_of_the_same_anchor() -> None: + """A pair labeled 0 with an identical first text still provides a negative for that anchor.""" + out_a = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) + out_b = torch.tensor([[0.9, 0.1], [0.8, 0.2]]) + loss_fn = PairInfoNCELoss(temperature=0.05) + + loss = loss_fn((out_a, out_b, torch.tensor([0, 0]), _distinct(2)), torch.tensor([1.0, 0.0])) + assert loss.item() > 0.1 + + +def test_pair_infonce_loss_keeps_distinct_but_parallel_negatives() -> None: + """Distinct second texts are negatives, even if their embeddings are nearly parallel.""" + out_a = torch.tensor([[1.0, 0.0], [0.0, 1.0]]) + out_b = torch.tensor([[1.0, 0.0], [1.0, 1e-4]]) + loss_fn = PairInfoNCELoss(temperature=0.05) + + loss = loss_fn((out_a, out_b, _distinct(2), _distinct(2)), torch.tensor([1.0, 0.0])) + assert loss.item() == pytest.approx(float(np.log(2)), abs=1e-3) + + +@pytest.mark.parametrize("temperature", [0.0, -0.05]) +def test_pair_infonce_loss_rejects_non_positive_temperature(temperature: float) -> None: + """The temperature must be positive.""" + with pytest.raises(ValueError): + PairInfoNCELoss(temperature=temperature) + + def test_pair_similarity_out_dim_defaults_to_embed_dim(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: """The output dimension defaults to the input embedding dimension when not specified.""" s = StaticModelForPairSimilarity(vectors=torch.from_numpy(mock_vectors).float(), tokenizer=mock_tokenizer) @@ -472,9 +514,11 @@ def test_pair_similarity_forward(mock_trained_pair_similarity_pipeline: StaticMo batch, _ = next(iter(dataset.to_dataloader(shuffle=False, batch_size=2))) with torch.no_grad(): - out_a, out_b = model(batch) + out_a, out_b, ids_a, ids_b = model(batch) assert out_a.shape == (2, model.out_dim) assert out_b.shape == (2, model.out_dim) + assert ids_a.tolist()[0] != ids_a.tolist()[1] + assert ids_b.tolist()[0] != ids_b.tolist()[1] def test_pair_similarity_mismatched_lengths( @@ -523,14 +567,64 @@ def test_pair_similarity_fit_with_explicit_val(mock_vectors: np.ndarray, mock_to text_a, text_b, labels=labels, - text_a_val=["word1"], - text_b_val=["word2"], - labels_val=[1], + text_a_val=["word1", "word3"], + text_b_val=["word2", "word1"], + labels_val=[1, 1], early_stopping_patience=1, max_epochs=1, ) +def test_pair_similarity_fit_rejects_splits_it_cannot_learn_from( + mock_vectors: np.ndarray, mock_tokenizer: Tokenizer +) -> None: + """Both splits need at least two pairs and a pair labeled 1.""" + model = StaticModelForPairSimilarity(vectors=torch.from_numpy(mock_vectors).float(), tokenizer=mock_tokenizer) + with pytest.raises(ValueError, match="training set needs at least two pairs"): + model.fit(["word1", "word2"], ["word2", "word3"], labels=[1, 0]) + with pytest.raises(ValueError, match="validation set needs at least one pair labeled 1"): + model.fit( + ["word1", "word2"], + ["word2", "word3"], + text_a_val=["word1", "word3"], + text_b_val=["word3", "word2"], + labels_val=[0, 0], + ) + + +def test_pair_similarity_fit_rejects_invalid_temperature_and_batch_size( + mock_vectors: np.ndarray, mock_tokenizer: Tokenizer +) -> None: + """A non-positive temperature and a batch size of 1 are rejected.""" + model = StaticModelForPairSimilarity(vectors=torch.from_numpy(mock_vectors).float(), tokenizer=mock_tokenizer) + text_a, text_b = ["word1", "word2", "word3", "word1 word2"], ["word2", "word3", "word1", "word3"] + with pytest.raises(ValueError, match="temperature"): + model.fit(text_a, text_b, test_size=0.5, temperature=0.0) + with pytest.raises(ValueError, match="batch_size"): + model.fit(text_a, text_b, test_size=0.5, batch_size=1) + + +def test_pairdataset_drops_single_pair_batches() -> None: + """A final batch with a single pair is dropped, unless it is the only pair.""" + dataset = PairDataset([[1], [2], [3]], [[1], [2], [3]]) + assert [len(y) for _, y in dataset.to_dataloader(shuffle=False, batch_size=2)] == [2] + single = PairDataset([[1]], [[1]]) + assert [len(y) for _, y in single.to_dataloader(shuffle=False, batch_size=2)] == [1] + + +def test_pair_similarity_forward_ids_identical_texts( + mock_trained_pair_similarity_pipeline: StaticModelForPairSimilarity, +) -> None: + """Identical texts in a batch get the same id.""" + model = mock_trained_pair_similarity_pipeline + dataset = model._prepare_pair_dataset(["dog", "cat", "dog"], ["puppy", "puppy", "kitten"], [1, 1, 1], None) + batch, _ = next(iter(dataset.to_dataloader(shuffle=False, batch_size=3))) + with torch.no_grad(): + _, _, ids_a, ids_b = model(batch) + assert ids_a[0] == ids_a[2] and ids_a[0] != ids_a[1] + assert ids_b[0] == ids_b[1] and ids_b[0] != ids_b[2] + + def test_pair_similarity_fit_with_labels(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: """A model can be fit with a mix of positive and negative pair labels.""" model = StaticModelForPairSimilarity(vectors=torch.from_numpy(mock_vectors).float(), tokenizer=mock_tokenizer) From 2bbab4652373d0dac16060ac8100e50cc4121a8f Mon Sep 17 00:00:00 2001 From: stephantul Date: Mon, 28 Sep 2026 07:58:57 +0200 Subject: [PATCH 3/3] remove labels --- model2vec/train/README.md | 6 --- model2vec/train/dataset.py | 22 +++----- model2vec/train/pairs.py | 86 ++++++++++--------------------- tests/test_trainable.py | 100 +++++-------------------------------- 4 files changed, 47 insertions(+), 167 deletions(-) diff --git a/model2vec/train/README.md b/model2vec/train/README.md index 5c95ae7..63ef8b2 100644 --- a/model2vec/train/README.md +++ b/model2vec/train/README.md @@ -111,12 +111,6 @@ model.fit(text_a=queries, text_b=documents) 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. -Pairs can also be labeled. If `labels` is omitted, every pair is treated as positive. Pairs labeled `0` are not used as anchors, but their `text_b` still serves as an additional negative for the other pairs in the batch, including pairs with the same `text_a`: - -```python -model.fit(text_a=queries, text_b=documents, labels=labels) -``` - The InfoNCE temperature can be set with `temperature` (default `0.05`). It must be positive. # Persistence diff --git a/model2vec/train/dataset.py b/model2vec/train/dataset.py index 65a5de3..ec760a8 100644 --- a/model2vec/train/dataset.py +++ b/model2vec/train/dataset.py @@ -45,49 +45,43 @@ 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 otherwise. - 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. A final batch with a single pair is dropped, unless it is the only pair.""" diff --git a/model2vec/train/pairs.py b/model2vec/train/pairs.py index e59087d..d788625 100644 --- a/model2vec/train/pairs.py +++ b/model2vec/train/pairs.py @@ -32,32 +32,23 @@ def __call__( ) -> torch.Tensor: """Returns the InfoNCE loss over a pair batch, using in-batch negatives. - Every first text labeled 1 is treated as an anchor, its paired second text as the positive, and - all other second texts in the batch as negatives. Pairs labeled 0 are not used as anchors, but - their second texts still serve as negatives for the other anchors. A second text is not used as a - negative for an anchor if it is identical to the anchor's positive, or if it belongs to another - pair labeled 1 with an identical first text. + 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: The label of each pair. + :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 - positive = y == 1 out_a = torch.nn.functional.normalize(out_a, dim=1) out_b = torch.nn.functional.normalize(out_b, dim=1) logits = (out_a @ out_b.T) / self.temperature - targets = torch.arange(len(logits), device=logits.device) - same_positive = ids_b[:, None] == ids_b[None, :] - other_positive = (ids_a[:, None] == ids_a[None, :]) & positive[None, :] - false_negatives = same_positive | other_positive - false_negatives[targets, targets] = False + 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")) - loss = torch.nn.functional.cross_entropy(logits, targets, reduction="none") - if not positive.any(): - return (loss * 0).sum() - return loss[positive].mean() + return torch.nn.functional.cross_entropy(logits, y) class StaticModelForPairSimilarity(BaseFinetuneable): @@ -129,65 +120,52 @@ 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(train_labels: list[int], val_labels: list[int]) -> None: - """Check that the training and validation sets each have at least two pairs, and a pair labeled 1. + def _check_pair_splits(n_train: int, n_val: int) -> None: + """Check that the training and validation sets each have at least two pairs. - :param train_labels: The labels of the training pairs. - :param val_labels: The labels of the validation pairs. - :raises ValueError: If either set has fewer than two pairs, or no pair labeled 1. + :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, split_labels in (("training", train_labels), ("validation", val_labels)): - if len(split_labels) < 2: + 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 {len(split_labels)}. Pass more pairs, " + f"The {name} set needs at least two pairs, got {n_pairs}. Pass more pairs, " "a different test_size, or an explicit validation set." ) - if not any(label == 1 for label in split_labels): - raise ValueError(f"The {name} set needs at least one pair labeled 1.") - def _prepare_pair_dataset( - self, text_a: list[str], text_b: list[str], labels: list[int], max_length: int | None - ) -> PairDataset: + 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, ) @@ -195,7 +173,6 @@ 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, @@ -205,7 +182,6 @@ 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, @@ -213,10 +189,9 @@ def fit( """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, and trained with an InfoNCE loss: each `text_a` labeled 1 is - pulled towards its paired `text_b` and pushed away from all other `text_b` in the batch. Pairs - labeled 0 are not used as anchors, but their `text_b` still serves as an in-batch negative. 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. @@ -226,8 +201,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 otherwise. - 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. @@ -239,7 +212,6 @@ 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. @@ -250,18 +222,14 @@ def fit( 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 - ) - self._check_pair_splits(train_labels, val_labels) + 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: diff --git a/tests/test_trainable.py b/tests/test_trainable.py index e0efa1e..947e640 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -396,30 +396,11 @@ def test_pairdataset_collate() -> None: dataset = PairDataset([[1], [1, 2]], [[1, 2, 3], [1]], pad_id=0) batch, y = next(iter(dataset.to_dataloader(shuffle=False, batch_size=2))) assert batch.shape == (2, 2, 3) - assert y.shape == (2,) + assert torch.equal(y, torch.tensor([0, 1])) assert torch.equal(batch[0], torch.tensor([[1, 0, 0], [1, 2, 0]])) assert torch.equal(batch[1], torch.tensor([[1, 2, 3], [1, 0, 0]])) -def test_pairdataset_default_labels_are_positive() -> None: - """Without explicit labels, every pair defaults to label 1.""" - dataset = PairDataset([[1], [2]], [[3], [4]]) - assert torch.equal(dataset.labels, torch.tensor([1.0, 1.0])) - - -def test_pairdataset_custom_labels() -> None: - """Custom labels are stored and returned by the collate function.""" - dataset = PairDataset([[1], [2]], [[3], [4]], labels=[1, 0]) - _, y = next(iter(dataset.to_dataloader(shuffle=False, batch_size=2))) - assert torch.equal(y, torch.tensor([1.0, 0.0])) - - -def test_pairdataset_labels_mismatched_length() -> None: - """Labels must have one entry per pair.""" - with pytest.raises(ValueError): - PairDataset([[1], [2]], [[3], [4]], labels=[1]) - - def _distinct(n: int) -> torch.Tensor: return torch.arange(n) @@ -432,60 +413,37 @@ def test_pair_infonce_loss_is_query_to_document() -> None: logits = torch.nn.functional.normalize(out_a, dim=1) @ torch.nn.functional.normalize(out_b, dim=1).T / 0.1 expected = torch.nn.functional.cross_entropy(logits, torch.arange(4)) - loss = loss_fn((out_a, out_b, _distinct(4), _distinct(4)), torch.ones(4)) + loss = loss_fn((out_a, out_b, _distinct(4), _distinct(4)), torch.arange(4)) assert loss.item() == pytest.approx(expected.item(), abs=1e-5) -def test_pair_infonce_loss_ignores_negative_anchors() -> None: - """Pairs labeled 0 are not used as anchors, but their second text is still a negative.""" - torch.manual_seed(0) - out_a, out_b = torch.randn(3, 3), torch.randn(3, 3) - loss_fn = PairInfoNCELoss(temperature=0.1) - logits = torch.nn.functional.normalize(out_a, dim=1) @ torch.nn.functional.normalize(out_b, dim=1).T / 0.1 - per_pair = torch.nn.functional.cross_entropy(logits, torch.arange(3), reduction="none") - - loss = loss_fn((out_a, out_b, _distinct(3), _distinct(3)), torch.tensor([1.0, 0.0, 1.0])) - assert loss.item() == pytest.approx(per_pair[[0, 2]].mean().item(), abs=1e-5) - assert loss_fn((out_a, out_b, _distinct(3), _distinct(3)), torch.zeros(3)).item() == 0.0 - - def test_pair_infonce_loss_masks_duplicate_positives() -> None: """Second texts identical to an anchor's positive are not used as negatives for that anchor.""" out_a = torch.tensor([[1.0, 0.0], [0.0, 1.0]]) out_b = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) loss_fn = PairInfoNCELoss(temperature=0.05) - loss = loss_fn((out_a, out_b, _distinct(2), torch.tensor([0, 0])), torch.tensor([1.0, 0.0])) + loss = loss_fn((out_a, out_b, _distinct(2), torch.tensor([0, 0])), torch.arange(2)) assert loss.item() == pytest.approx(0.0, abs=1e-6) def test_pair_infonce_loss_masks_other_positives_of_the_same_anchor() -> None: - """Pairs labeled 1 with an identical first text are all positives, so they don't compete.""" + """Pairs with an identical first text are all positives for it, so they don't compete.""" out_a = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) out_b = torch.tensor([[0.9, 0.1], [0.8, 0.2]]) loss_fn = PairInfoNCELoss(temperature=0.05) - loss = loss_fn((out_a, out_b, torch.tensor([0, 0]), _distinct(2)), torch.ones(2)) + loss = loss_fn((out_a, out_b, torch.tensor([0, 0]), _distinct(2)), torch.arange(2)) assert loss.item() == pytest.approx(0.0, abs=1e-6) -def test_pair_infonce_loss_keeps_negatives_of_the_same_anchor() -> None: - """A pair labeled 0 with an identical first text still provides a negative for that anchor.""" - out_a = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) - out_b = torch.tensor([[0.9, 0.1], [0.8, 0.2]]) - loss_fn = PairInfoNCELoss(temperature=0.05) - - loss = loss_fn((out_a, out_b, torch.tensor([0, 0]), _distinct(2)), torch.tensor([1.0, 0.0])) - assert loss.item() > 0.1 - - def test_pair_infonce_loss_keeps_distinct_but_parallel_negatives() -> None: """Distinct second texts are negatives, even if their embeddings are nearly parallel.""" out_a = torch.tensor([[1.0, 0.0], [0.0, 1.0]]) out_b = torch.tensor([[1.0, 0.0], [1.0, 1e-4]]) loss_fn = PairInfoNCELoss(temperature=0.05) - loss = loss_fn((out_a, out_b, _distinct(2), _distinct(2)), torch.tensor([1.0, 0.0])) + loss = loss_fn((out_a, out_b, _distinct(2), _distinct(2)), torch.arange(2)) assert loss.item() == pytest.approx(float(np.log(2)), abs=1e-3) @@ -510,7 +468,7 @@ def test_pair_similarity_out_dim_defaults_to_embed_dim(mock_vectors: np.ndarray, def test_pair_similarity_forward(mock_trained_pair_similarity_pipeline: StaticModelForPairSimilarity) -> None: """The forward pass should return one head output per half of the pair batch.""" model = mock_trained_pair_similarity_pipeline - dataset = model._prepare_pair_dataset(["dog cat", "dog"], ["puppy", "kitten cat"], [1, 1], max_length=None) + dataset = model._prepare_pair_dataset(["dog cat", "dog"], ["puppy", "kitten cat"], max_length=None) batch, _ = next(iter(dataset.to_dataloader(shuffle=False, batch_size=2))) with torch.no_grad(): @@ -541,35 +499,16 @@ def test_pair_similarity_val_split_errors(mock_trained_pair_similarity_pipeline: ) -def test_pair_similarity_labels_mismatched_length( - mock_trained_pair_similarity_pipeline: StaticModelForPairSimilarity, -) -> None: - """Labels must have one entry per training pair, and labels_val one entry per validation pair.""" - with pytest.raises(ValueError): - mock_trained_pair_similarity_pipeline.fit(["dog", "cat"], ["puppy", "kitten"], labels=[1]) - with pytest.raises(ValueError): - mock_trained_pair_similarity_pipeline.fit( - ["dog", "cat"], - ["puppy", "kitten"], - text_a_val=["dog"], - text_b_val=["puppy"], - labels_val=[1, 0], - ) - - def test_pair_similarity_fit_with_explicit_val(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: """A model can be fit with explicit validation pairs instead of an automatic split.""" model = StaticModelForPairSimilarity(vectors=torch.from_numpy(mock_vectors).float(), tokenizer=mock_tokenizer) text_a = ["word1", "word2", "word3", "word1 word2"] text_b = ["word2", "word3", "word1", "word3 word1"] - labels = [1, 1, 0, 0] model.fit( text_a, text_b, - labels=labels, text_a_val=["word1", "word3"], text_b_val=["word2", "word1"], - labels_val=[1, 1], early_stopping_patience=1, max_epochs=1, ) @@ -578,18 +517,12 @@ def test_pair_similarity_fit_with_explicit_val(mock_vectors: np.ndarray, mock_to def test_pair_similarity_fit_rejects_splits_it_cannot_learn_from( mock_vectors: np.ndarray, mock_tokenizer: Tokenizer ) -> None: - """Both splits need at least two pairs and a pair labeled 1.""" + """Both splits need at least two pairs.""" model = StaticModelForPairSimilarity(vectors=torch.from_numpy(mock_vectors).float(), tokenizer=mock_tokenizer) with pytest.raises(ValueError, match="training set needs at least two pairs"): - model.fit(["word1", "word2"], ["word2", "word3"], labels=[1, 0]) - with pytest.raises(ValueError, match="validation set needs at least one pair labeled 1"): - model.fit( - ["word1", "word2"], - ["word2", "word3"], - text_a_val=["word1", "word3"], - text_b_val=["word3", "word2"], - labels_val=[0, 0], - ) + model.fit(["word1", "word2"], ["word2", "word3"]) + with pytest.raises(ValueError, match="validation set needs at least two pairs"): + model.fit(["word1", "word2"], ["word2", "word3"], text_a_val=["word1"], text_b_val=["word3"]) def test_pair_similarity_fit_rejects_invalid_temperature_and_batch_size( @@ -617,7 +550,7 @@ def test_pair_similarity_forward_ids_identical_texts( ) -> None: """Identical texts in a batch get the same id.""" model = mock_trained_pair_similarity_pipeline - dataset = model._prepare_pair_dataset(["dog", "cat", "dog"], ["puppy", "puppy", "kitten"], [1, 1, 1], None) + dataset = model._prepare_pair_dataset(["dog", "cat", "dog"], ["puppy", "puppy", "kitten"], None) batch, _ = next(iter(dataset.to_dataloader(shuffle=False, batch_size=3))) with torch.no_grad(): _, _, ids_a, ids_b = model(batch) @@ -625,15 +558,6 @@ def test_pair_similarity_forward_ids_identical_texts( assert ids_b[0] == ids_b[1] and ids_b[0] != ids_b[2] -def test_pair_similarity_fit_with_labels(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: - """A model can be fit with a mix of positive and negative pair labels.""" - model = StaticModelForPairSimilarity(vectors=torch.from_numpy(mock_vectors).float(), tokenizer=mock_tokenizer) - text_a = ["word1", "word2", "word3", "word1 word2"] - text_b = ["word2", "word3", "word1", "word3 word1"] - labels = [1, 1, 0, 0] - model.fit(text_a, text_b, labels=labels, early_stopping_patience=1, max_epochs=1) - - def test_convert_to_pipeline_pair_similarity( mock_trained_pair_similarity_pipeline: StaticModelForPairSimilarity, ) -> None: