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 TransformerWeightTracker as 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

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 torchrun and ensure ltx-kernels is 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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →