How to Configure Sequence Parallel Training with Balanced Workload Distribution in LongLive
To configure sequence parallel training with balanced workload distribution in the NVlabs/LongLive framework, set sequence_parallel_size to divide the total world size, ensure the temporal dimension of your video data is divisible by the product of sequence_parallel_size and num_frame_per_block, and specify the Wan2.2-TI2V-5B model architecture.
The LongLive repository by NVlabs implements sequence parallel training (SP) to distribute video generation workloads across multiple GPUs. This technique splits the temporal dimension of video sequences across SP ranks while maintaining data parallel (DP) groups across different GPUs, ensuring each processing unit handles an identical frame count for balanced computation.
Prerequisites and Configuration Parameters
LongLive's sequence parallel implementation is specifically designed for the 5B parameter video model. Proper configuration requires setting four key parameters in your training config file (e.g., configs/wan_ti2v_5B.py).
Required Config Parameters
| Parameter | Requirement | Purpose |
|---|---|---|
model_kwargs.model_name |
Must be "Wan2.2-TI2V-5B" |
Enables SP-aware attention kernels |
sequence_parallel_size |
Integer ≥ 1, must divide world size | Number of SP ranks per DP replica |
num_frame_per_block |
Integer that divides frames evenly | Frames per transformer block |
image_or_video_shape |
List [B, T, C, H, W] where T is divisible by sequence_parallel_size × num_frame_per_block |
Defines input tensor dimensions |
A typical YAML configuration looks like this:
model_kwargs:
model_name: Wan2.2-TI2V-5B
sequence_parallel_size: 4 # 4 SP ranks per DP replica
num_frame_per_block: 8
image_or_video_shape: [1, 81, 4, 256, 256] # T = 81 frames
Validation Checks in trainer/diffusion.py
The framework validates constraints immediately after initializing the distributed environment. In trainer/diffusion.py, three critical assertions enforce balanced workload distribution:
# trainer/diffusion.py (lines 99-119)
self.sequence_parallel_size = getattr(config, "sequence_parallel_size", 1)
world_size = dist.get_world_size()
self.data_parallel_size = world_size // self.sequence_parallel_size if self.sequence_parallel_size > 1 else world_size
# Model compatibility check
assert config.model_kwargs.model_name == "Wan2.2-TI2V-5B", (
f"sequence_parallel_size is only supported for Wan2.2-TI2V-5B model, but got {config.model_kwargs.model_name}"
)
# World size divisibility check
assert world_size % self.sequence_parallel_size == 0, (
f"world_size ({world_size}) must be divisible by sequence_parallel_size ({self.sequence_parallel_size})"
)
# Temporal dimension balance check
assert list(config.image_or_video_shape)[1] % (self.sequence_parallel_size * config.num_frame_per_block) == 0, (
f"image_or_video_shape[1] ({list(config.image_or_video_shape)[1]}) must be divisible by the product of "
f"sequence_parallel_size ({self.sequence_parallel_size}) and num_frame_per_block ({config.num_frame_per_block})"
)
These checks prevent runtime errors by ensuring the total frame count splits evenly across all SP ranks.
Process Group Setup
When sequence_parallel_size > 1, LongLive automatically constructs two distinct process group types: SP groups for sequence sharding within a DP replica, and DP groups for replicating sequence chunks across replicas.
SP and DP Group Creation
In trainer/diffusion.py, the trainer initializes these groups using torch.distributed:
# trainer/diffusion.py (group creation logic)
sp_size = self.sequence_parallel_size
dp_size = self.data_parallel_size
# Create Sequence Parallel groups (all-to-all within DP replica)
sp_groups = []
for g in range(dp_size):
ranks_g = list(range(g * sp_size, (g + 1) * sp_size))
sp_groups.append(dist.new_group(ranks=ranks_g))
self.sp_group = sp_groups[global_rank // sp_size]
set_sequence_parallel_group(self.sp_group) # Registers globally
# Create Data Parallel groups (same sequence chunk across replicas)
dp_groups = []
for k in range(sp_size):
ranks_k = [g * sp_size + k for g in range(dp_size)]
dp_groups.append(dist.new_group(ranks=ranks_k))
self.dp_group = dp_groups[global_rank % sp_size]
set_data_parallel_group(self.dp_group) # Registers globally
The global registration functions reside in wan_5b/distributed/sp_training.py:
# wan_5b/distributed/sp_training.py (lines 43-49)
def set_sequence_parallel_group(group):
"""Set the SP group used by SP rank/world-size and all-to-all helpers."""
global _sp_group
_sp_group = group
def set_data_parallel_group(group):
"""Set the DP group for gradient synchronization."""
global _dp_group
_dp_group = group
This architecture ensures that all-to-all communication for sequence parallelism remains confined within each DP replica, avoiding unnecessary cross-replica traffic.
Balanced Tensor Sharding
The SequenceParallelHelper class in wan_5b/distributed/sp_training.py handles automatic tensor partitioning to guarantee perfect workload balance.
The SequenceParallelHelper Class
The helper shards tensors along the temporal dimension (dim=1) using the _chunk_tensor method:
# wan_5b/distributed/sp_training.py (tensor sharding)
def _chunk_tensor(self, tensor, dim):
"""Split tensor evenly across SP ranks."""
return tensor.chunk(self.sp_size, dim=dim)[self.local_sp_rank()].contiguous()
During the forward pass, partition_training_inputs applies this sharding to latents, conditional embeddings, and loss masks:
# wan_5b/distributed/sp_training.py (lines 91-96)
if clean_latent is not None and not clean_latent_is_sharded:
clean_latent = self._chunk_tensor(clean_latent, dim=1) # Split temporal dimension
clean_latent_is_sharded = True
# Adjust shape metadata to reflect local shard size
image_or_video_shape[1] = image_or_video_shape[1] // self.sp_size
Because sharding uses torch.chunk with sp_size, each rank receives exactly T / sequence_parallel_size frames. Combined with the earlier divisibility assertions, this ensures perfectly balanced compute where every GPU processes an identical number of frames.
Launching Sequence Parallel Training
Execute training using torchrun or mpirun with the appropriate configuration:
# Example: 8 GPUs total, SP size = 4, DP size = 2
torchrun --nproc_per_node=8 \
train.py \
--config configs/wan_ti2v_5B.py \
--sequence_parallel_size 4 \
--num_frame_per_block 8 \
--image_or_video_shape "[1, 81, 4, 256, 256]"
Upon startup, the trainer logs the parallel configuration:
[SP] Sequence Parallel enabled, sp_size=4, dp_size=2, world_size=8
If any configuration constraints are violated, the process aborts immediately with descriptive error messages indicating which divisibility check failed.
Summary
- Sequence parallel training in LongLive requires the
Wan2.2-TI2V-5Bmodel and specific configuration ofsequence_parallel_sizeandnum_frame_per_block. - The framework validates that
world_size % sequence_parallel_size == 0and that the temporal dimension divides evenly bysequence_parallel_size × num_frame_per_block. - Process groups are automatically created in
trainer/diffusion.pyto isolate SP communication within DP replicas. - The
SequenceParallelHelperclass inwan_5b/distributed/sp_training.pyuses_chunk_tensorto split sequences evenly, ensuring balanced workload distribution. - Launch with
torchrun, ensuring yourimage_or_video_shapetemporal dimension matches the divisibility requirements for your chosen parallel configuration.
Frequently Asked Questions
What models support sequence parallel training in LongLive?
Only the Wan2.2-TI2V-5B model supports sequence parallel training. The framework explicitly checks config.model_kwargs.model_name in trainer/diffusion.py and raises an assertion error if you attempt to use SP with other model architectures. This restriction exists because SP requires specialized attention kernels (sp_causal_attn_forward) implemented specifically for this model's architecture.
How do I verify the workload is balanced across SP ranks?
You can inspect the local shard size at runtime. Each rank should have image_or_video_shape[1] / sequence_parallel_size frames after the partition_training_inputs function executes. Additionally, check the console output for the [SP] log line confirming sp_size and dp_size. If the temporal dimension assertions in trainer/diffusion.py pass, the workload is mathematically guaranteed to be balanced across all SP ranks.
Can I use arbitrary values for sequence_parallel_size?
No. The sequence_parallel_size must divide the total world size evenly (e.g., with 8 GPUs, valid values are 1, 2, 4, or 8). Additionally, the product sequence_parallel_size × num_frame_per_block must divide the total temporal length T specified in image_or_video_shape. These constraints ensure that torch.chunk can split tensors into equal-sized pieces without remainder.
What happens if my temporal dimension isn't divisible by the required product?
The training process aborts during initialization with an AssertionError from trainer/diffusion.py. The error message explicitly states that image_or_video_shape[1] must be divisible by the product of sequence_parallel_size and num_frame_per_block. You must adjust your video length, SP size, or frames per block to satisfy this divisibility requirement before training can begin.
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 →