nemo_automodel.recipes.retrieval.train_bi_encoder

View as Markdown

Module Contents

Classes

NameDescription
TrainBiEncoderRecipeRecipe for training encoder models with contrastive learning.

Functions

NameDescription
_configure_sentence_transformer_exportBind static prompts and validate export against the runtime tokenizer or processor.
_get_autocast_ctxReturn the optional recipe-level autocast context.
_get_model_instantiate_kwargsReturn infrastructure kwargs forwarded to model instantiation.
_unpack_qpUnpack query and passage inputs from batch dictionary.
_unwrap_model_for_attrsReturn the underlying model object for configuration-style attribute reads.
_uses_multi_vector_scoringReturn whether the model emits token-level embeddings for MaxSim scoring.
contrastive_scores_and_labelsCompute contrastive scores and labels without in-batch negatives.
distributed_maxsim_scores_and_labelsCompute local-query multi-vector MaxSim scores against globally gathered passages.
main-
maxsim_scores_and_labelsCompute local multi-vector MaxSim scores and labels without in-batch negatives.

Data

logger

API

class nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe(
cfg
)

Bases: BaseRecipe

Recipe for training encoder models with contrastive learning.

cfg
temperature
= self.cfg.get('temperature', 1.0)
nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe._build_optimizer_param_groups() -> list[dict[str, typing.Any]]

Build optimizer parameter groups for trainable model parameters.

nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe._extract_scoring_reps(
model_output
)

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.

nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe._forward_backward_step(
idx,
batch,
loss_buffer,
num_batches,
is_train: bool = True
)

Forward and backward pass for a single micro-batch.

nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe._run_train_optim_step(
batches,
max_grad_norm = None
)

Run one optimization step with gradient accumulation.

nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe._run_validation_epoch(
val_dataloader
)

Run validation for one epoch and compute loss, accuracy@1, and MRR.

nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe._validate_model(
model: torch.nn.Module
) -> None

Validate recipe-specific model settings before constructing the optimizer.

Parameters:

model
torch.nn.Module

Constructed retrieval model, including infrastructure wrappers.

nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe.log_train_metrics(
)
nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe.log_val_metrics(
)
nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe.run_train_validation_loop()

Run the training loop over all epochs and batches.

nemo_automodel.recipes.retrieval.train_bi_encoder.TrainBiEncoderRecipe.setup()

Build all components needed for training/validation/logging/checkpointing.

nemo_automodel.recipes.retrieval.train_bi_encoder._configure_sentence_transformer_export(
model,
collate_fn,
tokenizer = None
) -> None

Bind static prompts and validate export against the runtime tokenizer or processor.

nemo_automodel.recipes.retrieval.train_bi_encoder._get_autocast_ctx(
distributed_config
)

Return the optional recipe-level autocast context.

nemo_automodel.recipes.retrieval.train_bi_encoder._get_model_instantiate_kwargs(
cfg,
distributed_setup,
peft_config
)

Return infrastructure kwargs forwarded to model instantiation.

nemo_automodel.recipes.retrieval.train_bi_encoder._unpack_qp(
inputs: dict[str, torch.Tensor]
) -> tuple

Unpack query and passage inputs from batch dictionary.

Parameters:

inputs
dict[str, torch.Tensor]

Dictionary containing query (q_) and passage (d_) tensors

Returns: tuple

Tuple of (query_batch_dict, doc_batch_dict)

nemo_automodel.recipes.retrieval.train_bi_encoder._unwrap_model_for_attrs(
model
)

Return the underlying model object for configuration-style attribute reads.

nemo_automodel.recipes.retrieval.train_bi_encoder._uses_multi_vector_scoring(
model
) -> bool

Return whether the model emits token-level embeddings for MaxSim scoring.

nemo_automodel.recipes.retrieval.train_bi_encoder.contrastive_scores_and_labels(
query: torch.Tensor,
key: torch.Tensor,
current_train_n_passages: int
) -> tuple[torch.Tensor, torch.Tensor]

Compute contrastive scores and labels without in-batch negatives.

Parameters:

query
torch.Tensor

Query embeddings [batch_size, hidden_dim]

key
torch.Tensor

Key/passage embeddings [batch_size * n_passages, hidden_dim]

current_train_n_passages
int

Number of passages per query

Returns: torch.Tensor

Tuple of (scores, labels) where scores is [batch_size, n_passages]

nemo_automodel.recipes.retrieval.train_bi_encoder.distributed_maxsim_scores_and_labels(
query: torch.Tensor,
key: torch.Tensor,
current_train_n_passages: int,
key_attention_mask: torch.Tensor,
rank: int
) -> tuple[torch.Tensor, torch.Tensor]

Compute local-query multi-vector MaxSim scores against globally gathered passages.

nemo_automodel.recipes.retrieval.train_bi_encoder.main(
default_config_path = 'examples/retrieval/bi_enco...
)
nemo_automodel.recipes.retrieval.train_bi_encoder.maxsim_scores_and_labels(
query: torch.Tensor,
key: torch.Tensor,
current_train_n_passages: int,
key_attention_mask: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]

Compute local multi-vector MaxSim scores and labels without in-batch negatives.

nemo_automodel.recipes.retrieval.train_bi_encoder.logger = logging.getLogger(__name__)