How to Configure Multi-GPU Inference with LTX-2 Pipelines: A Complete Guide
Configure multi-GPU inference in LTX-2 by subclassing MGPURunner, instantiating MGPUController with your GPU count, calling start() for setup, then streaming jobs through controller.stream() for lock-step SPMD execution across devices.
LTX-2 provides a built-in multi-GPU (MGPU) engine that enables any pipeline runner to execute in single-process, lock-step SPMD mode across multiple GPUs. This architecture eliminates Python pickling overhead for tensors through NCCL broadcasting and streamlines distributed inference for video generation workflows. This guide covers the complete configuration process based on the official Lightricks/LTX-2 source code.
Core Multi-GPU Components in LTX-2
Understanding the architecture helps you configure multi-GPU inference correctly. The engine consists of four primary components:
| Component | File Path | Purpose |
|---|---|---|
MGPUController |
ltx_pipelines/multigpu/controller.py |
Orchestrates worker fleet, dispatches single jobs, returns Stream for results |
_RunnersFleet / _RunnerShipper |
ltx_pipelines/multigpu/fleet.py |
Spawns one process per GPU, ships runner via cloud-pickle, manages shutdown |
MGPURunner |
ltx_pipelines/multigpu/runner.py |
Abstract base class requiring setup() and generator-based __call__() |
| Example implementation | ltx_pipelines/ti2vid_two_stages_mgpu.py |
Production-ready two-stage video generation runner |
The zero-copy tensor path is critical: top-level torch.Tensor arguments to stream() move from rank 0 to all ranks via NCCL, bypassing Python serialization.
Step 1: Implement Your MGPURunner Subclass
Every multi-GPU pipeline requires a concrete runner that subclasses MGPURunner. In ltx_pipelines/multigpu/runner.py, the base class defines two required methods:
from ltx_pipelines.multigpu.runner import MGPURunner, RunnerError
import torch
class SimpleImageGenerator(MGPURunner):
@torch.inference_mode()
def setup(self, model_path: str):
"""Runs once per GPU rank during controller.start()."""
self.model = torch.load(model_path, map_location="cpu")
@torch.inference_mode()
def __call__(self, *, latent: torch.Tensor, steps: int):
"""
Must be a generator (yield at least once).
Each yield becomes an element in the result Stream.
"""
for i in range(steps):
result = {"step": i, "image": latent * (i + 1)}
yield result # streamed immediately to host
Critical requirements for __call__:
- Must use
yieldat least once (generator requirement) - Receives broadcasted tensors already on the local GPU
- Raise
RunnerErroridentically on all ranks for symmetric failure handling
Step 2: Create and Configure MGPUController
In ltx_pipelines/multigpu/controller.py, the MGPUController accepts two mutually exclusive ways to specify GPUs:
from ltx_pipelines.multigpu.controller import MGPUController
# Option A: Use first N GPUs automatically
controller = MGPUController(
runner_cls=SimpleImageGenerator,
num_gpus=4
)
# Option B: Target specific GPU devices
controller = MGPUController(
runner_cls=SimpleImageGenerator,
devices=[2, 3, 4, 5] # for multiple controllers on disjoint subsets
)
Optional parameter: logs_specs controls Elastic launcher log verbosity for debugging worker crashes.
Step 3: Initialize the Fleet with start()
The start() method blocks until all ranks report ready via the internal ready queue. Pass all arguments required by your runner's setup():
controller.start(
model_path="models/dummy.pt",
# additional kwargs forwarded to setup() on every rank
)
In ltx_pipelines/multigpu/fleet.py, _RunnersFleet spawns processes, creates NCCL groups, and polls for readiness. A RuntimeError raises if any worker exits unexpectedly—check Elastic launcher logs for tracebacks.
Step 4: Dispatch Inference with stream()
The stream() method returns immediately with a Stream object. Only the calling thread may iterate it (enforced by thread-ID check):
import torch
# Top-level tensors are broadcast via NCCL automatically
latent = torch.randn(1, 3, 256, 256, device="cuda:0")
stream = controller.stream(
latent=latent,
steps=20,
# other kwargs passed to runner.__call__()
)
Tensor broadcast limitation: Only top-level tensor kwargs use NCCL. Nested tensors inside dicts/lists follow the pickle path.
Step 5: Consume Results and Cleanup
Always use try/finally with stream.drain() to prevent resource leaks:
try:
for rank_output in stream:
# Each item is a yield from one rank as soon as available
print(f"Rank result: {rank_output}")
# For symmetric pipelines, all ranks yield identical structure
# For asymmetric workloads, yields arrive as each rank completes
finally:
stream.drain() # mandatory cleanup before next stream() call
Error handling behavior:
- All ranks raise
RunnerError→SymmetricRunnerErrorre-raised - Mixed success/failure →
AsymmetricRunnerErrorraised
Step 6: Shutdown the Controller
Graceful shutdown with 60-second timeout (configurable):
controller.shutdown() # sends sentinel, waits for clean exit, kills stragglers
Multiple controllers on disjoint devices can coexist—ensure they don't overlap GPU assignments.
Complete Working Example
from ltx_pipelines.multigpu.controller import MGPUController, Stream
from ltx_pipelines.multigpu.runner import MGPURunner, RunnerError
import torch
import torch.multiprocessing as mp
class MinimalVideoRunner(MGPURunner):
def setup(self, spatial_upsampler_path: str):
self.spatial_path = spatial_upsampler_path
def __call__(self, *, height: int, width: int, steps: int):
for i in range(steps):
dummy_frame = torch.zeros(1, 3, height, width)
yield {"frame_idx": i, "tensor": dummy_frame}
def main():
controller = MGPUController(MinimalVideoRunner, num_gpus=2)
try:
controller.start(spatial_upsampler_path="models/upscaler.pt")
stream = controller.stream(height=512, width=768, steps=10)
try:
for item in stream:
print(f"Received frame {item['frame_idx']} from a rank")
finally:
stream.drain()
finally:
controller.shutdown()
if __name__ == "__main__":
mp.set_start_method("spawn", force=True) # required for NCCL
main()
Production Reference: TI2Vid Two-Stage Pipeline
The full-featured TI2VidTwoStagesRunner in ltx_pipelines/ti2vid_two_stages_mgpu.py demonstrates advanced multi-GPU patterns:
from ltx_pipelines.multigpu.controller import MGPUController
from ltx_pipelines.ti2vid_two_stages_mgpu import TI2VidTwoStagesRunner
import torch.multiprocessing as mp
vae_queue = mp.get_context("spawn").SimpleQueue()
controller = MGPUController(TI2VidTwoStagesRunner)
controller.start(
model_paths=model_paths,
prompt_enhancer_gemma_root=args.prompt_enhancer_gemma_root,
spatial_upsampler_path=args.spatial_upsampler_path,
vae_queue=vae_queue,
distilled_lora_path=args.distilled_lora[0].path,
compilation_config=args.compile,
diffvae_optimization=args.diffvae_optimization,
)
try:
for _ in controller.stream(
output_path=args.output_path,
prompt=args.prompt,
height=args.height,
width=args.width,
num_frames=args.num_frames,
num_inference_steps=args.num_inference_steps,
video_guider_params=video_guider_params,
images=args.images,
enhance_prompt=args.enhance_prompt,
):
pass # runner writes video as side-effect; iteration drives progress
finally:
controller.shutdown()
This runner uses weight_tracker.py for distributed transformer weight updates across NCCL groups.
Common Configuration Pitfalls
| Issue | Cause | Solution |
|---|---|---|
RuntimeError: stream already active |
Called stream() before drain() on previous stream |
Always stream.drain() in finally block |
AsymmetricRunnerError |
Ranks raised different exceptions | Raise identical RunnerError on all ranks for recoverable failures |
| Tensors not on expected GPU | Nested tensors not top-level | Flatten tensor arguments to stream() kwargs |
| Worker crash, no traceback | Elastic launcher logs not configured | Add logs_specs parameter to MGPUController |
__call__ not generator |
Missing yield statement |
Ensure at least one yield in __call__ implementation |
Performance Optimization Tips
- Batch tensor preparation on rank 0 — Construct all inputs on
cuda:0beforestream(); NCCL broadcast handles distribution - Minimize yields for throughput — Each yield incurs shared-memory queue overhead; batch intermediate results when possible
- Use
torch.inference_mode()— Decorator eliminates autograd overhead at runner entry points - Compile with
torch.compile— Passcompilation_configthroughstart()kwargs for optimized graphs per rank - Monitor NCCL health — Check
NVIDIA_NCCL_DEBUG=INFOif broadcasts hang; often indicates GPU memory fragmentation
Summary
- Subclass
MGPURunnerand implementsetup()(once per rank) and generator-based__call__()(must yield) - Instantiate
MGPUControllerwithnum_gpusor specificdevicesfor GPU allocation - Call
start()to spawn fleet, create NCCL groups, and initialize runners—blocks until ready - Dispatch with
stream()— top-level tensors auto-broadcast via NCCL; returnsStreamimmediately - Iterate and drain — consume yields with
for item in stream, alwaysstream.drain()infinally - Shutdown with
controller.shutdown()— graceful termination with timeout and cleanup
The LTX-2 multi-GPU engine in Lightricks/LTX-2 provides deterministic, single-process SPMD execution that scales video generation pipelines without the complexity of traditional distributed training frameworks.
Frequently Asked Questions
How does LTX-2 broadcast tensors across GPUs without pickling?
The MGPUController in ltx_pipelines/multigpu/controller.py identifies top-level torch.Tensor arguments to stream() and moves them from rank 0 to all ranks via NCCL collective operations. This zero-copy path avoids Python pickle overhead entirely. Nested tensors within lists or dictionaries still serialize through pickle, so flatten tensor arguments when possible.
Can I run multiple MGPUController instances on the same machine?
Yes—use the devices parameter instead of num_gpus to assign disjoint GPU subsets to each controller. For example, devices=[0,1,2,3] for one controller and devices=[4,5,6,7] for another allows concurrent multi-GPU pipelines on an 8-GPU server. Overlapping device assignments causes NCCL initialization failures.
What happens if my runner crashes on only some ranks?
The controller classifies failures into SymmetricRunnerError (all ranks raised identically) or AsymmetricRunnerError (mixed outcomes). For recoverable errors, raise RunnerError with identical messages on every rank to trigger symmetric handling. Asymmetric errors typically indicate bugs—check worker logs via logs_specs parameter.
Why must __call__ be a generator with at least one yield?
The streaming architecture in ltx_pipelines/multigpu/runner.py requires __call__ to be a generator so the controller can return a Stream object immediately and populate it asynchronously as workers yield results. Without yield, the controller cannot distinguish between setup completion and result availability, breaking the non-blocking streaming contract.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →