nemo_rl.models.policy.workers.checkpoint_engine#

Module Contents#

Classes#

DTensorCheckpointEngineSendMixin

Onload DTensor/FSDP2 policy weights for checkpoint-engine transfer.

MegatronCheckpointEngineSendMixin

Select destination-local Megatron weights for checkpoint-engine transfer.

PolicyCheckpointEngineMixin

Checkpoint-engine lifecycle shared by policy worker implementations.

Functions#

maybe_preinit_nixl_checkpoint_engine

Preinitialize NIXL when checkpoint-engine refit is configured.

API#

nemo_rl.models.policy.workers.checkpoint_engine.maybe_preinit_nixl_checkpoint_engine(
config: dict[str, Any],
) Any#

Preinitialize NIXL when checkpoint-engine refit is configured.

class nemo_rl.models.policy.workers.checkpoint_engine.DTensorCheckpointEngineSendMixin#

Onload DTensor/FSDP2 policy weights for checkpoint-engine transfer.

model: torch.nn.Module#

None

_prepare_checkpoint_engine_weight_send() None#
_finalize_checkpoint_engine_weight_send() None#
_checkpoint_engine_weight_iterator(
kv_scales: Optional[dict[str, float]] = None,
) collections.abc.Iterator[tuple[str, torch.Tensor]]#
class nemo_rl.models.policy.workers.checkpoint_engine.MegatronCheckpointEngineSendMixin#

Select destination-local Megatron weights for checkpoint-engine transfer.

_checkpoint_engine_weight_iterator(
kv_scales: Optional[dict[str, float]] = None,
) collections.abc.Iterator[tuple[str, torch.Tensor]]#
class nemo_rl.models.policy.workers.checkpoint_engine.PolicyCheckpointEngineMixin#

Checkpoint-engine lifecycle shared by policy worker implementations.

checkpoint_engine: nemo_rl.utils.checkpoint_engines.base.CheckpointEngine#

None

rank: int#

None

abstractmethod _checkpoint_engine_weight_iterator(
kv_scales: Optional[dict[str, float]] = None,
) collections.abc.Generator[tuple[str, torch.Tensor], None, None]#
_prepare_checkpoint_engine_weight_send() None#
_finalize_checkpoint_engine_weight_send() None#
async send_weights_via_checkpoint_engine(
kv_scales: Optional[dict[str, float]] = None,
) None#
async checkpoint_engine_rpc(
checkpoint_method: str,
method_kwargs: Optional[dict[str, Any]] = None,
) Any#