Distributed

Torch-native distributed helpers. AgileRL is single-device by default; multi-GPU LLM training initialises torch.distributed from the standard launcher environment variables set by torchrun (or an orchestration layer such as Ray), and these helpers no-op on a single device.

How to launch, when to use data parallel vs FSDP2, and every FSDPConfig field: Multi-GPU LLM training.

agilerl.distributed.init_distributed(timeout_seconds: int = 1800) → bool

Initialise torch.distributed from launcher env vars.

No-op (returns False) on a single device with no launcher env. Safe to call repeatedly; if a process group already exists it is reused.

Parameters:

timeout_seconds (int) – Collective timeout for the process group.

Returns:

True when distributed training is active.

Return type:

bool

agilerl.distributed.is_distributed() → bool

Whether torch.distributed is available and initialised.

agilerl.distributed.get_rank() → int

Global rank (0 on a single device).

agilerl.distributed.get_local_rank() → int

Rank within the node, used for device selection.

agilerl.distributed.get_world_size() → int

Number of processes (1 on a single device).

agilerl.distributed.is_main_process() → bool

Whether this is rank 0 (always True on a single device).

agilerl.distributed.barrier() → None

Synchronise all processes; no-op on a single device.

agilerl.distributed.broadcast_object_list(objects: list, src: int = 0) → list

Broadcast a list of picklable objects from src to all ranks.

Mutates objects in place on non-source ranks and returns it. No-op on a single device.

Parameters:
  • objects (list) – Objects to broadcast (same length on every rank).

  • src (int) – Source rank.

Returns:

The broadcast list.

Return type:

list

agilerl.distributed.gather_tensor(tensor: Tensor | ndarray | float) → Tensor

Gather a tensor from every rank (identity on a single device).

Parameters:

tensor – Tensor (or array/scalar convertible to one) to gather.

Returns:

Stacked / concatenated tensors from all ranks.

agilerl.distributed.gather_objects(objects: list) → list

Flatten object lists gathered from every rank.

Identity on a single device. Each rank contributes a picklable list; the result is concatenation in rank order.

Parameters:

objects (list) – Local objects to gather.

Returns:

Flattened objects from all ranks.

Return type:

list

agilerl.distributed.allreduce_minmax_int(value: int) → tuple[int, int]

Return (min, max) of value across ranks ((value, value) locally).

Parameters:

value (int) – This rank’s value.

Returns:

(min, max) across ranks.

Return type:

tuple[int, int]

agilerl.distributed.aggregate_metrics_across_gpus(metric_tensor: Tensor | ndarray | float) → float

Average a metric across ranks (local mean on a single device).

Parameters:

metric_tensor (torch.Tensor | np.ndarray | float) – Metric values on this rank.

Returns:

Mean across all ranks.

Return type:

float

agilerl.distributed.aggregate_metrics_dict(metrics: dict[str, Tensor | ndarray | float]) → dict[str, float]

Average every metric across ranks in one all-reduce (local mean on a single device).

Every rank must pass the same keys in the same order.

Parameters:

metrics (dict[str, torch.Tensor | np.ndarray | float]) – Metric values on this rank, by name.

Returns:

Mean of each metric across ranks.

Return type:

dict[str, float]

agilerl.distributed.sync_grads(params: Sequence[Parameter]) → None

Average gradients across data-parallel ranks.

One coalesced all-reduce of every .grad in params (SUM, then divide by world size so Gloo works). See all_reduce_grads() for params without a grad. No-op on a single device.

Parameters:

params – Optimizer parameters whose .grad should be averaged.

agilerl.distributed.materialize_dtensors(t0: torch.Tensor, t1: torch.Tensor | None, /) → _GeneratorContextManager[tuple[torch.Tensor, torch.Tensor | None], None, None]
agilerl.distributed.materialize_dtensors(*tensors: torch.Tensor | None) → _GeneratorContextManager[list[torch.Tensor | None], None, None]

All-gather DTensor shards to dense locals without swapping module params.

Prefer this for ephemeral matmuls (fused lm_head logprobs, Liger). Use gather_params() when in-module reads must see dense weights (state_dict / PEFT save_pretrained). All ranks must enter and exit together. Yields a list parallel to tensors.

