vllm_omni.diffusion.hooks.sequence_parallel ¶
Sequence Parallelism hooks for non-intrusive SP support.
This module implements the hook-based mechanism for applying sequence parallelism to models without modifying their forward() methods.
Usage
- Define _sp_plan on your model class (corresponds to diffusers' _cp_plan)
- Call apply_sequence_parallel(model, config, plan) to enable SP
- Call remove_sequence_parallel(model, plan) to disable SP
The hooks automatically shard inputs before forward and gather outputs after, based on the plan specification.
ModuleForwardMetadata dataclass ¶
Metadata for mapping forward() parameter names to args/kwargs positions.
This caches the inspection of a module's forward signature to efficiently locate parameters by name in subsequent calls.
SequenceParallelGatherHook ¶
Bases: ModelHook
Hook for gathering outputs after a module's forward pass.
This hook is registered to modules that need their outputs gathered from all sequence parallel ranks. It intercepts the output and gathers it according to the plan specification.
Note: This corresponds to ContextParallelGatherHook in diffusers.
SequenceParallelSplitHook ¶
Bases: ModelHook
Hook for splitting inputs before a module's forward pass.
This hook is registered to modules that need their inputs sharded across sequence parallel ranks. It intercepts the forward call, shards specified inputs according to the plan, and passes the sharded inputs to the original forward.
For split_output=True inputs, it shards the output instead.
Supports both SequenceParallelInput (full split) and SequenceParallelPartialInput (partial split for text/image separation).
Note: This corresponds to ContextParallelSplitHook in diffusers.
module_forward_metadata instance-attribute ¶
module_forward_metadata: ModuleForwardMetadata | None = None
post_forward ¶
Shard outputs for split_output=True entries.
apply_sequence_parallel ¶
apply_sequence_parallel(
module: Module,
config: SequenceParallelConfig,
plan: SequenceParallelModelPlan,
) -> None
Apply sequence parallel hooks to a model according to the plan.
This function registers hooks on the specified submodules to automatically shard inputs and gather outputs for sequence parallelism.
Note: This corresponds to apply_context_parallel in diffusers.
The complete SP flow is: 1. Input sharding (SequenceParallelSplitHook): Split sequence across SP ranks 2. Attention parallelism (handled by vLLM-Omni's Attention layer): - Ulysses: All-to-All over Q/K/V heads - Ring: K/V circulation in ring topology - Hybrid: Both (Ulysses handles head redistribution, Ring handles K/V) 3. Output gathering (SequenceParallelGatherHook): Gather sequence from SP ranks
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
module | Module | The model to apply SP to. | required |
config | SequenceParallelConfig | The sequence parallel configuration. | required |
plan | SequenceParallelModelPlan | Dictionary mapping module names to input/output specifications. | required |
Example
config = SequenceParallelConfig(ulysses_degree=2) plan = { "": {"hidden_states": SequenceParallelInput(split_dim=1, expected_dims=3)}, "proj_out": SequenceParallelOutput(gather_dim=1, expected_dims=3), } apply_sequence_parallel(model, config, plan)
Note
vLLM-Omni's Attention layer automatically handles the internal parallelism (Ulysses All-to-All or Ring attention) based on the forward_context configuration. This function only handles input/output sharding for the model as a whole.
remove_sequence_parallel ¶
remove_sequence_parallel(
module: Module, plan: SequenceParallelModelPlan
) -> None
Remove sequence parallel hooks from a model.
Note: This corresponds to remove_context_parallel in diffusers.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
module | Module | The model to remove SP from. | required |
plan | SequenceParallelModelPlan | The same plan used when applying SP. | required |