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: WeightsProvider manages GPU buffer allocation and weight copying
  • wrapper.py: BlockStreamingWrapper injects streaming logic via forward hooks

How WeightsProvider Works

The provider maintains an LRU-cached BufferPool. When block idx is requested:

  1. Cache check: Return existing GPU buffer if present
  2. Eviction: Release oldest buffer if pool is full
  3. Copy: Transfer weights from CPU WeightSource to 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: Calls provider.get(idx) and assigns weights via assign_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:

  1. Checks incoming batch size against max_batch_size
  2. For oversized batches, computes chunk sizes with _get_chunk_sizes
  3. Splits video/audio tensors with tensor.split(sizes)
  4. Slices perturbations via _split_perturbations
  5. 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.py limits static weight memory via on-demand GPU buffer loading from WeightsProvider
  • Batch splitting in ltx_core/batch_split.py caps activation memory by chunking large batches through BatchSplitAdapter
  • The mechanisms are orthogonal: streaming saves weight memory, splitting saves activation memory
  • Combine with StreamingModelBuilder(pool_capacity=1) wrapped in BatchSplitAdapter(max_batch_size=4-8) for 8-10 GB GPU inference
  • Tune pool_capacity against copy overhead and max_batch_size against 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:

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 →