Skip to main content

TensorCollective

Trait TensorCollective 

Source
pub trait TensorCollective<B>:
    Clone
    + Debug
    + Send
    + 'static
where B: Backend,
{ type Error: Debug; // Required methods fn world_size(&self) -> u32; fn all_gather_float( &self, value: <B as BackendTypes>::FloatTensorPrimitive, ) -> Result<<B as BackendTypes>::FloatTensorPrimitive, Self::Error>; fn reduce_scatter_sum( &self, value: <B as BackendTypes>::FloatTensorPrimitive, ) -> Result<<B as BackendTypes>::FloatTensorPrimitive, Self::Error>; // Provided method fn autodiff_context(&self) -> Option<&(dyn Any + Sync + Send + 'static)> { ... } }
Expand description

Matching rank-ordered collectives used by differentiable sharded tensors. All participating ranks must enter forward and backward in the same order.

Required Associated Types§

Source

type Error: Debug

Transport or tensor-contract failure.

Required Methods§

Source

fn world_size(&self) -> u32

Number of participating ranks, greater than zero.

Source

fn all_gather_float( &self, value: <B as BackendTypes>::FloatTensorPrimitive, ) -> Result<<B as BackendTypes>::FloatTensorPrimitive, Self::Error>

Gather equal leading-axis shards in rank order, retaining the input dtype.

Source

fn reduce_scatter_sum( &self, value: <B as BackendTypes>::FloatTensorPrimitive, ) -> Result<<B as BackendTypes>::FloatTensorPrimitive, Self::Error>

Sum tensors across ranks and return this rank’s equal leading-axis shard.

Provided Methods§

Source

fn autodiff_context(&self) -> Option<&(dyn Any + Sync + Send + 'static)>

Optional typed AD execution context carried by an explicit communicator wrapper. Native transports and ordinary AD calls retain no context by default.

Dyn Compatibility§

This trait is not dyn compatible.

In older versions of Rust, dyn compatibility was called "object safety".

Implementors§