[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 (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
cleanup_dist_env_and_memory,
|
cleanup_dist_env_and_memory,
|
||||||
|
get_decode_parallel_group_coordinator,
|
||||||
|
get_decode_parallel_rank,
|
||||||
|
get_decode_parallel_world_size,
|
||||||
get_dp_group,
|
get_dp_group,
|
||||||
get_dp_rank,
|
get_dp_rank,
|
||||||
get_dp_world_size,
|
get_dp_world_size,
|
||||||
@@ -51,6 +54,10 @@ __all__ = [
|
|||||||
"get_tp_group",
|
"get_tp_group",
|
||||||
"get_tp_rank",
|
"get_tp_rank",
|
||||||
"get_tp_world_size",
|
"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 torch device
|
||||||
"get_local_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())
|
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(
|
def init_dit_group(
|
||||||
dit_parallel_size: int,
|
dit_parallel_size: int,
|
||||||
backend: str,
|
backend: str,
|
||||||
|
|||||||
@@ -6,9 +6,9 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
get_sp_group,
|
get_decode_parallel_group_coordinator,
|
||||||
get_sp_parallel_rank,
|
get_decode_parallel_rank,
|
||||||
get_sp_world_size,
|
get_decode_parallel_world_size,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
|
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
|
||||||
from sglang.multimodal_gen.runtime.models.vaes.parallel.wan_common_utils import (
|
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):
|
def gather_and_trim_height(x: torch.Tensor, expected_height: int | None):
|
||||||
if expected_height is None:
|
if expected_height is None:
|
||||||
return x
|
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:
|
if x.shape[-2] != expected_height:
|
||||||
x = x[..., :expected_height, :].contiguous()
|
x = x[..., :expected_height, :].contiguous()
|
||||||
return x
|
return x
|
||||||
@@ -150,11 +152,11 @@ def halo_exchange(
|
|||||||
if height_halo_size == 0:
|
if height_halo_size == 0:
|
||||||
return x, recv_top_buf, recv_bottom_buf
|
return x, recv_top_buf, recv_bottom_buf
|
||||||
|
|
||||||
sp_group = get_sp_group()
|
decode_group = get_decode_parallel_group_coordinator()
|
||||||
rank = get_sp_parallel_rank()
|
rank = get_decode_parallel_rank()
|
||||||
world_size = get_sp_world_size()
|
world_size = get_decode_parallel_world_size()
|
||||||
group = sp_group.device_group
|
group = decode_group.device_group
|
||||||
group_ranks = sp_group.ranks
|
group_ranks = decode_group.ranks
|
||||||
|
|
||||||
top_row_ref = x[..., :height_halo_size, :]
|
top_row_ref = x[..., :height_halo_size, :]
|
||||||
bottom_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.padding = (0, self.padding[1])
|
||||||
self._halo_recv_top_buf: torch.Tensor | None = None
|
self._halo_recv_top_buf: torch.Tensor | None = None
|
||||||
self._halo_recv_bottom_buf: torch.Tensor | None = None
|
self._halo_recv_bottom_buf: torch.Tensor | None = None
|
||||||
self.rank = get_sp_parallel_rank()
|
self.rank = get_decode_parallel_rank()
|
||||||
self.world_size = get_sp_world_size()
|
self.world_size = get_decode_parallel_world_size()
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
if any(self._padding):
|
if any(self._padding):
|
||||||
@@ -324,8 +326,8 @@ class WanDistCausalConv3d(nn.Conv3d):
|
|||||||
self.padding = (0, 0, 0)
|
self.padding = (0, 0, 0)
|
||||||
self._halo_recv_top_buf: torch.Tensor | None = None
|
self._halo_recv_top_buf: torch.Tensor | None = None
|
||||||
self._halo_recv_bottom_buf: torch.Tensor | None = None
|
self._halo_recv_bottom_buf: torch.Tensor | None = None
|
||||||
self.rank = get_sp_parallel_rank()
|
self.rank = get_decode_parallel_rank()
|
||||||
self.world_size = get_sp_world_size()
|
self.world_size = get_decode_parallel_world_size()
|
||||||
|
|
||||||
def forward(self, x, cache_x=None):
|
def forward(self, x, cache_x=None):
|
||||||
padding = list(self._padding)
|
padding = list(self._padding)
|
||||||
@@ -386,8 +388,8 @@ class WanDistZeroPad2d(nn.Module):
|
|||||||
def __init__(self, padding: tuple[int, int, int, int]) -> None:
|
def __init__(self, padding: tuple[int, int, int, int]) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.padding = padding # (left, right, top, bottom)
|
self.padding = padding # (left, right, top, bottom)
|
||||||
self.rank = get_sp_parallel_rank()
|
self.rank = get_decode_parallel_rank()
|
||||||
self.world_size = get_sp_world_size()
|
self.world_size = get_decode_parallel_world_size()
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
left, right, top, bottom = self.padding
|
left, right, top, bottom = self.padding
|
||||||
@@ -512,13 +514,13 @@ class WanDistAttentionBlock(nn.Module):
|
|||||||
self.norm = WanRMS_norm(dim)
|
self.norm = WanRMS_norm(dim)
|
||||||
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
|
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
|
||||||
self.proj = nn.Conv2d(dim, dim, 1)
|
self.proj = nn.Conv2d(dim, dim, 1)
|
||||||
self.rank = get_sp_parallel_rank()
|
self.rank = get_decode_parallel_rank()
|
||||||
self.world_size = get_sp_world_size()
|
self.world_size = get_decode_parallel_world_size()
|
||||||
self.sp_group = get_sp_group()
|
self.decode_group = get_decode_parallel_group_coordinator()
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
if self.world_size > 1:
|
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 = x.contiguous()
|
||||||
x = attention_block_forward(self, x)
|
x = attention_block_forward(self, x)
|
||||||
if self.world_size > 1:
|
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.configs.models.vaes import WanVAEConfig
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
|
get_decode_parallel_rank,
|
||||||
|
get_decode_parallel_world_size,
|
||||||
get_sp_parallel_rank,
|
get_sp_parallel_rank,
|
||||||
get_sp_world_size,
|
get_sp_world_size,
|
||||||
)
|
)
|
||||||
@@ -623,7 +625,7 @@ class WanDecoder3d(nn.Module):
|
|||||||
|
|
||||||
world_size = 1
|
world_size = 1
|
||||||
if dist.is_initialized():
|
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:
|
if use_parallel_decode and world_size > 1:
|
||||||
CausalConv3d = WanDistCausalConv3d
|
CausalConv3d = WanDistCausalConv3d
|
||||||
@@ -692,8 +694,8 @@ class WanDecoder3d(nn.Module):
|
|||||||
self.world_size = 1
|
self.world_size = 1
|
||||||
self.rank = 0
|
self.rank = 0
|
||||||
if dist.is_initialized():
|
if dist.is_initialized():
|
||||||
self.world_size = get_sp_world_size()
|
self.world_size = get_decode_parallel_world_size()
|
||||||
self.rank = get_sp_parallel_rank()
|
self.rank = get_decode_parallel_rank()
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
expected_height = None
|
expected_height = None
|
||||||
|
|||||||
@@ -9,7 +9,11 @@ import weakref
|
|||||||
|
|
||||||
import torch
|
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.loader.component_loaders.vae_loader import VAELoader
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||||
ComponentUse,
|
ComponentUse,
|
||||||
@@ -114,10 +118,20 @@ class DecodingStage(PipelineStage):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def parallelism_type(self) -> StageParallelismType:
|
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.MAIN_RANK_ONLY
|
||||||
return StageParallelismType.REPLICATED
|
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:
|
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
|
||||||
"""Verify decoding stage inputs."""
|
"""Verify decoding stage inputs."""
|
||||||
result = VerificationResult()
|
result = VerificationResult()
|
||||||
|
|||||||
@@ -241,7 +241,9 @@ HFFN, WFFN = S, 1 # S spatial tokens per frame
|
|||||||
|
|
||||||
def _ffn():
|
def _ffn():
|
||||||
m = GLUMBConvTemp(C_FFN, HID, t_kernel_size=3).double().eval()
|
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))
|
m.t_conv.weight.copy_(torch.randn_like(m.t_conv.weight))
|
||||||
return m
|
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