diff --git a/model2vec/train/README.md b/model2vec/train/README.md index 98342f4..63ef8b2 100644 --- a/model2vec/train/README.md +++ b/model2vec/train/README.md @@ -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) ``` -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 diff --git a/model2vec/train/dataset.py b/model2vec/train/dataset.py index 0d7627e..ec760a8 100644 --- a/model2vec/train/dataset.py +++ b/model2vec/train/dataset.py @@ -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 + 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 672bf84..9cd8020 100644 --- a/model2vec/train/pairs.py +++ b/model2vec/train/pairs.py @@ -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 + 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): @@ -81,62 +102,71 @@ 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, ) @@ -144,7 +174,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, @@ -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. @@ -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. @@ -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}.") self._train( - loss_function=PairCosineLoss(), + 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 3953d05..69f6be2 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -17,7 +17,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 @@ -421,42 +421,62 @@ 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 _distinct(n: int) -> torch.Tensor: + return torch.arange(n) -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_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)) + 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_pairdataset_labels_mismatched_length() -> None: - """Labels must have one entry per pair.""" - with pytest.raises(ValueError): - PairDataset([[1], [2]], [[3], [4]], labels=[1]) +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.arange(2)) + assert loss.item() == pytest.approx(0.0, abs=1e-6) -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() + +def test_pair_infonce_loss_masks_other_positives_of_the_same_anchor() -> None: + """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.arange(2)) + assert loss.item() == pytest.approx(0.0, abs=1e-6) - identical = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) - orthogonal = torch.tensor([[0.0, 1.0], [0.0, 1.0]]) - 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) +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.arange(2)) + 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: @@ -530,13 +550,15 @@ def test_classifier_keeps_head_when_dimensions_match(mock_vectors: np.ndarray, m 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(): - 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( @@ -559,47 +581,63 @@ 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"], - text_b_val=["word2"], - labels_val=[1], + text_a_val=["word1", "word3"], + text_b_val=["word2", "word1"], early_stopping_patience=1, max_epochs=1, ) -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.""" +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.""" 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) + with pytest.raises(ValueError, match="training set needs at least two pairs"): + 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( + 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"], 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_convert_to_pipeline_pair_similarity(