nemo_automodel.components.distributed.config

View as Markdown

Strategy-specific distributed training configuration classes.

Design principle:

  • Size params (dp_size, dp_replicate_size, tp_size, pp_size, cp_size, ep_size) are grouped in ParallelismSizes.
  • dp_replicate_size is FSDP2-only: raises assertion if passed with non-FSDP2 config
  • Strategy-specific configs contain only additional flags unique to each strategy
  • Managers become normal classes that accept (config, device_mesh)

Usage: from nemo_automodel.components.distributed.config import FSDP2Config, MegatronFSDPConfig, DDPConfig

FSDP2 with custom options

config = FSDP2Config(sequence_parallel=True, activation_checkpointing=True)

MegatronFSDP with custom options

config = MegatronFSDPConfig(zero_dp_strategy=3, overlap_grad_reduce=True)

DDP with activation checkpointing

config = DDPConfig(activation_checkpointing=True)

Module Contents

Classes

NameDescription
DDPConfigAdditional configuration for DDP distributed training.
DistributedSetupResolved distributed topology and execution policies.
FSDP2ConfigAdditional configuration for FSDP2 distributed training.
MegatronFSDPConfigAdditional configuration for MegatronFSDP distributed training.
MoEParallelizerConfigConfiguration for MoE model parallelization (EP + FSDP settings).
MultimodalDistributedConfigDistributed policies for resolved multimodal modules.
MultimodalVisionConfigDistributed policies for resolved vision modules.

Functions

NameDescription
_resolve_strategy_configResolve a setup-level strategy name or config object.
normalize_activation_checkpointing_scopeValidate and normalize activation-checkpointing scope values.

Data

ActivationCheckpointingMode

ActivationCheckpointingScope

DistributedConfig

DistributedStrategyConfig

_STRATEGY_MAP

_StrategyConfigClass

_VALID_ACTIVATION_CHECKPOINTING_SCOPES

API

class nemo_automodel.components.distributed.config.DDPConfig(
broadcast_buffers: bool = False,
find_unused_parameters: bool = False,
static_graph: bool = False,
bucket_cap_mb: float | None = None,
gradient_as_bucket_view: bool = False,
autocast_dtype: torch.dtype | None = None
)
Dataclass

Additional configuration for DDP distributed training.

Note: DDP does not support tensor parallelism, pipeline parallelism, or expert parallelism. Only dp_size is relevant (inferred from world_size).

activation_checkpointing
ActivationCheckpointingMode = False

Enable activation checkpointing. True or "full" keeps the existing full activation checkpointing behavior. "selective" wraps transformer blocks with PyTorch selective activation checkpointing.

activation_checkpointing_scope
ActivationCheckpointingScope = 'all'

Which extracted layer groups activation checkpointing should wrap. "all" selects every extracted group. Scoped values such as "language", "vision", and "multimodal" are filtered to trainable layers before generic wrapping.

autocast_dtype
dtype | None = None

If set, recipes can wrap the forward pass in torch.autocast(device_type="cuda", dtype=autocast_dtype). Set to None to disable. Can be set from YAML as a string (e.g. autocast_dtype: bfloat16).

broadcast_buffers
bool = False

Synchronize module buffers before each forward.

bucket_cap_mb
float | None = None

DDP gradient bucket size in MiB. None uses PyTorch’s default.

find_unused_parameters
bool = False

Forwarded to PyTorch DDP for models with conditionally unused trainable parameters.

gradient_as_bucket_view
bool = False

Make gradients views into DDP buckets after the first iteration.

static_graph
bool = False

Tell DDP the used/unused parameter set is stable.

nemo_automodel.components.distributed.config.DDPConfig.__post_init__()
nemo_automodel.components.distributed.config.DDPConfig.to_dict() -> typing.Dict[str, typing.Any]

Convert config to dictionary.

class nemo_automodel.components.distributed.config.DistributedSetup(
mesh_context: 'MeshContext',
pipeline_config: 'PipelineConfig | None' = None,
moe_parallel_config: 'MoEParallelizerConfig | None' = None,
)
Dataclass

Resolved distributed topology and execution policies.

