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

> Overcome memory limits with LTX-2 batch splitting. Use BatchSplitAdapter to reduce GPU activation memory for large batches without changing model semantics.

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

---

**Use the `BatchSplitAdapter` class from [`ltx_core/batch_split.py`](https://github.com/Lightricks/LTX-2/blob/main/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`](https://github.com/Lightricks/LTX-2/blob/main/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:

```python
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`](https://github.com/Lightricks/LTX-2/blob/main/packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py) accepts a `max_batch_size` parameter that internally wraps the transformer:

```python
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`](https://github.com/Lightricks/LTX-2/blob/main/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`](https://github.com/Lightricks/LTX-2/blob/main/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`](https://github.com/Lightricks/LTX-2/blob/main/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.