# How to Set Up Multi-GPU Inference with Sequence Parallel and Tiled Data Parallel in LTX-2

> Learn how to set up multi-GPU inference with LTX-2 using Sequence Parallel and Tiled Data Parallel. Optimize transformer inference across multiple GPUs for faster results.

- Repository: [Lightricks/LTX-2](https://github.com/Lightricks/LTX-2)
- Tags: how-to-guide
- Published: 2026-08-20

---

**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/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/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/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

```python
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

```python
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

```python
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/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

```python
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

```python
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/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:

```bash
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

```python
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:

```python
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/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/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/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/runner.py)](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/runner.py) — `MGPURunner` contract
- [[`controller.py`](https://github.com/Lightricks/LTX-2/blob/main/controller.py)](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/multigpu/controller.py) — `MGPUController` orchestration
- [[`weight_tracker.py`](https://github.com/Lightricks/LTX-2/blob/main/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/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/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/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`](https://github.com/Lightricks/LTX-2/blob/main/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/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.