nemo_automodel.recipes.diffusion.train

View as Markdown

Module Contents

Classes

NameDescription
TrainDiffusionRecipeTraining recipe for diffusion models.

Functions

NameDescription
_build_diffusion_mesh_contextBuild the diffusion transformer’s sole parallelization input.
_build_transformer_engine_fp8_recipeBuild a Transformer Engine FP8 recipe from CLI-friendly config values.
_calculate_throughput_metricsCalculate directly measured training throughput metrics.
_count_local_batch_group_samplesCount local samples processed by one optimizer step.
_get_diffusion_microbatch_sizeReturn the number of samples in one local diffusion micro-batch.
_reject_removed_diffusion_keysRaise with old->new key mapping when a removed diffusion YAML key is present.
_resolve_model_dtypesResolve model storage and compute dtypes from the recipe config.
_resolve_transformer_engine_autocastResolve Transformer Engine’s quantization autocast context manager.
_validate_precision_configurationReject split storage/compute dtypes on paths without FSDP param casting.
build_diffusion_pipelineBuild the sharded diffusion pipeline (model + parallel scheme).

Data

_REMOVED_KEY_MIGRATIONS

API

class nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe(
cfg
)

Bases: BaseRecipe

Training recipe for diffusion models.

cfg
nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._autocast_context() -> typing.Any

Return the per-forward autocast context used when FSDP2 does not cast parameters.

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._count_global_samples(
local_samples: int
) -> int

Count samples processed across the data-parallel group.

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._elapsed_seconds_since(
start_time: float
) -> tuple[float, float]

Return the max elapsed wall-clock seconds across ranks since start_time.

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._get_collective_device() -> torch.device

Return a tensor device compatible with the active distributed backend.

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._get_dp_group_size(
include_cp: bool = False
) -> int

Get data parallel world size, handling DDP mode where device_mesh is None.

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._get_dp_rank(
include_cp: bool = False
) -> int

Get data parallel rank, handling DDP mode where device_mesh is None.

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._get_memory_metrics() -> typing.Dict[str, float]

Return PyTorch CUDA allocator memory counters, max-reduced across ranks.

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._run_validation_epoch(
global_step: int
) -> float

Score the held-out set with the training flow-matching objective.

Seeding by data rank mirrors training: every evaluation draws the same timesteps, noise, and CFG dropout, so val_loss tracks the model rather than the sampling, while data-parallel ranks stay decorrelated and context-parallel peers of one rank draw identical values for their shared batch. ScopedRNG restores the training RNG state afterwards, leaving the training trajectory unchanged.

The forward runs under the same compute-dtype autocast as training, so a single-rank split-dtype config evaluates the way it trains. FP8 autocast is left out: delayed-scaling amax history is training state and must not be updated from held-out batches.

Parameters:

global_step
int

Current optimizer step, forwarded to the flow-matching pipeline for its periodic diagnostics.

Returns: float

Mean loss over the validation batches of the data-parallel group, comparable to the

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._sync_device() -> None

Wait for queued CUDA work so timing reflects completed training work.

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe._transformer_engine_fp8_context() -> typing.Any

Return the per-forward Transformer Engine FP8 context.

nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe.run_train_validation_loop()
nemo_automodel.recipes.diffusion.train.TrainDiffusionRecipe.setup()
nemo_automodel.recipes.diffusion.train._build_diffusion_mesh_context(
fsdp_cfg: typing.Dict[str, typing.Any] | None,
ddp_cfg: typing.Dict[str, typing.Any] | None,
world_size: int,
dtype: torch.dtype,
compute_dtype: torch.dtype | None = None,
lora_enabled: bool

Build the diffusion transformer’s sole parallelization input.

nemo_automodel.recipes.diffusion.train._build_transformer_engine_fp8_recipe(
recipe_name: str,
amax_history_len: int,
amax_compute_algo: str
) -> typing.Any

Build a Transformer Engine FP8 recipe from CLI-friendly config values.

nemo_automodel.recipes.diffusion.train._calculate_throughput_metrics(
elapsed_seconds: float,
optimizer_steps: int,
global_samples: int,
world_size: int
) -> typing.Dict[str, float]

Calculate directly measured training throughput metrics.

nemo_automodel.recipes.diffusion.train._count_local_batch_group_samples(
batch_group: list[typing.Dict[str, typing.Any]]
) -> int

Count local samples processed by one optimizer step.

nemo_automodel.recipes.diffusion.train._get_diffusion_microbatch_size(
batch: typing.Dict[str, typing.Any]
) -> int

Return the number of samples in one local diffusion micro-batch.

nemo_automodel.recipes.diffusion.train._reject_removed_diffusion_keys(
cfg: typing.Any
) -> None

Raise with old->new key mapping when a removed diffusion YAML key is present.

Parameters:

cfg
Any

Raw recipe config (ConfigNode or RecipeConfig) as loaded from YAML.

Raises:

  • ValueError: If any key from _REMOVED_KEY_MIGRATIONS is present.
nemo_automodel.recipes.diffusion.train._resolve_model_dtypes(
cfg: typing.Any
) -> tuple[torch.dtype, torch.dtype]

Resolve model storage and compute dtypes from the recipe config.

nemo_automodel.recipes.diffusion.train._resolve_transformer_engine_autocast() -> typing.Any

Resolve Transformer Engine’s quantization autocast context manager.

nemo_automodel.recipes.diffusion.train._validate_precision_configuration(
dtype: torch.dtype,
compute_dtype: torch.dtype,
ddp_cfg: typing.Dict[str, typing.Any] | None,
peft_cfg: typing.Any
) -> None

Reject split storage/compute dtypes on paths without FSDP param casting.

nemo_automodel.recipes.diffusion.train.build_diffusion_pipeline(
model_id: str,
finetune_mode: bool,
device: torch.device,
dtype: torch.dtype,
compute_dtype: torch.dtype | None = None,
cpu_offload: bool = False,
fsdp_cfg: typing.Dict[str, typing.Any] | None = None,
ddp_cfg: typing.Dict[str, typing.Any] | None = None,
attention_backend: str | None = None,
transformer_engine_linear: bool = False,
transformer_engine_fp8_safe_only: bool = False,
fuse_qkv_projections: bool = False,
compact_fused_qkv_projections: bool = False,
pipeline_spec: typing.Dict[str, typing.Any] | None = None,
peft_cfg = None,
model_type = None,
active_transformer: str | None = None,
config_overrides: dict[str, nemo_automodel._diffusers.auto_diffusion_pipeline.ConfigFieldValue] | None = None

Build the sharded diffusion pipeline (model + parallel scheme).

The optimizer is built separately by the recipe via OptimizerConfig.build(...) on the returned pipeline’s transformer, so that parameters are collected after FSDP2 wrapping.

Parameters:

model_id
str

Pretrained model name or path.

finetune_mode
bool

Whether to load for finetuning (True) or pretraining (False).

device
torch.device

Target device.

dtype
torch.dtype

Model parameter storage dtype.

compute_dtype
torch.dtype | NoneDefaults to None

Forward/FSDP compute dtype. Defaults to dtype when unset.

cpu_offload
boolDefaults to False

Whether to enable CPU offload (FSDP only).

fsdp_cfg
Dict[str, Any] | NoneDefaults to None

FSDP configuration dict. Mutually exclusive with ddp_cfg.

ddp_cfg
Dict[str, Any] | NoneDefaults to None

DDP configuration dict. Mutually exclusive with fsdp_cfg.

attention_backend
str | NoneDefaults to None

Optional attention backend override.

transformer_engine_linear
boolDefaults to False

Whether to replace transformer torch.nn.Linear modules with Transformer Engine Linear.

transformer_engine_fp8_safe_only
boolDefaults to False

Whether to skip TE conversion for known FP8-incompatible modules.

fuse_qkv_projections
boolDefaults to False

Whether to call Diffusers QKV projection fusion on the transformer before FSDP.

compact_fused_qkv_projections
boolDefaults to False

Whether to remove original projection modules after QKV fusion.

pipeline_spec
Dict[str, Any] | NoneDefaults to None

Pipeline specification for pretraining (from_config). Required when finetune_mode is False. Should contain:

  • transformer_cls: str (e.g., “WanTransformer3DModel”, “FluxTransformer2DModel”)
  • subfolder: str (e.g., “transformer”)
  • Optional: pipeline_cls, load_full_pipeline
peft_cfg
Defaults to None

PeftConfig instance or None. When provided, only LoRA params are trained; base weights are frozen and sharded by FSDP2 for memory.

model_type
Defaults to None

“flux” | “flux2” | “wan” | “hunyuan” | “ltx2”. Required when peft_cfg is provided.

active_transformer
str | NoneDefaults to None

For two-transformer pipelines (Wan2.2), select which transformer to finetune. "transformer" (default for Wan2.2 = high-noise) or "transformer_2" (low-noise). The unused transformer is dropped before device placement so only one transformer lives on GPU.

backend
BackendConfig | NoneDefaults to None

BackendConfig for a transformer with a custom Automodel implementation (for example MoE diffusion transformers that need expert parallelism).

config_overrides
dict[str, ConfigFieldValue] | NoneDefaults to None

Config attributes overridden before building a custom-model transformer.

Returns: NeMoAutoDiffusionPipeline

Tuple of (pipeline, resolved MeshContext). The mesh context carries both the

Raises:

  • ValueError: If both fsdp_cfg and ddp_cfg are provided.
  • ValueError: If finetune_mode is False and pipeline_spec is not provided.
nemo_automodel.recipes.diffusion.train._REMOVED_KEY_MIGRATIONS = {'optim.learning_rate': 'optimizer.lr', 'optim.optimizer': 'optimizer (with an e...