How to Set Up Multi-GPU Inference with Sequence Parallel and Tiled Data Parallel in LTX-2
LTX-2 enables multi-GPU inference by composing Sequence Parallel (SP) for transformer stage 1 and Tiled Data Parallel (TDP) for stage 2, orchestrated through MGPUController and a custom MGPURunner subclass.
Setting up multi-GPU inference with Sequence Parallel and Tiled Data Parallel in LTX-2 requires understanding three core building blocks: SequenceParallelBuilder for splitting sequence dimensions across GPUs, TiledDataParallelBuilder for spatial tile distribution, and MGPURunner as the execution contract. These components work together in the TI2VidTwoStagesRunner implementation to maximize throughput on high-resolution video generation tasks.
Core Multi-GPU Components
SequenceParallelBuilder: Distributing Sequence Dimensions
The SequenceParallelBuilder wraps a single-GPU transformer and injects sequence-parallel operations. Located in [sp_builder.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/sp_builder.py#L25-L62), it returns a SequenceParallelModelWrapper that coordinates all-to-all communication through an AttentionManager.
Key responsibilities:
- Splits the sequence dimension across the transformer process group
- Manages attention kernel dispatch via
AttentionManager - Rebinds weights through a shared
TransformerWeightTracker
TiledDataParallelBuilder: Distributing Spatial Tiles
The TiledDataParallelBuilder handles spatial parallelism for stage 2. Found in [tdp_builder.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/tdp_builder.py#L25-L75), it creates a TiledDataParallelModelWrapper where each GPU processes a height × width tile with configurable overlap boundaries.
Key responsibilities:
- Computes balanced 2D tile splits across available GPUs
- Overlaps tile boundaries (default 5 pixels) for seamless stitching
- Shares the same
TransformerWeightTrackeras SP for zero-copy weight rebinding
MGPURunner: The Execution Contract
The abstract MGPURunner class in [runner.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/runner.py#L29-L63) defines the interface:
setup(): Initializes models, process groups, and builders on each rank__call__(): Generator that yields results; only driver rank returns actual outputs
The concrete implementation TI2VidTwoStagesRunner ties SP and TDP together with optional Gemma parallelism and distributed VAE decoding.
Step-by-Step Setup in TI2VidTwoStagesRunner
The setup method in TI2VidTwoStagesRunner performs six critical steps to configure multi-GPU inference with Sequence Parallel and Tiled Data Parallel:
1. Create the Weight Tracker
tracker = TransformerWeightTracker(group=self.groups.transformer_group)
The TransformerWeightTracker keeps model weights resident on each GPU and automatically rebinds them when builders swap the inner model.
2. Instantiate the Attention Manager
attn_mgr = AttentionManager(
max_tokens=sp_max_tokens,
num_heads=model_cfg["num_attention_heads"],
head_dim=model_cfg["attention_head_dim"],
tensor_dtype=pipeline.dtype,
group=self.groups.transformer_group,
)
The AttentionManager from ltx_core.multigpu.transformer.attention manages SP attention kernels and all-to-all communication patterns.
3. Configure Sequence Parallel for Stage 1
pipeline.stage_1._transformer_builder = SequenceParallelBuilder(
inner=pipeline.stage_1._transformer_builder,
attn_mgr=attn_mgr,
registry=registry,
tracker=tracker,
)
This replaces the default transformer builder with one that injects SP operations through [sequence_parallel.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-core/src/ltx_core/multigpu/transformer/sequence_parallel.py).
4. Compute Balanced Tile Split
tdp_height_tiles, tdp_width_tiles = balanced_tile_split(
dist.get_world_size(self.groups.transformer_group)
)
tdp_tiling = TileCountConfig(
height=DimensionTilingConfig(num_tiles=tdp_height_tiles, overlap=5),
width=DimensionTilingConfig(num_tiles=tdp_width_tiles, overlap=5),
)
The balanced_tile_split function distributes tiles optimally across the transformer group size.
5. Configure Tiled Data Parallel for Stage 2
pipeline.stage_2._transformer_builder = TiledDataParallelBuilder(
inner=pipeline.stage_2._transformer_builder,
group=self.groups.transformer_group,
tiling=tdp_tiling,
registry=registry,
tracker=tracker,
)
This wraps stage 2 with the TDP implementation in [tiled_data_parallel.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-core/src/ltx_core/multigpu/transformer/tiled_data_parallel.py).
6. Attach Optional Components
Gemma-based prompt enhancement and distributed VAE decoding can be added independently. These use separate process groups but integrate with the same controller.
Running Multi-GPU Inference
Launch Command
Use torchrun to spawn the distributed processes:
torchrun --nproc_per_node=4 -m ltx_pipelines.ti2vid_two_stages_mgpu \
--model_paths /path/to/stage1 /path/to/stage2 \
--prompt "your prompt here" \
--height 1024 --width 1024 --num_frames 121
Complete Python Example
import torch
from ltx_pipelines.multigpu.controller import MGPUController
from ltx_pipelines.ti2vid_two_stages_mgpu import TI2VidTwoStagesRunner
from ltx_pipelines.utils.args import default_2_stage_arg_parser, resolve_cli_params
# 1️⃣ Parse CLI arguments
params = resolve_cli_params()
parser = default_2_stage_arg_parser(params=params, supports_auto_duration=True)
args = parser.parse_args()
# 2️⃣ Create multiprocessing queue for distributed VAE
vae_queue = torch.multiprocessing.get_context("spawn").SimpleQueue()
# 3️⃣ Instantiate controller with concrete runner
controller = MGPUController(TI2VidTwoStagesRunner)
# 4️⃣ Start controller — spawns one process per GPU
controller.start(
model_paths=args.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,
)
# 5️⃣ Stream inference — yields output path on driver rank only
for _ in controller.stream(
output_path=args.output_path,
prompt=args.prompt,
negative_prompt=args.negative_prompt,
seed=args.seed,
height=args.height,
width=args.width,
num_frames=args.num_frames,
frame_rate=args.frame_rate,
num_inference_steps=args.num_inference_steps,
video_guider_params=args.video_guidance,
audio_guider_params=args.audio_guidance,
images=args.images,
enhance_prompt=args.enhance_prompt,
hdr=args.hdr,
):
pass # video written by driver rank
# 6️⃣ Clean shutdown
controller.shutdown()
Manual Builder Construction (Custom Pipelines)
For pipelines beyond TI2Vid, construct builders directly:
from ltx_core.multigpu.transformer.attention import AttentionManager
from ltx_pipelines.multigpu.sp_builder import SequenceParallelBuilder
from ltx_pipelines.multigpu.tdp_builder import TiledDataParallelBuilder
from ltx_pipelines.multigpu.weight_tracker import TransformerWeightTracker
from ltx_core.loader.registry import Registry
registry = Registry()
tracker = TransformerWeightTracker(group=transformer_group)
# ---- Sequence Parallel for stage 1 ----
attn_mgr = AttentionManager(
max_tokens=32768,
num_heads=12,
head_dim=64,
tensor_dtype=torch.float16,
group=transformer_group,
)
sp_builder = SequenceParallelBuilder(
inner=inner_builder,
attn_mgr=attn_mgr,
registry=registry,
tracker=tracker,
)
# ---- Tiled Data Parallel for stage 2 ----
from ltx_core.multigpu.transformer.tiled_data_parallel import TileCountConfig, DimensionTilingConfig
tdp_tiling = TileCountConfig(
height=DimensionTilingConfig(num_tiles=2, overlap=5),
width=DimensionTilingConfig(num_tiles=2, overlap=5),
)
tdp_builder = TiledDataParallelBuilder(
inner=inner_builder,
group=transformer_group,
tiling=tdp_tiling,
registry=registry,
tracker=tracker,
)
System Requirements
| Component | Requirement |
|---|---|
| ltx-kernels | Install via pip install ltx-kernels — provides low-level SP attention kernels (create_video_self_attention_module_ops) |
| CUDA + NCCL | Required for inter-GPU communication; NCCLGroups manages process group creation |
| torchrun | Use PyTorch distributed launcher with consistent world size across all groups |
| Process groups | Controller creates transformer_group, gemma_group, and vae_group automatically |
Process groups in MGPUController use the same world size internally. Do not manually move model weights — the TransformerWeightTracker handles residency across SP/TDP swaps.
Key Source Files
- [
ti2vid_two_stages_mgpu.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages_mgpu.py) — Complete two-stage pipeline implementation - [
sp_builder.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/sp_builder.py) — Sequence Parallel builder - [
tdp_builder.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/tdp_builder.py) — Tiled Data Parallel builder - [
runner.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/runner.py) —MGPURunnercontract - [
controller.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/controller.py) —MGPUControllerorchestration - [
weight_tracker.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/weight_tracker.py) — Weight management - [
attention.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-core/src/ltx_core/multigpu/transformer/attention.py) — Low-level attention utilities - [
sequence_parallel.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-core/src/ltx_core/multigpu/transformer/sequence_parallel.py) — SP core wrapper - [
tiled_data_parallel.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-core/src/ltx_core/multigpu/transformer/tiled_data_parallel.py) — TDP core wrapper
Summary
- SequenceParallelBuilder splits sequence dimensions across GPUs for stage 1 transformers, coordinated by
AttentionManager - TiledDataParallelBuilder distributes spatial tiles with configurable overlap for stage 2, using
TileCountConfig - TransformerWeightTracker maintains GPU-resident weights across both parallelism strategies without manual movement
- TI2VidTwoStagesRunner demonstrates the complete integration pattern in
ti2vid_two_stages_mgpu.py - MGPUController handles NCCL process groups, broadcast/collect operations, and result streaming
- Launch with
torchrunand ensureltx-kernelsis installed for optimized SP attention
Frequently Asked Questions
What is the difference between Sequence Parallel and Tiled Data Parallel in LTX-2?
Sequence Parallel splits the token sequence dimension across GPUs, reducing memory per rank for long sequences. Tiled Data Parallel splits the spatial dimensions (height × width) across GPUs, enabling higher resolution processing. According to the LTX-2 source code, SP applies to stage 1 transformers while TDP applies to stage 2, with both using the same TransformerWeightTracker to avoid weight duplication.
How do I choose the number of tiles for TDP?
The balanced_tile_split function in TI2VidTwoStagesRunner automatically computes a 2D grid based on dist.get_world_size(self.groups.transformer_group). For 4 GPUs, this typically yields 2×2 tiles. You can override this by constructing TileCountConfig manually with custom num_tiles values per dimension and adjusting overlap (default 5 pixels) based on your model's receptive field requirements.
Why does only the driver rank return results during inference?
The MGPURunner.__call__ contract specifies a generator pattern where all ranks participate in computation but only rank yields actual outputs. According to [runner.py](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/runner.py#L29-L63), other ranks yield None to maintain synchronization. The MGPUController.stream() method handles NCCL-based result collection automatically, ensuring the driver rank has full outputs for encoding and writing.
Can I use Sequence Parallel or Tiled Data Parallel independently?
Yes. Both SequenceParallelBuilder and TiledDataParallelBuilder can wrap any SingleGPUModelBuilder independently. The TI2VidTwoStagesRunner demonstrates combined usage for maximum throughput, but you can apply either strategy to single-stage pipelines by importing the respective builder from ltx_pipelines.multigpu and following the manual construction pattern shown above.
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 →