activation_checkpointing
ActivationCheckpointingMode = False
mesh_context
'MeshContext'
moe_parallel_config
'MoEParallelizerConfig | None' = None
pipeline_config
'PipelineConfig | None' = None
strategy_config
DistributedStrategyConfig | None = None
nemo_automodel.components.distributed.config.DistributedSetup.__post_init__() -> None

Keep the compatibility bundle and its MeshContext policy in sync.

nemo_automodel.components.distributed.config.DistributedSetup.build(
parallelism_sizes: 'ParallelismSizes | None' = None,
pipeline_config: 'PipelineConfig | dict | None' = None,
moe_parallel_config: 'MoEParallelizerConfig | dict | None' = None,
world_size: int | None = None,
timeout_minutes: int | None = None,
ranks: list[int] | tuple[int, ...] | None = None
) -> 'DistributedSetup'
classmethod

Create a resolved distributed setup from sizes and policy configs.

Intentionally, this function is forgiving wrt the input types, allowing strings for the strategy and dicts for the pipeline and MoE configs.

class nemo_automodel.components.distributed.config.FSDP2Config(
sequence_parallel: bool = False,
tp_plan: dict | str | None = None,
patch_is_packed_sequence: bool = False,
mp_policy: torch.distributed.fsdp.MixedPrecisionPolicy | None = (lambda: MixedPrecisionPoli...,
offload_policy: torch.distributed.fsdp.CPUOffloadPolicy | None = None,
autocast_dtype: torch.dtype | None = None,
defer_fsdp_grad_sync: bool = True,
reshard_after_forward: bool | None = None,
enable_async_tensor_parallel: bool = False,
enable_compile: bool = False,
enable_fsdp2_prefetch: bool = False,
fsdp2_backward_prefetch_depth: int = 2,
fsdp2_forward_prefetch_depth: int = 1,
)
Dataclass

Additional configuration for FSDP2 distributed training.

Note: Size parameters (dp_size, dp_replicate_size, tp_size, pp_size, cp_size, ep_size) are grouped separately in ParallelismSizes.

activation_checkpointing
ActivationCheckpointingMode = False

Enable activation checkpointing. True or "full" keeps the existing full activation checkpointing behavior. "selective" wraps transformer blocks with PyTorch selective activation checkpointing.

activation_checkpointing_scope
ActivationCheckpointingScope = 'all'

Which extracted layer groups activation checkpointing should wrap. "all" selects every extracted group. Scoped values such as "language", "vision", and "multimodal" are filtered to trainable layers before generic wrapping.

autocast_dtype
dtype | None = None

If set, wraps the forward pass in torch.autocast(device_type="cuda", dtype=autocast_dtype). Use with output_dtype=float32 in mp_policy to keep the residual stream in fp32 while running matmuls in lower precision. Set to None to disable. Can be set from YAML as a string (e.g. autocast_dtype: bfloat16).

defer_fsdp_grad_sync
bool = True

Defer FSDP gradient sync to final micro-batch.

enable_async_tensor_parallel
bool = False

Enable async tensor parallelism via torch._inductor.config._micro_pipeline_tp. Overlaps ReduceScatter with compute in row-parallel layers. Requires sequence_parallel=True (forced automatically with a warning if not set). Also enables symmetric memory for the TP group.

enable_compile
bool = False

Apply per-layer torch.compile to transformer decoder layers (with NO_REENTRANT activation checkpointing inside each compiled layer). Skips whole-model compile so that checkpoint loading does not produce _orig_mod key-prefix mismatches.

enable_fsdp2_prefetch
bool = False

Enable explicit forward/backward prefetch chains between FSDP2 sharded layers. Default True.

fsdp2_backward_prefetch_depth
int = 2

Number of FSDP units to prefetch during backward pass. 2 hides AllGather behind compute; 1 reduces peak memory at a small throughput cost. Default 2.

fsdp2_forward_prefetch_depth
int = 1

Number of FSDP units to prefetch during forward pass. Default 1.

mp_policy
MixedPrecisionPolicy | None

MixedPrecisionPolicy for FSDP2. If None (default), uses bf16 forward/backward compute with fp32 gradient reduction. Pair this with model.torch_dtype: float32 for the Megatron-style fp32 master-weights pattern. Override from YAML using the _target_ pattern::

mp_policy: target: torch.distributed.fsdp.MixedPrecisionPolicy param_dtype: bfloat16 reduce_dtype: float32 output_dtype: bfloat16

See docs/guides/mixed-precision-training.md for the full set of recommended patterns and the bf16-storage trap.

multimodal
MultimodalDistributedConfig = field(default_factory=MultimodalDistributedConfig)

Policies for resolved multimodal modules that use the text model’s distributed axes.

offload_policy
CPUOffloadPolicy | None = None

CPUOffloadPolicy for CPU offloading.

patch_is_packed_sequence
bool = False

Patch transformers._is_packed_sequence to always return Python False. This does two things: (1) removes a CPU-GPU sync per attention layer (aten::is_nonzero triggered by HF when batch_size==1), and (2) ensures static attention shapes for torch.compile. Safe for standard (non-packed) training only. Disable if using packed-sequence training (position_ids that reset to 0 mid-sequence). Default False.

reshard_after_forward
bool | None = None

Override layer-level FSDP2 resharding. None preserves AutoModel’s heuristic: pipeline-parallel layers do not reshard after forward, while non-pipeline layers reshard all but the last layer. Set False for a ZeRO-2-like benchmark where gathered parameters stay resident after forward. Set True to force resharding everywhere, including pipeline-parallel layers, which may reduce throughput by adding per-microbatch all-gathers.

sequence_parallel
bool = False

Enable sequence parallelism in TP plan.

tp_plan
dict | str | None = None

Custom TP plan. If None, auto-selected based on model type.

nemo_automodel.components.distributed.config.FSDP2Config.__post_init__() -> None
nemo_automodel.components.distributed.config.FSDP2Config.to_dict() -> typing.Dict[str, typing.Any]

Convert config to dictionary (shallow, preserves policy objects).

class nemo_automodel.components.distributed.config.MegatronFSDPConfig(
megatron_fsdp_unit_modules: list[str] | None = None,
zero_dp_strategy: int = 3,
init_fsdp_with_meta_device: bool = False,
grad_reduce_in_fp32: bool = False,
preserve_fp32_weights: bool = False,
overlap_grad_reduce: bool = True,
overlap_param_gather: bool = True,
check_for_nan_in_grad: bool = True,
report_nan_in_param_grad: bool = False,
average_in_collective: bool = False,
disable_bucketing: bool = False,
calculate_per_token_loss: bool = False,
keep_fp8_transpose_cache: bool = False,
nccl_ub: bool = False,
fsdp_double_buffer: bool = False,
activation_checkpointing: bool = False
)
Dataclass

Additional configuration for MegatronFSDP distributed training.

Note: Size parameters (dp_size, tp_size, cp_size) are grouped separately in ParallelismSizes. MegatronFSDP does not support pp_size, dp_replicate_size, or ep_size.

activation_checkpointing
bool = False

Enable activation checkpointing for transformer MLP layers to save memory.

average_in_collective
bool = False

Average in collective if True.

calculate_per_token_loss
bool = False

Calculate per token loss if True.

check_for_nan_in_grad
bool = True

Legacy buffer-level gradient NaN check. BREAKING CHANGE on megatron-fsdp 0.5.0: this flag is a no-op, preserved only for config compatibility. 0.5.0 removed the buffer-level NaN check entirely, so gradient NaN checking is now OFF regardless of this value; a truthy value is dropped with a one-time warning. The default is kept True for config compatibility, but has no effect. To restore gradient NaN checking, enable report_nan_in_param_grad instead.

disable_bucketing
bool = False

Disable bucketing if True.

fsdp_double_buffer
bool = False

Use double buffer if True.

grad_reduce_in_fp32
bool = False

Reduce gradients in fp32 if True.

init_fsdp_with_meta_device
bool = False

Initialize MegatronFSDP with meta device if True.

keep_fp8_transpose_cache
bool = False

Keep fp8 transpose cache when using custom FSDP if True.

megatron_fsdp_unit_modules
list[str] | None = None

Class paths of the submodules to wrap as individual MegatronFSDP units. When None (the default), the wrap classes are auto-derived from the model’s _no_split_modules so the real instantiated block classes are used regardless of backend (HF or NeMo-custom).

nccl_ub
bool = False

Use NCCL UBs if True.

overlap_grad_reduce
bool = True

Overlap gradient reduction if True.

overlap_param_gather
bool = True

Overlap parameter gathering if True.

preserve_fp32_weights
bool = False

Preserve fp32 weights if True.

report_nan_in_param_grad
bool = False

Enable megatron-fsdp 0.5.0’s precise per-parameter gradient NaN check. This is the replacement for the removed check_for_nan_in_grad and is OFF by default; enabling it can significantly reduce training throughput.

zero_dp_strategy
int = 3

Data parallel sharding strategy.

nemo_automodel.components.distributed.config.MegatronFSDPConfig.to_dict() -> typing.Dict[str, typing.Any]

Convert config to dictionary (shallow, preserves objects).

class nemo_automodel.components.distributed.config.MoEParallelizerConfig(
ignore_router_for_ac: bool = True,
reshard_after_forward: bool = False,
lm_head_precision: typing.Union[str, torch.dtype] | None = None,
wrap_outer_model: bool = True,
mp_policy: torch.distributed.fsdp.MixedPrecisionPolicy | None = None
)
Dataclass

Configuration for MoE model parallelization (EP + FSDP settings).

ignore_router_for_ac
bool = True
lm_head_precision
Union[str, dtype] | None = None
mp_policy
MixedPrecisionPolicy | None = None
reshard_after_forward
bool = False
wrap_outer_model
bool = True
nemo_automodel.components.distributed.config.MoEParallelizerConfig.to_dict() -> typing.Dict[str, typing.Any]
class nemo_automodel.components.distributed.config.MultimodalDistributedConfig(
)
Dataclass

Distributed policies for resolved multimodal modules.

frozen_sharding
FrozenMultimodalSharding = 'root'

Controls fully frozen multimodal modules such as vision/audio towers and projectors. "root" (default) keeps their parameters in an always-run outer FSDP root, which is safe when modality execution differs across ranks. "per_layer" uses normal layer/container FSDP units and requires every rank in the FSDP group to execute or skip the module identically on every microbatch. "replicate" excludes the frozen parameters from FSDP roots so each rank keeps a full copy. Modules with any trainable parameters use normal layer/container sharding regardless of this setting.

vision
MultimodalVisionConfig = field(default_factory=MultimodalVisionConfig)

Policies that apply specifically to resolved vision modules.

nemo_automodel.components.distributed.config.MultimodalDistributedConfig.__post_init__() -> None
class nemo_automodel.components.distributed.config.MultimodalVisionConfig(
)
Dataclass

Distributed policies for resolved vision modules.

frame_sharding
CpVisionFrameShardingConfig = field(default_factory=CpVisionFrameShardingConfig)

Controls frame-level encoder compute sharding over the selected text-model mesh dimensions.

nemo_automodel.components.distributed.config.MultimodalVisionConfig.__post_init__() -> None
nemo_automodel.components.distributed.config._resolve_strategy_config(
strategy_kwargs: typing.Any = {}

Resolve a setup-level strategy name or config object.

nemo_automodel.components.distributed.config.normalize_activation_checkpointing_scope(
value: typing.Any
) -> typing.Tuple[str, ...]

Validate and normalize activation-checkpointing scope values.

nemo_automodel.components.distributed.config.ActivationCheckpointingMode = Union[bool, Literal['full', 'selective']]
nemo_automodel.components.distributed.config.ActivationCheckpointingScope = Union[str, List[str], Tuple[str, ...]]
nemo_automodel.components.distributed.config.DistributedConfig = DistributedStrategyConfig
nemo_automodel.components.distributed.config.DistributedStrategyConfig = Union['FSDP2Config', 'MegatronFSDPConfig', 'DDPConfig']
nemo_automodel.components.distributed.config._STRATEGY_MAP: Dict[str, _StrategyConfigClass] = {'fsdp2': FSDP2Config, 'megatron_fsdp': MegatronFSDPConfig, 'megatron-fsdp': Meg...
nemo_automodel.components.distributed.config._StrategyConfigClass = type[FSDP2Config] | type[MegatronFSDPConfig] | type[DDPConfig]
nemo_automodel.components.distributed.config._VALID_ACTIVATION_CHECKPOINTING_SCOPES = {'all', 'language', 'vision', 'audio', 'multimodal'}