pub trait TensorCollective<B>:
Clone
+ Debug
+ Send
+ 'staticwhere
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§
Required Methods§
Sourcefn world_size(&self) -> u32
fn world_size(&self) -> u32
Number of participating ranks, greater than zero.
Sourcefn all_gather_float(
&self,
value: <B as BackendTypes>::FloatTensorPrimitive,
) -> Result<<B as BackendTypes>::FloatTensorPrimitive, Self::Error>
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.
Sourcefn reduce_scatter_sum(
&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>
Sum tensors across ranks and return this rank’s equal leading-axis shard.
Provided Methods§
Dyn Compatibility§
This trait is not dyn compatible.
In older versions of Rust, dyn compatibility was called "object safety".