credit.parallel.collectives#

Shared gradient all-reduce machinery for the parallel package.

Used by both sync_domain_gradients (domain.py) and sync_replicated_gradients (tensor_parallel.py) so the subtle DTensor handling lives in exactly one place. Also home to the mixed-mesh-safe gradient clipping used by the gen2 trainer.

Functions#

all_reduce_avg(→ None)

In-place all_reduce average that also works on gloo.

allreduce_grads_avg(→ None)

Average gradients across a process group with minimal NCCL calls.

total_grad_norm(grads[, norm_type])

Total p-norm of a set of gradients that may live on different meshes.

clip_grad_norm_(parameters, max_norm[, norm_type])

Mixed-mesh-safe drop-in for torch.nn.utils.clip_grad_norm_.

Module Contents#

credit.parallel.collectives.all_reduce_avg(tensor, group=None) → None#

In-place all_reduce average that also works on gloo.

ReduceOp.AVG is NCCL-only; gloo (CPU multi-rank runs, –backend gloo) needs SUM + divide.

credit.parallel.collectives.allreduce_grads_avg(grads, group) → None#

Average gradients across a process group with minimal NCCL calls.

Plain dense grads are flattened into one bucket per (dtype, device) so the sync is a handful of large all_reduces instead of one per parameter. DTensor grads (FSDP2) are reduced in place on their local shards — shards can be 0-sized or oddly strided per rank, which breaks the flatten/unflatten round-trip — but the per-shard all_reduces are issued async and waited together, so they cost one latency, not one per param.

Parameters:
  • grads – iterable of gradient tensors (dense or DTensor).

  • group – process group to average over.

credit.parallel.collectives.total_grad_norm(grads, norm_type=2.0)#

Total p-norm of a set of gradients that may live on different meshes.

torch.nn.utils.get_total_norm stacks every per-grad norm with foreach ops, and DTensor dispatch rejects operands on different meshes (and plain/DTensor mixes). With native TP composed under FSDP2, the TP’d projections carry grads on the 2D (dp, tp) mesh while every other param lives on the 1D dp mesh, so the stack raises. Here grads are grouped by mesh; each homogeneous group goes through get_total_norm unchanged, a DTensor group norm is made global with full_tensor() (the correct reduction over its sharded dims), and the group norms combine as (sum_i n_i**p)**(1/p) (max for p=inf) via vector_norm. With a single group this is exactly torch’s computation.

Parameters:
  • grads – iterable of gradient tensors (dense or DTensor).

  • norm_type – p-norm exponent (any float, or inf).

Returns:

The total norm as a plain scalar tensor (never a DTensor).

credit.parallel.collectives.clip_grad_norm_(parameters, max_norm, norm_type=2.0)#

Mixed-mesh-safe drop-in for torch.nn.utils.clip_grad_norm_.

The total norm comes from total_grad_norm (mesh-grouped, see above); the scaling reuses torch.nn.utils.clip_grads_with_norm_ per mesh group, because torch’s foreach multiply also rejects mixed plain/DTensor lists. The clip coefficient depends only on the already-global total norm, so per-group scaling is exact. For a homogeneous parameter set (plain DDP tensors, or all DTensors on one mesh under FSDP2/domain) both steps reduce to exactly torch’s clip_grad_norm_ computation.

Parameters:
  • parameters – tensor or iterable of tensors whose .grad gets clipped in place.

  • max_norm – max norm of the gradients (float or scalar tensor).

  • norm_type – p-norm exponent (any float, or inf).

Returns:

The total norm of the gradients (before clipping) as a plain scalar tensor, matching the torch.nn.utils.clip_grad_norm_ contract.