diff --git a/python/sglang/multimodal_gen/runtime/distributed/__init__.py b/python/sglang/multimodal_gen/runtime/distributed/__init__.py index c91de13d3..276615cce 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/__init__.py +++ b/python/sglang/multimodal_gen/runtime/distributed/__init__.py @@ -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", ] diff --git a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py index 320422e02..5dd033169 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py +++ b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_dist_utils.py b/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_dist_utils.py index 096bf3095..e590a1bb5 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_dist_utils.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_dist_utils.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py index 0343754a9..2a96c33ad 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py index 1a130cd87..d6b0dc9bd 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py @@ -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() diff --git a/python/sglang/multimodal_gen/test/unit/sana_wm/test_streaming_cached.py b/python/sglang/multimodal_gen/test/unit/sana_wm/test_streaming_cached.py index 48f14bab5..c96087527 100644 --- a/python/sglang/multimodal_gen/test/unit/sana_wm/test_streaming_cached.py +++ b/python/sglang/multimodal_gen/test/unit/sana_wm/test_streaming_cached.py @@ -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 diff --git a/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py b/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py new file mode 100644 index 000000000..9b308545b --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py @@ -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()