latest (dev)
Copy
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 aDTensorwith 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 aDTensorout 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:
Split one logical tensor dimension across a mesh dimension. |
|
Keep a complete copy of the logical tensor on every rank. |
|
Store values that still need a reduction across one mesh dimension. |
|
Base class for a tensor layout on one mesh dimension. |
Shardsplits the tensor along a logical dimension across the ranks of the corresponding mesh dimension.Replicatestores a full copy of the tensor on every rank of the corresponding mesh dimension.Partialmarks a tensor that has been only partially reduced on each rank; thereduce_opnames the operation (e.g."sum") that completes the reduction when a full value is needed.
The DTensor class
A logical tensor represented by a local value and a mesh placement. |
|
Split one logical tensor dimension across a mesh dimension. |
|
Keep a complete copy of the logical tensor on every rank. |
|
Store values that still need a reduction across one mesh dimension. |
|
Base class for a tensor layout on one mesh dimension. |
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
device mesh — the mesh abstraction that
DTensorruns on.tensor parallelism — sharding entire modules with
parallelize_module.the distributed package — process groups, collectives, and the key-value stores underneath.
Help improve this page
Found an error, an unclear step, or a missing example?
tensorplay.distributed.rpc
tensorplay.distributed.rpc lets you call functions on a worker running in another process — possibly on another machine — and get the result back. It is the primitive behind the higher-level distributed model APIs (remot
tensorplay.distributed.tensor.parallel
Tensor parallelism shards the parameters of a module across the ranks of a mesh dimension, so that each rank computes on part of the weights and the results are combined with collectives. The entry point is parallelize_m

