This document is relevant for: Trn2, Trn3
Context Parallelism in vLLM Omni Neuron#
Overview#
Context Parallelism (CP) in vLLM Omni Neuron enables efficient processing of long sequences by distributing the sequence dimension across multiple ranks. This implementation leverages vLLM Omni’s sequence parallelism infrastructure (sequence_parallel_size) to shard the input tokens while maintaining full attention computation. Because vLLM Omni enforces sequence_parallel_size = ulysses_degree * ring_degree, the CP degree is configured via ring_degree (with ulysses_degree left at its default of 1) — see Configuration.
Unlike traditional tensor parallelism that shards model parameters, context parallelism shards the sequence dimension, allowing each rank to process a subset of tokens while still computing attention over the full sequence. The communication strategy for attention can vary — ring attention (K/V chunks passed between ranks in a ring, so no rank materializes the full K/V), Full K/V AllGather (K/V gathered before attention so each rank computes local Q × full K/V), or DeepSpeed-Ulysses (All-to-All redistribution from sequence-partitioned to head-partitioned).
The current implementation uses ring-attention CP via a vendored ring-attention NKI kernel, which drives its own collective_permute around the CP group and keeps K/V memory at 1/cp_size. It falls back to Full K/V AllGather only where the ring kernel cannot run (CPU mode, fake-tensor tracing, NKI disabled).
Architecture#
Sequence Parallelism Integration#
Context parallelism is implemented through vLLM Omni’s sequence parallelism (SP) framework:
from vllm_omni.diffusion.distributed.parallel_state import get_sp_group
sp_group = get_sp_group()
self.cp_size = sp_group.world_size
self.cp_rank = sp_group.rank_in_group
self.cp_group = sp_group if self.cp_size > 1 else None
Process Groups#
Context parallelism operates alongside tensor parallelism (TP):
TP Group: Shards model parameters (attention heads, FFN dimensions)
CP Group: Shards the sequence dimension across ranks
Combined: Each rank has
(tp_rank, cp_rank)coordinates
Example with TP=4, CP=8 (32 ranks total):
Rank 0: tp_rank=0, cp_rank=0 (heads 0-9, tokens 0-127)
Rank 1: tp_rank=0, cp_rank=1 (heads 0-9, tokens 128-255)
...
Rank 7: tp_rank=0, cp_rank=7 (heads 0-9, tokens 896-1023)
Rank 8: tp_rank=1, cp_rank=0 (heads 10-19, tokens 0-127)
...
Rank 31: tp_rank=3, cp_rank=7 (heads 30-39, tokens 896-1023)
Model Implementation#
WanTransformer3DModel#
The main transformer model implements context parallelism through sequence sharding:
# Split sequence across CP ranks before transformer blocks
if self.cp_size > 1:
S = hidden_states.shape[1] # Total sequence length
if S % self.cp_size != 0:
raise ValueError(
f"Sequence length {S} is not divisible by cp_size {self.cp_size}. "
f"Choose a resolution/frame count that yields a divisible patch sequence length."
)
local_S = S // self.cp_size
start = self.cp_rank * local_S
hidden_states = hidden_states[:, start : start + local_S, :]
# Slice rotary embeddings to match local positions
freqs_cos, freqs_sin = rotary_emb
freqs_cos = freqs_cos[:, start : start + local_S, :, :]
freqs_sin = freqs_sin[:, start : start + local_S, :, :]
rotary_emb = (freqs_cos, freqs_sin)
# Process through transformer blocks
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb)
# Gather full sequence after transformer blocks
if self.cp_size > 1:
hidden_states = self.cp_group.all_gather(hidden_states.contiguous(), dim=1)
Self-Attention with Context Parallelism#
The WanSelfAttention module projects and RoPEs its local Q/K/V shard, then hands the CP policy to the shared wan_cp_self_attention core.
Ring attention (current implementation)#
K/V stay sharded (local_S tokens per rank). The ring-attention NKI kernel drives a collective_permute around the CP group, streaming each rank’s K/V chunk through every other rank so local Q attends the full sequence without ever materializing full K/V. Q remains sequence-partitioned — no output redistribution is needed.
Full K/V AllGather (fallback)#
Where the ring kernel cannot run (CPU mode, fake-tensor tracing, NKI disabled), each rank AllGathers K/V from all CP ranks to materialize the full sequence, then computes local Q × full K/V attention:
class WanSelfAttention(nn.Module):
def __init__(self, ...):
# CP group setup
sp_group = get_sp_group()
self.cp_size = sp_group.world_size
self.cp_group = sp_group if self.cp_size > 1 else None
def forward(self, hidden_states: torch.Tensor, rotary_emb=None) -> torch.Tensor:
# QKV projection (NKI kernel or matmul fallback)
if self._use_nki_qkv:
qkv = NF.qkv_proj(hidden_states, self.qkv_proj_weight, bias=self.qkv_proj_bias.unsqueeze(0))
else:
qkv = torch.matmul(hidden_states, self.qkv_proj_weight) + self.qkv_proj_bias
q, k, v = torch.tensor_split(qkv, self.qkv_split, dim=-1)
# QK-norm + reshape to multi-head: [B, S, H] -> [B, S, N, D]
query = self.norm_q(q).unflatten(2, (self.num_heads, self.head_dim))
key = self.norm_k(k).unflatten(2, (self.num_heads, self.head_dim))
value = v.unflatten(2, (self.num_heads, self.head_dim))
# Apply rotary embeddings to local tokens
if rotary_emb is not None:
freqs_cos, freqs_sin = rotary_emb
query = apply_rotary_emb_wan(query, freqs_cos, freqs_sin)
key = apply_rotary_emb_wan(key, freqs_cos, freqs_sin)
query = query.transpose(1, 2) # [B, N, local_S, D]
key = key.transpose(1, 2)
value = value.transpose(1, 2)
# CP: AllGather K/V across SP group to get full sequence
# key/value before: [B, N, local_S, D] -> after: [B, N, S, D]
if self.cp_size > 1:
key = self.cp_group.all_gather(key.contiguous(), dim=2)
value = self.cp_group.all_gather(value.contiguous(), dim=2)
# Attention: local Q [B, N, local_S, D] × full K/V [B, N, S, D]
hidden_states = _nf_attend(query, key, value, self.scale)
# Output projection via NKI kernel: [B, N, local_S, D] -> [B, local_S, H]
output = NF.o_proj(
hidden_states.transpose(2, 3), # [B, N, D, local_S]
self.o_proj_weight,
self.o_proj_bias.unsqueeze(0),
)
if self.tp_size > 1:
dist.all_reduce(output, op=dist.ReduceOp.SUM, group=self.tp_group)
return output
Comparison with Other CP Strategies#
Aspect |
Ring Attention (current) |
Full K/V AllGather (fallback) |
DeepSpeed-Ulysses |
|---|---|---|---|
Communication primitive |
|
AllGather on K/V |
All-to-All on Q,K,V + All-to-All on output |
Rounds per layer |
cp_size send/recv |
1 AllGather |
2 All-to-All |
K/V memory per rank |
1/cp_size of sequence |
Full sequence (redundant) |
Full sequence for local heads only |
Head constraint from CP |
None |
None |
|
Compute-comm overlap |
Overlaps attention with P2P |
None |
None |
Ring attention minimizes peak K/V memory (its main draw for long sequences) at the cost of cp_size communication rounds; Full K/V AllGather is a single collective with no head-count constraints, kept as the always-correct fallback; DeepSpeed-Ulysses is more memory-efficient than AllGather at large CP degrees but adds head-count constraints and is not used here.
Communication Pattern#
Forward Pass Flow#
Input Embedding: Full sequence processed locally (e.g., patch embedding for DiT)
Sequence Sharding: Split tokens across CP ranks
Transformer Blocks:
Each rank processes its local token subset
Self-attention: ring
collective_permuteof local K/V shards → local Q attends full sequence (Full K/V AllGather only on the fallback path)Cross-attention uses full encoder states (no sharding)
FFN operates on local tokens
Sequence Gathering: Reconstruct full sequence
Output Projection: Full sequence for downstream processing (e.g., unpatchify)
Memory and Compute Benefits#
Memory: Reduces activation memory by
1/cp_sizefor transformer blocksCompute: Maintains full attention quality while distributing sequence processing
Communication: Inter-rank communication scales with sequence length, not model size
Generality: Applicable to any model with a sequence dimension (DiT, LLM, etc.)
Configuration#
Context parallelism is configured through vLLM Omni’s sequence parallelism. vLLM Omni enforces sequence_parallel_size = ulysses_degree * ring_degree, so rather than setting sequence_parallel_size directly, the stage config sets ring_degree and leaves the other two unset: ulysses_degree defaults to 1, and sequence_parallel_size is derived as ulysses_degree * ring_degree = ring_degree. ring_degree is therefore the single CP knob in the stage config, and the CP degree (cp_size) equals it:
# engine_args.parallel_config in the stage config
parallel_config:
tensor_parallel_size: 4
ring_degree: 8 # CP degree — sets sequence_parallel_size = 1 * 8 = 8
cfg_parallel_size: 2
Note that despite the name, ring_degree only sizes the sequence-parallel group — it does not necessarily select a ring-attention algorithm. The SP group it produces is what get_sp_group() returns and what this implementation shards the sequence across.
Constraints#
sequence_length % cp_size == 0: Sequence must be evenly divisiblelocal_S % 2 == 0: NKI MLP kernel requires even local sequence length (enforced by padding within the MLP kernel)Compatible with tensor parallelism:
total_ranks = tp_size * cp_sizePositional embeddings must support local position slicing
CP group ranks must map to directly connected devices on the target hardware
Performance Characteristics#
Scaling Properties#
Sequence Length: Linear memory reduction with CP degree
Model Size: Orthogonal to tensor parallelism scaling
Communication: Ring
collective_permutemoves one K/V shard (1/cp_sizeof the sequence) per round overcp_sizerounds, overlapped with attention compute; the AllGather fallback instead materializes full K/V on every rank
Optimal Use Cases#
Long sequences that exceed single-device memory capacity
Long video sequences (high frame count) in diffusion models
High-resolution images (large spatial token count)
Memory-constrained scenarios with long contexts
Balanced with tensor parallelism for model parameter distribution
Implementation Notes#
Positional Embedding Handling#
Local positional embeddings (e.g., rotary embeddings) are sliced to match the local token positions:
# Slice rotary_emb to match local token positions
freqs_cos = freqs_cos[:, start : start + local_S, :, :]
freqs_sin = freqs_sin[:, start : start + local_S, :, :]
Distributed Normalization#
DistributedRMSNorm computes global statistics across TP ranks only, ensuring correct normalization within each CP group:
if self.tp_size > 1:
global_sum_sq = local_sum_sq.clone()
dist.all_reduce(global_sum_sq, group=self.tp_group) # TP group only
This design ensures that context parallelism provides efficient long-sequence processing while maintaining compatibility with existing tensor parallelism and attention mechanisms.
Sequence Length Constraints#
The sequence length must be divisible by cp_size. If it is not, the model raises a ValueError at runtime. Users must choose input dimensions (resolution, frame count) that produce a compatible sequence length.
Problem#
For example, height=240, width=432, num_frames=33 yields:
latent_frames = (33-1)//4 + 1 = 9post_patch_h = (240//8) // 2 = 15post_patch_w = (432//8) // 2 = 27S_total = 9 * 15 * 27 = 36453645 % 8 = 5— not divisible bycp_size=8
This will raise:
ValueError: Sequence length 3645 is not divisible by cp_size 8.
Choose a resolution/frame count that yields a divisible patch sequence length.
NKI MLP Kernel Padding#
Separately, the NKI MLP kernel uses grid=(2,) which requires the sequence dimension to be even. This is handled locally within WanFeedForward.forward — the input is padded by 1 token before the kernel call and trimmed after, so no other layer sees padded tokens:
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
input_shape = hidden_states.shape
hidden_2d = hidden_states.reshape(-1, input_shape[-1])
# NKI MLP kernel uses grid=(2,) requiring T to be even.
T = hidden_2d.shape[0]
pad_T = T % 2
if pad_T:
hidden_2d = F.pad(hidden_2d, (0, 0, 0, 1))
output = NF.mlp(hidden_2d, ...)
if pad_T:
output = output[:T, :]
...
This localized padding has no impact on accuracy since the padded position is discarded immediately after the kernel.
vLLM Omni Integration#
The CP implementation depends on vLLM Omni’s get_sp_group() for process group management. In production, vLLM Omni’s initialize_model_parallel() sets up the SP group automatically. In unit tests, it must be initialized manually:
import vllm_omni.diffusion.distributed.parallel_state as omni_ps
# SP group (required by WanTransformer3DModel.__init__ and WanSelfAttention)
ulysses_pg, ring_pg = omni_ps.set_seq_parallel_pg(
sp_ulysses_degree=1, sp_ring_degree=sp_size,
rank=rank, world_size=world_size, sp_group_ranks=sp_group_ranks,
)
omni_ps._SP = omni_ps.init_model_parallel_group(
group_ranks=sp_group_ranks, ..., parallel_mode="sequence",
ulysses_group=ulysses_pg, ring_group=ring_pg,
)
Known Issues#
Issue |
Description |
|---|---|
NKI QKV kernel disabled under TP4 |
The |
Sequence length must be divisible by |
Arbitrary input resolutions/frame counts that produce a non-divisible sequence length will raise a |
NKI MLP kernel requires even sequence length |
The NKI MLP kernel uses |