How to Optimize Memory Usage During AR Training with Sequence Parallelism in LongLive

Enable Sequence Parallelism by setting sequence_parallel_size to shard the temporal dimension across GPUs, which reduces activation memory from 40 GB+ to 8‑16 GB per device without changing hyper‑parameters.

LongLive trains large‑scale autoregressive video diffusion models that routinely exceed single‑GPU memory capacity. The repository implements Sequence Parallelism (SP)—a tensor parallelism strategy that distributes video frames across multiple GPUs to optimize memory usage during AR training with sequence parallelism in LongLive. Each device stores only a fraction of the token sequence, dramatically cutting activation memory while preserving batch size and model quality.

How Sequence Parallelism Reduces Memory in LongLive

Sequence Parallelism in LongLive operates by shredding the temporal dimension of video data across a group of GPUs. In wan_5b/distributed/sequence_parallel.py, the sp_dit_causal_forward_train function (lines 45‑135) implements the core sharding logic:

  • Activation sharding: Each SP rank holds local_frames = total_frames / sp_world_size, reducing per‑device activation storage from B × L × C to B × (L / sp_size) × C
  • KV‑cache optimization: The AllToAllWithGrad class in wan_5b/distributed/sp_training.py exchanges only necessary attention pieces during the forward pass, avoiding full‑sequence storage on each device
  • VAE efficiency: The chunk_halo_meta utility ensures only receptive‑field "halo" frames transfer between ranks during encoding, rather than full‑resolution copies

This architecture achieves approximately 4×‑6× reduction in peak memory, allowing a 5B‑parameter model to train on consumer GPUs (RTX 3090/4090) that would otherwise require enterprise‑grade hardware.

Key Components for Memory Optimization

Temporal Chunking in sp_dit_causal_forward_train

The forward pass partitions tensors along the temporal axis (dimension 1) immediately after patch embedding. In wan_5b/distributed/sequence_parallel.py (lines 89‑119), the implementation pads the global sequence once, then distributes slices:


# Patch embedding and global padding

x = torch.cat([
    torch.cat([u, u.new_zeros(1, max_len - u.size(1), u.size(2))], dim=1)
    for u in x
])

# Shard across SP ranks

x = torch.chunk(x, get_world_size(), dim=1)[get_rank()]
e = torch.chunk(e, get_world_size(), dim=1)[get_rank()]

Each rank processes only its local chunk, keeping the rest of the sequence off‑device.

All‑to‑All Communication with Gradients

The distributed_flex_attention function in wan_5b/distributed/sp_training.py (lines 61‑80) enables full‑sequence attention while maintaining gradients. Before the FlexAttention kernel, the all_to_all_with_grad helper exchanges sharded heads across the SP group:

x = distributed_flex_attention(
    roped_q, roped_k, v, block_mask
)

The scatter_dim=2, gather_dim=1 parameters ensure that query/key/value tensors are correctly gathered for global attention computation, then scattered back to their original layout for the backward pass.

Halo Exchange for VAE Latents

When encoding raw video through the VAE, encode_raw_video_latents (lines 84‑107 in wan_5b/distributed/sp_training.py) uses chunk_halo_meta to compute minimal frame windows. Only the overlapping "halo" regions required by the VAE receptive field transfer between ranks via scatter_frame_windows_for_chunk_halo, eliminating redundant full‑tensor transfers.

Gradient Checkpointing Integration

Inside the transformer block loop, conditional checkpointing further reduces memory:

if torch.is_grad_enabled() and self.gradient_checkpointing:
    x = torch.utils.checkpoint.checkpoint(
        create_custom_forward(block), x, **kwargs, use_reentrant=False
    )
else:
    x = block(x, **kwargs)

This pattern, found in sp_dit_causal_forward_train, keeps only active block tensors in memory during the forward pass.

Implementation Guide

Configuration Setup

To activate Sequence Parallelism, set the following in your config file (e.g., configs/wan_ti2v_5B.py):

model_kwargs:
  model_name: "wan_ti2v_5B"
training:
  sequence_parallel_size: 4      # Must equal nproc_per_node

  gradient_checkpointing: true   # Optional, for additional memory savings

The SequenceParallelHelper class reads sequence_parallel_size and initializes the SP group automatically when sp_size > 1.

