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

> Optimize AR training memory with LongLive sequence parallelism. Shard temporal data across GPUs to slash activation memory from 40GB+ to 8-16GB per device.

- Repository: [NVIDIA Research Projects/LongLive](https://github.com/NVlabs/LongLive)
- Tags: performance
- Published: 2026-05-24

---

**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`](https://github.com/NVlabs/LongLive/blob/main/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`](https://github.com/NVlabs/LongLive/blob/main/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`](https://github.com/NVlabs/LongLive/blob/main/wan_5b/distributed/sequence_parallel.py) (lines 89‑119), the implementation pads the global sequence once, then distributes slices:

```python

# 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`](https://github.com/NVlabs/LongLive/blob/main/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:

```python
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`](https://github.com/NVlabs/LongLive/blob/main/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:

```python
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`](https://github.com/NVlabs/LongLive/blob/main/configs/wan_ti2v_5B.py)):

```python
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`:

```bash
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:

```python
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`](https://github.com/NVlabs/LongLive/blob/main/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`](https://github.com/NVlabs/LongLive/blob/main/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`](https://github.com/NVlabs/LongLive/blob/main/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`](https://github.com/NVlabs/LongLive/blob/main/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.