nemo_automodel.recipes.retrieval.train_bi_encoder
nemo_automodel.recipes.retrieval.train_bi_encoder
Module Contents
Classes
Functions
Data
API
Bases: BaseRecipe
Recipe for training encoder models with contrastive learning.
Build optimizer parameter groups for trainable model parameters.
Return the embedding tensor used for validation scoring from a forward output.
The base bi-encoder forward returns an embedding tensor directly. Subclasses whose
forward returns a richer structure (e.g. the distillation student, which returns
(pooled, projected, intermediate_outputs)) should override this to select the
tensor to score with.
Forward and backward pass for a single micro-batch.
Run one optimization step with gradient accumulation.
Run validation for one epoch and compute loss, accuracy@1, and MRR.
Validate recipe-specific model settings before constructing the optimizer.
Parameters:
Constructed retrieval model, including infrastructure wrappers.
Run the training loop over all epochs and batches.
Build all components needed for training/validation/logging/checkpointing.
Bind static prompts and validate export against the runtime tokenizer or processor.
Return the optional recipe-level autocast context.
Return infrastructure kwargs forwarded to model instantiation.
Unpack query and passage inputs from batch dictionary.
Parameters:
Dictionary containing query (q_) and passage (d_) tensors
Returns: tuple
Tuple of (query_batch_dict, doc_batch_dict)
Return the underlying model object for configuration-style attribute reads.
Return whether the model emits token-level embeddings for MaxSim scoring.
Compute contrastive scores and labels without in-batch negatives.
Parameters:
Query embeddings [batch_size, hidden_dim]
Key/passage embeddings [batch_size * n_passages, hidden_dim]
Number of passages per query
Returns: torch.Tensor
Tuple of (scores, labels) where scores is [batch_size, n_passages]
Compute local-query multi-vector MaxSim scores against globally gathered passages.
Compute local multi-vector MaxSim scores and labels without in-batch negatives.