ReferenceFull Library ReferenceNemo AutomodelNemo AutomodelComponentsModelsNemotron V3nemo_automodel.components.models.nemotron_v3.parallelization

nemo_automodel.components.models.nemotron_v3.parallelization

View as Markdown

Model-owned distributed parallelization for Nemotron-H and Nemotron-V3.

Module Contents

Classes

NameDescription
NemotronHModelParallelizerApply Nemotron-H’s specialized TP, CP, AC, and FSDP policy.

Functions

NameDescription
_decoder_blocksReturn the mutable decoder container and its ordered blocks.

Data

PARALLELIZER

__all__

logger

API

class nemo_automodel.components.models.nemotron_v3.parallelization.NemotronHModelParallelizer()

Bases: ModelParallelizer

Apply Nemotron-H’s specialized TP, CP, AC, and FSDP policy.

nemo_automodel.components.models.nemotron_v3.parallelization.NemotronHModelParallelizer._apply(
model: torch.nn.Module,
device_mesh: torch.distributed.device_mesh.DeviceMesh,
mp_policy: torch.distributed.fsdp.MixedPrecisionPolicy | None = None,
offload_policy: torch.distributed.fsdp.OffloadPolicy | None = None,
sequence_parallel: bool = False,
activation_checkpointing: bool = False,
tp_shard_plan: typing.Union[typing.Dict[str, torch.distributed.tensor.parallel.ParallelStyle], str] | None = None,
dp_replicate_mesh_name: str = 'dp_replicate',
dp_shard_cp_mesh_name: str = 'dp_shard_cp',
tp_mesh_name: str = 'tp',
reshard_after_forward: bool | None = None,
reapply_trainability: collections.abc.Callable[[nn.Module], None] | None = None,
kwargs = {}
) -> torch.nn.Module

Apply every requested parallelism to a Nemotron-H model.

nemo_automodel.components.models.nemotron_v3.parallelization._decoder_blocks(
model: torch.nn.Module
) -> tuple[torch.nn.Module, list[torch.nn.Module]]

Return the mutable decoder container and its ordered blocks.

nemo_automodel.components.models.nemotron_v3.parallelization.PARALLELIZER = NemotronHModelParallelizer()
nemo_automodel.components.models.nemotron_v3.parallelization.__all__ = ['PARALLELIZER']
nemo_automodel.components.models.nemotron_v3.parallelization.logger = logging.getLogger(__name__)