latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
tensorplay.distributed.pipelining
tensorplay.distributed.pipelining splits one model so its layers run on
different workers, one stage of the pipeline per worker, and orchestrates the
forward and backward passes as streams of microbatches flowing through the
stages. It is a form of model parallelism: instead of replicating the whole
model on every rank, each rank holds a different slice of the layers, so the
model whose parameters do not fit on one device still fits across the group.
Training with a pipeline split has three parts:
Splitting — describe where the model graph is cut into stages. The
pipeline()function traces the module with example micro-batched inputs and produces aPipe— the intermediate representation of the split model. You choose the cut points either with asplit_spec(a mapping of submodule names toSplitPointmarkers) or with asplit_policycallable, or by placingpipe_split()calls inside the module’s forward.Building stages —
build_stage()turns thePipeand a stage index into aPipelineStage: the worker-local module plus the send/recv channels for the tensors crossing its two boundaries. Each rank builds its own stage.Scheduling — a schedule class runs the training step over the pipeline, chunking each full batch into
n_microbatchesmicrobatches so the stages can work on different microbatches at the same time.
import tensorplay as tp
import tensorplay.distributed as dist
from tensorplay.distributed.pipelining import (
SplitPoint, pipeline, build_stage, Schedule1F1B,
)
dist.init_process_group("nccl")
class Block(tp.nn.Module):
def __init__(self, dim):
super().__init__()
self.net = tp.nn.Sequential(tp.nn.Linear(dim, dim), tp.nn.ReLU())
def forward(self, x):
return self.net(x)
class Model(tp.nn.Module):
def __init__(self):
super().__init__()
self.layer0 = Block(128) # stage 0 on rank 0
self.layer1 = Block(128) # stage 1 on rank 1
def forward(self, x):
x = self.layer0(x)
x = self.layer1(x)
return x
model = Model()
mb_args = (tp.randn(8, 128),) # one microbatch of inputs
pipe = pipeline(model, mb_args, split_spec={"layer1": SplitPoint.BEGINNING})
stage = build_stage(
pipe, stage_index=dist.get_rank(),
device=tp.device(f"cuda:{dist.get_rank()}"),
)
schedule = Schedule1F1B(stage, n_microbatches=4)
for x, y in dataloader:
out = schedule.step(x.to(dev))
loss = out.sum()
loss.backward()
stage.optimizer.step()
Describing the split
pipeline()traces the module withmb_args/mb_kwargsshaped like one microbatch, applies the split specification, and returns aPipe. Passsplit_spec(submodule name →SplitPoint) orsplit_policy(a callable that rewrites the traced graph), but not both.Pipeis the split representation: a module that knows the stage count, the per-stage submodules and their parameters, and the shapes of the tensors that cross stage boundaries. It is whatbuild_stageconsumes.pipe_split()is a marker you call inside a module’sforward; tracing records “cut here”. Combined with a nested call topipeline, this is the least intrusive way to annotate split points.SplitPointis the enum describing where a stage boundary falls relative to a submodule:BEGINNING(before the submodule runs) orEND(after it).
Building a stage
build_stage() takes the Pipe, a
stage index, the active device and the process group, and returns the
PipelineStage for this rank. The
stage owns the submodule slice, holds the input/output buffers and gradient
buffers that bridge the stage boundaries, and can wrap a weight-grad
communication callback (dw_builder) for overlapping all-reduce with
computation. For the device-mesh style of scheduling (a 2D mesh where the
pipeline dimension is separate from the data-parallel dimension), build the
stage with get_mesh so the schedule routes communication over the mesh’s
pipeline groups.
Scheduling
Execute all forward microbatches before draining their backwards. |
|
Overlap forward and backward microbatches after warmup. |
|
|
|
Run the bidirectional local stage schedule. |
The schedule is what actually runs the step: it chunks the input into
microbatches, walks them through the stage(s), and returns the per-microbatch
losses. PipelineScheduleSingle
drives one stage (used when every rank runs one stage); for interleaved
schedules where a rank owns several stages, use the multi-stage
PipelineScheduleMulti base, which
schedules the same microbatch over the local stages.
The concrete schedules differ in how forward and backward microbatches are ordered:
ScheduleGPiperuns all forward microbatches first, then drains their backwards — simple, but leaves the pipeline idle until the last forward finishes.Schedule1F1B(one forward, one backward) overlaps a microbatch’s backward with the next microbatch’s forward after a warm-up phase, keeping more stages busy.ScheduleInterleaved1F1Bextends 1F1B to interleaved stage assignment (several stages per rank), so a rank alternates between its chunks of the model.ScheduleLoopedBFSand the ZeroBubble family (ScheduleInterleavedZeroBubble,ScheduleZBVZeroBubble) apply increasingly aggressive recomputation and bubble-elimination schemes to approach the theoretical pipeline throughput.ScheduleDualPipeVruns the bidirectional local-stage schedule for the DualPipe V variant.
All schedules take n_microbatches (how finely the batch is chunked), an
optional loss_fn applied to the stage’s output, and chunk specs describing
which input dimensions are sharded per microbatch (args_chunk_spec /
kwargs_chunk_spec), so a tensor of shape (16, 128) with
n_microbatches=4 can be split along dim 0 into four (4, 128) chunks.
Where to go next
the distributed package — process groups and collectives, the messaging layer pipeline communication runs on.
device mesh — 2D meshes for combining the pipeline dimension with data or tensor parallelism.
FSDP — sharding the parameter memory within each stage’s rank, complementary to splitting the model across stages.
Help improve this page
Found an error, an unclear step, or a missing example?
tensorplay.distributed.optim
tensorplay.distributed.optim provides optimizers that are distributed-aware: their internal state is sharded, averaged, or moved, and the update is communicated across the process group rather than replicated on every ra
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

