[diffusion] optimize: enable vae parallel decode with cfg-parallel (#27875)
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user