How to Optimize Memory with Block Streaming and Batch Splitting in LTX-2
Combine LTX-2's block streaming (for weight memory) and batch splitting (for activation memory) to run large transformer models on GPUs with 8-10 GB of VRAM.
LTX-2's transformer architecture can overwhelm GPU memory with static weights and large batch activations. The repository provides two orthogonal mechanisms—block streaming and batch splitting—that work together to dramatically reduce memory footprint without sacrificing model capability. This guide explains exactly how to implement both techniques using the actual source code from Lightricks/LTX-2.
Understanding the Memory Problem in LTX-2
LTX-2 consumes GPU memory through two distinct channels:
- Static weight residency: All transformer blocks remain loaded on GPU throughout inference
- Large activation batches: Forward passes over big batches create massive intermediate tensors
The solution addresses each problem separately. Block streaming limits resident weights, while batch splitting caps activation size per forward pass.
Block Streaming: On-Demand Weight Loading
Block streaming keeps only a small number of transformer blocks in GPU memory at any moment, streaming others from CPU storage as needed.
Core Components
The implementation spans two files in ltx_core/block_streaming/:
provider.py:WeightsProvidermanages GPU buffer allocation and weight copyingwrapper.py:BlockStreamingWrapperinjects streaming logic via forward hooks
How WeightsProvider Works
The provider maintains an LRU-cached BufferPool. When block idx is requested:
- Cache check: Return existing GPU buffer if present
- Eviction: Release oldest buffer if pool is full
- Copy: Transfer weights from CPU
WeightSourceto GPU via_copy_to_gpu
# From ltx_core/block_streaming/provider.py
# The provider.get(idx) call implements this flow
provider.get(idx) # Fetches block weights, evicting old blocks as needed
How BlockStreamingWrapper Works
The wrapper registers hooks on each transformer block:
_pre_hook: Callsprovider.get(idx)and assigns weights viaassign_tensor_to_module_post_hook: Signals completion, allowing buffer reuse
Only the currently executing block (plus prefetch) occupies GPU memory. With pool_capacity=1, weight memory drops to ~1/N of the full model where N is total blocks.
Batch Splitting: Chunked Activation Processing
Batch splitting divides large input batches into smaller chunks, processing each separately to bound peak activation memory.
BatchSplitAdapter Implementation
Located in ltx_core/batch_split.py, the adapter:
- Checks incoming batch size against
max_batch_size - For oversized batches, computes chunk sizes with
_get_chunk_sizes - Splits video/audio tensors with
tensor.split(sizes) - Slices perturbations via
_split_perturbations - Processes chunks sequentially, merging outputs with
_merge_tensors
# From ltx_core/batch_split.py conceptual flow
if batch_size <= max_batch_size:
return model(video, audio, perturbations) # Direct pass
else:
chunks = video.split(chunk_sizes)
outputs = [model(chunk, audio_chunk, pert_chunk) for chunk, ... in zip(...)]
return _merge_tensors(outputs) # Concatenate results
Each chunk allocates activations only for its samples, keeping peak memory proportional to max_batch_size rather than total batch size.
Complete Memory-Optimized Setup
Combine both mechanisms for maximum memory reduction:
from ltx_core.block_streaming.builder import StreamingModelBuilder
from ltx_core.batch_split import BatchSplitAdapter
from ltx_core.model.transformer.modality import Modality
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
import torch
# Step 1: Build block-streaming model
# pool_capacity=1 keeps only ONE block resident (maximum weight savings)
builder = StreamingModelBuilder.from_checkpoint(
checkpoint_path="path/to/model.safetensors",
target_device=torch.device("cuda:0"),
pool_capacity=1, # Adjust: higher reduces copy overhead
)
streaming_model = builder.build() # Returns BlockStreamingWrapper
# Step 2: Wrap with batch splitting
# max_batch_size controls activation memory per forward pass
batch_adapter = BatchSplitAdapter(streaming_model, max_batch_size=4)
# Step 3: Run inference on large batch
video = Modality(latent=torch.randn(32, 3, 64, 64, device="cuda"))
audio = None
perturb = BatchedPerturbationConfig(...) # Optional guidance
denoised_video, denoised_audio = batch_adapter(
video=video,
audio=audio,
perturbations=perturb
)
# 32 samples processed as 8 chunks of 4, with only 1 transformer block resident
Tuning Parameters for Your Hardware
| Parameter | Location | Effect | Typical Values |
|---|---|---|---|
pool_capacity |
StreamingModelBuilder |
Resident blocks (weight memory) | 1 for max savings; 2-3 if copy bandwidth limits throughput |
max_batch_size |
BatchSplitAdapter |
Samples per forward (activation memory) | 4-8 for 16 GB GPUs; 16-32 for 24 GB GPUs |
PYTORCH_CUDA_ALLOC_CONF |
Environment | Allows expandable CUDA segments | expandable_segments:True recommended |
Increase pool_capacity if GPU utilization drops due to weight copying overhead. Decrease max_batch_size if you hit OOM during forward passes.
Key Implementation Files
| File | Path | Purpose |
|---|---|---|
provider.py |
ltx_core/block_streaming/provider.py |
WeightsProvider with buffer pool and LoRA fusion |
wrapper.py |
ltx_core/block_streaming/wrapper.py |
BlockStreamingWrapper with pre/post hooks |
batch_split.py |
ltx_core/batch_split.py |
BatchSplitAdapter for chunked processing |
builder.py |
ltx_core/block_streaming/builder.py |
Convenience constructor from Safetensors checkpoints |
Summary
- Block streaming in
ltx_core/block_streaming/wrapper.pylimits static weight memory via on-demand GPU buffer loading fromWeightsProvider - Batch splitting in
ltx_core/batch_split.pycaps activation memory by chunking large batches throughBatchSplitAdapter - The mechanisms are orthogonal: streaming saves weight memory, splitting saves activation memory
- Combine with
StreamingModelBuilder(pool_capacity=1)wrapped inBatchSplitAdapter(max_batch_size=4-8)for 8-10 GB GPU inference - Tune
pool_capacityagainst copy overhead andmax_batch_sizeagainst activation limits for your hardware
Frequently Asked Questions
What GPU memory is needed to run LTX-2 with block streaming and batch splitting?
With pool_capacity=1 and max_batch_size=4, LTX-2 runs on GPUs with 8-10 GB of VRAM. The full model typically requires 24+ GB. Weight memory drops to ~1/N of baseline (N = number of transformer blocks), while activation memory scales with max_batch_size rather than total batch size.
Does block streaming slow down inference?
Block streaming adds weight copy overhead from CPU to GPU. With pool_capacity=1, this overhead is highest—each block transfer happens during forward pass. Increasing to pool_capacity=2 or 3 hides latency via prefetching without substantially increasing memory. The default WeightsProvider LRU cache minimizes redundant copies for repeated block access patterns.
How do I choose the right max_batch_size for batch splitting?
Start with max_batch_size=4 and increase until you observe OOM errors or reach your GPU's comfortable utilization. Larger values improve throughput (fewer forward passes) but increase peak activation memory. For 16 GB GPUs, 4-8 works well; for 24 GB, try 16-32. Profile with torch.cuda.max_memory_allocated() to find your ceiling.
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 →