# Source code for tensorplay.distributed.nn.functional Source: https://www.tensorplay.cn/docs/_modules/tensorplay/distributed/nn/functional.html ``` from __future__ import annotations import tensorplay as tp import tensorplay.distributed as dist from tensorplay.autograd.function import Function ReduceOp = dist.ReduceOp __all__ = [ "broadcast", "gather", "scatter", "reduce", "reduce_scatter", "all_gather", "all_gather_single", "all_to_all", "all_to_all_single", "all_reduce", ] def _not_supported_under_compile(name: str, suggestion: str | None = None) -> None: message = f"tensorplay.distributed.nn.functional.{name} is not available during graph compilation" if suggestion: message += f"; use {suggestion}" raise RuntimeError(message) [docs] def broadcast(tensor: tp.Tensor, src: int, group=None): return _Broadcast.apply(src, group, tensor) [docs] def gather(tensor: tp.Tensor, dst: int = 0, group=None): return _Gather.apply(dst, group, tensor) [docs] def scatter(tensors, src: int = 0, group=None): if tensors is None: raise ValueError("scatter requires a tensor sequence") return _Scatter.apply(src, group, *tuple(tensors)) [docs] def reduce(tensor: tp.Tensor, dst: int, op: int = ReduceOp.SUM, group=None): return _Reduce.apply(dst, op, group, tensor) [docs] def reduce_scatter(output, input_list, op: int = ReduceOp.SUM, group=None): return _Reduce_Scatter.apply(op, group, output, *tuple(input_list)) [docs] def all_gather(tensor: tp.Tensor, group=None): return _AllGather.apply(group, tensor) [docs] def all_gather_single(output_tensor, input_tensor, group=None): return _AllGatherSingle.apply(output_tensor, input_tensor, group) [docs] def all_to_all(output_tensor_list, input_tensor_list, group=None): return _AlltoAll.apply(group, output_tensor_list, *tuple(input_tensor_list)) [docs] def all_to_all_single( output, input, output_split_sizes=None, input_split_sizes=None, group=None, ): return _AlltoAllSingle.apply( group, output, output_split_sizes, input_split_sizes, input ) [docs] def all_reduce(tensor: tp.Tensor, op: int = ReduceOp.SUM, group=None): return _AllReduce.apply(op, group, tensor) class _Broadcast(Function): @staticmethod def forward(ctx, src, group, tensor): ctx.src = src ctx.group = group ctx.global_rank = dist.get_rank(group) result = tensor.clone() dist.broadcast(result, group=group, group_src=src) return result @staticmethod def backward(ctx, grad_output): grad = _Reduce.apply(ctx.src, ReduceOp.SUM, ctx.group, grad_output) if ctx.src != ctx.global_rank: grad.zero_() return None, None, grad class _Gather(Function): @staticmethod def forward(ctx, dst, group, tensor): ctx.dst = dst ctx.group = group outputs = [ tp.zeros_like(tensor) for _ in range(dist.get_world_size(group=group)) ] dist.gather( tensor.contiguous(), outputs if dist.get_rank(group=group) == dst else None, group=group, group_dst=dst, ) return tuple(outputs) @staticmethod def backward(ctx, *grad_outputs): return None, None, _Scatter.apply(ctx.dst, ctx.group, *grad_outputs) class _Scatter(Function): @staticmethod def forward(ctx, src, group, *tensors): if not tensors: raise ValueError("scatter requires at least one tensor") ctx.src = src ctx.group = group first = tensors[0] if any(tensor.shape != first.shape for tensor in tensors[1:]): raise ValueError("scatter tensors must have equal shapes") output = tp.zeros_like(first) dist.scatter( output, list(tensors) if dist.get_rank(group=group) == src else None, group=group, group_src=src, ) return output @staticmethod def backward(ctx, grad_output): return None, None, *_Gather.apply(ctx.src, ctx.group, grad_output) class _Reduce(Function): @staticmethod def forward(ctx, src, op, group, tensor): ctx.src = src ctx.op = op ctx.group = group result = tensor.clone() dist.reduce(result, op=op, group=group, group_dst=src) return result @staticmethod def backward(ctx, grad_output): return None, None, None, _Broadcast.apply(ctx.src, ctx.group, grad_output) class _Reduce_Scatter(Function): @staticmethod def forward(ctx, op, group, output, *input_tensors): ctx.op = op ctx.group = group dist.reduce_scatter( output, [tensor.contiguous() for tensor in input_tensors], op=op, group=group, ) return output @staticmethod def backward(ctx, grad_output): return None, None, None, *_AllGather.apply(ctx.group, grad_output) class _AllGather(Function): @staticmethod def forward(ctx, group, tensor): ctx.group = group ctx.input_shape = tuple(tensor.shape) ctx.input_dtype = tensor.dtype ctx.input_device = tensor.device outputs = [ tp.empty_like(tensor) for _ in range(dist.get_world_size(group=group)) ] dist.all_gather(outputs, tensor.contiguous(), group=group) return tuple(outputs) @staticmethod def backward(ctx, *grad_outputs): rank = dist.get_rank(group=ctx.group) world_size = dist.get_world_size(group=ctx.group) if len(grad_outputs) != world_size: raise RuntimeError( "all_gather backward received an invalid number of output gradients" ) reduced: Any = None for index, grad_output in enumerate(grad_outputs): if grad_output is None: current = tp.zeros( ctx.input_shape, dtype=ctx.input_dtype, device=ctx.input_device, ) else: current = grad_output.clone() dist.all_reduce(current, op=ReduceOp.SUM, group=ctx.group) if index == rank: reduced = current return None, reduced class _AllGatherSingle(Function): @staticmethod def forward(ctx, output_tensor, input_tensor, group): ctx.group = group dist.all_gather_single( output_tensor, input_tensor.contiguous(), group=group ) return output_tensor @staticmethod def backward(ctx, grad_output): world_size = dist.get_world_size(group=ctx.group) output_shape = list(grad_output.shape) if not output_shape or output_shape[0] % world_size != 0: raise RuntimeError( "all_gather_single backward requires the leading dimension " "to be divisible by the group size" ) output_shape[0] //= world_size grad_input = tp.empty( output_shape, dtype=grad_output.dtype, device=grad_output.device, ) dist.reduce_scatter_single( grad_input, grad_output.contiguous(), op=ReduceOp.SUM, group=ctx.group, ) return None, grad_input, None class _AlltoAll(Function): @staticmethod def forward(ctx, group, output_tensor_list, *tensors): ctx.group = group ctx.input_sizes = [tensor.shape for tensor in tensors] dist.all_to_all( output_tensor_list, [tensor.contiguous() for tensor in tensors], group=group, ) return tuple(output_tensor_list) @staticmethod def backward(ctx, *grad_outputs): outputs = [ tp.empty(size, dtype=grad_outputs[0].dtype, device=grad_outputs[0].device) for size in ctx.input_sizes ] return (None, None, *_AlltoAll.apply(ctx.group, outputs, *grad_outputs)) class _AlltoAllSingle(Function): @staticmethod def forward(ctx, group, output, output_split_sizes, input_split_sizes, input): ctx.group = group ctx.input_size = input.shape ctx.output_split_sizes = input_split_sizes ctx.input_split_sizes = output_split_sizes dist.all_to_all_single( output, input, output_split_sizes=output_split_sizes, input_split_sizes=input_split_sizes, group=group, ) return output @staticmethod def backward(ctx, grad_output): tensor = tp.empty_like(grad_output) return ( None, None, None, None, _AlltoAllSingle.apply( ctx.group, tensor, ctx.output_split_sizes, ctx.input_split_sizes, grad_output.contiguous(), ), ) class _AllReduce(Function): @staticmethod def forward(ctx, op, group, tensor): ctx.op = op ctx.group = group result = tensor.clone() dist.all_reduce(result, op=op, group=group) return result @staticmethod def backward(ctx, grad_output): return None, None, _AllReduce.apply(ctx.op, ctx.group, grad_output) ```