feat: train pair models with InfoNCE - #381
stephantul wants to merge 3 commits into
Conversation
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.
Codecov Report✅ All modified and coverable lines are covered by tests.
🚀 New features to boost your workflow:
|
|
|
/greptile review |
| """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.
Sole positive pair gets dropped When a validation set has three pairs labeled
[0, 0, 1] and batch_size=2, this drops the only positive pair. The remaining negative-only batch reports zero loss, so checkpoint selection and early stopping never assess a positive match. Shuffling can also drop the only positive training pair for an epoch.
| if batch_size < 2: | ||
| raise ValueError(f"batch_size must be at least 2, got {batch_size}.") |
|
|
||
| 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) |
There was a problem hiding this comment.
Example uses undefined variables The example calls
model.fit with queries and documents but never defines them. Readers cannot run the snippet as shown; sample lists or clearly marked placeholders would make it usable.
Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!
This PR replaces 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.