Parameters:

tensors (torch.Tensor | None) – Tensors to densify; None passes through.

agilerl.distributed.gather_params(root: Module, params: Sequence[Tensor | None]) → Generator[list[Tensor | None], None, None]

Materialize full (unsharded) views of params for the duration of the context.

Plain tensors are left unchanged. For FSDP2 DTensor parameters, each tensor is all-gathered with full_tensor() and temporarily installed on its owning module so in-module reads (e.g. state_dict / PEFT save_pretrained) see dense weights. Original shards are restored on exit.

Yields a list parallel to params: dense locals for any gathered DTensor, and the original handles otherwise. Callers that hold pre-gather tensor references must use the yielded list for math — those references still point at the shard. For matmul-only gathers prefer materialize_dtensors() (no module Parameter install).

Gathered parameters are read-only: writes are discarded when the sharded DTensor is restored. Write into a sharded model with set_full_model_state_dict(). All ranks must enter and exit together.

Parameters:
  • root (nn.Module) – Module whose submodules own params.

  • params (Sequence[torch.Tensor | None]) – Parameters to gather; None entries pass through.

agilerl.distributed.full_shape_views(root: Module, params: Sequence[Tensor | None]) → Generator[None, None, None]

Expose global-shape views on FSDP2 DTensor params for shape-only reads.

Each sharded param is temporarily replaced on its owning module with a zero-storage view (a scalar expanded to the DTensor global shape) so shape, dtype and device reads see the full tensor without full_tensor(). Values must not be read inside the block; the original DTensor is restored on exit when the installed view is still the live parameter. Plain tensors pass through untouched. Duplicate references are processed once (by identity).

Parameters:
  • root (nn.Module) – Module whose submodules own params.

  • params (Sequence[torch.Tensor | None]) – Parameters to expose; None entries are skipped.

agilerl.distributed.resolve_device(requested: str | device | None = None) → str

Pick the training device.

Distributed runs are pinned to cuda:<local_rank>, or CPU without CUDA; otherwise the requested device (or the best available) is used.

Parameters:

requested (str | torch.device | None) – Device requested by the caller, if any.

Returns:

Device string.

Return type:

str

