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-5B model and specific configuration of sequence_parallel_size and num_frame_per_block.
  • The framework validates that world_size % sequence_parallel_size == 0 and that the temporal dimension divides evenly by sequence_parallel_size × num_frame_per_block.
  • Process groups are automatically created in trainer/diffusion.py to isolate SP communication within DP replicas.
  • The SequenceParallelHelper class in wan_5b/distributed/sp_training.py uses _chunk_tensor to split sequences evenly, ensuring balanced workload distribution.
  • Launch with torchrun, ensuring your image_or_video_shape temporal 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:

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 →