How to Optimize Memory in LTX-2 Using BlockStreamingWrapper and BatchSplitAdapter

Use BlockStreamingWrapper to stream transformer block weights on-demand from CPU/disk, and BatchSplitAdapter to split large batches into smaller sub-batches, keeping parameter and activation memory respectively within GPU limits.

LTX-2's inference pipelines handle massive transformer models with hundreds of millions of parameters and high-resolution video/audio tensors. Without optimization, these workloads exhaust GPU memory through both static weights (model parameters) and dynamic activations (intermediate tensors). The BlockStreamingWrapper and BatchSplitAdapter classes in ltx-core provide complementary solutions: weight streaming caps the parameter footprint, while batch splitting limits activation peaks.


What BlockStreamingWrapper Does for Weight Memory

The BlockStreamingWrapper class in ltx_core/block_streaming/wrapper.py solves the parameter memory problem by ensuring only 1-2 transformer blocks reside on GPU at any moment. The rest remain in pinned CPU memory or on disk until needed.

How Weight Streaming Works

When StreamingModelBuilder constructs a transformer, it creates a WeightsProvider that knows how to load block parameters from checkpoints. The wrapper orchestrates the handoff:

  1. Pre-hook registration (_pre_hook): Before a block's forward pass, the provider loads that block's weights into GPU buffers via assign_tensor_to_module.
  2. Post-hook cleanup (_post_hook): After the CUDA kernel finishes, the buffer is released for the next block.

A 30-block LTX-2 model that would normally require ~24 GB for parameters can run on an 8 GB GPU because the static memory cost becomes constant—roughly one block plus a prefetch buffer, regardless of total model size.

Key Implementation Details

  • Lines 20-30 of wrapper.py handle hook registration
  • provider.get(block_idx) fetches weights on-demand
  • Only current block weights + optional prefetch buffer occupy cuda:0

What BatchSplitAdapter Does for Activation Memory

The BatchSplitAdapter class in ltx_core/batch_split.py solves the activation memory problem. Guided denoisers combine multiple guidance passes (CFG, STG, modality scaling) into a single forward call with enlarged batch dimensions (B=2-4). Without intervention, activation memory grows linearly with batch size and quickly exhausts VRAM.

How Batch Splitting Works

The adapter transparently partitions incoming tensors:

sizes = self._get_chunk_sizes(batch_size)   # e.g., [1, 1, 1, 1] for B=4, max_batch=1

v_chunks = video.split(sizes)              # split video tensor

a_chunks = audio.split(sizes)              # split audio tensor

p_chunks = _split_perturbations(perturbations, sizes)   # split masks

Each chunk processes sequentially through the underlying model (which may itself be a BlockStreamingWrapper). Results are concatenated via _merge_tensors so the external API returns the original batch shape unchanged.

Key Implementation Details

  • Lines 54-98 of batch_split.py implement chunking logic
  • max_batch_size controls the trade-off between memory and throughput
  • The adapter preserves the forward(video, audio, perturbations) signature of X0Model

Combining Both Optimizations

Weight streaming and batch splitting are orthogonal—they address different memory pools and can be stacked. The LTX-2 pipelines use exactly this composition in ltx_pipelines/utils/blocks.py (lines 645-666).

Manual Setup Example

from ltx_core.block_streaming import StreamingModelBuilder
from ltx_core.batch_split import BatchSplitAdapter

# Step 1: Build streaming transformer (weights on CPU/disk)

builder = StreamingModelBuilder(
    model_path=checkpoint_path,
    model_class_configurator=MyTransformerConfigurator(),
    model_sd_ops=sd_ops,
    module_ops=module_ops,
    registry=model_registry,
)
streaming_transformer = builder.build()  # Returns BlockStreamingWrapper

# Step 2: Wrap for batch size limiting

max_batch = 1  # Tune per GPU capacity

transformer = BatchSplitAdapter(streaming_transformer, max_batch_size=max_batch)

# Step 3: Use normally—the composition is transparent

denoised_video, denoised_audio = transformer(
    video=video_modality,
    audio=audio_modality,
    perturbations=perturb_cfg,
)

