[diffusion] optimize: enable vae parallel decode with cfg-parallel (#27875)

This commit is contained in:
Mick
2026-06-13 13:52:27 +08:00
committed by GitHub
parent bcd45d3ac7
commit cb4933b22e
7 changed files with 168 additions and 25 deletions
@@ -6,6 +6,9 @@ from sglang.multimodal_gen.runtime.distributed.group_coordinator import (
)
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
cleanup_dist_env_and_memory,
get_decode_parallel_group_coordinator,
get_decode_parallel_rank,
get_decode_parallel_world_size,
get_dp_group,
get_dp_rank,
get_dp_world_size,
@@ -51,6 +54,10 @@ __all__ = [
"get_tp_group",
"get_tp_rank",
"get_tp_world_size",
# Decode parallel group
"get_decode_parallel_group_coordinator",
"get_decode_parallel_rank",
"get_decode_parallel_world_size",
# Get torch device
"get_local_torch_device",
]
@@ -812,6 +812,22 @@ def get_vae_parallel_rank() -> int:
return torch.distributed.get_rank(group=get_vae_parallel_group())
def get_decode_parallel_group_coordinator() -> GroupCoordinator:
sp_group = get_sp_group()
cfg_group = get_cfg_group()
if sp_group.world_size == 1 and cfg_group.world_size > 1:
return cfg_group
return sp_group
def get_decode_parallel_world_size() -> int:
return get_decode_parallel_group_coordinator().world_size
def get_decode_parallel_rank() -> int:
return get_decode_parallel_group_coordinator().rank_in_group
def init_dit_group(
dit_parallel_size: int,
backend: str,
@@ -6,9 +6,9 @@ import torch.nn as nn
import torch.nn.functional as F
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_sp_group,
get_sp_parallel_rank,
get_sp_world_size,
get_decode_parallel_group_coordinator,
get_decode_parallel_rank,
get_decode_parallel_world_size,
)
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
from sglang.multimodal_gen.runtime.models.vaes.parallel.wan_common_utils import (
@@ -115,7 +115,9 @@ def _halo_memory_format(reference: torch.Tensor) -> torch.memory_format:
def gather_and_trim_height(x: torch.Tensor, expected_height: int | None):
if expected_height is None:
return x
x = get_sp_group().all_gather(_maybe_contiguous_for_sp_gather(x), dim=-2)
x = get_decode_parallel_group_coordinator().all_gather(
_maybe_contiguous_for_sp_gather(x), dim=-2
)
if x.shape[-2] != expected_height:
x = x[..., :expected_height, :].contiguous()
return x
@@ -150,11 +152,11 @@ def halo_exchange(
if height_halo_size == 0:
return x, recv_top_buf, recv_bottom_buf
sp_group = get_sp_group()
rank = get_sp_parallel_rank()
world_size = get_sp_world_size()
group = sp_group.device_group
group_ranks = sp_group.ranks
decode_group = get_decode_parallel_group_coordinator()
rank = get_decode_parallel_rank()
world_size = get_decode_parallel_world_size()
group = decode_group.device_group
group_ranks = decode_group.ranks
top_row_ref = x[..., :height_halo_size, :]
bottom_row_ref = x[..., -height_halo_size:, :]
@@ -234,8 +236,8 @@ class WanDistConv2d(nn.Conv2d):
self.padding = (0, self.padding[1])
self._halo_recv_top_buf: torch.Tensor | None = None
self._halo_recv_bottom_buf: torch.Tensor | None = None
self.rank = get_sp_parallel_rank()
self.world_size = get_sp_world_size()
self.rank = get_decode_parallel_rank()
self.world_size = get_decode_parallel_world_size()
def forward(self, x):
if any(self._padding):
@@ -324,8 +326,8 @@ class WanDistCausalConv3d(nn.Conv3d):
self.padding = (0, 0, 0)
self._halo_recv_top_buf: torch.Tensor | None = None
self._halo_recv_bottom_buf: torch.Tensor | None = None
self.rank = get_sp_parallel_rank()
self.world_size = get_sp_world_size()
self.rank = get_decode_parallel_rank()
self.world_size = get_decode_parallel_world_size()
def forward(self, x, cache_x=None):
padding = list(self._padding)
@@ -386,8 +388,8 @@ class WanDistZeroPad2d(nn.Module):
def __init__(self, padding: tuple[int, int, int, int]) -> None:
super().__init__()
self.padding = padding # (left, right, top, bottom)
self.rank = get_sp_parallel_rank()
self.world_size = get_sp_world_size()
self.rank = get_decode_parallel_rank()
self.world_size = get_decode_parallel_world_size()
def forward(self, x: torch.Tensor) -> torch.Tensor:
left, right, top, bottom = self.padding
@@ -512,13 +514,13 @@ class WanDistAttentionBlock(nn.Module):
self.norm = WanRMS_norm(dim)
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
self.proj = nn.Conv2d(dim, dim, 1)
self.rank = get_sp_parallel_rank()
self.world_size = get_sp_world_size()
self.sp_group = get_sp_group()
self.rank = get_decode_parallel_rank()
self.world_size = get_decode_parallel_world_size()
self.decode_group = get_decode_parallel_group_coordinator()
def forward(self, x):
if self.world_size > 1:
x = self.sp_group.all_gather(_maybe_contiguous_for_sp_gather(x), dim=-2)
x = self.decode_group.all_gather(_maybe_contiguous_for_sp_gather(x), dim=-2)
x = x.contiguous()
x = attention_block_forward(self, x)
if self.world_size > 1:
@@ -26,6 +26,8 @@ from einops import rearrange
from sglang.multimodal_gen.configs.models.vaes import WanVAEConfig
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_decode_parallel_rank,
get_decode_parallel_world_size,
get_sp_parallel_rank,
get_sp_world_size,
)
@@ -623,7 +625,7 @@ class WanDecoder3d(nn.Module):
world_size = 1
if dist.is_initialized():
world_size = get_sp_world_size()
world_size = get_decode_parallel_world_size()
if use_parallel_decode and world_size > 1:
CausalConv3d = WanDistCausalConv3d
@@ -692,8 +694,8 @@ class WanDecoder3d(nn.Module):
self.world_size = 1
self.rank = 0
if dist.is_initialized():
self.world_size = get_sp_world_size()
self.rank = get_sp_parallel_rank()
self.world_size = get_decode_parallel_world_size()
self.rank = get_decode_parallel_rank()
def forward(self, x):
expected_height = None
@@ -9,7 +9,11 @@ import weakref
import torch
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.distributed import (
get_decode_parallel_world_size,
get_local_torch_device,
model_parallel_is_initialized,
)
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import VAELoader
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
ComponentUse,
@@ -114,10 +118,20 @@ class DecodingStage(PipelineStage):
@property
def parallelism_type(self) -> StageParallelismType:
if get_global_server_args().enable_cfg_parallel:
server_args = get_global_server_args()
if server_args.enable_cfg_parallel:
if self._can_use_parallel_decode():
return StageParallelismType.REPLICATED
return StageParallelismType.MAIN_RANK_ONLY
return StageParallelismType.REPLICATED
def _can_use_parallel_decode(self) -> bool:
return (
model_parallel_is_initialized()
and get_decode_parallel_world_size() > 1
and self.vae.use_parallel_decode
)
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
"""Verify decoding stage inputs."""
result = VerificationResult()
@@ -241,7 +241,9 @@ HFFN, WFFN = S, 1 # S spatial tokens per frame
def _ffn():
m = GLUMBConvTemp(C_FFN, HID, t_kernel_size=3).double().eval()
with torch.no_grad(): # zero-init t_conv -> randomize for a non-trivial temporal filter
with (
torch.no_grad()
): # zero-init t_conv -> randomize for a non-trivial temporal filter
m.t_conv.weight.copy_(torch.randn_like(m.t_conv.weight))
return m
@@ -0,0 +1,100 @@
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
StageParallelismType,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import DecodingStage
class TestDecodingStageParallelism(unittest.TestCase):
def test_cfg_parallel_uses_replicated_decode_when_decode_group_has_multiple_ranks(
self,
):
stage = object.__new__(DecodingStage)
stage.vae = SimpleNamespace(use_parallel_decode=True)
with (
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.get_global_server_args",
return_value=SimpleNamespace(enable_cfg_parallel=True),
),
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.model_parallel_is_initialized",
return_value=True,
),
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.get_decode_parallel_world_size",
return_value=2,
),
):
self.assertEqual(
stage.parallelism_type,
StageParallelismType.REPLICATED,
)
def test_cfg_parallel_keeps_main_rank_decode_without_parallel_decode(self):
stage = object.__new__(DecodingStage)
stage.vae = SimpleNamespace(use_parallel_decode=False)
with (
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.get_global_server_args",
return_value=SimpleNamespace(enable_cfg_parallel=True),
),
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.model_parallel_is_initialized",
return_value=True,
),
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.get_decode_parallel_world_size",
return_value=2,
),
):
self.assertEqual(
stage.parallelism_type,
StageParallelismType.MAIN_RANK_ONLY,
)
def test_cfg_parallel_keeps_main_rank_decode_when_decode_group_is_single_rank(
self,
):
stage = object.__new__(DecodingStage)
stage.vae = SimpleNamespace(use_parallel_decode=True)
with (
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.get_global_server_args",
return_value=SimpleNamespace(enable_cfg_parallel=True),
),
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.model_parallel_is_initialized",
return_value=True,
),
patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.get_decode_parallel_world_size",
return_value=1,
),
):
self.assertEqual(
stage.parallelism_type,
StageParallelismType.MAIN_RANK_ONLY,
)
def test_non_cfg_parallel_keeps_replicated_decode(self):
stage = object.__new__(DecodingStage)
stage.vae = SimpleNamespace(use_parallel_decode=True)
with patch(
"sglang.multimodal_gen.runtime.pipelines_core.stages.decoding.get_global_server_args",
return_value=SimpleNamespace(enable_cfg_parallel=False),
):
self.assertEqual(
stage.parallelism_type,
StageParallelismType.REPLICATED,
)
if __name__ == "__main__":
unittest.main()