How to Handle Large Batch Sizes with Memory Constraints Using Batch Splitting in LTX-2

Use the BatchSplitAdapter class from ltx_core/batch_split.py to automatically split large input batches into smaller sub-batches during forward passes, reducing peak GPU activation memory while preserving the original model semantics and API.

LTX-2's transformer-based video generation models can rapidly exhaust GPU memory when processing large batches, creating a bottleneck for both training and inference workflows. To solve this without modifying model architecture or resorting to multi-GPU setups, the Lightricks/LTX-2 codebase provides a transparent batch-splitting mechanism that wraps any X0-model interface compliant object. This adapter intercepts forward calls, divides the work into memory-safe chunks, and reconstructs the results, allowing you to scale batch sizes beyond your hardware's native capacity.

Why Large Batches Exhaust GPU Memory in LTX-2

Transformer models in LTX-2 store extensive activation data during forward passes, causing memory consumption to scale linearly with batch size. When processing video and audio modalities simultaneously, a single forward pass with batch size B can exceed VRAM limits, triggering out-of-memory errors. The BatchSplitAdapter addresses this by decomposing the computation into sequential sub-batches of size M, where only M samples reside in memory at any given moment.

How BatchSplitAdapter Works

Chunking Strategy and Memory Management

The adapter implements a transparent splitting algorithm in packages/ltx-core/src/ltx_core/batch_split.py. When invoked, it first computes optimal chunk sizes via _get_chunk_sizes (lines 57-62), dividing the total batch size B into ⌊B/M⌋ full chunks plus one remainder chunk. This ensures no sample is processed with unnecessary padding.

Splitting Modalities and Perturbations

For each chunk, the adapter splits video and audio tensors using torch.split (lines 78-80) and decomposes the BatchedPerturbationConfig along the batch dimension via _split_perturbations. This maintains alignment between data and configuration across all sub-batches.

Execution and Reconstruction

The wrapped model processes each chunk independently (lines 82-85), and the adapter concatenates results using _merge_tensors (lines 87-88). The final output preserves the original batch dimension, making the splitting transparent to downstream code. Attribute access remains proxied to the underlying model, so methods like state_dict() and eval() function normally.

Implementing Batch Splitting in LTX-2

Wrapping a Model Manually

To use the adapter directly, instantiate BatchSplitAdapter with your transformer model and a max_batch_size that fits your GPU memory:

from ltx_core.batch_split import BatchSplitAdapter
from ltx_core.model.transformer import X0Model
import torch

# Initialize your transformer

model = X0Model(...)

# Wrap with batch splitting (max 2 samples per forward pass)

adapter = BatchSplitAdapter(model, max_batch_size=2)

# Prepare inputs (batch size 5)

video = Modality(latent=torch.randn(5, 16, 64, 64))  # (B, C, H, W)

audio = None
perturb = BatchedPerturbationConfig([])

# Automatic splitting: processes as 2 + 2 + 1 samples

denoised_video, denoised_audio = adapter(
    video=video, 
    audio=audio, 
    perturbations=perturb
)

Integration with BlockRun Pipelines

For high-level inference, the BlockRun class in packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py accepts a max_batch_size parameter that internally wraps the transformer:

from ltx_pipelines.utils.blocks import BlockRun

# Configure pipeline with batch splitting

pipeline = BlockRun(
    checkpoint_path="model.ckpt",
    offload_mode=OffloadMode.NONE,
    # ... other parameters

)

# Process 8 frames in chunks of 4

video_state, audio_state = pipeline(
    denoiser=my_denoiser,
    sigmas=my_sigmas,
    noiser=my_noiser,
    width=640,
    height=360,
    frames=32,
    fps=30.0,
    video=my_video_tensor,
    audio=None,
    max_batch_size=4,  # Activates BatchSplitAdapter internally

)

When max_batch_size=4 is specified, the pipeline processes an 8-frame batch as two sequential 4-frame passes, cutting peak activation memory by approximately half while maintaining output consistency. This pattern is also utilized in packages/ltx-trainer/src/ltx_trainer/trainer.py for memory-constrained training scenarios.

Performance Characteristics and Constraints

The adapter only reduces peak activation memory; the total computational cost increases proportionally with the number of chunks since each sub-batch requires a separate forward pass. The design is deliberately transparent to the X0-model interface, meaning the wrapped object maintains the signature (video, audio, perturbations) → (denoised_video, denoised_audio) regardless of how many internal splits occur.

Summary

  • Wrap any X0-model with BatchSplitAdapter from ltx_core/batch_split.py to enable automatic batch splitting without API changes.
  • Set max_batch_size to a value that fits safely within GPU memory (typically 1-4 for high-resolution video generation).
  • Preserve semantics—the adapter maintains the original forward signature and produces mathematically identical outputs to full-batch processing.
  • Integrate seamlessly with BlockRun by passing max_batch_size to the pipeline constructor, or use directly in custom training loops.
  • Expect trade-offs—memory savings come at the cost of sequential forward passes, increasing total computation time proportional to the number of chunks.

Frequently Asked Questions

Does batch splitting change the model's output quality?

No. The BatchSplitAdapter preserves mathematical equivalence by processing identical input data in smaller chunks and concatenating results along the batch dimension. Since transformer forward passes in LTX-2 are independent across batch dimensions, the final output tensors match those produced by a full-batch forward pass exactly.

Can I use batch splitting during training as well as inference?

Yes. The adapter is utilized in packages/ltx-trainer/src/ltx_trainer/trainer.py for training scenarios where large batch sizes would otherwise exceed GPU capacity. You can wrap your training model with BatchSplitAdapter to handle large effective batch sizes while keeping the per-step memory footprint manageable.

What happens if my batch size is not evenly divisible by max_batch_size?

The _get_chunk_sizes method handles remainders automatically. For a batch size of 5 and max_batch_size of 2, it creates chunk sizes [2, 2, 1], processing the remainder chunk separately without padding or truncation, ensuring all samples are processed exactly once.

Is there performance overhead beyond the extra forward passes?

Minimal. The splitting and merging operations (torch.split and concatenation via _merge_tensors) execute on the CPU and introduce negligible latency compared to the GPU computation time. The primary trade-off is sequential execution of chunks versus parallel processing, which is unavoidable when memory constraints prevent full batch processing.

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 →