Pipeline Integration Example

The DiffusionStage class automatically composes these layers:


# Inside ltx_pipelines/utils/blocks.py, line 665

wrapped = BatchSplitAdapter(transformer, max_batch_size=max_batch_size)

video_state, audio_state = loop(
    sigmas=sigmas,
    video_state=video_state,
    audio_state=audio_state,
    stepper=stepper,
    transformer=wrapped,  # Denoiser receives the adapted model

    denoiser=denoiser,
)

Pipeline authors only adjust max_batch_size. The rest of the codebase remains unchanged because both wrapper and adapter expose identical interfaces to plain X0Model.


Tuning max_batch_size for Your Hardware

GPU VRAM Recommended max_batch_size Expected Behavior
4-8 GB 1 Serial processing, minimal memory pressure
12-16 GB 1-2 Balance throughput and safety margin
24+ GB 2-4 Higher throughput, streaming still beneficial for largest models

Higher max_batch_size reduces kernel launch overhead but increases activation memory. Weight streaming remains valuable even on large GPUs when running the full LTX-2 variants.


Debugging Memory Usage

Verify optimization effectiveness with PyTorch's memory tools:

import torch

# Baseline measurement

torch.cuda.empty_cache()
print(torch.cuda.memory_summary(device=None, abbreviated=False))

# Run inference with your configuration

# ...

# Compare allocated memory before/after adding BatchSplitAdapter

The allocated metric should drop proportionally to max_batch_size reduction. Streaming effectiveness appears in the reserved section—parameter memory stays flat regardless of model depth.


Summary

  • BlockStreamingWrapper streams transformer block weights from WeightsProvider on-demand, keeping GPU parameter memory constant at ~1-2 blocks regardless of total model size
  • BatchSplitAdapter splits large guidance batches into sequential sub-batches, capping activation memory at the per-sample peak
  • Stack both optimizations by wrapping a streaming transformer with the batch adapter, as implemented in DiffusionStage (lines 645-666 of blocks.py)
  • Tune max_batch_size to your GPU's safe capacity—default 1 works on 4-8 GB cards, higher values trade memory for speed on larger hardware
  • Source files: ltx_core/block_streaming/wrapper.py, ltx_core/batch_split.py, ltx_core/block_streaming/builder.py, and ltx_pipelines/utils/blocks.py

Frequently Asked Questions

What is the difference between weight streaming and batch splitting in LTX-2?

Weight streaming (via BlockStreamingWrapper) controls static memory—the model parameters themselves. Only active blocks load into GPU memory, so a 30-block transformer uses the same VRAM as a 2-block model. Batch splitting (via BatchSplitAdapter) controls dynamic memory—the intermediate activations produced during forward passes. Large batches fragment into smaller sub-batches, so peak activation memory equals one sub-batch rather than the full batch. They address orthogonal memory pools and work best together.

Can I use BatchSplitAdapter without BlockStreamingWrapper?

Yes. BatchSplitAdapter wraps any X0Model conforming object, including standard (non-streaming) transformers. However, without streaming, you still pay the full parameter memory cost. The combination is most powerful: streaming keeps weights minimal, while splitting keeps activations bounded.

How does WeightsProvider know where to load weights from?

The provider is constructed by StreamingModelBuilder in ltx_core/block_streaming/builder.py and initialized from your checkpoint path. It manages file handles to the checkpoint (CPU memory or disk) and maintains a pool of GPU buffers. When _pre_hook requests block_idx, the provider copies that block's tensors into an available buffer and returns it.

What happens if max_batch_size exceeds the actual batch size?

The adapter detects this via _get_chunk_sizes and returns a single chunk containing the full batch. No splitting occurs, and overhead remains minimal. The optimization activates automatically only when needed.

Does batch splitting change inference results?

No. BatchSplitAdapter is deterministic and mathematically equivalent to full-batch processing. Sub-batches process sequentially through identical weights, and _merge_tensors concatenates outputs in the original order. The external API is unchanged—code consuming the wrapped model sees no difference except reduced memory pressure.

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 →