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

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. 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. 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 Collective communication for tensor-parallel weight sharding
Sequence utilities ltx_pipelines/utils/res2s.py Split/merge long token sequences across GPUs
Model path resolver ltx_pipelines/utils/model_paths.py Load correct checkpoint shards per rank
Accelerate config ltx-trainer/configs/accelerate/fsdp_compile.yaml Activates FSDP, tensor, and sequence parallelism
Inference script ltx-trainer/scripts/serve_captioner.py Entry point for distributed video generation
Quantization support ltx_pipelines/utils/quantization_factory.py 8-bit/4-bit quantizers compatible with parallel shards

Configuration: The Accelerate YAML File

The 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:

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:

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 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 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.

  • Sequence parallelism splits long sequences via utilities in ltx_pipelines/utils/res2s.py.

  • Launch via accelerate with fsdp_compile.yaml, setting TENSOR_PARALLEL_SIZE and SEQ_PARALLEL_SIZE environment variables.

  • Each rank loads only required checkpoint shards through 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 utilities automatically handle the splitting when the parallelism is configured.

Does multi-GPU inference work with quantized models?

Yes. The 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.

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 →