Launching Distributed Training

Launch training with torchrun, ensuring nproc_per_node matches your sequence_parallel_size:

torchrun --nnodes=1 --nproc_per_node=4 \
    -m main.train \
    --config_path=wan_5b/configs/wan_ti2v_5B.py \
    --output_dir=./output_sp

Manual Tensor Partitioning

For custom training loops, manually partition inputs using the helper utilities:

from wan_5b.trainer.sp_helper import SequenceParallelHelper

# Initialize helper with trainer instance

sp_helper = SequenceParallelHelper(trainer)

# Partition batch across temporal dimension

clean_latent, cond, shape = sp_helper.partition_training_inputs(
    image_or_video_shape=cfg.image_or_video_shape,
    clean_latent=latent_tensor,
    conditional_dict=cond_dict,
)

# Forward with SP-enabled model

output = model.sp_dit_causal_forward_train(
    x=clean_latent,
    t=timestep_tensor,
    context=cond,
    seq_len=cfg.max_seq_len,
    clean_x=None,
    aug_t=None,
)

The helper automatically shards tensors via _chunk_tensor (dimension 1) and updates sequence length metadata.

Memory Profiling and Optimization Tips

Profile communication overhead by setting NCCL_DEBUG=INFO during training to verify all‑to‑all efficiency across SP ranks.

Capture GPU timelines by enabling NVTX_ENABLED=1 before launching training. The repository includes NVTXRange markers that appear in Nsight Systems profiles, helping identify bottlenecks in distributed_flex_attention or halo exchange operations.

Tune the VAE halo size if using custom VAE architectures. Adjust DEFAULT_SP_VAE_HALO_LATENTS (default = 28) in wan_5b/distributed/sp_training.py to match your encoder's receptive field—larger halos increase memory slightly but ensure correct boundary handling.

Summary

  • Sequence Parallelism shards the temporal dimension of video tensors across GPUs, reducing per‑device activation memory by the factor of sp_world_size.
  • AllToAllWithGrad in wan_5b/distributed/sp_training.py enables full‑sequence FlexAttention while keeping gradients flow‑compatible with distributed training.
  • Halo exchange minimizes VAE communication overhead by transferring only necessary boundary frames rather than full latents.
  • Gradient checkpointing integrates seamlessly with SP blocks to trade computation for additional memory savings.
  • Together, these techniques reduce AR training memory requirements from 40 GB+ to 8‑16 GB per GPU as implemented in NVlabs/LongLive.

Frequently Asked Questions

What is the minimum GPU memory required for AR training with Sequence Parallelism in LongLive?

With sequence_parallel_size set to 4 and gradient checkpointing enabled, LongLive can train 5B‑parameter autoregressive video models on GPUs with 8‑16 GB of VRAM, such as the RTX 3090 or RTX 4090. Without Sequence Parallelism, the same configuration requires 40 GB+ per device.

How does Sequence Parallelism differ from FSDP in LongLive?

Sequence Parallelism shards the temporal dimension of activations across GPUs (tensor parallelism), while FSDP (Fully Sharded Data Parallel) shards model parameters and optimizer states across data parallel ranks. According to the source code in wan_5b/distributed/fsdp.py, these strategies are orthogonal and can be combined—FSDP handles weight sharding while SP handles sequence‑length sharding.

Can I combine Sequence Parallelism with gradient checkpointing?

Yes. The sp_dit_causal_forward_train function in wan_5b/distributed/sequence_parallel.py conditionally applies torch.utils.checkpoint.checkpoint when self.gradient_checkpointing is enabled. This stacks seamlessly with Sequence Parallelism because checkpointing operates within the local SP rank's slice, further reducing activation memory without interfering with all‑to‑all communication patterns.

What is the purpose of the halo exchange in the VAE encoding step?

The chunk_halo_meta utility computes exact frame indices needed for VAE encoding across SP boundaries. Since VAE encoders require overlapping receptive fields, scatter_frame_windows_for_chunk_halo transfers only these "halo" frames (approximately 28 latents by default) between ranks during encode_raw_video_latents. This avoids materializing full‑resolution video tensors on every GPU, preserving the memory benefits of Sequence Parallelism during the initial encoding phase.

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 →