TensorPlay
Reference guides
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

tensorplay.distributed.tensor

Distributed (sharded) tensors built on top of a device mesh. A DTensor is a logical tensor whose data is split across the ranks of a mesh: each rank stores a local shard, and the shards are expected to behave like a single tensor for the operations the library implements. The layout of a DTensor is described by its mesh together with a placement for every mesh dimension.

The module is the natural building block for tensor parallelism: shard the weights of a module across a mesh with distribute_module(), then feed it inputs produced by the distributed factory functions below.

Creating distributed tensors

  • distribute_tensor() wraps an existing tensor into a DTensor with the given placements, redistributing the data if the current layout on the mesh differs from the requested one.

  • distribute_module() converts an entire module into its distributed form: a partition function decides, per parameter, whether it is sharded or left replicated, and optional input and output functions transform the tensors that enter and leave the module.

  • from_local() builds a DTensor out of the local shards already present on each rank, without any cross-rank movement of data (optionally checking that the local shards match the requested global shape).

Factory functions

The factory functions create DTensor values directly on a mesh. They accept the usual construction arguments (dtype, layout, requires_grad) plus device_mesh and placements:

For example, creating an 8-element replicated vector on a 4-rank mesh keeps a full copy on every rank, while Shard(0) distributes it one element per pair of ranks:

from tensorplay.distributed.tensor import Shard, ones

# requires an initialized process group, e.g. via dist.init_process_group()
mesh = ...   # a DeviceMesh, see distributed.device_mesh
local = ones(8, device_mesh=mesh, placements=[Shard(0)])

Placements

A placement tells the runtime how one dimension of the mesh participates in storing the tensor:

tensorplay.distributed.tensor.Shard

Split one logical tensor dimension across a mesh dimension.

tensorplay.distributed.tensor.Replicate

Keep a complete copy of the logical tensor on every rank.

tensorplay.distributed.tensor.Partial

Store values that still need a reduction across one mesh dimension.

tensorplay.distributed.tensor.Placement

Base class for a tensor layout on one mesh dimension.

  • Shard splits the tensor along a logical dimension across the ranks of the corresponding mesh dimension.

  • Replicate stores a full copy of the tensor on every rank of the corresponding mesh dimension.

  • Partial marks a tensor that has been only partially reduced on each rank; the reduce_op names the operation (e.g. "sum") that completes the reduction when a full value is needed.

The DTensor class

The class itself exposes the shard on the current rank through to_local() (moving the tensor across devices if the shard is on a different accelerator than the local one), the fully materialized value through full_tensor(), and a re-layout operation through redistribute() that converts between placements, possibly asynchronously. Its metadata follows a plain tensor: shape, size, stride, ndim, numel, dtype, device, plus the device_mesh and placements that describe its layout.

Where to go next

On this page

Ask DeepWiki