Skip to content

feat: train pair models with InfoNCE - #381

Open
stephantul wants to merge 3 commits into
mainfrom
pair-infonce
Open

stephantul wants to merge 3 commits into
mainfrom
pair-infonce

Conversation

@stephantul

Copy link
Copy Markdown
Contributor

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.

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.
@stephantul
stephantul marked this pull request as ready for review September 27, 2026 17:52
@stephantul
stephantul requested a review from Pringled September 27, 2026 17:52
@codecov

codecov Bot commented Sep 27, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.

Files with missing lines Coverage Δ
model2vec/train/dataset.py 100.00% <100.00%> (ø)
model2vec/train/pairs.py 100.00% <100.00%> (ø)
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@greptile-apps

greptile-apps Bot commented Sep 27, 2026 •

Copy link
Copy Markdown

RetriggerConfidence Score: 3/5

[Medium risk] Changes pair training loss function and dataset handling.

The PR does not appear safe to merge while batching can silently remove the only positive pair and invalid input can reset an existing model.

Reviews (2) · Last reviewed commit: "address reviewer comments"

Comment thread model2vec/train/pairs.py
Comment thread model2vec/train/pairs.py Outdated
Comment thread model2vec/train/pairs.py
Comment thread model2vec/train/pairs.py Outdated
Comment thread model2vec/train/pairs.py Outdated
Comment thread model2vec/train/pairs.py Outdated
@stephantul

Copy link
Copy Markdown
Contributor Author

/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

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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

Comment thread model2vec/train/pairs.py
Comment on lines +267 to +268
if batch_size < 2:
raise ValueError(f"batch_size must be at least 2, got {batch_size}.")

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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

Comment thread model2vec/train/README.md

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

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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

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

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant