ReferenceFull Library ReferenceNemo AutomodelNemo AutomodelComponentsModelsBagelnemo_automodel.components.models.bagel.parallelization

nemo_automodel.components.models.bagel.parallelization

View as Markdown

Model-owned distributed parallelization for BAGEL.

Module Contents

Classes

NameDescription
BagelModelParallelizerApply BAGEL whole-layer checkpointing before the generic FSDP2 flow.

Functions

Data

PARALLELIZER

_FULL_LAYER_CONTAINERS

__all__

logger

API

class nemo_automodel.components.models.bagel.parallelization.BagelModelParallelizer()

Bases: ModelParallelizer

Apply BAGEL whole-layer checkpointing before the generic FSDP2 flow.

nemo_automodel.components.models.bagel.parallelization.BagelModelParallelizer._apply(
model: torch.nn.Module,
args = (),
kwargs = {}
) -> torch.nn.Module
nemo_automodel.components.models.bagel.parallelization._apply_full_layer_checkpointing(
model: torch.nn.Module
) -> None
nemo_automodel.components.models.bagel.parallelization._module_by_path(
module: torch.nn.Module,
path: str
) -> torch.nn.Module | None
nemo_automodel.components.models.bagel.parallelization.PARALLELIZER = BagelModelParallelizer()
nemo_automodel.components.models.bagel.parallelization._FULL_LAYER_CONTAINERS = ('model.language_model.model.layers', 'model.vit_model.vision_model.encoder.laye...
nemo_automodel.components.models.bagel.parallelization.__all__ = ['PARALLELIZER']
nemo_automodel.components.models.bagel.parallelization.logger = logging.getLogger(__name__)