latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
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_module(), which
applies a style — a plan describing how each submodule is partitioned — to
an existing module. The styles in this module target the classic linear /
embedding layout: column-splitting the output dimension, row-splitting the
input dimension, and sequence-tensor parallelism for transformer blocks.
Parallelizing a module
parallelize_module() takes the
module, a device mesh, and either a single
style or a mapping from submodule names to styles. Inside the transformed
module, parameters and inputs are distributed tensors
with the placements the style chose; the runtime inserts the collectives that
assemble partial results where needed.
Loss with sharded logits
loss_parallel() is a context
manager that activates the distributed cross-entropy path for logits sharded
along the class dimension. When the model’s logits stay sharded through the
loss (instead of being gathered into a replicated tensor first), the
cross-entropy reduces the sharded logits directly and the loss gradient is
scaled for the shard size, saving the all-gather that would otherwise
materialize the full logits tensor on every rank. Wrap the loss computation
in with loss_parallel(): to enable it.
Where to go next
distributed tensors — the
DTensortype these styles operate on.device mesh — creating the mesh that
parallelize_moduleruns on.the distributed package — process groups and collectives underneath.
Help improve this page
Found an error, an unclear step, or a missing example?
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 sin
tensorplay.distributions
tensorplay.distributions.AbsTransform

