nemo_automodel.recipes.llm.train_domino

View as Markdown

Domino draft-model training recipe (Qwen3-style targets).

Domino (sgl-project/SpecForge#571) extends the DFlash parallel draft backbone with a lightweight causal correction head (a GRU state plus a low-rank logit correction; see nemo_automodel.components.speculative.dflash.domino_core). This recipe reuses every piece of the DFlash recipe — online target hidden-state capture, anchor sampling, the block attention mask, gradient accumulation, and checkpointing — and only swaps in the Domino trainer wrapper, enables the Domino head on the draft via dflash_config, and drives the base-anchor lambda_base curriculum.

Module Contents

Classes

NameDescription
TrainDominoRecipeRecipe for Domino draft-model training: DFlash backbone + causal correction head.

Functions

NameDescription
mainEntrypoint for TrainDominoRecipe.

Data

logger

API

class nemo_automodel.recipes.llm.train_domino.TrainDominoRecipe()

Bases: TrainDFlashRecipe

Recipe for Domino draft-model training: DFlash backbone + causal correction head.

nemo_automodel.recipes.llm.train_domino.TrainDominoRecipe._build_dflash_config(
recipe_cfg,
target_layer_ids: list[int]
) -> dict

Extend the DFlash draft config with the Domino head fields.

nemo_automodel.recipes.llm.train_domino.TrainDominoRecipe._build_trainer_module(
attention_backend: str,
recipe_cfg
)

Build the Domino trainer wrapper on the (Domino-head-enabled) DFlash draft.

nemo_automodel.recipes.llm.train_domino.TrainDominoRecipe._empty_extra_eval_metric_sums() -> dict[str, list[torch.Tensor]]

Create rank-symmetric Domino validation accumulators.

nemo_automodel.recipes.llm.train_domino.TrainDominoRecipe._extra_eval_metric_sums(
metrics
) -> dict[str, tuple[torch.Tensor, torch.Tensor]]

Return additive Domino base-head validation statistics.

Every returned value is a pair of scalar tensors on the trainer device. The shared validation loop SUM-reduces each pair before division.

nemo_automodel.recipes.llm.train_domino.TrainDominoRecipe._extra_train_metric_sums(
metrics
) -> dict[str, tuple[float, float]]

Return Domino head and curriculum diagnostics as window sums.

lambda_base is a schedule value rather than a statistic, so its denominator is the micro-batch count: averaging it over the window is what makes it comparable with the loss curves beside it.

nemo_automodel.recipes.llm.train_domino.TrainDominoRecipe._log_extra_train_metrics(
epoch_idx: int
) -> None

Log the Domino-specific diagnostics for the most recent step (rank-0 local).

nemo_automodel.recipes.llm.train_domino.TrainDominoRecipe._run_trainer_step(
target_batch
)

Forward through the Domino wrapper, injecting the current curriculum weight.

nemo_automodel.recipes.llm.train_domino.TrainDominoRecipe.setup()

Build everything via the DFlash recipe, then read the lambda_base schedule.

nemo_automodel.recipes.llm.train_domino.main(
config_path: str | None = None
)

Entrypoint for TrainDominoRecipe.

nemo_automodel.recipes.llm.train_domino.logger = logging.getLogger(__name__)