class agilerl.distributed.FSDPConfig(reshard_after_forward: Annotated[bool, FieldInfo(annotation=NoneType, required=True, description="Free gathered parameters after each module's forward. False keeps them gathered: faster, more VRAM.")] = True, cpu_offload: Annotated[bool, FieldInfo(annotation=NoneType, required=True, description='Offload sharded parameters and gradients to CPU. Needs colocated vLLM. Mutually exclusive with optim_cpu_offload.')] = False, optim_cpu_offload: Annotated[bool, FieldInfo(annotation=NoneType, required=True, description='Keep Adam state on CPU and move it to GPU only for step(). Parameters and gradients stay on GPU.')] = True, defer_grad_sync: Annotated[bool, FieldInfo(annotation=NoneType, required=True, description='Reduce-scatter gradients only on the last micro-batch of an optimizer step. Saves communication; holds unsharded grads in between.')] = True, param_dtype: Annotated[str, FieldInfo(annotation=NoneType, required=True, description="Mixed-precision parameter dtype, e.g. 'bfloat16'. A torch.dtype is accepted and stored by name.")] = 'bfloat16', reduce_dtype: Annotated[str, FieldInfo(annotation=NoneType, required=True, description='Dtype for gradient reduce-scatter and all-reduce.')] = 'float32', prefetch_units: Annotated[int, FieldInfo(annotation=NoneType, required=True, description='Neighbouring FSDP units to all-gather ahead during forward. 1 gathers unit i+1 while unit i runs.', metadata=[Ge(ge=1)])] = 1, backward_prefetch_units: Annotated[int, FieldInfo(annotation=NoneType, required=True, description='Neighbouring FSDP units to all-gather ahead during backward. 1 gathers unit i-1 while unit i runs backward.', metadata=[Ge(ge=1)])] = 1, checkpoint_skip_layer_types: Annotated[tuple[str, ...], FieldInfo(annotation=NoneType, required=True, description="Transformer block kinds left out of activation checkpointing. A block's kind is its block_type (hybrid models, e.g. 'linear_attention', 'full_attention', 'moe') or else its class name. Empty checkpoints every block. Applies when gradient_checkpointing is on.")] = (), checkpoint_every_n_blocks: Annotated[int, FieldInfo(annotation=NoneType, required=True, description='Checkpoint one in every n blocks not skipped by checkpoint_skip_layer_types. 1 checkpoints all of them.', metadata=[Ge(ge=1)])] = 1, wrap_every_n_blocks: Annotated[int, FieldInfo(annotation=NoneType, required=True, description='Consecutive transformer blocks per FSDP unit. Larger values mean fewer, bigger collectives.', metadata=[Ge(ge=1)])] = 1, param_persistence_threshold: Annotated[int, FieldInfo(annotation=NoneType, required=True, description='Parameters with fewer elements than this stay unsharded. 0 shards every parameter.', metadata=[Ge(ge=0)])] = 100000, ep: Annotated[int, FieldInfo(annotation=NoneType, required=True, description='Expert-parallel degree: packed MoE experts per layer are split across this many GPUs. 1 keeps data parallel plus FSDP sharding with no expert split.', metadata=[Ge(ge=1)])] = 1, ep_token_blocks: Annotated[int, FieldInfo(annotation=NoneType, required=True, description="Token blocks per routed MoE layer under expert parallel. Each block runs dispatch, experts, and combine; on GPU the next block's all-to-all overlaps this block's experts. 1 moves every token in one all-to-all.", metadata=[Ge(ge=1)])] = 1, tp: Annotated[int, FieldInfo(annotation=NoneType, required=True, description='Tensor-parallel degree for dense layers. Ranks in one TP group share a batch shard. 1 gives every rank its own batch.', metadata=[Ge(ge=1)])] = 1, shard_group_size: Annotated[int | None, FieldInfo(annotation=NoneType, required=True, description='Ranks that shard weights between them; groups replicate (HSDP). Set to GPUs per node so weight gathers stay inside a node. None shards across all trainer ranks.', metadata=[Ge(ge=1)])] | None = None, compile_blocks: Annotated[bool, FieldInfo(annotation=NoneType, required=True, description='torch.compile the dense submodules of each transformer block (norms, MLPs) in place. MoE experts, routers, attention and Mamba mixers stay eager.')] = False, compile_backend: Annotated[str, FieldInfo(annotation=NoneType, required=True, description='torch.compile backend for compile_blocks.')] = 'inductor')

Settings for sharding the LLM actor with PyTorch FSDP2 (fully_shard).

agilerl.distributed.apply_fsdp2(model: Module, config: FSDPConfig | None = None, mesh: DeviceMesh | None = None, expert_mesh: DeviceMesh | None = None, gradient_checkpointing: bool = False) → Module

Shard model with FSDP2: blocks, embed/(untied) lm_head, root, prefetch.

Parameters become DTensors in place, so any optimizer must be (re)built after this call. Callers must not pass a dense full-model replica on CUDA — use materialize_fsdp2_from_cpu_state() so weights stay on CPU/meta until only local shards are allocated on the compute device.

When mesh is provided, non-expert fully_shard uses that mesh; when None, FSDP uses the default process group (flat path). Packed expert modules are sharded on expert_mesh when that is provided.

When gradient_checkpointing is on, the transformer blocks chosen by config.checkpoint_skip_layer_types and config.checkpoint_every_n_blocks are wrapped with checkpoint_wrapper before fully_shard so the checkpoint boundary sits inside the FSDP unit.

Parameters:
  • model (nn.Module) – Model to shard (CPU or meta parameters).

  • config (FSDPConfig | None) – Sharding settings; defaults to FSDPConfig’s defaults.

  • mesh (DeviceMesh | None) – Optional FSDP device mesh for non-expert units. Ignored when None.

  • expert_mesh (DeviceMesh | None) – Optional FSDP mesh for packed-expert modules (the leftover data-parallel axis). Ignored when None.

  • gradient_checkpointing (bool) – Wrap the transformer blocks selected by config with non-reentrant activation checkpointing before sharding.

Returns:

The sharded model (same object).

Return type:

nn.Module