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.distributedfrom 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.
- agilerl.distributed.is_main_process() bool¶
Whether this is rank 0 (always
Trueon a single device).
- agilerl.distributed.broadcast_object_list(objects: list, src: int = 0) list¶
Broadcast a list of picklable objects from
srcto all ranks.Mutates
objectsin place on non-source ranks and returns it. No-op on a single device.
- 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.
- agilerl.distributed.allreduce_minmax_int(value: int) tuple[int, int]¶
Return
(min, max)ofvalueacross ranks ((value, value)locally).
- agilerl.distributed.aggregate_metrics_across_gpus(metric_tensor: Tensor | ndarray | float) float¶
Average a metric across ranks (local mean on a single device).
- 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.
- agilerl.distributed.sync_grads(params: Sequence[Parameter]) None¶
Average gradients across data-parallel ranks.
One coalesced all-reduce of every
.gradinparams(SUM, then divide by world size so Gloo works). Seeall_reduce_grads()for params without a grad. No-op on a single device.- Parameters:
params – Optimizer parameters whose
.gradshould 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
DTensorshards 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/ PEFTsave_pretrained). All ranks must enter and exit together. Yields a list parallel totensors.- Parameters:
tensors (torch.Tensor | None) – Tensors to densify;
Nonepasses through.
- agilerl.distributed.gather_params(root: Module, params: Sequence[Tensor | None]) Generator[list[Tensor | None], None, None]¶
Materialize full (unsharded) views of
paramsfor the duration of the context.Plain tensors are left unchanged. For FSDP2
DTensorparameters, each tensor is all-gathered withfull_tensor()and temporarily installed on its owning module so in-module reads (e.g.state_dict/ PEFTsave_pretrained) see dense weights. Original shards are restored on exit.Yields a list parallel to
params: dense locals for any gatheredDTensor, 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 prefermaterialize_dtensors()(no module Parameter install).Gathered parameters are read-only: writes are discarded when the sharded
DTensoris restored. Write into a sharded model withset_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;
Noneentries pass through.
- agilerl.distributed.full_shape_views(root: Module, params: Sequence[Tensor | None]) Generator[None, None, None]¶
Expose global-shape views on FSDP2
DTensorparams 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 originalDTensoris 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;
Noneentries 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.
- 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
modelwith 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
meshis provided, non-expertfully_sharduses that mesh; whenNone, FSDP uses the default process group (flat path). Packed expert modules are sharded onexpert_meshwhen that is provided.When
gradient_checkpointingis on, the transformer blocks chosen byconfig.checkpoint_skip_layer_typesandconfig.checkpoint_every_n_blocksare wrapped withcheckpoint_wrapperbeforefully_shardso 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
configwith non-reentrant activation checkpointing before sharding.
- Returns:
The sharded model (same object).
- Return type:
nn.Module