nemo_rl.utils.checkpoint_engines.nixl#
Module Contents#
Classes#
Wrap a NIXL agent and its peer-control messaging. |
|
Transfer checkpoint weight buckets through a NIXL backend. |
Functions#
Resolve |
|
Data#
API#
- nemo_rl.utils.checkpoint_engines.nixl.NixlAgentMetadata#
None
- nemo_rl.utils.checkpoint_engines.nixl.NIXL_DEFAULT_BACKEND_NAME#
‘UCX’
- nemo_rl.utils.checkpoint_engines.nixl.NIXL_TRANSFER_BUFFER_COUNT#
2
- nemo_rl.utils.checkpoint_engines.nixl._source_rank_for_rollout(
- rollout_rank: int,
- *,
- train_world_size: int,
- rollout_world_size: int,
- nemo_rl.utils.checkpoint_engines.nixl._create_nixl_agent(
- agent_name: str,
- backend_name: str,
- backend_init_params: dict[str, Any] | None = None,
- nemo_rl.utils.checkpoint_engines.nixl.resolve_nixl_backend_kwargs(
- nixl_kwargs: dict[str, Any],
Resolve
(backend_name, backend_init_params)fromengine_kwargs.nixl.Single source for the NIXL backend-name default so preinit call sites don’t each repeat
.get("backend_name", NIXL_DEFAULT_BACKEND_NAME).
- nemo_rl.utils.checkpoint_engines.nixl.preinit_nixl_agent(
- *,
- backend_name: str = NIXL_DEFAULT_BACKEND_NAME,
- backend_init_params: dict[str, Any] | None = None,
- nemo_rl.utils.checkpoint_engines.nixl._sync_device(device: torch.device) None#
- class nemo_rl.utils.checkpoint_engines.nixl.NixlAgent(
- backend_name: str = NIXL_DEFAULT_BACKEND_NAME,
- backend_init_params: dict[str, Any] | None = None,
Wrap a NIXL agent and its peer-control messaging.
Initialization
- get_agent_metadata() nemo_rl.utils.checkpoint_engines.nixl.NixlAgentMetadata#
- add_remote_agent( ) str#
- remove_remote_agent(agent_name: str) None#
- send_message(agent_name: str, message: dict[str, Any]) None#
- async read_message(agent_name: str) dict[str, Any]#
- async wait_notification(agent_name: str, notify_key: bytes) None#
- class nemo_rl.utils.checkpoint_engines.nixl.NIXLCheckpointEngine(
- bucket_size: int,
- device: str | torch.device = 'cuda',
- backend_name: str = NIXL_DEFAULT_BACKEND_NAME,
- backend_init_params: dict[str, Any] | None = None,
- shard_expert_weights: bool = False,
- release_after_refit: bool = False,
Bases:
nemo_rl.utils.checkpoint_engines.base.CheckpointEngineTransfer checkpoint weight buckets through a NIXL backend.
Initialization
- _allocate_transfer_buffer() torch.Tensor#
- get_target_weight_layout() dict[str, Any] | None#
- init_policy_process_group(
- *,
- worker_rank: int,
- train_world_size: int,
- rollout_world_size: int,
- metadata: list[nemo_rl.utils.checkpoint_engines.nixl.NixlAgentMetadata],
- init_rollout_process_group(
- *,
- rollout_rank: int,
- train_world_size: int,
- rollout_world_size: int,
- metadata: list[nemo_rl.utils.checkpoint_engines.nixl.NixlAgentMetadata],
- _disconnect_peers() None#
- _release_transfer_buffers() None#
- finalize() None#
- async send_weights(
- weights: collections.abc.Generator[tuple[str, torch.Tensor], None, None],
- async receive_weight_batches() collections.abc.AsyncGenerator[list[tuple[str, torch.Tensor]], None]#
- async _wait_read(xfer_handle: Any, remote_agent: str) None#
- async _receive_weight_chunk_batches() collections.abc.AsyncGenerator[list[tuple[nemo_rl.utils.checkpoint_engines.base.TensorMeta, torch.Tensor]], None]#