ReferenceFull Library ReferenceNemo AutomodelNemo AutomodelRecipesRetrievalnemo_automodel.recipes.retrieval.distill_bi_encoder

nemo_automodel.recipes.retrieval.distill_bi_encoder

View as Markdown

Module Contents

Classes

NameDescription
EmbeddingDistillRecipeRecipe for Stage-1 embedding distillation on bi-encoder backbones.

Functions

NameDescription
_build_or_none-
_cfg_get_pathRead either OmegaConf-style dotted keys or nested dict/config values.
_clean_path-
_copy_checkpoint_metadataCopy HF metadata/tokenizer files needed by AutoModel.from_pretrained.
_dp_group_src_rank-
_export_hf_student_checkpointMaterialize an evaluator-facing HF checkpoint from Automodel wrapper weights.
_mirror_hf_metadataMirror non-weight HF artifacts from the original student checkpoint.
_move_to_device-
_strip_student_prefixReturn the HF backbone tensors from a RetrieverStudentWithProjection checkpoint.
_unpack_qpn-
mainEntry point: load config, build the distillation recipe, and run training.

Data

logger

API

class nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe()

Bases: TrainBiEncoderRecipe

Recipe for Stage-1 embedding distillation on bi-encoder backbones.

nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe._build_optimizer_param_groups() -> list[dict[str, typing.Any]]

Build optimizer groups with projection params isolated before checkpoint restore.

nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe._extract_scoring_reps(
model_output
)

Select the pooled student embedding for validation scoring.

RetrieverStudentWithProjection.forward returns (pooled, projected, intermediate_outputs). The pooled embedding is the student’s native retrieval representation (and what training’s InfoNCE terms score with), so it is what the inherited validation loop should compare against.

nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe._forward_backward_step(
idx,
batch,
loss_buffer,
num_batches,
is_train: bool = True
)
nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe._projection_parameters() -> list[torch.nn.Parameter]
nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe._run_train_optim_step(
batches,
max_grad_norm = None
)
nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe._sync_projection_gradients() -> None

Average projection gradients that are outside FSDP/DDP wrapping.

nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe._sync_projection_parameters() -> None

Keep the rank-local projection head replicated across DP ranks.

nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe.save_checkpoint(
epoch: int,
step: int,
train_loss: float,
val_loss: dict[str, float] | None = None,
best_metric_key: str = 'default'
)
nemo_automodel.recipes.retrieval.distill_bi_encoder.EmbeddingDistillRecipe.setup()
nemo_automodel.recipes.retrieval.distill_bi_encoder._build_or_none(
cfg_section
)
nemo_automodel.recipes.retrieval.distill_bi_encoder._cfg_get_path(
cfg,
path: str,
default = None
)

Read either OmegaConf-style dotted keys or nested dict/config values.

nemo_automodel.recipes.retrieval.distill_bi_encoder._clean_path(
path: pathlib.Path
) -> None
nemo_automodel.recipes.retrieval.distill_bi_encoder._copy_checkpoint_metadata(
src_dir: pathlib.Path,
dst_dir: pathlib.Path
) -> None

Copy HF metadata/tokenizer files needed by AutoModel.from_pretrained.

nemo_automodel.recipes.retrieval.distill_bi_encoder._dp_group_src_rank(
group
) -> int
nemo_automodel.recipes.retrieval.distill_bi_encoder._export_hf_student_checkpoint(
src_dir: pathlib.Path,
dst_dir: pathlib.Path
) -> None

Materialize an evaluator-facing HF checkpoint from Automodel wrapper weights.

nemo_automodel.recipes.retrieval.distill_bi_encoder._mirror_hf_metadata(
src_model_dir: pathlib.Path,
dst_dir: pathlib.Path,
overwrite: bool = False
) -> None

Mirror non-weight HF artifacts from the original student checkpoint.

nemo_automodel.recipes.retrieval.distill_bi_encoder._move_to_device(
batch: dict,
device: torch.device
) -> dict
nemo_automodel.recipes.retrieval.distill_bi_encoder._strip_student_prefix(
state: dict[str, torch.Tensor]
) -> dict[str, torch.Tensor]

Return the HF backbone tensors from a RetrieverStudentWithProjection checkpoint.

nemo_automodel.recipes.retrieval.distill_bi_encoder._unpack_qpn(
batch: dict[str, torch.Tensor]
)
nemo_automodel.recipes.retrieval.distill_bi_encoder.main(
default_config_path = 'examples/retrieval/distill...
)

Entry point: load config, build the distillation recipe, and run training.

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