Multi-Vector-Encoder Losses (ColBERT / late-interaction)
All losses live in sentence_transformers.multi_vector_encoder.losses.
Multi-vector models emit one embedding per token and score (query, document) via MaxSim (for each query token, take the max similarity to any document token, then sum across query tokens). This is fundamentally different from single-vector cosine, which changes what "temperature" means: MaxSim is an unbounded sum over query-token similarities, so for it a scale near 1.0 makes sense. The default scale=1.0 (temperature=1.0) is correct for unnormalized MaxSim: do not copy scale=20.0 from bi-encoder MNRL. With MeanMaxSim scoring (similarity_fct=mean_colbert_scores, and set model.similarity_fn_name = "meanmaxsim" so evaluation matches), each score is divided by its query's token count, so a scale of roughly the average length is a reasonable start, although some works use larger scales.
Top-line decision table
| You have | Use |
|---|---|
(anchor, positive) or (anchor, positive, negative) triplets |
MultiVectorMultipleNegativesRankingLoss |
| Same, want effective batch size of 128+ | CachedMultiVectorMultipleNegativesRankingLoss |
Cross-encoder teacher scores, (query, positive, negative, score_diff) |
MultiVectorMarginMSELoss |
Listwise distillation (query, [doc_1..doc_N], scores) |
MultiVectorDistillKLDivLoss |
Hard-negative mining is essential for competitive results, because random in-batch negatives leave a lot on the table for late-interaction models. See dataset_formats.md (Hard-negative mining section) and ../scripts/mine_hard_negatives.py.
Contrastive losses
MultiVectorMultipleNegativesRankingLoss
The default late-interaction contrastive loss. In-batch positives plus explicit hard negatives, scored with MaxSim (or XTR).
from sentence_transformers.multi_vector_encoder.losses import MultiVectorMultipleNegativesRankingLoss
loss = MultiVectorMultipleNegativesRankingLoss(model=model) # scale=1.0 (temperature=1.0) is the correct default- Data:
(anchor, positive)or(anchor, positive, negative_1, ..., negative_n). The collator stamps each column's task (column 0 becomes query, others become document). Passtask=...on a column to override. - Set
batch_sampler=BatchSamplers.NO_DUPLICATESon training args (same reason as bi-encoder MNRL). similarity_fct=colbert_scoresby default. PassXTRScores(top_k=...)for the sparser XTR-style scoring (fromsentence_transformers.multi_vector_encoder.scoring). XTR applies to training only: evaluation always scores with MaxSim (see Gotchas).scale=1.0(temperature=1.0) matches PyLate and is correct for unnormalized MaxSim. Do NOT copyscale=20.0from bi-encoder MNRL. With MeanMaxSim (similarity_fct=mean_colbert_scores), raise the scale to roughly the average query length instead, which is the same objective at the same strength.
CachedMultiVectorMultipleNegativesRankingLoss
GradCache variant: chunked embedding forward, cached gradients, and a second re-embedding pass. Decouples per-device batch size from effective in-batch negatives.
loss = CachedMultiVectorMultipleNegativesRankingLoss(
model=model,
mini_batch_size=8, # or use mini_batch_num_tokens=... for token-budgeted packing
score_mini_batch_size=4, # optional, smaller trims the transient (Q, Q*N, q_tok, d_tok) buffer
)- Incompatible with
gradient_checkpointing=True(same as everyCached*loss). mini_batch_num_tokens=Npacks sequences until N real tokens per mini-batch instead of a fixedmini_batch_size. Big win on variable-length data with flash-attention / input flattening.score_mini_batch_sizechunks the SCORING phase (which builds(Q, Q*N, q_tokens, d_tokens)intermediates) independently from the embedding phase. Drop it first when hitting OOM in the loss stage.chunk_elementsis the second lever on that same phase, and it cuts the other axis:score_mini_batch_sizechunks queries,chunk_elementschunks documents inside each MaxSim call. Pass it through the scorer, e.g.similarity_fct=partial(colbert_scores, chunk_elements=1_000_000). Scores and gradients are unchanged (it is a pure memory knob), and it composes withscore_mini_batch_size.gather_across_devices=Truegathers document embeddings across DDP ranks. Enables cross-rank in-batch negatives.
Distillation losses
Distillation is where multi-vector models learn most efficiently: cross-encoder teachers (e.g. gte-modernbert-base) score (query, doc) pairs offline, and the student MaxSim model regresses to that signal.
For the standard MS MARCO KD format (LightOn's ms-marco-en-bge and similar: (query_id, document_ids, scores) with separate queries / documents splits), you can use sentence_transformers.util.resolve_ids to join the IDs against the text splits at load time. It's a factory that returns a batched transform: pass a mapping from each input ID column to its lookup dataset, e.g. dataset.set_transform(resolve_ids({"query_id": queries, "document_ids": documents}, max_list_length=32)) (lazy, no caching) or via dataset.map(..., batched=True, remove_columns=["query_id", "document_ids"]) (eager, cached: map keeps input columns unless told to drop them, so pass the ID columns to remove_columns), yielding (query, document_1, ..., document_32, scores) rows ready for MultiVectorDistillKLDivLoss. The same mapping shape covers triplets with IDs ({"query_id": queries, "positive_id": documents, "negative_id": documents}) and ID-only contrastive datasets without a scores column.
MultiVectorMarginMSELoss
Regress the margin between positive and negative MaxSim scores against the teacher's margin.
loss = MultiVectorMarginMSELoss(model=model)- Data:
(query, positive, negative, score_diff)wherescore_diff = teacher_score(query, positive) - teacher_score(query, negative). Raw teacher scores[score(q, pos), score(q, neg_1), ...]also work, converted to margins internally. - Popular recipe from PyLate / colpali-engine.
- Teacher scores are precomputed once, stored as the label column. The loss does not run the teacher inline.
- Same
similarity_fctconvention as the other losses, but pairwise-shaped: defaults tocolbert_scores_pairwise, or passxtr_scores_pairwisefor XTR.
MultiVectorDistillKLDivLoss
Listwise KL-div: student's softmax distribution over N candidates should match the teacher's.
loss = MultiVectorDistillKLDivLoss(model=model)- Data:
(query, document_1, ..., document_N, scores)wherescoresis a list of N teacher scores per row. One column per candidate document (the standard positional multi-column convention), and the label column must use a recognized name (label,labels,score,scores).resolve_idsexpands a stored ID list into the numbered columns. - Stronger training signal than
MarginMSEwhen you have fullN-way teacher scores (not just positive/negative margins). similarity_fctdefaults tocolbert_kd_scores, the listwise KD variant, not thecolbert_scoresused by the MNRL family.XTRKDScoresis the XTR counterpart.temperaturesoftens both distributions before the softmax, and can be overridden per side withstudent_temperature/teacher_temperature. Setteacher_temperaturefrom the spread of your teacher's scores rather than from a value copied out of a paper: a float32 softmax underflows to exact zeros once a row's spread divided by the temperature exceeds about 100, and every candidate that underflows drops out of the KL, which is the ranking information you are distilling. Sharpeningstudent_temperaturefar below the student's own score spread collapses that side too, and once both are one-hot the loss and its gradient are exactly zero. Start at 1.0 / 1.0 and only move a side you have a measured reason to move.- The loss is scaled by the student temperature squared, which keeps gradient magnitudes comparable only for a shared temperature at or above 1.0 (the regime Hinton et al. cover). Below that the factor shrinks the loss instead, by 1e-6 at
student_temperature=0.001: when combining KD with a contrastive loss, scale the KD weight back up accordingly. - OOM: drop
per_device_train_batch_sizefirst and raisegradient_accumulation_stepsto hold the effective batch. Only reduceN(candidate-list length) as a last resort, since lowering N changes the experiment.
Data-shape gotchas
- Column 0 is always the query. Losses call
self.model(sf, task="query" if idx==0 else "document"). Use the collator'srouter_mappingto override if you need a non-standard column layout. - Cross-column varlen
Tis handled: the batch's positive and negative columns can have different token counts (Qwen2-VL family etc.).stack_padded_token_embeddingspads to the per-batch max before MaxSim. - Mask flow: MVE reads the SCORING mask from the model OUTPUT dict (
MultiVectorMaskrewrites the input mask). Custom loss code that readssentence_features[i]["attention_mask"]directly instead ofoutputs[i]["attention_mask"]will silently score against skiplisted tokens.
Gotchas
scale=20.0copied from bi-encoder MNRL: misscaled for the unnormalized MaxSim sum. Keepscale=1.0there. Only length-normalized MeanMaxSim wants a larger scale, roughly the average query length.- Missing
Normalizeat the token level in the pipeline:colbert_scoresassumes L2-normalized token embeddings. If your custom pipeline drops the token-levelNormalize, either add one or passnormalize_embeddings=Truesemantics via a wrapper. CachedMultiVectorMultipleNegativesRankingLoss+gradient_checkpointing=True: crash. Pick one.MultiVectorMarginMSELosswithout precomputedscore_diff: label column must be populated from a teacher pass ahead of training. The loss does not compute the teacher inline.- Expecting XTR scoring at eval:
XTRScoresis a train-onlysimilarity_fct. XTR takes its top-k across the whole candidate set, so a(query, document)pair has no standalone score, which means it cannot be the model'ssimilarity_fn_nameand the evaluators reject it outright. Evaluation and inference score with MaxSim, also for XTR-trained models. That is by design, not a mismatch to fix. To score a fixed candidate set ad hoc, callxtr_scoresdirectly. - Distillation with a weak teacher: multi-vector students easily match a small-model teacher and then plateau. Use a strong cross-encoder (e.g.
gte-modernbert-base,mxbai-rerank-large-v2) for the teacher pass.