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 fromB × L × CtoB × (L / sp_size) × C - KV‑cache optimization: The
AllToAllWithGradclass inwan_5b/distributed/sp_training.pyexchanges only necessary attention pieces during the forward pass, avoiding full‑sequence storage on each device - VAE efficiency: The
chunk_halo_metautility 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.pyenables 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →