Distributed primitives#

Low-level communication primitives used internally by the distributed layers. They are documented here so that advanced users can build custom distributed operators on top of the same infrastructure.

All primitives are autograd-compatible: each one defines both a forward and backward communication pattern so that gradients flow correctly through distributed computations.

Transpose (all-to-all)#

distributed_transpose_azimuth

All-to-all transpose across the azimuth process group.

distributed_transpose_polar

All-to-all transpose across the polar process group.

Copy / Reduce#

copy_to_polar_region

Identity in the forward pass; all-reduce across polar ranks in the backward pass.

copy_to_azimuth_region

Identity in the forward pass; all-reduce across azimuth ranks in the backward pass.

reduce_from_polar_region

All-reduce across polar ranks in the forward pass; identity in the backward pass.

reduce_from_azimuth_region

All-reduce across azimuth ranks in the forward pass; identity in the backward pass.

Scatter / Gather#

scatter_to_polar_region

Split input_ along dim_ and keep only the local polar rank's chunk.

gather_from_polar_region

All-gather along dim_ across polar ranks to reconstruct the full tensor.

gather_from_copy_to_polar_region

All-gather along dim_ across polar ranks; reduce-scatter in the backward pass.

reduce_from_scatter_to_polar_region

Fused reduce-scatter across polar ranks along dim_.

reduce_from_scatter_to_azimuth_region

Fused reduce-scatter across azimuth ranks along dim_.

Halo exchange#

polar_halo_exchange

Exchange r_lat halo rows with neighbouring polar ranks.

get_group_neighbors

Return the (prev_rank, next_rank) global ranks of the immediate neighbours in group.

Tensor reshaping#

flatten_and_pad_leading_dims

Collapse all but the trailing num_trailing_dims dims into a single leading dim, padding it to at least min_leading_size.

unpad_and_unflatten_leading_dims

Inverse of flatten_and_pad_leading_dims(): drop the padding rows and restore the leading dims.