-
Notifications
You must be signed in to change notification settings - Fork 128
feat: train pair models with InfoNCE #381
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| return DataLoader(self, collate_fn=self.collate_fn, shuffle=shuffle, batch_size=batch_size, drop_last=drop_last) | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 | ||
|
stephantul marked this conversation as resolved.
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): | ||
|
|
@@ -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, | ||
|
|
@@ -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}.") | ||
|
Comment on lines
+236
to
+237
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. |
||
|
|
||
| self._train( | ||
| loss_function=PairCosineLoss(), | ||
| loss_function=loss_function, | ||
| learning_rate=learning_rate, | ||
| train_dataset=train_dataset, | ||
| val_dataset=val_dataset, | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
model.fitwithqueriesanddocumentsbut 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!