# Enabling Multi-GPU Inference with Tensor and Sequence Parallelism in LTX-2

> Scale video generation with LTX-2 multi-GPU inference. Learn how tensor and sequence parallelism optimize performance using PyTorch distributed primitives in ltx-pipelines.

- Repository: [Lightricks/LTX-2](https://github.com/Lightricks/LTX-2)
- Tags: performance
- Published: 2026-08-15

---

**LTX-2 supports scaling video generation across multiple GPUs through tensor parallelism (sharding model weights) and sequence parallelism (splitting token sequences), implemented in the `ltx-pipelines` package using PyTorch distributed primitives.**

The LTX-2 video generation model from Lightricks can process longer videos and larger batches by distributing computation across multiple GPUs. This article explains how to activate and configure **multi-GPU inference with tensor and sequence parallelism** based on the actual source code implementation.

## Understanding the Two Parallelism Strategies

LTX-2 combines two complementary approaches to parallelize inference:

### Tensor Parallelism

**Tensor parallelism** shards model layers along the feature dimension. Each GPU holds a slice of the weights and performs its forward pass on the full input, with **All-to-All collective operations** synchronizing activations between layers.

The core communication primitives live in [`packages/ltx-kernels/src/ltx_kernels/all_to_all.py`](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-kernels/src/ltx_kernels/all_to_all.py). This module provides the low-level kernels that exchange tensor shards across the tensor-parallel group at each layer boundary.

### Sequence Parallelism

**Sequence parallelism** addresses memory constraints when generating long videos. When token or latent sequences exceed single-GPU memory, the sequence is split into contiguous chunks across GPUs. Each rank processes its chunk independently, with minimal cross-rank synchronization.

The splitting and merging logic is implemented in [`packages/ltx-pipelines/src/ltx_pipelines/utils/res2s.py`](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/utils/res2s.py). This utility handles the decomposition of long latent sequences and reconstruction of the final output.

## Key Source Files and Their Roles

| Component | File Path | Purpose |
|-----------|-----------|---------|
| All-to-All kernels | [`ltx_kernels/all_to_all.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_kernels/all_to_all.py) | Collective communication for tensor-parallel weight sharding |
| Sequence utilities | [`ltx_pipelines/utils/res2s.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_pipelines/utils/res2s.py) | Split/merge long token sequences across GPUs |
| Model path resolver | [`ltx_pipelines/utils/model_paths.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_pipelines/utils/model_paths.py) | Load correct checkpoint shards per rank |
| Accelerate config | [`ltx-trainer/configs/accelerate/fsdp_compile.yaml`](https://github.com/Lightricks/LTX-2/blob/main/ltx-trainer/configs/accelerate/fsdp_compile.yaml) | Activates FSDP, tensor, and sequence parallelism |
| Inference script | [`ltx-trainer/scripts/serve_captioner.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx-trainer/scripts/serve_captioner.py) | Entry point for distributed video generation |
| Quantization support | [`ltx_pipelines/utils/quantization_factory.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_pipelines/utils/quantization_factory.py) | 8-bit/4-bit quantizers compatible with parallel shards |

## Configuration: The Accelerate YAML File

The [`fsdp_compile.yaml`](https://github.com/Lightricks/LTX-2/blob/main/fsdp_compile.yaml) configuration file in `packages/ltx-trainer/configs/accelerate/` is the control center for multi-GPU inference. It combines **FSDP (Fully Sharded Data Parallel)** with tensor and sequence parallelism options.

Key environment variables read from this configuration:

- `WORLD_SIZE` — total number of GPUs
- `TENSOR_PARALLEL_SIZE` — number of GPUs in each tensor-parallel group
- `SEQ_PARALLEL_SIZE` — number of GPUs in each sequence-parallel group

The constraint `WORLD_SIZE % (TENSOR_PARALLEL_SIZE * SEQ_PARALLEL_SIZE) == 0` must hold, as the product defines the size of a complete parallel replica.

## Launching Distributed Inference

Here is a complete Python launch script that configures and runs multi-GPU inference with both parallelism modes:

```python
import os
import subprocess


def launch_multi_gpu_inference(
    gpu_count: int,
    tensor_parallel: int,
    sequence_parallel: int,
    checkpoint_dir: str,
    prompt: str,
    output_path: str,
    config_path: str = "packages/ltx-trainer/configs/accelerate/fsdp_compile.yaml",
):
    """
    Launch LTX-2 inference with tensor and sequence parallelism.
    
    Parameters
    ----------
    gpu_count : int
        Total GPUs to use (WORLD_SIZE).
    tensor_parallel : int
        GPUs per tensor-parallel group. Must divide gpu_count.
    sequence_parallel : int
        GPUs per sequence-parallel group. Must divide gpu_count / tensor_parallel.
    checkpoint_dir : str
        Path to sharded model checkpoint.
    prompt : str
        Text prompt for video generation.
    output_path : str
        Where to save the generated video.
    """
    # Validate parallelism configuration

    replica_size = tensor_parallel * sequence_parallel
    if gpu_count % replica_size != 0:
        raise ValueError(
            f"WORLD_SIZE ({gpu_count}) must be divisible by "
            f"tensor_parallel * sequence_parallel ({replica_size})"
        )
    
    # Set environment for distributed launch

    env = os.environ.copy()
    env["WORLD_SIZE"] = str(gpu_count)
    env["TENSOR_PARALLEL_SIZE"] = str(tensor_parallel)
    env["SEQ_PARALLEL_SIZE"] = str(sequence_parallel)
    
    # Build accelerate launch command

    cmd = [
        "accelerate", "launch",
        "--config_file", config_path,
        "--num_processes", str(gpu_count),
        "packages/ltx-trainer/scripts/serve_captioner.py",
        "--ckpt_dir", checkpoint_dir,
        "--prompt", prompt,
        "--output", output_path,
    ]
    
    subprocess.run(cmd, check=True, env=env)


# Example: 8 GPUs, 2-way tensor parallel, 2-way sequence parallel

# This creates 2 model replicas, each with 4 GPUs (2x2)

launch_multi_gpu_inference(
    gpu_count=8,
    tensor_parallel=2,
    sequence_parallel=2,
    checkpoint_dir="/mnt/models/ltx2-sharded",
    prompt="A spacecraft docking with a ringed station above Mars",
    output_path="mars_docking.mp4",
)

```

## Initializing Distributed Groups in Inference Code

Inside the inference script, the parallelism environment variables are read to construct appropriate **process groups**. Here is how the distributed topology is established:

```python
import os
import torch
import torch.distributed as dist


def init_tensor_sequence_parallel():
    """
    Initialize torch.distributed with tensor and sequence parallel groups.
    Called by each rank on startup.
    """
    rank = int(os.environ["RANK"])
    world_size = int(os.environ["WORLD_SIZE"])
    
    # Initialize global process group

    dist.init_process_group(
        backend="nccl",
        init_method="env://",
        rank=rank,
        world_size=world_size,
    )
    
    # Read parallelism sizes from environment

    tp_size = int(os.environ["TENSOR_PARALLEL_SIZE"])
    sp_size = int(os.environ["SEQ_PARALLEL_SIZE"])
    
    # Validate configuration

    assert world_size % (tp_size * sp_size) == 0, "Invalid parallel group sizing"
    
    # Calculate position in parallel topology

    # replica_id: which model replica this rank belongs to

    replica_id = rank // (tp_size * sp_size)
    
    # Tensor-parallel group: ranks that share the same position within a replica

    tp_group_id = rank % tp_size
    tp_ranks = [
        replica_id * tp_size * sp_size + i * sp_size + tp_group_id
        for i in range(tp_size)
    ]
    tp_group = dist.new_group(ranks=tp_ranks)
    
    # Sequence-parallel group: ranks adjacent in sequence dimension

    sp_group_id = (rank // tp_size) % sp_size
    sp_ranks = [
        replica_id * tp_size * sp_size + i * tp_size + sp_group_id
        for i in range(sp_size)
    ]
    sp_group = dist.new_group(ranks=sp_ranks)
    
    return tp_group, sp_group, replica_id


def load_sharded_checkpoint(ckpt_dir: str, tp_group, sp_group):
    """
    Load only the checkpoint shards needed by this rank.
    Uses model_paths.py utilities for shard discovery.
    """
    from ltx_pipelines.utils.model_paths import resolve_rank_checkpoint
    
    # Resolve which .safetensors files this rank should load

    local_shards = resolve_rank_checkpoint(
        base_dir=ckpt_dir,
        rank=dist.get_rank(),
        tensor_parallel_group=tp_group,
        sequence_parallel_group=sp_group,
    )
    return local_shards

```

## Memory and Performance Trade-offs

**Tensor parallelism** reduces per-GPU memory by sharding layer weights, but increases communication volume between layers. Best for **large models** that don't fit on single devices.

**Sequence parallelism** reduces activation memory for long sequences without weight sharding overhead. Best for **long videos** where activations dominate memory usage.

The optimal configuration depends on your hardware:

| Scenario | Recommended Configuration |
|----------|---------------------------|
| 70B+ parameter model, short videos | High tensor parallel (4-8), sequence parallel = 1 |
| 7B parameter model, 10+ second videos | Tensor parallel = 1-2, high sequence parallel (4-8) |
| Balanced scaling | Tensor parallel = 2, sequence parallel = 2-4 |

## Checkpoint Sharding and Loading

The [`model_paths.py`](https://github.com/Lightricks/LTX-2/blob/main/model_paths.py) utility ensures each GPU loads only its required weight shards. For a model with tensor-parallel size 4, the checkpoint directory contains files like:

```

model-00001-of-00004.safetensors  # TP rank 0

model-00002-of-00004.safetensors  # TP rank 1

model-00003-of-00004.safetensors  # TP rank 2

model-00004-of-00004.safetensors  # TP rank 3

```

At load time, `resolve_rank_checkpoint()` maps the current distributed rank to its corresponding shard files, avoiding redundant memory usage.

## Quantization with Parallelism

The [`quantization_factory.py`](https://github.com/Lightricks/LTX-2/blob/main/quantization_factory.py) module provides 8-bit and 4-bit quantizers that work correctly with tensor-parallel shards. When quantization is enabled, each rank quantizes its local weight slice independently, preserving the parallel communication patterns.

## Summary

- **Tensor parallelism** shards model weights across GPUs using All-to-All kernels in [`ltx_kernels/all_to_all.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_kernels/all_to_all.py).

- **Sequence parallelism** splits long sequences via utilities in [`ltx_pipelines/utils/res2s.py`](https://github.com/Lightricks/LTX-2/blob/main/ltx_pipelines/utils/res2s.py).

- Launch via `accelerate` with [`fsdp_compile.yaml`](https://github.com/Lightricks/LTX-2/blob/main/fsdp_compile.yaml), setting `TENSOR_PARALLEL_SIZE` and `SEQ_PARALLEL_SIZE` environment variables.

- Each rank loads only required checkpoint shards through [`model_paths.py`](https://github.com/Lightricks/LTX-2/blob/main/model_paths.py).

- Combine both parallelism modes to scale LTX-2 to large models and long videos simultaneously.

## Frequently Asked Questions

### What hardware is required for multi-GPU inference in LTX-2?

LTX-2 multi-GPU inference requires **NVIDIA GPUs with NVLink or high-bandwidth interconnect** for efficient All-to-All communication in tensor parallelism. The code uses NCCL backend exclusively. Inference has been validated on A100, H100, and RTX 4090 clusters.

### Can I use tensor parallelism without sequence parallelism?

Yes. Set `SEQ_PARALLEL_SIZE=1` to disable sequence parallelism while keeping tensor parallelism active. This is the default configuration for single-GPU sequence lengths that fit in memory.

### How do I know if my sequence is long enough to benefit from sequence parallelism?

Enable sequence parallelism when generating videos longer than **4 seconds at 24fps** (roughly 96 frames) or when you encounter **out-of-memory errors during the denoising loop**. The [`res2s.py`](https://github.com/Lightricks/LTX-2/blob/main/res2s.py) utilities automatically handle the splitting when the parallelism is configured.

### Does multi-GPU inference work with quantized models?

Yes. The [`quantization_factory.py`](https://github.com/Lightricks/LTX-2/blob/main/quantization_factory.py) utilities provide per-rank quantizers that operate on sharded weights. Apply quantization flags in the accelerate config or pass `--load-in-8bit` to the inference script.