nemo_rl.utils.checkpoint_engines.nixl#

Module Contents#

Classes#

NixlAgent

Wrap a NIXL agent and its peer-control messaging.

NIXLCheckpointEngine

Transfer checkpoint weight buckets through a NIXL backend.

Functions#

_source_rank_for_rollout

_create_nixl_agent

resolve_nixl_backend_kwargs

Resolve (backend_name, backend_init_params) from engine_kwargs.nixl.

preinit_nixl_agent

_sync_device

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,
) int#
nemo_rl.utils.checkpoint_engines.nixl._create_nixl_agent(
agent_name: str,
backend_name: str,
backend_init_params: dict[str, Any] | None = None,
) Any#
nemo_rl.utils.checkpoint_engines.nixl.resolve_nixl_backend_kwargs(
nixl_kwargs: dict[str, Any],
) tuple[str, dict[str, Any] | None]#

Resolve (backend_name, backend_init_params) from engine_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,
) Any#
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(
metadata: nemo_rl.utils.checkpoint_engines.nixl.NixlAgentMetadata,
) 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.CheckpointEngine

Transfer checkpoint weight buckets through a NIXL backend.

Initialization

_allocate_transfer_buffer() torch.Tensor#
prepare() nemo_rl.utils.checkpoint_engines.nixl.NixlAgentMetadata#
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],
) None#
init_rollout_process_group(
*,
rollout_rank: int,
train_world_size: int,
rollout_world_size: int,
metadata: list[nemo_rl.utils.checkpoint_engines.nixl.NixlAgentMetadata],
) None#
_disconnect_peers() None#
_release_transfer_buffers() None#
finalize() None#
async send_weights(
weights: collections.abc.Generator[tuple[str, torch.Tensor], None, 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]#