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 GPUsTENSOR_PARALLEL_SIZE— number of GPUs in each tensor-parallel groupSEQ_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
acceleratewithfsdp_compile.yaml, settingTENSOR_PARALLEL_SIZEandSEQ_PARALLEL_SIZEenvironment 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →