From 582bd23f71be92337776fb99e4c9c6784a146f60 Mon Sep 17 00:00:00 2001 From: Mick Date: Sun, 14 Jun 2026 20:19:44 +0800 Subject: [PATCH] [diffusion] feat: enable spatial-shard vae decode across GPUs (#28071) --- .../configs/models/vaes/base.py | 96 ++- .../configs/models/vaes/hunyuan3d.py | 2 + .../configs/models/vaes/hunyuanvae.py | 4 + .../configs/models/vaes/ltx_audio.py | 2 + .../configs/models/vaes/ltx_video.py | 4 + .../configs/models/vaes/qwenimage.py | 2 - .../configs/models/vaes/wanvae.py | 4 +- .../runtime/distributed/parallel_state.py | 30 +- .../runtime/layers/parallel_conv.py | 717 ++++++++++++++++ .../runtime/models/vaes/autoencoder.py | 23 +- .../runtime/models/vaes/autoencoder_dc.py | 37 +- .../models/vaes/autoencoder_kl_flux2.py | 23 +- .../models/vaes/autoencoder_kl_qwenimage.py | 221 +++-- .../runtime/models/vaes/common.py | 62 +- .../runtime/models/vaes/hunyuanvae.py | 203 ++++- .../runtime/models/vaes/ltx_2_vae.py | 68 +- .../models/vaes/parallel/diffusers_spatial.py | 109 +++ .../models/vaes/parallel/wan_common_utils.py | 480 ----------- .../models/vaes/parallel/wan_dist_utils.py | 716 ---------------- .../runtime/models/vaes/wanvae.py | 788 ++++++++++++++++-- .../ltx_2/denoising_av.py | 13 +- .../sglang/multimodal_gen/test/test_utils.py | 2 +- .../unit/test_decoding_stage_parallelism.py | 43 + .../test/unit/test_vae_loader.py | 22 +- .../unit/test_vae_spatial_parallel_decode.py | 512 ++++++++++++ 25 files changed, 2767 insertions(+), 1416 deletions(-) create mode 100644 python/sglang/multimodal_gen/runtime/layers/parallel_conv.py create mode 100644 python/sglang/multimodal_gen/runtime/models/vaes/parallel/diffusers_spatial.py delete mode 100644 python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_common_utils.py delete mode 100644 python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_dist_utils.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py diff --git a/python/sglang/multimodal_gen/configs/models/vaes/base.py b/python/sglang/multimodal_gen/configs/models/vaes/base.py index 64c3653cd..5d55e92ee 100644 --- a/python/sglang/multimodal_gen/configs/models/vaes/base.py +++ b/python/sglang/multimodal_gen/configs/models/vaes/base.py @@ -4,6 +4,7 @@ import argparse import dataclasses from dataclasses import dataclass, field +from functools import lru_cache from typing import Any import torch @@ -11,6 +12,54 @@ import torch from sglang.multimodal_gen.configs.models.base import ArchConfig, ModelConfig from sglang.multimodal_gen.utils import StoreBoolean +AUTO_PARALLEL_DECODE_MODE = "auto" +SPATIAL_SHARD_PARALLEL_DECODE_MODES = ("spatial_shard", "spatial") + + +@lru_cache(maxsize=8) +def is_spatial_shard_parallel_decode_mode(mode: str) -> bool: + return mode in SPATIAL_SHARD_PARALLEL_DECODE_MODES + + +@lru_cache(maxsize=8) +def is_auto_parallel_decode_mode(mode: str) -> bool: + return mode == AUTO_PARALLEL_DECODE_MODE + + +@lru_cache(maxsize=128) +def _should_use_auto_spatial_shard_parallel_decode( + z_shape: tuple[int, ...], + world_size: int, + min_latent_elements_per_rank: int, +) -> bool: + if world_size <= 1 or z_shape[-2] < world_size: + return False + latent_elements_per_rank = ( + z_shape[0] * z_shape[-3] * z_shape[-2] * z_shape[-1] + ) // world_size + return latent_elements_per_rank >= min_latent_elements_per_rank + + +def should_use_spatial_shard_parallel_decode( + config: Any, z: torch.Tensor | None = None, world_size: int = 1 +) -> bool: + if not config.use_parallel_decode: + return False + + if is_spatial_shard_parallel_decode_mode(config.parallel_decode_mode): + return True + + if not is_auto_parallel_decode_mode(config.parallel_decode_mode): + return False + + if not config.auto_parallel_decode_prefers_spatial_shard(): + return False + + if z is None: + return True + + return config.should_use_auto_spatial_shard_parallel_decode(z, world_size) + @dataclass class VAEArchConfig(ArchConfig): @@ -41,8 +90,9 @@ class VAEConfig(ModelConfig): use_temporal_tiling: bool = True use_parallel_tiling: bool = True use_temporal_scaling_frames: bool = True - use_parallel_decode: bool = False - parallel_decode_mode: str = "tiled" + use_parallel_decode: bool = True + parallel_decode_mode: str = AUTO_PARALLEL_DECODE_MODE + auto_parallel_decode_min_latent_elements_per_rank: int = 4096 def __post_init__(self): self.blend_num_frames = ( @@ -52,6 +102,18 @@ class VAEConfig(ModelConfig): def post_init(self): pass + def auto_parallel_decode_prefers_spatial_shard(self) -> bool: + return False + + def should_use_auto_spatial_shard_parallel_decode( + self, z: torch.Tensor, world_size: int + ) -> bool: + return _should_use_auto_spatial_shard_parallel_decode( + tuple(z.shape), + world_size, + self.auto_parallel_decode_min_latent_elements_per_rank, + ) + @staticmethod def add_cli_args(parser: Any, prefix: str = "vae-config") -> Any: """Add CLI arguments for VAEConfig fields""" @@ -59,98 +121,98 @@ class VAEConfig(ModelConfig): f"--{prefix}.load-encoder", action=StoreBoolean, dest=f"{prefix.replace('-', '_')}.load_encoder", - default=VAEConfig.load_encoder, + default=None, help="Whether to load the VAE encoder", ) parser.add_argument( f"--{prefix}.load-decoder", action=StoreBoolean, dest=f"{prefix.replace('-', '_')}.load_decoder", - default=VAEConfig.load_decoder, + default=None, help="Whether to load the VAE decoder", ) parser.add_argument( f"--{prefix}.tile-sample-min-height", type=int, dest=f"{prefix.replace('-', '_')}.tile_sample_min_height", - default=VAEConfig.tile_sample_min_height, + default=None, help="Minimum height for VAE tile sampling", ) parser.add_argument( f"--{prefix}.tile-sample-min-width", type=int, dest=f"{prefix.replace('-', '_')}.tile_sample_min_width", - default=VAEConfig.tile_sample_min_width, + default=None, help="Minimum width for VAE tile sampling", ) parser.add_argument( f"--{prefix}.tile-sample-min-num-frames", type=int, dest=f"{prefix.replace('-', '_')}.tile_sample_min_num_frames", - default=VAEConfig.tile_sample_min_num_frames, + default=None, help="Minimum number of frames for VAE tile sampling", ) parser.add_argument( f"--{prefix}.tile-sample-stride-height", type=int, dest=f"{prefix.replace('-', '_')}.tile_sample_stride_height", - default=VAEConfig.tile_sample_stride_height, + default=None, help="Stride height for VAE tile sampling", ) parser.add_argument( f"--{prefix}.tile-sample-stride-width", type=int, dest=f"{prefix.replace('-', '_')}.tile_sample_stride_width", - default=VAEConfig.tile_sample_stride_width, + default=None, help="Stride width for VAE tile sampling", ) parser.add_argument( f"--{prefix}.tile-sample-stride-num-frames", type=int, dest=f"{prefix.replace('-', '_')}.tile_sample_stride_num_frames", - default=VAEConfig.tile_sample_stride_num_frames, + default=None, help="Stride number of frames for VAE tile sampling", ) parser.add_argument( f"--{prefix}.blend-num-frames", type=int, dest=f"{prefix.replace('-', '_')}.blend_num_frames", - default=VAEConfig.blend_num_frames, + default=None, help="Number of frames to blend for VAE tile sampling", ) parser.add_argument( f"--{prefix}.use-tiling", action=StoreBoolean, dest=f"{prefix.replace('-', '_')}.use_tiling", - default=VAEConfig.use_tiling, + default=None, help="Whether to use tiling for VAE", ) parser.add_argument( f"--{prefix}.use-temporal-tiling", action=StoreBoolean, dest=f"{prefix.replace('-', '_')}.use_temporal_tiling", - default=VAEConfig.use_temporal_tiling, + default=None, help="Whether to use temporal tiling for VAE", ) parser.add_argument( f"--{prefix}.use-parallel-tiling", action=StoreBoolean, dest=f"{prefix.replace('-', '_')}.use_parallel_tiling", - default=VAEConfig.use_parallel_tiling, + default=None, help="Whether to use parallel tiling for VAE", ) parser.add_argument( f"--{prefix}.use-parallel-decode", action=StoreBoolean, dest=f"{prefix.replace('-', '_')}.use_parallel_decode", - default=VAEConfig.use_parallel_decode, + default=None, help="Whether to use parallel decode for VAE", ) parser.add_argument( f"--{prefix}.parallel-decode-mode", - choices=("tiled", "patch", "auto"), + choices=("tiled", "patch", "spatial_shard", "spatial", "auto"), dest=f"{prefix.replace('-', '_')}.parallel_decode_mode", - default=VAEConfig.parallel_decode_mode, + default=None, help="Parallel decode mode for VAE", ) diff --git a/python/sglang/multimodal_gen/configs/models/vaes/hunyuan3d.py b/python/sglang/multimodal_gen/configs/models/vaes/hunyuan3d.py index fd6d72adb..a5c0deeb0 100644 --- a/python/sglang/multimodal_gen/configs/models/vaes/hunyuan3d.py +++ b/python/sglang/multimodal_gen/configs/models/vaes/hunyuan3d.py @@ -20,3 +20,5 @@ class Hunyuan3DVAEConfig(VAEConfig): subfolder: str = "hunyuan3d-dit-v2-0" load_encoder: bool = False load_decoder: bool = True + use_parallel_decode: bool = False + parallel_decode_mode: str = "tiled" diff --git a/python/sglang/multimodal_gen/configs/models/vaes/hunyuanvae.py b/python/sglang/multimodal_gen/configs/models/vaes/hunyuanvae.py index 601b72d57..30a1b9f1f 100644 --- a/python/sglang/multimodal_gen/configs/models/vaes/hunyuanvae.py +++ b/python/sglang/multimodal_gen/configs/models/vaes/hunyuanvae.py @@ -39,3 +39,7 @@ class HunyuanVAEArchConfig(VAEArchConfig): @dataclass class HunyuanVAEConfig(VAEConfig): arch_config: VAEArchConfig = field(default_factory=HunyuanVAEArchConfig) + auto_parallel_decode_min_latent_elements_per_rank: int = 32768 + + def auto_parallel_decode_prefers_spatial_shard(self) -> bool: + return True diff --git a/python/sglang/multimodal_gen/configs/models/vaes/ltx_audio.py b/python/sglang/multimodal_gen/configs/models/vaes/ltx_audio.py index 54a2b5a0e..a61331947 100644 --- a/python/sglang/multimodal_gen/configs/models/vaes/ltx_audio.py +++ b/python/sglang/multimodal_gen/configs/models/vaes/ltx_audio.py @@ -30,3 +30,5 @@ class LTXAudioVAEArchConfig(VAEArchConfig): @dataclass class LTXAudioVAEConfig(VAEConfig): arch_config: LTXAudioVAEArchConfig = field(default_factory=LTXAudioVAEArchConfig) + use_parallel_decode: bool = False + parallel_decode_mode: str = "tiled" diff --git a/python/sglang/multimodal_gen/configs/models/vaes/ltx_video.py b/python/sglang/multimodal_gen/configs/models/vaes/ltx_video.py index 92d29537d..240214639 100644 --- a/python/sglang/multimodal_gen/configs/models/vaes/ltx_video.py +++ b/python/sglang/multimodal_gen/configs/models/vaes/ltx_video.py @@ -62,3 +62,7 @@ class LTXVideoVAEArchConfig(VAEArchConfig): @dataclass class LTXVideoVAEConfig(VAEConfig): arch_config: LTXVideoVAEArchConfig = field(default_factory=LTXVideoVAEArchConfig) + auto_parallel_decode_min_latent_elements_per_rank: int = 1024 + + def auto_parallel_decode_prefers_spatial_shard(self) -> bool: + return True diff --git a/python/sglang/multimodal_gen/configs/models/vaes/qwenimage.py b/python/sglang/multimodal_gen/configs/models/vaes/qwenimage.py index fae449fe7..af9fa9d2a 100644 --- a/python/sglang/multimodal_gen/configs/models/vaes/qwenimage.py +++ b/python/sglang/multimodal_gen/configs/models/vaes/qwenimage.py @@ -38,8 +38,6 @@ class QwenImageVAEConfig(VAEConfig): use_temporal_tiling: bool = False use_parallel_tiling: bool = False - use_parallel_decode: bool = False - def get_vae_scale_factor(self): return 2 ** len(self.arch_config.temperal_downsample) diff --git a/python/sglang/multimodal_gen/configs/models/vaes/wanvae.py b/python/sglang/multimodal_gen/configs/models/vaes/wanvae.py index caf59c407..27597a32f 100644 --- a/python/sglang/multimodal_gen/configs/models/vaes/wanvae.py +++ b/python/sglang/multimodal_gen/configs/models/vaes/wanvae.py @@ -83,13 +83,15 @@ class WanVAEConfig(VAEConfig): use_parallel_tiling: bool = False use_parallel_encode: bool = True - use_parallel_decode: bool = True def __post_init__(self): self.blend_num_frames = ( self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames ) * 2 + def auto_parallel_decode_prefers_spatial_shard(self) -> bool: + return True + def get_vae_scale_factor(self): # Wan VAE does not expose block_out_channels like SD-style VAEs. # Its spatial downsample factor is explicitly defined by scale_factor_spatial. diff --git a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py index 5dd033169..7f4a06d12 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py +++ b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py @@ -66,6 +66,7 @@ _SP: SequenceParallelGroupCoordinator | None = None _PP: PipelineGroupCoordinator | None = None _CFG: GroupCoordinator | None = None _DP: GroupCoordinator | None = None +_VAE_DECODE: GroupCoordinator | None = None _DIT: ProcessGroup | None = None _VAE: ProcessGroup | None = None @@ -152,6 +153,7 @@ def init_parallel_group_coordinator( "tensor", "sequence", "classifier_free_guidance", + "vae_decode", ], f"parallel_mode {parallel_mode} is not supported" if parallel_mode == "pipeline": return PipelineGroupCoordinator( @@ -169,12 +171,13 @@ def init_parallel_group_coordinator( **kwargs, ) else: - # fallback to GroupCoordinator return GroupCoordinator( group_ranks=group_ranks, local_rank=local_rank, torch_distributed_backend=backend, - group_name="cfg_group", + group_name=( + "vae_decode_group" if parallel_mode == "vae_decode" else "cfg_group" + ), ) @@ -434,6 +437,15 @@ def initialize_model_parallel( parallel_mode="tensor", ) + global _VAE_DECODE + assert _VAE_DECODE is None, "VAE decode parallel group is already initialized" + _VAE_DECODE = init_parallel_group_coordinator( + group_ranks=rank_generator.get_ranks("tp-sp-pp-cfg"), + local_rank=get_world_group().local_rank, + backend=backend, + parallel_mode="vae_decode", + ) + if vae_parallel_size > 0: init_vae_group(dit_parallel_size, vae_parallel_size, backend) init_dit_group(dit_parallel_size, backend) @@ -537,6 +549,7 @@ def model_parallel_is_initialized() -> bool: and _SP is not None and _PP is not None and _TP is not None + and _VAE_DECODE is not None ) @@ -813,11 +826,8 @@ def get_vae_parallel_rank() -> int: 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 + assert _VAE_DECODE is not None, "VAE decode parallel group is not initialized" + return _VAE_DECODE def get_decode_parallel_world_size() -> int: @@ -858,9 +868,9 @@ def init_vae_group( def destroy_model_parallel() -> None: """Set the groups to none and destroy them.""" - global _TP, _SP, _DP, _CFG, _PP, _DIT, _VAE + global _TP, _SP, _DP, _CFG, _PP, _VAE_DECODE, _DIT, _VAE - for group in (_TP, _SP, _DP, _CFG, _PP): + for group in (_TP, _SP, _DP, _CFG, _PP, _VAE_DECODE): if group is not None: group.destroy() @@ -868,4 +878,4 @@ def destroy_model_parallel() -> None: if group is not None: torch.distributed.destroy_process_group(group) - _TP, _SP, _DP, _CFG, _PP, _DIT, _VAE = (None,) * 7 + _TP, _SP, _DP, _CFG, _PP, _VAE_DECODE, _DIT, _VAE = (None,) * 8 diff --git a/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py b/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py new file mode 100644 index 000000000..068a6b925 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py @@ -0,0 +1,717 @@ +import contextvars +import math +from contextlib import contextmanager + +import torch +import torch.distributed as dist +import torch.nn as nn +import torch.nn.functional as F + +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_decode_parallel_group_coordinator, + get_decode_parallel_rank, + get_decode_parallel_world_size, +) +from sglang.multimodal_gen.runtime.platforms import current_platform + +_SPATIAL_PARALLEL_DECODE_DISABLED = contextvars.ContextVar( + "spatial_parallel_decode_disabled", default=False +) + + +@contextmanager +def disable_spatial_parallel_decode(): + token = _SPATIAL_PARALLEL_DECODE_DISABLED.set(True) + try: + yield + finally: + _SPATIAL_PARALLEL_DECODE_DISABLED.reset(token) + + +def spatial_parallel_decode_disabled() -> bool: + return _SPATIAL_PARALLEL_DECODE_DISABLED.get() + + +def _tensor_pad(x: torch.Tensor, len_to_pad: int, dim: int = -2): + return torch.cat( + [ + x, + torch.zeros( + *x.shape[:dim], + len_to_pad, + *x.shape[dim + 1 :], + dtype=x.dtype, + device=x.device, + ), + ], + dim=dim, + ) + + +def _tensor_chunk(x: torch.Tensor, dim: int = -2, world_size: int = 1, rank: int = 0): + if x is None: + return x + if world_size <= 1: + return x + return torch.tensor_split(x, world_size, dim=dim)[rank].contiguous( + memory_format=_halo_memory_format(x) + ) + + +def split_for_parallel_decode( + x: torch.Tensor, upsample_count: int, world_size: int, rank: int +): + return split_height_for_parallel_decode( + x, + expected_height=x.shape[-2] * (2**upsample_count), + world_size=world_size, + rank=rank, + ) + + +def split_height_for_parallel_decode( + x: torch.Tensor, expected_height: int, world_size: int, rank: int +): + if spatial_parallel_decode_disabled(): + return x, None + x = _tensor_chunk(x, dim=-2, world_size=world_size, rank=rank) + return x, expected_height + + +def _maybe_contiguous_for_sp_gather(x: torch.Tensor) -> torch.Tensor: + if ( + x.dim() == 5 + and hasattr(torch, "channels_last_3d") + and x.is_contiguous(memory_format=torch.channels_last_3d) + and not x.is_contiguous() + ): + return x.contiguous() + if ( + x.dim() == 4 + and x.is_contiguous(memory_format=torch.channels_last) + and not x.is_contiguous() + ): + return x.contiguous() + return x + + +def gather_and_trim_height(x: torch.Tensor, expected_height: int | None): + if spatial_parallel_decode_disabled(): + return x + if expected_height is None: + return x + x, _ = gather_variable_height(x) + if x.shape[-2] != expected_height: + x = x[..., :expected_height, :].contiguous() + return x + + +def gather_height_for_global_op(x: torch.Tensor) -> torch.Tensor: + if spatial_parallel_decode_disabled(): + return x + return gather_variable_height(x)[0] + + +def chunk_height_for_parallel_decode(x: torch.Tensor) -> torch.Tensor: + if spatial_parallel_decode_disabled(): + return x + return _tensor_chunk( + x, + dim=-2, + world_size=get_decode_parallel_world_size(), + rank=get_decode_parallel_rank(), + ) + + +def chunk_height_by_sizes(x: torch.Tensor, heights: list[int]) -> torch.Tensor: + if spatial_parallel_decode_disabled(): + return x + rank = get_decode_parallel_rank() + start = sum(heights[:rank]) + return x[..., start : start + heights[rank], :].contiguous( + memory_format=_halo_memory_format(x) + ) + + +def gather_height_sizes(x: torch.Tensor) -> list[int]: + """gather heights of sharded feature_maps from peers""" + if spatial_parallel_decode_disabled(): + return [x.shape[-2]] + world_size = get_decode_parallel_world_size() + if world_size <= 1: + return [x.shape[-2]] + local_height = torch.tensor([x.shape[-2]], device=x.device, dtype=torch.int64) + gathered = [torch.empty_like(local_height) for _ in range(world_size)] + dist.all_gather( + gathered, + local_height, + group=get_decode_parallel_group_coordinator().device_group, + ) + return [int(height.item()) for height in gathered] + + +def gather_variable_height(x: torch.Tensor) -> tuple[torch.Tensor, list[int]]: + if spatial_parallel_decode_disabled(): + return x, [x.shape[-2]] + world_size = get_decode_parallel_world_size() + if world_size <= 1: + return x, [x.shape[-2]] + + heights = gather_height_sizes(x) + max_height = max(heights) + if x.shape[-2] < max_height: + x = _tensor_pad(x, max_height - x.shape[-2], dim=-2) + + gathered = get_decode_parallel_group_coordinator().all_gather( + _maybe_contiguous_for_sp_gather(x), dim=-2 + ) + chunks = torch.split(gathered, max_height, dim=-2) + return ( + torch.cat( + [chunk[..., :height, :] for chunk, height in zip(chunks, heights)], dim=-2 + ), + heights, + ) + + +def _halo_memory_format(reference: torch.Tensor) -> torch.memory_format: + if reference.dim() > 1 and reference.stride(1) == 1: + if reference.dim() == 5 and hasattr(torch, "channels_last_3d"): + return torch.channels_last_3d + if reference.dim() == 4: + return torch.channels_last + return torch.contiguous_format + + +def _ensure_recv_buf( + recv_buf: torch.Tensor | None, reference: torch.Tensor +) -> torch.Tensor: + memory_format = _halo_memory_format(reference) + if ( + recv_buf is None + or recv_buf.shape != reference.shape + or recv_buf.dtype != reference.dtype + or recv_buf.device != reference.device + or not recv_buf.is_contiguous(memory_format=memory_format) + ): + return torch.empty( + reference.shape, + dtype=reference.dtype, + device=reference.device, + memory_format=memory_format, + ) + return recv_buf + + +def halo_exchange( + x: torch.Tensor, + height_halo_size: int = 1, + recv_top_buf: torch.Tensor | None = None, + recv_bottom_buf: torch.Tensor | None = None, + height_pad_mode: str = "zeros", +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """exchange(send and recv) top/bottom conv-input halos with adjacent spatial ranks""" + if spatial_parallel_decode_disabled(): + return x, recv_top_buf, recv_bottom_buf + if height_halo_size == 0: + return x, recv_top_buf, recv_bottom_buf + + 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:, :] + + recv_top_buf = _ensure_recv_buf(recv_top_buf, top_row_ref) + recv_bottom_buf = _ensure_recv_buf(recv_bottom_buf, bottom_row_ref) + p2p_ops = [] + + if rank > 0: + prev_rank = group_ranks[rank - 1] + top_row = top_row_ref.contiguous(memory_format=_halo_memory_format(top_row_ref)) + p2p_ops.append(dist.P2POp(dist.irecv, recv_top_buf, prev_rank, group)) + p2p_ops.append(dist.P2POp(dist.isend, top_row, prev_rank, group)) + if rank < world_size - 1: + next_rank = group_ranks[rank + 1] + bottom_row = bottom_row_ref.contiguous( + memory_format=_halo_memory_format(bottom_row_ref) + ) + p2p_ops.append(dist.P2POp(dist.isend, bottom_row, next_rank, group)) + p2p_ops.append(dist.P2POp(dist.irecv, recv_bottom_buf, next_rank, group)) + + if p2p_ops: + reqs = dist.batch_isend_irecv(p2p_ops) + for req in reqs: + req.wait() + + if rank == 0: + recv_top_buf.copy_( + _make_boundary_halo( + x, + recv_bottom_buf if world_size > 1 else None, + height_halo_size=height_halo_size, + is_top=True, + mode=height_pad_mode, + ) + ) + if rank == world_size - 1: + recv_bottom_buf.copy_( + _make_boundary_halo( + x, + recv_top_buf if world_size > 1 else None, + height_halo_size=height_halo_size, + is_top=False, + mode=height_pad_mode, + ) + ) + + return ( + torch.concat([recv_top_buf, x, recv_bottom_buf], dim=-2), + recv_top_buf, + recv_bottom_buf, + ) + + +def _make_boundary_halo( + x: torch.Tensor, + neighbor: torch.Tensor | None, + *, + height_halo_size: int, + is_top: bool, + mode: str, +) -> torch.Tensor: + if mode == "zeros": + shape = list(x.shape) + shape[-2] = height_halo_size + return torch.zeros(shape, dtype=x.dtype, device=x.device) + if mode == "replicate": + edge = x[..., :1, :] if is_top else x[..., -1:, :] + return edge.expand(*edge.shape[:-2], height_halo_size, edge.shape[-1]) + if mode == "reflect": + source = x + if is_top and neighbor is not None: + source = torch.cat([x, neighbor], dim=-2) + elif not is_top and neighbor is not None: + source = torch.cat([neighbor, x], dim=-2) + if is_top: + index = torch.arange( + height_halo_size, 0, -1, device=x.device, dtype=torch.long + ) + else: + index = torch.arange( + source.shape[-2] - 2, + source.shape[-2] - 2 - height_halo_size, + -1, + device=x.device, + dtype=torch.long, + ) + return source.index_select(-2, index) + raise ValueError(f"Unsupported spatial padding mode for parallel decode: {mode}") + + +def _pad_with_mode( + x: torch.Tensor, padding: tuple[int, ...], mode: str +) -> torch.Tensor: + if mode == "zeros": + return F.pad(x, padding) + return F.pad(x, padding, mode=mode) + + +def _set_conv_padding(module: nn.Module, padding: tuple[int, ...]) -> None: + module.padding = padding + module._reversed_padding_repeated_twice = tuple( + value for pad in reversed(padding) for value in (pad, pad) + ) + + +def _conv_preserves_local_height( + *, + height_halo_size: int, + height_pad_top: int, + height_pad_bottom: int, + kernel_height: int, + dilation_height: int, + stride_height: int, +) -> bool: + kernel_span = dilation_height * (kernel_height - 1) + return ( + stride_height == 1 + and 2 * height_halo_size == kernel_span + and height_pad_top == height_halo_size + and height_pad_bottom == height_halo_size + ) + + +def _conv3d_weight_is_channels_last_3d(weight: torch.Tensor) -> bool: + return ( + weight.dim() == 5 + and hasattr(torch, "channels_last_3d") + and (current_platform.is_cuda() or current_platform.is_rocm()) + and weight.is_contiguous(memory_format=torch.channels_last_3d) + ) + + +def _match_conv3d_input_format(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + if x.dim() == 5 and _conv3d_weight_is_channels_last_3d(weight): + return x.contiguous(memory_format=torch.channels_last_3d) + return x + + +def _spatial_parallel_conv_forward( + module: nn.Module, + x: torch.Tensor, + conv_forward, + *, + height_pad_mode: str, + match_conv3d_format: bool = False, +) -> torch.Tensor: + # send and recv halo + # x_padded: concatenated input + x_padded, module._halo_recv_top_buf, module._halo_recv_bottom_buf = halo_exchange( + x, + height_halo_size=module.height_halo_size, + recv_top_buf=module._halo_recv_top_buf, + recv_bottom_buf=module._halo_recv_bottom_buf, + height_pad_mode=height_pad_mode, + ) + if match_conv3d_format: + x_padded = _match_conv3d_input_format(x_padded, module.weight) + if module.height_halo_size == 0: + return conv_forward(x_padded) + + stride = module.stride[-2] + if _conv_preserves_local_height( + height_halo_size=module.height_halo_size, + height_pad_top=module.height_pad_top, + height_pad_bottom=module.height_pad_bottom, + kernel_height=module.kernel_size[-2], + dilation_height=module.dilation[-2], + stride_height=stride, + ): + return conv_forward(x_padded) + + heights = gather_height_sizes(x) + global_start = sum(heights[: module.rank]) + global_height = sum(heights) + if stride > 1: + shift = ( + global_start - module.height_halo_size + module.height_pad_top + ) % stride + if shift: + x_padded = x_padded[..., shift:, :] + global_start += shift + if match_conv3d_format: + x_padded = _match_conv3d_input_format(x_padded, module.weight) + + out = conv_forward(x_padded) + + # trim the output to original shape + return _trim_conv_output_height( + out, + local_height=x.shape[-2], + global_height=global_height, + global_start=global_start, + height_halo_size=module.height_halo_size, + height_pad_top=module.height_pad_top, + height_pad_bottom=module.height_pad_bottom, + kernel_height=module.kernel_size[-2], + dilation_height=module.dilation[-2], + stride_height=stride, + ) + + +class SpatialParallelConv2d(nn.Conv2d): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int | tuple[int, int], + stride: int | tuple[int, int] = 1, + padding: int | tuple[int, int] = 0, + dilation: int | tuple[int, int] = 1, + groups: int = 1, + bias: bool = True, + padding_mode: str = "zeros", + height_padding: tuple[int, int] | None = None, + ): + super().__init__( + in_channels=in_channels, + out_channels=out_channels, + kernel_size=kernel_size, + stride=stride, + padding=padding, + dilation=dilation, + groups=groups, + bias=bias, + padding_mode=padding_mode, + ) + self.height_halo_size = (self.dilation[-2] * (self.kernel_size[-2] - 1)) // 2 + if height_padding is None: + height_padding = (self.padding[-2], self.padding[-2]) + self.height_pad_top, self.height_pad_bottom = height_padding + + self.padding: tuple[int, int] + if self.height_halo_size > 0: + self._padding = (0, 0, 0, 0) + else: + self._padding = (0, 0, self.padding[0], self.padding[0]) + + _set_conv_padding(self, (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_decode_parallel_rank() + self.world_size = get_decode_parallel_world_size() + + def forward(self, x): + if spatial_parallel_decode_disabled(): + return self._direct_forward(x) + + if any(self._padding): + x = _pad_with_mode(x, self._padding, self.padding_mode) + + return _spatial_parallel_conv_forward( + self, + x, + super().forward, + height_pad_mode=self.padding_mode, + ) + + def _direct_forward(self, x): + width_pad = self.padding[-1] + padding = ( + width_pad, + width_pad, + self.height_pad_top, + self.height_pad_bottom, + ) + if any(padding): + x = _pad_with_mode(x, padding, self.padding_mode) + return F.conv2d( + x, + self.weight, + self.bias, + self.stride, + (0, 0), + self.dilation, + self.groups, + ) + + +class SpatialParallelCausalConv3d(nn.Conv3d): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int | tuple[int, int, int], + stride: int | tuple[int, int, int] = 1, + padding: int | tuple[int, int, int] = 0, + ): + super().__init__( + in_channels=in_channels, + out_channels=out_channels, + kernel_size=kernel_size, + stride=stride, + padding=padding, + ) + + self.height_pad_top = self.padding[1] + self.height_pad_bottom = self.padding[1] + self.height_halo_size = (self.kernel_size[-2] - 1) // 2 + + self.padding: tuple[int, int, int] + if self.height_halo_size > 0: + self._padding = ( + self.padding[2], + self.padding[2], + 0, + 0, + 2 * self.padding[0], + 0, + ) + else: + self._padding = ( + self.padding[2], + self.padding[2], + self.padding[1], + self.padding[1], + 2 * self.padding[0], + 0, + ) + 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_decode_parallel_rank() + self.world_size = get_decode_parallel_world_size() + + def forward(self, x, cache_x=None): + padding = list(self._padding) + if spatial_parallel_decode_disabled(): + padding[2] = self.height_pad_top + padding[3] = self.height_pad_bottom + if cache_x is not None and self._padding[4] > 0: + cache_x = cache_x.to(x.device) + x = torch.cat([cache_x, x], dim=2) + padding[4] -= cache_x.shape[2] + + x = F.pad(x, padding) + x = x if current_platform.is_amp_supported() else x.to(self.weight.dtype) + + if spatial_parallel_decode_disabled(): + x = _match_conv3d_input_format(x, self.weight) + return F.conv3d( + x, + self.weight, + self.bias, + self.stride, + (0, 0, 0), + self.dilation, + self.groups, + ) + + return _spatial_parallel_conv_forward( + self, + x, + super().forward, + height_pad_mode="zeros", + match_conv3d_format=True, + ) + + +class SpatialParallelConv3d(nn.Conv3d): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int | tuple[int, int, int], + stride: int | tuple[int, int, int] = 1, + padding: int | tuple[int, int, int] = 0, + dilation: int | tuple[int, int, int] = 1, + groups: int = 1, + bias: bool = True, + padding_mode: str = "zeros", + height_padding: tuple[int, int] | None = None, + ): + super().__init__( + in_channels=in_channels, + out_channels=out_channels, + kernel_size=kernel_size, + stride=stride, + padding=padding, + dilation=dilation, + groups=groups, + bias=bias, + padding_mode=padding_mode, + ) + self.height_halo_size = (self.dilation[-2] * (self.kernel_size[-2] - 1)) // 2 + if height_padding is None: + height_padding = (self.padding[-2], self.padding[-2]) + self.height_pad_top, self.height_pad_bottom = height_padding + + self.padding: tuple[int, int, int] + if self.height_halo_size > 0: + self._padding = (0, 0, 0, 0, 0, 0) + else: + self._padding = ( + 0, + 0, + self.padding[1], + self.padding[1], + 0, + 0, + ) + + _set_conv_padding(self, (self.padding[0], 0, self.padding[2])) + self._halo_recv_top_buf: torch.Tensor | None = None + self._halo_recv_bottom_buf: torch.Tensor | None = None + self.rank = get_decode_parallel_rank() + self.world_size = get_decode_parallel_world_size() + + def forward(self, x): + if spatial_parallel_decode_disabled(): + return self._direct_forward(x) + + if any(self._padding): + x = _pad_with_mode(x, self._padding, self.padding_mode) + + return _spatial_parallel_conv_forward( + self, + x, + super().forward, + height_pad_mode=self.padding_mode, + match_conv3d_format=True, + ) + + def _direct_forward(self, x): + time_pad = self.padding[0] + width_pad = self.padding[-1] + padding = ( + width_pad, + width_pad, + self.height_pad_top, + self.height_pad_bottom, + time_pad, + time_pad, + ) + if any(padding): + x = _pad_with_mode(x, padding, self.padding_mode) + x = _match_conv3d_input_format(x, self.weight) + return F.conv3d( + x, + self.weight, + self.bias, + self.stride, + (0, 0, 0), + self.dilation, + self.groups, + ) + + +class SpatialParallelZeroPad2d(nn.Module): + def __init__(self, padding: tuple[int, int, int, int]) -> None: + super().__init__() + self.padding = padding + self.rank = get_decode_parallel_rank() + self.world_size = get_decode_parallel_world_size() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if spatial_parallel_decode_disabled(): + return F.pad(x, self.padding) + left, right, top, bottom = self.padding + top = top if self.rank == 0 else 0 + bottom = bottom if self.rank == self.world_size - 1 else 0 + return F.pad(x, (left, right, top, bottom)) + + +def _trim_conv_output_height( + out: torch.Tensor, + *, + local_height: int, + global_height: int, + global_start: int, + height_halo_size: int, + height_pad_top: int, + height_pad_bottom: int, + kernel_height: int, + dilation_height: int, + stride_height: int, +) -> torch.Tensor: + kernel_span = dilation_height * (kernel_height - 1) + min_i = math.ceil( + ((-height_pad_top) - (global_start - height_halo_size)) / stride_height + ) + max_i = math.floor( + ( + (global_height - 1 + height_pad_bottom) + - kernel_span + - (global_start - height_halo_size) + ) + / stride_height + ) + start = max(min_i, 0) + end = min(max_i + 1, out.shape[-2]) + if start != 0 or end != out.shape[-2]: + out = out[..., start:end, :] + return out diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder.py index ecb780458..3e79b5b76 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder.py @@ -25,6 +25,13 @@ from sglang.multimodal_gen.configs.models.vaes.flux import FluxVAEConfig from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( LayerwiseOffloadableModuleMixin, ) +from sglang.multimodal_gen.runtime.models.vaes.common import ( + can_install_spatial_shard_parallel_decode, +) +from sglang.multimodal_gen.runtime.models.vaes.parallel.diffusers_spatial import ( + enable_diffusers_decoder_spatial_parallel, + spatial_parallel_diffusers_decode, +) class AutoencoderKL(nn.Module, LayerwiseOffloadableModuleMixin): @@ -127,6 +134,15 @@ class AutoencoderKL(nn.Module, LayerwiseOffloadableModuleMixin): self.use_slicing = False self.use_tiling = False + self.use_parallel_decode = config.use_parallel_decode + self.parallel_decode_mode = config.parallel_decode_mode + self._spatial_parallel_decode_enabled = False + self._spatial_parallel_upsample_count = 0 + if can_install_spatial_shard_parallel_decode(self.config): + self._spatial_parallel_upsample_count = ( + enable_diffusers_decoder_spatial_parallel(self.decoder) + ) + self._spatial_parallel_decode_enabled = True # only relevant if vae tiling is enabled self.tile_sample_min_size = sample_size @@ -311,7 +327,12 @@ class AutoencoderKL(nn.Module, LayerwiseOffloadableModuleMixin): if self.post_quant_conv is not None: z = self.post_quant_conv(z) - dec = self.decoder(z) + if self._spatial_parallel_decode_enabled: + dec = spatial_parallel_diffusers_decode( + self.decoder, z, self._spatial_parallel_upsample_count + ) + else: + dec = self.decoder(z) if not return_dict: return (dec,) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_dc.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_dc.py index 8259a5b7c..dc5e3a367 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_dc.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_dc.py @@ -6,9 +6,23 @@ import torch from torch import nn from sglang.multimodal_gen.configs.models.vaes.sana import SanaVAEConfig +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_decode_parallel_rank, + get_decode_parallel_world_size, +) +from sglang.multimodal_gen.runtime.layers.parallel_conv import ( + gather_and_trim_height, + split_height_for_parallel_decode, +) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( LayerwiseOffloadableModuleMixin, ) +from sglang.multimodal_gen.runtime.models.vaes.common import ( + can_install_spatial_shard_parallel_decode, +) +from sglang.multimodal_gen.runtime.models.vaes.parallel.diffusers_spatial import ( + enable_diffusers_decoder_spatial_parallel, +) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) @@ -25,6 +39,7 @@ class AutoencoderDC(nn.Module, LayerwiseOffloadableModuleMixin): self._config = config self._inner_model = None self._loaded_state_dict: dict[str, torch.Tensor] = {} + self._spatial_parallel_decode_enabled = False def _ensure_inner_model(self, state_dict: dict[str, torch.Tensor] | None = None): if self._inner_model is not None: @@ -72,6 +87,9 @@ class AutoencoderDC(nn.Module, LayerwiseOffloadableModuleMixin): self._loaded_state_dict.clear() self._inner_model = self._inner_model.to(device) + if can_install_spatial_shard_parallel_decode(self._config): + enable_diffusers_decoder_spatial_parallel(self._inner_model.decoder) + self._spatial_parallel_decode_enabled = True @property def config(self): @@ -98,7 +116,24 @@ class AutoencoderDC(nn.Module, LayerwiseOffloadableModuleMixin): def decode(self, z: torch.Tensor, **kwargs): self._ensure_inner_model() z = z.to(dtype=self.dtype) - return self._inner_model.decode(z, **kwargs) + if not self._spatial_parallel_decode_enabled: + return self._inner_model.decode(z, **kwargs) + + expected_height = ( + z.shape[-2] * self._config.arch_config.spatial_compression_ratio + ) + z, expected_height = split_height_for_parallel_decode( + z, + expected_height=expected_height, + world_size=get_decode_parallel_world_size(), + rank=get_decode_parallel_rank(), + ) + decoded = self._inner_model.decode(z, **kwargs) + if isinstance(decoded, tuple): + sample = gather_and_trim_height(decoded[0], expected_height) + return (sample, *decoded[1:]) + sample = gather_and_trim_height(decoded.sample, expected_height) + return decoded.__class__(sample=sample) def forward(self, x: torch.Tensor, **kwargs): self._ensure_inner_model() diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_flux2.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_flux2.py index 5c1bf65ea..c46238dc6 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_flux2.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_flux2.py @@ -19,7 +19,14 @@ from diffusers.models.autoencoders.vae import ( from diffusers.models.modeling_outputs import AutoencoderKLOutput from sglang.multimodal_gen.configs.models.vaes.flux import Flux2VAEConfig -from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE +from sglang.multimodal_gen.runtime.models.vaes.common import ( + ParallelTiledVAE, + can_install_spatial_shard_parallel_decode, +) +from sglang.multimodal_gen.runtime.models.vaes.parallel.diffusers_spatial import ( + enable_diffusers_decoder_spatial_parallel, + spatial_parallel_diffusers_decode, +) class AutoencoderKLFlux2(ParallelTiledVAE): @@ -110,6 +117,13 @@ class AutoencoderKLFlux2(ParallelTiledVAE): self.use_slicing = False self.use_tiling = False + self._spatial_parallel_decode_enabled = False + self._spatial_parallel_upsample_count = 0 + if can_install_spatial_shard_parallel_decode(self.config): + self._spatial_parallel_upsample_count = ( + enable_diffusers_decoder_spatial_parallel(self.decoder) + ) + self._spatial_parallel_decode_enabled = True # only relevant if vae tiling is enabled self.tile_sample_min_size = self.config.sample_size @@ -266,7 +280,12 @@ class AutoencoderKLFlux2(ParallelTiledVAE): if self.post_quant_conv is not None: z = self.post_quant_conv(z) - dec = self.decoder(z) + if self._spatial_parallel_decode_enabled: + dec = spatial_parallel_diffusers_decode( + self.decoder, z, self._spatial_parallel_upsample_count + ) + else: + dec = self.decoder(z) if not return_dict: return (dec,) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py index 2578b4c42..cad3530f1 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage.py @@ -3,6 +3,7 @@ from typing import Optional, Tuple, Union import torch +import torch.distributed as dist import torch.nn as nn import torch.nn.functional as F from diffusers.models.activations import get_activation @@ -13,11 +14,28 @@ from diffusers.models.autoencoders.vae import ( from diffusers.models.modeling_outputs import AutoencoderKLOutput from sglang.multimodal_gen.configs.models.vaes.qwenimage import QwenImageVAEConfig -from sglang.multimodal_gen.runtime.distributed import ( - get_local_torch_device, - get_sp_world_size, +from sglang.multimodal_gen.runtime.distributed import get_local_torch_device +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_decode_parallel_rank, + get_decode_parallel_world_size, + model_parallel_is_initialized, +) +from sglang.multimodal_gen.runtime.layers.parallel_conv import ( + SpatialParallelCausalConv3d, + SpatialParallelConv2d, + SpatialParallelZeroPad2d, + chunk_height_for_parallel_decode, + disable_spatial_parallel_decode, + gather_and_trim_height, + gather_height_for_global_op, + split_for_parallel_decode, +) +from sglang.multimodal_gen.runtime.models.vaes.common import ( + ParallelTiledVAE, + can_install_spatial_shard_parallel_decode, + has_decode_parallel_world, + should_run_spatial_shard_parallel_decode, ) -from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name @@ -140,7 +158,14 @@ class QwenImageResample(nn.Module): - 'downsample3d': 3D downsampling with zero-padding, convolution, and causal 3D convolution. """ - def __init__(self, dim: int, mode: str) -> None: + def __init__( + self, + dim: int, + mode: str, + conv2d_cls: type[nn.Conv2d] = nn.Conv2d, + causal_conv3d_cls: type[nn.Conv3d] = QwenImageCausalConv3d, + zero_pad2d_cls: type[nn.Module] = nn.ZeroPad2d, + ) -> None: super().__init__() self.dim = dim self.mode = mode @@ -149,26 +174,40 @@ class QwenImageResample(nn.Module): if mode == "upsample2d": self.resample = nn.Sequential( QwenImageUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - nn.Conv2d(dim, dim // 2, 3, padding=1), + conv2d_cls(dim, dim // 2, 3, padding=1), ) elif mode == "upsample3d": self.resample = nn.Sequential( QwenImageUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - nn.Conv2d(dim, dim // 2, 3, padding=1), + conv2d_cls(dim, dim // 2, 3, padding=1), ) - self.time_conv = QwenImageCausalConv3d( + self.time_conv = causal_conv3d_cls( dim, dim * 2, (3, 1, 1), padding=(1, 0, 0) ) elif mode == "downsample2d": - self.resample = nn.Sequential( - nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)) - ) + if conv2d_cls is SpatialParallelConv2d: + self.resample = nn.Sequential( + zero_pad2d_cls((0, 1, 0, 0)), + conv2d_cls(dim, dim, 3, stride=(2, 2), height_padding=(0, 1)), + ) + else: + self.resample = nn.Sequential( + zero_pad2d_cls((0, 1, 0, 1)), + conv2d_cls(dim, dim, 3, stride=(2, 2)), + ) elif mode == "downsample3d": - self.resample = nn.Sequential( - nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)) - ) - self.time_conv = QwenImageCausalConv3d( + if conv2d_cls is SpatialParallelConv2d: + self.resample = nn.Sequential( + zero_pad2d_cls((0, 1, 0, 0)), + conv2d_cls(dim, dim, 3, stride=(2, 2), height_padding=(0, 1)), + ) + else: + self.resample = nn.Sequential( + zero_pad2d_cls((0, 1, 0, 1)), + conv2d_cls(dim, dim, 3, stride=(2, 2)), + ) + self.time_conv = causal_conv3d_cls( dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0) ) @@ -257,6 +296,7 @@ class QwenImageResidualBlock(nn.Module): out_dim: int, dropout: float = 0.0, non_linearity: str = "silu", + causal_conv3d_cls: type[nn.Conv3d] = QwenImageCausalConv3d, ) -> None: super().__init__() self.in_dim = in_dim @@ -265,12 +305,17 @@ class QwenImageResidualBlock(nn.Module): # layers self.norm1 = QwenImageRMS_norm(in_dim, images=False) - self.conv1 = QwenImageCausalConv3d(in_dim, out_dim, 3, padding=1) + self.conv1 = causal_conv3d_cls(in_dim, out_dim, 3, padding=1) self.norm2 = QwenImageRMS_norm(out_dim, images=False) self.dropout = nn.Dropout(dropout) - self.conv2 = QwenImageCausalConv3d(out_dim, out_dim, 3, padding=1) + self.conv2 = causal_conv3d_cls(out_dim, out_dim, 3, padding=1) + shortcut_conv3d_cls = ( + QwenImageCausalConv3d + if causal_conv3d_cls is SpatialParallelCausalConv3d + else causal_conv3d_cls + ) self.conv_shortcut = ( - QwenImageCausalConv3d(in_dim, out_dim, 1) + shortcut_conv3d_cls(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity() ) @@ -338,9 +383,10 @@ class QwenImageAttentionBlock(nn.Module): dim (int): The number of channels in the input tensor. """ - def __init__(self, dim): + def __init__(self, dim, spatial_parallel: bool = False): super().__init__() self.dim = dim + self.spatial_parallel = spatial_parallel # layers self.norm = QwenImageRMS_norm(dim) @@ -348,6 +394,8 @@ class QwenImageAttentionBlock(nn.Module): self.proj = nn.Conv2d(dim, dim, 1) def forward(self, x): + if self.spatial_parallel: + x = gather_height_for_global_op(x).contiguous() identity = x batch_size, channels, time, height, width = x.size() @@ -376,7 +424,10 @@ class QwenImageAttentionBlock(nn.Module): x = x.view(batch_size, time, channels, height, width) x = x.permute(0, 2, 1, 3, 4) - return x + identity + x = x + identity + if self.spatial_parallel: + x = chunk_height_for_parallel_decode(x) + return x class QwenImageMidBlock(nn.Module): @@ -395,16 +446,24 @@ class QwenImageMidBlock(nn.Module): dropout: float = 0.0, non_linearity: str = "silu", num_layers: int = 1, + spatial_parallel: bool = False, + causal_conv3d_cls: type[nn.Conv3d] = QwenImageCausalConv3d, ): super().__init__() self.dim = dim # Create the components - resnets = [QwenImageResidualBlock(dim, dim, dropout, non_linearity)] + resnets = [ + QwenImageResidualBlock(dim, dim, dropout, non_linearity, causal_conv3d_cls) + ] attentions = [] for _ in range(num_layers): - attentions.append(QwenImageAttentionBlock(dim)) - resnets.append(QwenImageResidualBlock(dim, dim, dropout, non_linearity)) + attentions.append(QwenImageAttentionBlock(dim, spatial_parallel)) + resnets.append( + QwenImageResidualBlock( + dim, dim, dropout, non_linearity, causal_conv3d_cls + ) + ) self.attentions = nn.ModuleList(attentions) self.resnets = nn.ModuleList(resnets) @@ -452,6 +511,8 @@ class QwenImageEncoder3d(nn.Module): input_channels: int = 3, ): super().__init__() + dim_mult = list(dim_mult) + temperal_downsample = list(temperal_downsample) # dim = config.arch_config.dim # z_dim = config.arch_config.z_dim # dim_mult = config.arch_config.dim_mult @@ -577,6 +638,9 @@ class QwenImageUpBlock(nn.Module): dropout: float = 0.0, upsample_mode: Optional[str] = None, non_linearity: str = "silu", + conv2d_cls: type[nn.Conv2d] = nn.Conv2d, + causal_conv3d_cls: type[nn.Conv3d] = QwenImageCausalConv3d, + zero_pad2d_cls: type[nn.Module] = nn.ZeroPad2d, ): super().__init__() self.in_dim = in_dim @@ -588,7 +652,13 @@ class QwenImageUpBlock(nn.Module): current_dim = in_dim for _ in range(num_res_blocks + 1): resnets.append( - QwenImageResidualBlock(current_dim, out_dim, dropout, non_linearity) + QwenImageResidualBlock( + current_dim, + out_dim, + dropout, + non_linearity, + causal_conv3d_cls, + ) ) current_dim = out_dim @@ -598,7 +668,15 @@ class QwenImageUpBlock(nn.Module): self.upsamplers = None if upsample_mode is not None: self.upsamplers = nn.ModuleList( - [QwenImageResample(out_dim, mode=upsample_mode)] + [ + QwenImageResample( + out_dim, + mode=upsample_mode, + conv2d_cls=conv2d_cls, + causal_conv3d_cls=causal_conv3d_cls, + zero_pad2d_cls=zero_pad2d_cls, + ) + ] ) self.gradient_checkpointing = False @@ -655,8 +733,11 @@ class QwenImageDecoder3d(nn.Module): dropout=0.0, non_linearity: str = "silu", input_channels=3, + use_parallel_decode: bool = False, ): super().__init__() + dim_mult = list(dim_mult) + temperal_upsample = list(temperal_upsample) self.dim = dim self.z_dim = z_dim self.dim_mult = dim_mult @@ -665,20 +746,44 @@ class QwenImageDecoder3d(nn.Module): self.temperal_upsample = temperal_upsample self.nonlinearity = get_activation(non_linearity) + self.use_parallel_decode = use_parallel_decode + self.upsample_count = 0 + self.world_size = 1 + self.rank = 0 + if dist.is_initialized() and model_parallel_is_initialized(): + self.world_size = get_decode_parallel_world_size() + self.rank = get_decode_parallel_rank() + + if use_parallel_decode and self.world_size > 1: + CausalConv3d = SpatialParallelCausalConv3d + Conv2d = SpatialParallelConv2d + ZeroPad2d = SpatialParallelZeroPad2d + spatial_parallel = True + else: + CausalConv3d = QwenImageCausalConv3d + Conv2d = nn.Conv2d + ZeroPad2d = nn.ZeroPad2d + spatial_parallel = False # dimensions dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]] scale = 1.0 / 2 ** (len(dim_mult) - 2) # init block - self.conv_in = QwenImageCausalConv3d(z_dim, dims[0], 3, padding=1) + self.conv_in = CausalConv3d(z_dim, dims[0], 3, padding=1) # middle blocks self.mid_block = QwenImageMidBlock( - dims[0], dropout, non_linearity, num_layers=1 + dims[0], + dropout, + non_linearity, + num_layers=1, + spatial_parallel=spatial_parallel, + causal_conv3d_cls=CausalConv3d, ) # upsample blocks + self.upsample_count = 0 self.up_blocks = nn.ModuleList([]) for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): # residual (+attention) blocks @@ -698,8 +803,13 @@ class QwenImageDecoder3d(nn.Module): dropout=dropout, upsample_mode=upsample_mode, non_linearity=non_linearity, + conv2d_cls=Conv2d, + causal_conv3d_cls=CausalConv3d, + zero_pad2d_cls=ZeroPad2d, ) self.up_blocks.append(up_block) + if upsample_mode is not None: + self.upsample_count += 1 # Update scale for next iteration if upsample_mode is not None: @@ -707,11 +817,17 @@ class QwenImageDecoder3d(nn.Module): # output blocks self.norm_out = QwenImageRMS_norm(out_dim, images=False) - self.conv_out = QwenImageCausalConv3d(out_dim, input_channels, 3, padding=1) + self.conv_out = CausalConv3d(out_dim, input_channels, 3, padding=1) self.gradient_checkpointing = False def forward(self, x, feat_cache=None, feat_idx=[0]): + expected_height = None + if self.use_parallel_decode and self.world_size > 1: + x, expected_height = split_for_parallel_decode( + x, self.upsample_count, self.world_size, self.rank + ) + ## conv1 if feat_cache is not None: idx = feat_idx[0] @@ -758,6 +874,8 @@ class QwenImageDecoder3d(nn.Module): feat_idx[0] += 1 else: x = self.conv_out(x) + if self.use_parallel_decode and self.world_size > 1: + x = gather_and_trim_height(x, expected_height) return x @@ -791,8 +909,12 @@ class AutoencoderKLQwenImage(ParallelTiledVAE): self.temperal_upsample = temperal_downsample[::-1] self.input_channels = config.arch_config.input_channels self.latents_mean = config.arch_config.latents_mean + self.vae_config = config self.config = config.arch_config self.use_parallel_decode = config.use_parallel_decode + self._spatial_parallel_decode_enabled = ( + can_install_spatial_shard_parallel_decode(config) + ) self.encoder = QwenImageEncoder3d( base_dim, z_dim * 2, dim_mult, num_res_blocks, attn_scales, self.temperal_downsample, dropout, @@ -803,7 +925,8 @@ class AutoencoderKLQwenImage(ParallelTiledVAE): self.decoder = QwenImageDecoder3d( base_dim, z_dim, dim_mult, num_res_blocks, attn_scales, self.temperal_upsample, dropout, - input_channels=self.input_channels + input_channels=self.input_channels, + use_parallel_decode=self._spatial_parallel_decode_enabled, ) # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension @@ -825,7 +948,10 @@ class AutoencoderKLQwenImage(ParallelTiledVAE): # Precompute and cache conv counts for encoder and decoder for clear_cache speedup self._cached_conv_counts = { - "decoder": sum(isinstance(m, QwenImageCausalConv3d) for m in self.decoder.modules()) + "decoder": sum( + isinstance(m, (QwenImageCausalConv3d, SpatialParallelCausalConv3d)) + for m in self.decoder.modules() + ) if self.decoder is not None else 0, "encoder": sum(isinstance(m, QwenImageCausalConv3d) for m in self.encoder.modules()) @@ -902,7 +1028,7 @@ class AutoencoderKLQwenImage(ParallelTiledVAE): def _count_conv3d(model): count = 0 for m in model.modules(): - if isinstance(m, QwenImageCausalConv3d): + if isinstance(m, (QwenImageCausalConv3d, SpatialParallelCausalConv3d)): count += 1 return count @@ -962,30 +1088,29 @@ class AutoencoderKLQwenImage(ParallelTiledVAE): return posterior + def _should_use_spatial_parallel_decode(self, z: torch.Tensor) -> bool: + return ( + self._spatial_parallel_decode_enabled + and should_run_spatial_shard_parallel_decode(self.vae_config, z) + ) + def _decode_with_parallel_dispatch(self, z: torch.Tensor) -> DecoderOutput: - if self.use_parallel_decode and get_sp_world_size() > 1: + if self.use_parallel_decode and has_decode_parallel_world(): num_frame = z.shape[2] num_sample_frames = (num_frame - 1) * self.temporal_compression_ratio + 1 - tile_latent_min_height = ( - self.tile_sample_min_height // self.spatial_compression_ratio - ) - tile_latent_min_width = ( - self.tile_sample_min_width // self.spatial_compression_ratio - ) mode = self.parallel_decode_mode - if mode == "auto": - if ( - z.shape[-2] > tile_latent_min_height - or z.shape[-1] > tile_latent_min_width - ): - mode = "tiled" - else: - mode = "patch" - if mode == "patch": + if self._should_use_spatial_parallel_decode(z): + decoded = self._decode(z)[:, :, :num_sample_frames] + elif self._spatial_parallel_decode_enabled: + with disable_spatial_parallel_decode(): + decoded = self._decode(z)[:, :, :num_sample_frames] + elif mode == "patch": decoded = super().parallel_patch_decode(z)[:, :, :num_sample_frames] - else: + elif mode == "tiled": decoded = super().parallel_tiled_decode(z)[:, :, :num_sample_frames] + else: + decoded = self._decode(z)[:, :, :num_sample_frames] return DecoderOutput(sample=decoded) return DecoderOutput(sample=self._decode(z)) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/common.py b/python/sglang/multimodal_gen/runtime/models/vaes/common.py index d83d417d1..e7e4edbeb 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/common.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/common.py @@ -3,6 +3,7 @@ # SPDX-License-Identifier: Apache-2.0 from abc import ABC, abstractmethod +from functools import lru_cache from math import isqrt, prod from typing import Optional, cast @@ -14,15 +15,66 @@ from diffusers.utils.torch_utils import randn_tensor from torch import nn from sglang.multimodal_gen.configs.models import VAEConfig +from sglang.multimodal_gen.configs.models.vaes.base import ( + should_use_spatial_shard_parallel_decode, +) from sglang.multimodal_gen.runtime.distributed import ( + get_decode_parallel_group_coordinator, + get_decode_parallel_world_size, get_sp_parallel_rank, get_sp_world_size, + model_parallel_is_initialized, ) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( LayerwiseOffloadableModuleMixin, ) +@lru_cache(maxsize=1) +def _cached_decode_parallel_world_size( + is_dist_initialized: bool, is_model_parallel_initialized: bool, group_id: int +) -> int: + if not is_dist_initialized or not is_model_parallel_initialized: + return 1 + return get_decode_parallel_world_size() + + +def _decode_parallel_world_size() -> int: + is_dist_initialized = dist.is_initialized() + is_model_parallel_initialized = model_parallel_is_initialized() + if not is_dist_initialized or not is_model_parallel_initialized: + return _cached_decode_parallel_world_size( + is_dist_initialized, is_model_parallel_initialized, 0 + ) + return _cached_decode_parallel_world_size( + is_dist_initialized, + is_model_parallel_initialized, + id(get_decode_parallel_group_coordinator()), + ) + + +def has_decode_parallel_world() -> bool: + return _decode_parallel_world_size() > 1 + + +def can_install_spatial_shard_parallel_decode(config: VAEConfig | None) -> bool: + world_size = _decode_parallel_world_size() + return ( + config is not None + and world_size > 1 + and should_use_spatial_shard_parallel_decode(config, world_size=world_size) + ) + + +def should_run_spatial_shard_parallel_decode( + config: VAEConfig, z: torch.Tensor +) -> bool: + world_size = _decode_parallel_world_size() + return world_size > 1 and should_use_spatial_shard_parallel_decode( + config, z, world_size + ) + + class ParallelTiledVAE(ABC, nn.Module, LayerwiseOffloadableModuleMixin): layerwise_offload_dit_group_enabled = False layer_names = [ @@ -115,7 +167,15 @@ class ParallelTiledVAE(ABC, nn.Module, LayerwiseOffloadableModuleMixin): ) num_sample_frames = (num_frames - 1) * self.temporal_compression_ratio + 1 - if self.use_tiling and self.use_parallel_tiling and get_sp_world_size() > 1: + if should_run_spatial_shard_parallel_decode(self.config, z): + return self._decode(z)[:, :, :num_sample_frames] + + if ( + self.parallel_decode_mode == "tiled" + and self.use_tiling + and self.use_parallel_tiling + and get_sp_world_size() > 1 + ): return self.parallel_tiled_decode(z)[:, :, :num_sample_frames] if ( self.use_tiling diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/hunyuanvae.py b/python/sglang/multimodal_gen/runtime/models/vaes/hunyuanvae.py index 3e3f62ecd..bf2d39f01 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/hunyuanvae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/hunyuanvae.py @@ -24,8 +24,24 @@ import torch.nn.functional as F from sglang.jit_kernel.diffusion.group_norm_silu import apply_group_norm_silu from sglang.multimodal_gen.configs.models.vaes import HunyuanVAEConfig +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + 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.common import ParallelTiledVAE +from sglang.multimodal_gen.runtime.layers.parallel_conv import ( + SpatialParallelConv3d, + chunk_height_by_sizes, + disable_spatial_parallel_decode, + gather_and_trim_height, + gather_variable_height, + split_height_for_parallel_decode, +) +from sglang.multimodal_gen.runtime.models.vaes.common import ( + ParallelTiledVAE, + can_install_spatial_shard_parallel_decode, + should_run_spatial_shard_parallel_decode, +) def prepare_causal_attention_mask( @@ -45,6 +61,44 @@ def prepare_causal_attention_mask( return mask +def _make_spatial_parallel_conv3d( + conv: nn.Conv3d, + *, + height_pad: int, + width_pad: int, + padding_mode: str, +) -> SpatialParallelConv3d: + spatial_conv = SpatialParallelConv3d( + in_channels=conv.in_channels, + out_channels=conv.out_channels, + kernel_size=conv.kernel_size, + stride=conv.stride, + padding=(0, height_pad, width_pad), + dilation=conv.dilation, + groups=conv.groups, + bias=conv.bias is not None, + padding_mode=padding_mode, + ) + spatial_conv.weight = conv.weight + spatial_conv.bias = conv.bias + return spatial_conv + + +def _apply_group_norm_silu( + hidden_states: torch.Tensor, + norm: nn.GroupNorm, + nonlinearity: nn.Module, + spatial_parallel: bool, +) -> torch.Tensor: + if not spatial_parallel: + return apply_group_norm_silu(hidden_states, norm, nonlinearity) + hidden_states, heights = gather_variable_height(hidden_states) + hidden_states = apply_group_norm_silu( + hidden_states.contiguous(), norm, nonlinearity + ) + return chunk_height_by_sizes(hidden_states, heights) + + class HunyuanVAEAttention(nn.Module): def __init__( @@ -142,15 +196,43 @@ class HunyuanVideoCausalConv3d(nn.Module): kernel_size[2] - 1, 0, ) + self.spatial_parallel_time_padding = ( + 0, + 0, + 0, + 0, + kernel_size[2] - 1, + 0, + ) + self.spatial_parallel = False self.conv = nn.Conv3d( in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias ) - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - hidden_states = F.pad( - hidden_states, self.time_causal_padding, mode=self.pad_mode + def enable_spatial_parallel(self) -> None: + if isinstance(self.conv, SpatialParallelConv3d): + self.spatial_parallel = True + return + if self.conv.kernel_size == (1, 1, 1): + self.spatial_parallel = True + return + self.conv = _make_spatial_parallel_conv3d( + self.conv, + height_pad=self.time_causal_padding[2], + width_pad=self.time_causal_padding[0], + padding_mode=self.pad_mode, ) + self.spatial_parallel = True + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + padding = ( + self.spatial_parallel_time_padding + if self.spatial_parallel + else self.time_causal_padding + ) + if any(padding): + hidden_states = F.pad(hidden_states, padding, mode=self.pad_mode) return self.conv(hidden_states) @@ -254,18 +336,19 @@ class HunyuanVideoResnetBlockCausal3D(nn.Module): self.conv_shortcut = HunyuanVideoCausalConv3d( in_channels, out_channels, 1, 1, 0 ) + self.spatial_parallel = False def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: hidden_states = hidden_states.contiguous() residual = hidden_states - hidden_states = apply_group_norm_silu( - hidden_states, self.norm1, self.nonlinearity + hidden_states = _apply_group_norm_silu( + hidden_states, self.norm1, self.nonlinearity, self.spatial_parallel ) hidden_states = self.conv1(hidden_states) - hidden_states = apply_group_norm_silu( - hidden_states, self.norm2, self.nonlinearity + hidden_states = _apply_group_norm_silu( + hidden_states, self.norm2, self.nonlinearity, self.spatial_parallel ) hidden_states = self.dropout(hidden_states) hidden_states = self.conv2(hidden_states) @@ -338,8 +421,33 @@ class HunyuanVideoMidBlock3D(nn.Module): self.attentions = nn.ModuleList(attentions) self.resnets = nn.ModuleList(resnets) + self.spatial_parallel = False self.gradient_checkpointing = False + def _run_attention( + self, attn: HunyuanVAEAttention, hidden_states: torch.Tensor + ) -> torch.Tensor: + heights = None + if self.spatial_parallel: + hidden_states, heights = gather_variable_height(hidden_states) + + batch_size, num_channels, num_frames, height, width = hidden_states.shape + hidden_states = hidden_states.permute(0, 2, 3, 4, 1).flatten(1, 3) + attention_mask = prepare_causal_attention_mask( + num_frames, + height * width, + hidden_states.dtype, + hidden_states.device, + batch_size=batch_size, + ) + hidden_states = attn(hidden_states, attention_mask=attention_mask) + hidden_states = hidden_states.unflatten(1, (num_frames, height, width)).permute( + 0, 4, 1, 2, 3 + ) + if heights is not None: + hidden_states = chunk_height_by_sizes(hidden_states, heights) + return hidden_states + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: if torch.is_grad_enabled() and self.gradient_checkpointing: hidden_states = self._gradient_checkpointing_func( @@ -348,21 +456,7 @@ class HunyuanVideoMidBlock3D(nn.Module): for attn, resnet in zip(self.attentions, self.resnets[1:], strict=True): if attn is not None: - batch_size, num_channels, num_frames, height, width = ( - hidden_states.shape - ) - hidden_states = hidden_states.permute(0, 2, 3, 4, 1).flatten(1, 3) - attention_mask = prepare_causal_attention_mask( - num_frames, - height * width, - hidden_states.dtype, - hidden_states.device, - batch_size=batch_size, - ) - hidden_states = attn(hidden_states, attention_mask=attention_mask) - hidden_states = hidden_states.unflatten( - 1, (num_frames, height, width) - ).permute(0, 4, 1, 2, 3) + hidden_states = self._run_attention(attn, hidden_states) hidden_states = self._gradient_checkpointing_func(resnet, hidden_states) @@ -371,21 +465,7 @@ class HunyuanVideoMidBlock3D(nn.Module): for attn, resnet in zip(self.attentions, self.resnets[1:], strict=True): if attn is not None: - batch_size, num_channels, num_frames, height, width = ( - hidden_states.shape - ) - hidden_states = hidden_states.permute(0, 2, 3, 4, 1).flatten(1, 3) - attention_mask = prepare_causal_attention_mask( - num_frames, - height * width, - hidden_states.dtype, - hidden_states.device, - batch_size=batch_size, - ) - hidden_states = attn(hidden_states, attention_mask=attention_mask) - hidden_states = hidden_states.unflatten( - 1, (num_frames, height, width) - ).permute(0, 4, 1, 2, 3) + hidden_states = self._run_attention(attn, hidden_states) hidden_states = resnet(hidden_states) @@ -736,6 +816,7 @@ class HunyuanVideoDecoder3D(nn.Module): block_out_channels[0], out_channels, kernel_size=3 ) + self.spatial_parallel = False self.gradient_checkpointing = False def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: @@ -757,14 +838,29 @@ class HunyuanVideoDecoder3D(nn.Module): hidden_states = up_block(hidden_states) # post-process - hidden_states = apply_group_norm_silu( - hidden_states, self.conv_norm_out, self.conv_act + hidden_states = _apply_group_norm_silu( + hidden_states, + self.conv_norm_out, + self.conv_act, + self.spatial_parallel, ) hidden_states = self.conv_out(hidden_states) return hidden_states +def _enable_hunyuan_decoder_spatial_parallel(decoder: nn.Module) -> None: + for module in decoder.modules(): + if isinstance(module, HunyuanVideoCausalConv3d): + module.enable_spatial_parallel() + elif isinstance(module, HunyuanVideoResnetBlockCausal3D): + module.spatial_parallel = True + elif isinstance(module, HunyuanVideoMidBlock3D): + module.spatial_parallel = True + elif isinstance(module, HunyuanVideoDecoder3D): + module.spatial_parallel = True + + class AutoencoderKLHunyuanVideo(ParallelTiledVAE): r""" A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. @@ -821,6 +917,16 @@ class AutoencoderKLHunyuanVideo(ParallelTiledVAE): self.post_quant_conv = nn.Conv3d( config.latent_channels, config.latent_channels, kernel_size=1 ) + self._spatial_parallel_decode_enabled = False + if can_install_spatial_shard_parallel_decode(self.config): + _enable_hunyuan_decoder_spatial_parallel(self.decoder) + self._spatial_parallel_decode_enabled = True + + def _should_use_spatial_parallel_decode(self, z: torch.Tensor) -> bool: + return ( + self._spatial_parallel_decode_enabled + and should_run_spatial_shard_parallel_decode(self.config, z) + ) def _encode(self, x: torch.Tensor) -> torch.Tensor: x = self.encoder(x) @@ -829,7 +935,22 @@ class AutoencoderKLHunyuanVideo(ParallelTiledVAE): def _decode(self, z: torch.Tensor) -> torch.Tensor: z = self.post_quant_conv(z) - dec = self.decoder(z) + if self._should_use_spatial_parallel_decode(z): + expected_height = ( + z.shape[-2] * self.config.arch_config.spatial_compression_ratio + ) + z, expected_height = split_height_for_parallel_decode( + z, + expected_height=expected_height, + world_size=get_decode_parallel_world_size(), + rank=get_decode_parallel_rank(), + ) + dec = gather_and_trim_height(self.decoder(z), expected_height) + elif self._spatial_parallel_decode_enabled: + with disable_spatial_parallel_decode(): + dec = self.decoder(z) + else: + dec = self.decoder(z) return dec def forward( diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_vae.py b/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_vae.py index cc652a66f..a985ad871 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_vae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_vae.py @@ -12,7 +12,21 @@ from diffusers.models.embeddings import PixArtAlphaCombinedTimestepSizeEmbedding from diffusers.models.modeling_outputs import AutoencoderKLOutput from sglang.multimodal_gen.configs.models.vaes.ltx_video import LTXVideoVAEConfig -from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_decode_parallel_rank, + get_decode_parallel_world_size, +) +from sglang.multimodal_gen.runtime.layers.parallel_conv import ( + SpatialParallelConv3d, + disable_spatial_parallel_decode, + gather_and_trim_height, + split_height_for_parallel_decode, +) +from sglang.multimodal_gen.runtime.models.vaes.common import ( + ParallelTiledVAE, + can_install_spatial_shard_parallel_decode, + should_run_spatial_shard_parallel_decode, +) @lru_cache(maxsize=128) @@ -218,6 +232,29 @@ class LTX2VideoCausalConv3d(nn.Module): return hidden_states +def _make_spatial_parallel_conv3d(conv: nn.Conv3d) -> SpatialParallelConv3d: + spatial_conv = SpatialParallelConv3d( + in_channels=conv.in_channels, + out_channels=conv.out_channels, + kernel_size=conv.kernel_size, + stride=conv.stride, + padding=conv.padding, + dilation=conv.dilation, + groups=conv.groups, + bias=conv.bias is not None, + padding_mode=conv.padding_mode, + ) + spatial_conv.weight = conv.weight + spatial_conv.bias = conv.bias + return spatial_conv + + +def _enable_ltx_decoder_spatial_parallel(decoder: nn.Module) -> None: + for module in decoder.modules(): + if isinstance(module, LTX2VideoCausalConv3d) and type(module.conv) is nn.Conv3d: + module.conv = _make_spatial_parallel_conv3d(module.conv) + + # Like LTXVideoResnetBlock3d, but uses new causal Conv3d, normal Conv3d for the conv_shortcut, and the spatial padding # mode is configurable class LTX2VideoResnetBlock3d(nn.Module): @@ -1683,6 +1720,10 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE): latents_std = torch.ones((latent_channels,), requires_grad=False) self.register_buffer("latents_mean", latents_mean, persistent=True) self.register_buffer("latents_std", latents_std, persistent=True) + self._spatial_parallel_decode_enabled = False + if can_install_spatial_shard_parallel_decode(self.config): + _enable_ltx_decoder_spatial_parallel(self.decoder) + self._spatial_parallel_decode_enabled = True # When decoding a batch of video latents at a time, one can save memory by slicing across the batch dimension # to perform decoding of a single video latent at a time. @@ -1714,6 +1755,12 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE): self.tile_sample_stride_width = 448 self.tile_sample_stride_num_frames = 8 + def _should_use_spatial_parallel_decode(self, z: torch.Tensor) -> bool: + return ( + self._spatial_parallel_decode_enabled + and should_run_spatial_shard_parallel_decode(self.config, z) + ) + def enable_tiling( self, tile_sample_min_height: Optional[int] = None, @@ -1829,7 +1876,24 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE): ): return self.tiled_decode(z, temb, causal=causal, return_dict=return_dict) - dec = self.decoder(z, temb, causal=causal) + if self._should_use_spatial_parallel_decode(z): + expected_height = ( + z.shape[-2] * self.config.arch_config.spatial_compression_ratio + ) + z, expected_height = split_height_for_parallel_decode( + z, + expected_height=expected_height, + world_size=get_decode_parallel_world_size(), + rank=get_decode_parallel_rank(), + ) + dec = gather_and_trim_height( + self.decoder(z, temb, causal=causal), expected_height + ) + elif self._spatial_parallel_decode_enabled: + with disable_spatial_parallel_decode(): + dec = self.decoder(z, temb, causal=causal) + else: + dec = self.decoder(z, temb, causal=causal) if not return_dict: return (dec,) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/parallel/diffusers_spatial.py b/python/sglang/multimodal_gen/runtime/models/vaes/parallel/diffusers_spatial.py new file mode 100644 index 000000000..3a31f75cb --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/vaes/parallel/diffusers_spatial.py @@ -0,0 +1,109 @@ +import torch +import torch.nn as nn + +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_decode_parallel_rank, + get_decode_parallel_world_size, +) +from sglang.multimodal_gen.runtime.layers.parallel_conv import ( + SpatialParallelConv2d, + chunk_height_by_sizes, + gather_and_trim_height, + gather_variable_height, + split_for_parallel_decode, +) + + +def count_decoder_spatial_upsamples(decoder: nn.Module) -> int: + return sum( + len(upsamplers) + for block in getattr(decoder, "up_blocks", []) + if (upsamplers := getattr(block, "upsamplers", None)) is not None + ) + + +def enable_diffusers_decoder_spatial_parallel(decoder: nn.Module) -> int: + _replace_conv2d_modules(decoder) + _patch_groupnorm_modules(decoder) + _patch_attention_modules(decoder) + return count_decoder_spatial_upsamples(decoder) + + +def spatial_parallel_diffusers_decode( + decoder: nn.Module, z: torch.Tensor, upsample_count: int +) -> torch.Tensor: + z, expected_height = split_for_parallel_decode( + z, + upsample_count=upsample_count, + world_size=get_decode_parallel_world_size(), + rank=get_decode_parallel_rank(), + ) + return gather_and_trim_height(decoder(z), expected_height) + + +def _replace_conv2d_modules(module: nn.Module) -> None: + for name, child in list(module.named_children()): + if type(child) is nn.Conv2d: + setattr(module, name, _make_spatial_conv2d(child)) + else: + _replace_conv2d_modules(child) + + +def _make_spatial_conv2d(conv: nn.Conv2d) -> SpatialParallelConv2d: + spatial_conv = SpatialParallelConv2d( + in_channels=conv.in_channels, + out_channels=conv.out_channels, + kernel_size=conv.kernel_size, + stride=conv.stride, + padding=conv.padding, + dilation=conv.dilation, + groups=conv.groups, + bias=conv.bias is not None, + padding_mode=conv.padding_mode, + ) + spatial_conv.weight = conv.weight + spatial_conv.bias = conv.bias + return spatial_conv + + +def _patch_attention_modules(module: nn.Module) -> None: + for child in module.children(): + if child.__class__.__name__ == "Attention": + _patch_attention_forward(child) + _patch_attention_modules(child) + + +def _patch_groupnorm_modules(module: nn.Module) -> None: + for child in module.children(): + if type(child) is nn.GroupNorm: + _patch_groupnorm_forward(child) + else: + _patch_groupnorm_modules(child) + + +def _patch_groupnorm_forward(norm: nn.GroupNorm) -> None: + original_forward = norm.forward + + def spatial_parallel_forward(hidden_states): + if hidden_states.dim() < 4: + return original_forward(hidden_states) + hidden_states, heights = gather_variable_height(hidden_states) + hidden_states = hidden_states.contiguous() + hidden_states = original_forward(hidden_states) + return chunk_height_by_sizes(hidden_states, heights) + + norm.forward = spatial_parallel_forward + + +def _patch_attention_forward(attn: nn.Module) -> None: + original_forward = attn.forward + + def spatial_parallel_forward(hidden_states, *args, **kwargs): + if hidden_states.dim() != 4: + return original_forward(hidden_states, *args, **kwargs) + hidden_states, heights = gather_variable_height(hidden_states) + hidden_states = hidden_states.contiguous() + hidden_states = original_forward(hidden_states, *args, **kwargs) + return chunk_height_by_sizes(hidden_states, heights) + + attn.forward = spatial_parallel_forward diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_common_utils.py b/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_common_utils.py deleted file mode 100644 index 2fdb933b5..000000000 --- a/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_common_utils.py +++ /dev/null @@ -1,480 +0,0 @@ -from __future__ import annotations - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from sglang.multimodal_gen.runtime.platforms import current_platform - - -def _channels_last_3d_supported_by_platform() -> bool: - return hasattr(torch, "channels_last_3d") and ( - current_platform.is_cuda() or current_platform.is_rocm() - ) - - -def _conv3d_weight_is_channels_last_3d(weight: torch.Tensor) -> bool: - return ( - weight.dim() == 5 - and _channels_last_3d_supported_by_platform() - and weight.is_contiguous(memory_format=torch.channels_last_3d) - ) - - -def match_conv3d_input_format(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: - if x.dim() == 5 and _conv3d_weight_is_channels_last_3d(weight): - return x.contiguous(memory_format=torch.channels_last_3d) - return x - - -class AvgDown3D(nn.Module): - def __init__( - self, - in_channels, - out_channels, - factor_t, - factor_s=1, - ): - super().__init__() - self.in_channels = in_channels - self.out_channels = out_channels - self.factor_t = factor_t - self.factor_s = factor_s - self.factor = self.factor_t * self.factor_s * self.factor_s - - assert in_channels * self.factor % out_channels == 0 - self.group_size = in_channels * self.factor // out_channels - - def forward(self, x: torch.Tensor) -> torch.Tensor: - pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t - pad = (0, 0, 0, 0, pad_t, 0) - x = F.pad(x, pad) - B, C, T, H, W = x.shape - x = x.view( - B, - C, - T // self.factor_t, - self.factor_t, - H // self.factor_s, - self.factor_s, - W // self.factor_s, - self.factor_s, - ) - x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous() - x = x.view( - B, - C * self.factor, - T // self.factor_t, - H // self.factor_s, - W // self.factor_s, - ) - x = x.view( - B, - self.out_channels, - self.group_size, - T // self.factor_t, - H // self.factor_s, - W // self.factor_s, - ) - x = x.mean(dim=2) - return x - - -class DupUp3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - factor_t, - factor_s=1, - ): - super().__init__() - self.in_channels = in_channels - self.out_channels = out_channels - - self.factor_t = factor_t - self.factor_s = factor_s - self.factor = self.factor_t * self.factor_s * self.factor_s - - assert out_channels * self.factor % in_channels == 0 - self.repeats = out_channels * self.factor // in_channels - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = x.repeat_interleave(self.repeats, dim=1) - x = x.view( - x.size(0), - self.out_channels, - self.factor_t, - self.factor_s, - self.factor_s, - x.size(2), - x.size(3), - x.size(4), - ) - x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous() - x = x.view( - x.size(0), - self.out_channels, - x.size(2) * self.factor_t, - x.size(4) * self.factor_s, - x.size(6) * self.factor_s, - ) - - _first_chunk = first_chunk.get() if first_chunk is not None else None - if _first_chunk: - x = x[:, :, self.factor_t - 1 :, :, :] - return x - - -class WanCausalConv3d(nn.Conv3d): - r""" - A custom 3D causal convolution layer with feature caching support. - - This layer extends the standard Conv3D layer by ensuring causality in the time dimension and handling feature - caching for efficient inference. - """ - - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int], - stride: int | tuple[int, int, int] = 1, - padding: int | tuple[int, int, int] = 0, - ) -> None: - super().__init__( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=kernel_size, - stride=stride, - padding=padding, - ) - self.padding: tuple[int, int, int] - # Set up causal padding - self._padding: tuple[int, ...] = ( - self.padding[2], - self.padding[2], - self.padding[1], - self.padding[1], - 2 * self.padding[0], - 0, - ) - self.padding = (0, 0, 0) - - def forward(self, x, cache_x=None): - padding = list(self._padding) - if cache_x is not None and self._padding[4] > 0: - cache_x = cache_x.to(x.device) - x = torch.cat([cache_x, x], dim=2) - padding[4] -= cache_x.shape[2] - x = F.pad(x, padding) - x = ( - x if current_platform.is_amp_supported() else x.to(self.weight.dtype) - ) # casting needed if amp isn't supported - x = match_conv3d_input_format(x, self.weight) - return super().forward(x) - - -class WanRMS_norm(nn.Module): - r""" - A custom RMS normalization layer. - """ - - def __init__( - self, - dim: int, - channel_first: bool = True, - images: bool = True, - bias: bool = False, - ) -> None: - super().__init__() - broadcastable_dims = (1, 1, 1) if not images else (1, 1) - shape = (dim, *broadcastable_dims) if channel_first else (dim,) - - self.channel_first = channel_first - self.scale = dim**0.5 - self.gamma = nn.Parameter(torch.ones(shape)) - self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 - - def forward(self, x): - return ( - F.normalize(x, dim=(1 if self.channel_first else -1)) - * self.scale - * self.gamma - + self.bias - ) - - -class WanUpsample(nn.Upsample): - r""" - Perform upsampling while ensuring the output tensor has the same data type as the input. - """ - - def forward(self, x): - if current_platform.is_amp_supported(): - return super().forward(x) - return super().forward(x.float()).type_as(x) - - -is_first_frame = None -feat_cache = None -feat_idx = None -cache_t = None -first_chunk = None - - -def bind_context( - is_first_frame_var, - feat_cache_var, - feat_idx_var, - cache_t_value, - first_chunk_var, -): - global is_first_frame - global feat_cache - global feat_idx - global cache_t - global first_chunk - is_first_frame = is_first_frame_var - feat_cache = feat_cache_var - feat_idx = feat_idx_var - cache_t = cache_t_value - first_chunk = first_chunk_var - - -def _ensure_bound(): - if ( - is_first_frame is None - or feat_cache is None - or feat_idx is None - or cache_t is None - or first_chunk is None - ): - raise RuntimeError("common_utils.bind_context() must be called before use.") - - -def resample_forward(self, x): - _ensure_bound() - b, c, t, h, w = x.size() - first_frame = is_first_frame.get() - if first_frame: - assert t == 1 - _feat_cache = feat_cache.get() - _feat_idx = feat_idx.get() - if self.mode == "upsample3d": - if _feat_cache is not None: - idx = _feat_idx - if _feat_cache[idx] is None: - _feat_cache[idx] = "Rep" - _feat_idx += 1 - else: - cache_x = x[:, :, -cache_t:, :, :].clone() - if ( - cache_x.shape[2] < 2 - and _feat_cache[idx] is not None - and _feat_cache[idx] != "Rep" - ): - # cache last frame of last two chunk - cache_x = torch.cat( - [ - _feat_cache[idx][:, :, -1, :, :] - .unsqueeze(2) - .to(cache_x.device), - cache_x, - ], - dim=2, - ) - if ( - cache_x.shape[2] < 2 - and _feat_cache[idx] is not None - and _feat_cache[idx] == "Rep" - ): - cache_x = torch.cat( - [torch.zeros_like(cache_x).to(cache_x.device), cache_x], - dim=2, - ) - if _feat_cache[idx] == "Rep": - x = self.time_conv(x) - else: - x = self.time_conv(x, _feat_cache[idx]) - _feat_cache[idx] = cache_x - _feat_idx += 1 - - x = x.reshape(b, 2, c, t, h, w) - x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3) - x = x.reshape(b, c, t * 2, h, w) - feat_cache.set(_feat_cache) - feat_idx.set(_feat_idx) - elif not first_frame and hasattr(self, "time_conv"): - x = self.time_conv(x) - x = x.reshape(b, 2, c, t, h, w) - x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3) - x = x.reshape(b, c, t * 2, h, w) - t = x.shape[2] - x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) - x = self.resample(x) - x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4) - - _feat_cache = feat_cache.get() - _feat_idx = feat_idx.get() - if self.mode == "downsample3d": - if _feat_cache is not None: - idx = _feat_idx - if _feat_cache[idx] is None: - _feat_cache[idx] = x.clone() - _feat_idx += 1 - else: - cache_x = x[:, :, -1:, :, :].clone() - x = self.time_conv(torch.cat([_feat_cache[idx][:, :, -1:, :, :], x], 2)) - _feat_cache[idx] = cache_x - _feat_idx += 1 - feat_cache.set(_feat_cache) - feat_idx.set(_feat_idx) - elif not first_frame and hasattr(self, "time_conv"): - x = self.time_conv(x) - return x - - -def residual_block_forward(self, x): - _ensure_bound() - # Apply shortcut connection - h = self.conv_shortcut(x) - - # First normalization and activation - x = self.norm1(x) - x = self.nonlinearity(x) - - _feat_cache = feat_cache.get() - _feat_idx = feat_idx.get() - if _feat_cache is not None: - idx = _feat_idx - cache_x = x[:, :, -cache_t:, :, :].clone() - if cache_x.shape[2] < 2 and _feat_cache[idx] is not None: - cache_x = torch.cat( - [ - _feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), - cache_x, - ], - dim=2, - ) - - x = self.conv1(x, _feat_cache[idx]) - _feat_cache[idx] = cache_x - _feat_idx += 1 - feat_cache.set(_feat_cache) - feat_idx.set(_feat_idx) - else: - x = self.conv1(x) - - # Second normalization and activation - x = self.norm2(x) - x = self.nonlinearity(x) - - # Dropout - x = self.dropout(x) - - _feat_cache = feat_cache.get() - _feat_idx = feat_idx.get() - if _feat_cache is not None: - idx = _feat_idx - cache_x = x[:, :, -cache_t:, :, :].clone() - if cache_x.shape[2] < 2 and _feat_cache[idx] is not None: - cache_x = torch.cat( - [ - _feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), - cache_x, - ], - dim=2, - ) - - x = self.conv2(x, _feat_cache[idx]) - _feat_cache[idx] = cache_x - _feat_idx += 1 - feat_cache.set(_feat_cache) - feat_idx.set(_feat_idx) - else: - x = self.conv2(x) - - # Add residual connection - return x + h - - -def attention_block_forward(self, x): - identity = x - batch_size, channels, num_frames, height, width = x.size() - x = x.permute(0, 2, 1, 3, 4).reshape( - batch_size * num_frames, channels, height, width - ) - x = self.norm(x) - - # compute query, key, value - qkv = self.to_qkv(x) - qkv = qkv.reshape(batch_size * num_frames, 1, channels * 3, -1) - qkv = qkv.permute(0, 1, 3, 2).contiguous() - q, k, v = qkv.chunk(3, dim=-1) - - x = torch.nn.functional.scaled_dot_product_attention(q, k, v) - - x = ( - x.squeeze(1) - .permute(0, 2, 1) - .reshape(batch_size * num_frames, channels, height, width) - ) - - # output projection - x = self.proj(x) - - # Reshape back: [(b*t), c, h, w] -> [b, c, t, h, w] - x = x.view(batch_size, num_frames, channels, height, width) - x = x.permute(0, 2, 1, 3, 4) - - return x + identity - - -def mid_block_forward(self, x): - # First residual block - x = self.resnets[0](x) - - # Process through attention and residual blocks - for attn, resnet in zip(self.attentions, self.resnets[1:], strict=True): - if attn is not None: - x = attn(x) - - x = resnet(x) - - return x - - -def residual_down_block_forward(self, x): - x_copy = x - for resnet in self.resnets: - x = resnet(x) - if self.downsampler is not None: - x = self.downsampler(x) - - return x + self.avg_shortcut(x_copy) - - -def residual_up_block_forward(self, x): - if self.avg_shortcut is not None: - x_copy = x - - for resnet in self.resnets: - x = resnet(x) - - if self.upsampler is not None: - x = self.upsampler(x) - - if self.avg_shortcut is not None: - x = x + self.avg_shortcut(x_copy) - - return x - - -def up_block_forward(self, x): - for resnet in self.resnets: - x = resnet(x) - - if self.upsamplers is not None: - x = self.upsamplers[0](x) - return x 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 deleted file mode 100644 index e590a1bb5..000000000 --- a/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_dist_utils.py +++ /dev/null @@ -1,716 +0,0 @@ -import math - -import torch -import torch.distributed as dist -import torch.nn as nn -import torch.nn.functional as F - -from sglang.multimodal_gen.runtime.distributed.parallel_state import ( - 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 ( - AvgDown3D, - DupUp3D, - WanCausalConv3d, - WanRMS_norm, - WanUpsample, - attention_block_forward, - match_conv3d_input_format, - mid_block_forward, - resample_forward, - residual_block_forward, - residual_down_block_forward, - residual_up_block_forward, - up_block_forward, -) -from sglang.multimodal_gen.runtime.platforms import current_platform - - -def tensor_pad(x: torch.Tensor, len_to_pad: int, dim: int = -2): - x = torch.cat( - [ - x, - torch.zeros( - *x.shape[:dim], - len_to_pad, - *x.shape[dim + 1 :], - dtype=x.dtype, - device=x.device, - ), - ], - dim=dim, - ) - return x - - -def tensor_chunk(x: torch.Tensor, dim: int = -2, world_size: int = 1, rank: int = 0): - if x is None: - return None - if world_size <= 1: - return x - len_to_padding = (int(math.ceil(x.shape[dim] / world_size)) * world_size) - x.shape[ - dim - ] - if len_to_padding != 0: - x = tensor_pad(x, len_to_padding, dim=dim) - return torch.chunk(x, world_size, dim=dim)[rank] - - -def split_for_parallel_encode( - x: torch.Tensor, downsample_count: int, world_size: int, rank: int -): - orig_height = x.shape[-2] - expected_height = orig_height // (2**downsample_count) - factor = world_size * (2**downsample_count) - pad_h = (factor - orig_height % factor) % factor - if pad_h: - x = F.pad(x, (0, 0, 0, pad_h, 0, 0)) - expected_local_height = (orig_height + pad_h) // (2**downsample_count) // world_size - x = tensor_chunk(x, dim=-2, world_size=world_size, rank=rank) - return x, expected_height, expected_local_height - - -def ensure_local_height(x: torch.Tensor, expected_local_height: int | None): - if expected_local_height is None: - return x - if x.shape[-2] < expected_local_height: - pad = expected_local_height - x.shape[-2] - return F.pad(x, (0, 0, 0, pad, 0, 0)) - if x.shape[-2] > expected_local_height: - return x[..., :expected_local_height, :].contiguous() - return x - - -def split_for_parallel_decode( - x: torch.Tensor, upsample_count: int, world_size: int, rank: int -): - expected_height = x.shape[-2] * (2**upsample_count) - x = tensor_chunk(x, dim=-2, world_size=world_size, rank=rank) - return x, expected_height - - -def _maybe_contiguous_for_sp_gather(x: torch.Tensor) -> torch.Tensor: - if ( - x.dim() == 5 - and hasattr(torch, "channels_last_3d") - and x.is_contiguous(memory_format=torch.channels_last_3d) - and not x.is_contiguous() - ): - return x.contiguous() - return x - - -def _halo_memory_format(reference: torch.Tensor) -> torch.memory_format: - if reference.dim() > 1 and reference.stride(1) == 1: - if reference.dim() == 5 and hasattr(torch, "channels_last_3d"): - return torch.channels_last_3d - if reference.dim() == 4: - return torch.channels_last - return torch.contiguous_format - - -def gather_and_trim_height(x: torch.Tensor, expected_height: int | None): - if expected_height is None: - return x - 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 - - -def _ensure_recv_buf( - recv_buf: torch.Tensor | None, reference: torch.Tensor -) -> torch.Tensor: - memory_format = _halo_memory_format(reference) - if ( - recv_buf is None - or recv_buf.shape != reference.shape - or recv_buf.dtype != reference.dtype - or recv_buf.device != reference.device - or not recv_buf.is_contiguous(memory_format=memory_format) - ): - return torch.empty( - reference.shape, - dtype=reference.dtype, - device=reference.device, - memory_format=memory_format, - ) - return recv_buf - - -def halo_exchange( - x: torch.Tensor, - height_halo_size: int = 1, - recv_top_buf: torch.Tensor | None = None, - recv_bottom_buf: torch.Tensor | None = None, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - if height_halo_size == 0: - return x, recv_top_buf, recv_bottom_buf - - 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:, :] - - recv_top_buf = _ensure_recv_buf(recv_top_buf, top_row_ref) - recv_bottom_buf = _ensure_recv_buf(recv_bottom_buf, bottom_row_ref) - - # use batched P2P operations - p2p_ops = [] - - if rank > 0: - # has previous neighbor, recv previous rank's data to recv_top_buf and send top_row to it. - prev_rank = group_ranks[rank - 1] - top_row = top_row_ref.contiguous(memory_format=_halo_memory_format(top_row_ref)) - p2p_ops.append(dist.P2POp(dist.irecv, recv_top_buf, prev_rank, group)) - p2p_ops.append(dist.P2POp(dist.isend, top_row, prev_rank, group)) - if rank < world_size - 1: - # has next neighbor, send bottom_row to next rank and recv next rank's data to recv_bottom_buf. - next_rank = group_ranks[rank + 1] - bottom_row = bottom_row_ref.contiguous( - memory_format=_halo_memory_format(bottom_row_ref) - ) - p2p_ops.append(dist.P2POp(dist.isend, bottom_row, next_rank, group)) - p2p_ops.append(dist.P2POp(dist.irecv, recv_bottom_buf, next_rank, group)) - - if rank == 0: - recv_top_buf.zero_() - if rank == world_size - 1: - recv_bottom_buf.zero_() - - if p2p_ops: - reqs = dist.batch_isend_irecv(p2p_ops) - for req in reqs: - req.wait() - - return ( - torch.concat([recv_top_buf, x, recv_bottom_buf], dim=-2), - recv_top_buf, - recv_bottom_buf, - ) - - -class WanDistConv2d(nn.Conv2d): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int], - stride: int | tuple[int, int, int] = 1, - padding: int | tuple[int, int, int] = 0, - height_padding: tuple[int, int] | None = None, - ): - super().__init__( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=kernel_size, - stride=stride, - padding=padding, - ) - - self.height_halo_size = (self.kernel_size[-2] - 1) // 2 - if height_padding is None: - height_padding = (self.padding[-2], self.padding[-2]) - self.height_pad_top, self.height_pad_bottom = height_padding - - self.padding: tuple[int, int] - if self.height_halo_size > 0: - self._padding = (0, 0, 0, 0) - else: - self._padding = ( - 0, - 0, - self.padding[0], - self.padding[0], - ) - - 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_decode_parallel_rank() - self.world_size = get_decode_parallel_world_size() - - def forward(self, x): - if any(self._padding): - x = F.pad(x, self._padding) - - x_padded, self._halo_recv_top_buf, self._halo_recv_bottom_buf = halo_exchange( - x, - height_halo_size=self.height_halo_size, - recv_top_buf=self._halo_recv_top_buf, - recv_bottom_buf=self._halo_recv_bottom_buf, - ) - - pad_top = self.height_pad_top - stride = self.stride[-2] - global_start = self.rank * x.shape[-2] - if self.height_halo_size > 0 and stride > 1: - shift = (global_start - self.height_halo_size + pad_top) % stride - if shift: - x_padded = x_padded[..., shift:, :] - global_start += shift - - out = super().forward(x_padded) - - if self.height_halo_size == 0: - return out - - local_height = x.shape[-2] - global_height = local_height * self.world_size - halo = self.height_halo_size - pad_bottom = self.height_pad_bottom - kernel = self.kernel_size[-2] - min_i = math.ceil(((-pad_top) - (global_start - halo)) / stride) - max_i = math.floor( - ((global_height - 1 + pad_bottom) - (kernel - 1) - (global_start - halo)) - / stride - ) - start = max(min_i, 0) - end = min(max_i + 1, out.shape[-2]) - if start != 0 or end != out.shape[-2]: - out = out[..., start:end, :] - - return out - - -class WanDistCausalConv3d(nn.Conv3d): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int | tuple[int, int, int], - stride: int | tuple[int, int, int] = 1, - padding: int | tuple[int, int, int] = 0, - ): - super().__init__( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=kernel_size, - stride=stride, - padding=padding, - ) - - self.height_pad_top = self.padding[1] - self.height_pad_bottom = self.padding[1] - self.height_halo_size = (self.kernel_size[-2] - 1) // 2 - - self.padding: tuple[int, int, int] - # Set up causal padding, let the halo to control height padding - if self.height_halo_size > 0: - self._padding: tuple[int, ...] = ( - self.padding[2], - self.padding[2], - 0, - 0, - 2 * self.padding[0], - 0, - ) - else: - self._padding: tuple[int, ...] = ( - self.padding[2], - self.padding[2], - self.padding[1], - self.padding[1], - 2 * self.padding[0], - 0, - ) - 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_decode_parallel_rank() - self.world_size = get_decode_parallel_world_size() - - def forward(self, x, cache_x=None): - padding = list(self._padding) - if cache_x is not None and self._padding[4] > 0: - cache_x = cache_x.to(x.device) - x = torch.cat([cache_x, x], dim=2) - padding[4] -= cache_x.shape[2] - - x = F.pad(x, padding) - - x = ( - x if current_platform.is_amp_supported() else x.to(self.weight.dtype) - ) # casting needed if amp isn't supported - - x_padded, self._halo_recv_top_buf, self._halo_recv_bottom_buf = halo_exchange( - x, - height_halo_size=self.height_halo_size, - recv_top_buf=self._halo_recv_top_buf, - recv_bottom_buf=self._halo_recv_bottom_buf, - ) - - pad_top = self.height_pad_top - stride = self.stride[-2] - global_start = self.rank * x.shape[-2] - if self.height_halo_size > 0 and stride > 1: - shift = (global_start - self.height_halo_size + pad_top) % stride - if shift: - x_padded = x_padded[..., shift:, :] - global_start += shift - - x_padded = match_conv3d_input_format(x_padded, self.weight) - out = super().forward(x_padded) - - if self.height_halo_size == 0: - return out - - local_height = x.shape[-2] - global_height = local_height * self.world_size - halo = self.height_halo_size - pad_bottom = self.height_pad_bottom - kernel = self.kernel_size[-2] - min_i = math.ceil(((-pad_top) - (global_start - halo)) / stride) - max_i = math.floor( - ((global_height - 1 + pad_bottom) - (kernel - 1) - (global_start - halo)) - / stride - ) - start = max(min_i, 0) - end = min(max_i + 1, out.shape[-2]) - if start != 0 or end != out.shape[-2]: - out = out[..., start:end, :] - - return out - - -class WanDistZeroPad2d(nn.Module): - """Apply 2D padding once globally across sequence-parallel height splits.""" - - def __init__(self, padding: tuple[int, int, int, int]) -> None: - super().__init__() - self.padding = padding # (left, right, top, bottom) - 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 - if self.world_size <= 1: - return F.pad(x, (left, right, top, bottom)) - # Only the first/last rank should contribute global top/bottom padding. - top = top if self.rank == 0 else 0 - bottom = bottom if self.rank == self.world_size - 1 else 0 - return F.pad(x, (left, right, top, bottom)) - - -class WanDistResample(nn.Module): - r""" - A custom resampling module for 2D and 3D data used for parallel decoding. - - Args: - dim (int): The number of input/output channels. - mode (str): The resampling mode. Must be one of: - - 'none': No resampling (identity operation). - - 'upsample2d': 2D upsampling with nearest-exact interpolation and convolution. - - 'upsample3d': 3D upsampling with nearest-exact interpolation, convolution, and causal 3D convolution. - - 'downsample2d': 2D downsampling with zero-padding and convolution. - - 'downsample3d': 3D downsampling with zero-padding, convolution, and causal 3D convolution. - """ - - def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None: - super().__init__() - self.dim = dim - self.mode = mode - - # default to dim //2 - if upsample_out_dim is None: - upsample_out_dim = dim // 2 - - # layers - # We support parallel encode/decode; downsample uses halo exchange as well. - if mode == "upsample2d": - self.resample = nn.Sequential( - WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - WanDistConv2d(dim, upsample_out_dim, 3, padding=1), - ) - elif mode == "upsample3d": - self.resample = nn.Sequential( - WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - WanDistConv2d(dim, upsample_out_dim, 3, padding=1), - ) - self.time_conv = WanCausalConv3d(dim, dim * 2, (3, 1, 1), padding=(1, 0, 0)) - - elif mode == "downsample2d": - self.resample = nn.Sequential( - WanDistZeroPad2d((0, 1, 0, 0)), - WanDistConv2d(dim, dim, 3, stride=(2, 2), height_padding=(0, 1)), - ) - elif mode == "downsample3d": - self.resample = nn.Sequential( - WanDistZeroPad2d((0, 1, 0, 0)), - WanDistConv2d(dim, dim, 3, stride=(2, 2), height_padding=(0, 1)), - ) - self.time_conv = WanCausalConv3d( - dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0) - ) - - else: - self.resample = nn.Identity() - - def forward(self, x): - return resample_forward(self, x) - - -class WanDistResidualBlock(nn.Module): - r""" - A custom residual block module. - - Args: - in_dim (int): Number of input channels. - out_dim (int): Number of output channels. - dropout (float, optional): Dropout rate for the dropout layer. Default is 0.0. - non_linearity (str, optional): Type of non-linearity to use. Default is "silu". - """ - - def __init__( - self, - in_dim: int, - out_dim: int, - dropout: float = 0.0, - non_linearity: str = "silu", - ) -> None: - super().__init__() - self.in_dim = in_dim - self.out_dim = out_dim - self.nonlinearity = get_act_fn(non_linearity) - - # layers - self.norm1 = WanRMS_norm(in_dim, images=False) - self.conv1 = WanDistCausalConv3d(in_dim, out_dim, 3, padding=1) - self.norm2 = WanRMS_norm(out_dim, images=False) - self.dropout = nn.Dropout(dropout) - self.conv2 = WanDistCausalConv3d(out_dim, out_dim, 3, padding=1) - self.conv_shortcut = ( - WanDistCausalConv3d(in_dim, out_dim, 1) - if in_dim != out_dim - else nn.Identity() - ) - - def forward(self, x): - return residual_block_forward(self, x) - - -class WanDistAttentionBlock(nn.Module): - r""" - Causal self-attention with a single head. - - Args: - dim (int): The number of channels in the input tensor. - """ - - def __init__(self, dim) -> None: - super().__init__() - self.dim = dim - - # layers - self.norm = WanRMS_norm(dim) - self.to_qkv = nn.Conv2d(dim, dim * 3, 1) - self.proj = nn.Conv2d(dim, dim, 1) - 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.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: - x = torch.chunk(x, self.world_size, dim=-2)[self.rank] - - return x - - -class WanDistMidBlock(nn.Module): - """ - Middle block for WanVAE encoder and decoder. - - Args: - dim (int): Number of input/output channels. - dropout (float): Dropout rate. - non_linearity (str): Type of non-linearity to use. - """ - - def __init__( - self, - dim: int, - dropout: float = 0.0, - non_linearity: str = "silu", - num_layers: int = 1, - ): - super().__init__() - self.dim = dim - - # Create the components - resnets = [WanDistResidualBlock(dim, dim, dropout, non_linearity)] - attentions = [] - for _ in range(num_layers): - attentions.append(WanDistAttentionBlock(dim)) - resnets.append(WanDistResidualBlock(dim, dim, dropout, non_linearity)) - self.attentions = nn.ModuleList(attentions) - self.resnets = nn.ModuleList(resnets) - - self.gradient_checkpointing = False - - def forward(self, x): - return mid_block_forward(self, x) - - -class WanDistResidualDownBlock(nn.Module): - def __init__( - self, - in_dim, - out_dim, - dropout, - num_res_blocks, - temperal_downsample=False, - down_flag=False, - ): - super().__init__() - - # Shortcut path with downsample - self.avg_shortcut = AvgDown3D( - in_dim, - out_dim, - factor_t=2 if temperal_downsample else 1, - factor_s=2 if down_flag else 1, - ) - - # Main path with residual blocks and downsample - resnets = [] - for _ in range(num_res_blocks): - resnets.append(WanDistResidualBlock(in_dim, out_dim, dropout)) - in_dim = out_dim - self.resnets = nn.ModuleList(resnets) - - # Add the final downsample block - if down_flag: - mode = "downsample3d" if temperal_downsample else "downsample2d" - self.downsampler = WanDistResample(out_dim, mode=mode) - else: - self.downsampler = None - - def forward(self, x): - return residual_down_block_forward(self, x) - - -class WanDistResidualUpBlock(nn.Module): - """ - A block that handles upsampling for the WanVAE decoder. - Args: - in_dim (int): Input dimension - out_dim (int): Output dimension - num_res_blocks (int): Number of residual blocks - dropout (float): Dropout rate - temperal_upsample (bool): Whether to upsample on temporal dimension - up_flag (bool): Whether to upsample or not - non_linearity (str): Type of non-linearity to use - """ - - def __init__( - self, - in_dim: int, - out_dim: int, - num_res_blocks: int, - dropout: float = 0.0, - temperal_upsample: bool = False, - up_flag: bool = False, - non_linearity: str = "silu", - ): - super().__init__() - self.in_dim = in_dim - self.out_dim = out_dim - - if up_flag: - self.avg_shortcut = DupUp3D( - in_dim, - out_dim, - factor_t=2 if temperal_upsample else 1, - factor_s=2, - ) - else: - self.avg_shortcut = None - - # create residual blocks - resnets = [] - current_dim = in_dim - for _ in range(num_res_blocks + 1): - resnets.append( - WanDistResidualBlock(current_dim, out_dim, dropout, non_linearity) - ) - current_dim = out_dim - - self.resnets = nn.ModuleList(resnets) - - # Add upsampling layer if needed - if up_flag: - upsample_mode = "upsample3d" if temperal_upsample else "upsample2d" - self.upsampler = WanDistResample( - out_dim, mode=upsample_mode, upsample_out_dim=out_dim - ) - else: - self.upsampler = None - - self.gradient_checkpointing = False - - def forward(self, x): - return residual_up_block_forward(self, x) - - -class WanDistUpBlock(nn.Module): - """ - A block that handles upsampling for the WanVAE decoder. - - Args: - in_dim (int): Input dimension - out_dim (int): Output dimension - num_res_blocks (int): Number of residual blocks - dropout (float): Dropout rate - upsample_mode (str, optional): Mode for upsampling ('upsample2d' or 'upsample3d') - non_linearity (str): Type of non-linearity to use - """ - - def __init__( - self, - in_dim: int, - out_dim: int, - num_res_blocks: int, - dropout: float = 0.0, - upsample_mode: str | None = None, - non_linearity: str = "silu", - ): - super().__init__() - self.in_dim = in_dim - self.out_dim = out_dim - - # Create layers list - resnets = [] - # Add residual blocks and attention if needed - current_dim = in_dim - for _ in range(num_res_blocks + 1): - resnets.append( - WanDistResidualBlock(current_dim, out_dim, dropout, non_linearity) - ) - current_dim = out_dim - - self.resnets = nn.ModuleList(resnets) - - # Add upsampling layer if needed - self.upsamplers = None - if upsample_mode is not None: - self.upsamplers = nn.ModuleList( - [WanDistResample(out_dim, mode=upsample_mode)] - ) - - self.gradient_checkpointing = False - - def forward(self, x): - return up_block_forward(self, x) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py index 2a96c33ad..0f3e2f424 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py @@ -17,14 +17,18 @@ # limitations under the License. import contextvars -from contextlib import contextmanager +from contextlib import contextmanager, nullcontext import torch import torch.distributed as dist import torch.nn as nn +import torch.nn.functional as F from einops import rearrange from sglang.multimodal_gen.configs.models.vaes import WanVAEConfig +from sglang.multimodal_gen.configs.models.vaes.base import ( + should_use_spatial_shard_parallel_decode, +) from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_decode_parallel_rank, get_decode_parallel_world_size, @@ -32,39 +36,22 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_sp_world_size, ) from sglang.multimodal_gen.runtime.layers.activation import get_act_fn +from sglang.multimodal_gen.runtime.layers.parallel_conv import ( + SpatialParallelCausalConv3d, + SpatialParallelConv2d, + SpatialParallelZeroPad2d, + chunk_height_for_parallel_decode, + disable_spatial_parallel_decode, + gather_and_trim_height, + gather_height_for_global_op, + split_for_parallel_decode, +) from sglang.multimodal_gen.runtime.models.vaes.common import ( DiagonalGaussianDistribution, ParallelTiledVAE, + should_run_spatial_shard_parallel_decode, ) -from sglang.multimodal_gen.runtime.models.vaes.parallel.wan_common_utils import ( - AvgDown3D, - DupUp3D, - WanCausalConv3d, - WanRMS_norm, - WanUpsample, - attention_block_forward, - bind_context, - mid_block_forward, - resample_forward, - residual_block_forward, - residual_down_block_forward, - residual_up_block_forward, - up_block_forward, -) -from sglang.multimodal_gen.runtime.models.vaes.parallel.wan_dist_utils import ( - WanDistAttentionBlock, - WanDistCausalConv3d, - WanDistMidBlock, - WanDistResample, - WanDistResidualBlock, - WanDistResidualDownBlock, - WanDistResidualUpBlock, - WanDistUpBlock, - ensure_local_height, - gather_and_trim_height, - split_for_parallel_decode, - split_for_parallel_encode, -) +from sglang.multimodal_gen.runtime.platforms import current_platform CACHE_T = 2 @@ -73,7 +60,464 @@ feat_cache = contextvars.ContextVar("feat_cache", default=None) feat_idx = contextvars.ContextVar("feat_idx", default=0) first_chunk = contextvars.ContextVar("first_chunk", default=None) -bind_context(is_first_frame, feat_cache, feat_idx, CACHE_T, first_chunk) + +def _channels_last_3d_supported_by_platform() -> bool: + return hasattr(torch, "channels_last_3d") and ( + current_platform.is_cuda() or current_platform.is_rocm() + ) + + +def _conv3d_weight_is_channels_last_3d(weight: torch.Tensor) -> bool: + return ( + weight.dim() == 5 + and _channels_last_3d_supported_by_platform() + and weight.is_contiguous(memory_format=torch.channels_last_3d) + ) + + +def match_conv3d_input_format(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + if x.dim() == 5 and _conv3d_weight_is_channels_last_3d(weight): + return x.contiguous(memory_format=torch.channels_last_3d) + return x + + +class AvgDown3D(nn.Module): + def __init__( + self, + in_channels, + out_channels, + factor_t, + factor_s=1, + ): + super().__init__() + self.in_channels = in_channels + self.out_channels = out_channels + self.factor_t = factor_t + self.factor_s = factor_s + self.factor = self.factor_t * self.factor_s * self.factor_s + + assert in_channels * self.factor % out_channels == 0 + self.group_size = in_channels * self.factor // out_channels + + def forward(self, x: torch.Tensor) -> torch.Tensor: + pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t + pad = (0, 0, 0, 0, pad_t, 0) + x = F.pad(x, pad) + B, C, T, H, W = x.shape + x = x.view( + B, + C, + T // self.factor_t, + self.factor_t, + H // self.factor_s, + self.factor_s, + W // self.factor_s, + self.factor_s, + ) + x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous() + x = x.view( + B, + C * self.factor, + T // self.factor_t, + H // self.factor_s, + W // self.factor_s, + ) + x = x.view( + B, + self.out_channels, + self.group_size, + T // self.factor_t, + H // self.factor_s, + W // self.factor_s, + ) + x = x.mean(dim=2) + return x + + +class DupUp3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + factor_t, + factor_s=1, + ): + super().__init__() + self.in_channels = in_channels + self.out_channels = out_channels + + self.factor_t = factor_t + self.factor_s = factor_s + self.factor = self.factor_t * self.factor_s * self.factor_s + + assert out_channels * self.factor % in_channels == 0 + self.repeats = out_channels * self.factor // in_channels + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = x.repeat_interleave(self.repeats, dim=1) + x = x.view( + x.size(0), + self.out_channels, + self.factor_t, + self.factor_s, + self.factor_s, + x.size(2), + x.size(3), + x.size(4), + ) + x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous() + x = x.view( + x.size(0), + self.out_channels, + x.size(2) * self.factor_t, + x.size(4) * self.factor_s, + x.size(6) * self.factor_s, + ) + + _first_chunk = first_chunk.get() if first_chunk is not None else None + if _first_chunk: + x = x[:, :, self.factor_t - 1 :, :, :] + return x + + +class WanCausalConv3d(nn.Conv3d): + r""" + A custom 3D causal convolution layer with feature caching support. + + This layer extends the standard Conv3D layer by ensuring causality in the time dimension and handling feature + caching for efficient inference. + """ + + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int | tuple[int, int, int], + stride: int | tuple[int, int, int] = 1, + padding: int | tuple[int, int, int] = 0, + ) -> None: + super().__init__( + in_channels=in_channels, + out_channels=out_channels, + kernel_size=kernel_size, + stride=stride, + padding=padding, + ) + self.padding: tuple[int, int, int] + # Set up causal padding + self._padding: tuple[int, ...] = ( + self.padding[2], + self.padding[2], + self.padding[1], + self.padding[1], + 2 * self.padding[0], + 0, + ) + self.padding = (0, 0, 0) + + def forward(self, x, cache_x=None): + padding = list(self._padding) + if cache_x is not None and self._padding[4] > 0: + cache_x = cache_x.to(x.device) + x = torch.cat([cache_x, x], dim=2) + padding[4] -= cache_x.shape[2] + x = F.pad(x, padding) + x = ( + x if current_platform.is_amp_supported() else x.to(self.weight.dtype) + ) # casting needed if amp isn't supported + x = match_conv3d_input_format(x, self.weight) + return super().forward(x) + + +class WanRMS_norm(nn.Module): + r""" + A custom RMS normalization layer. + """ + + def __init__( + self, + dim: int, + channel_first: bool = True, + images: bool = True, + bias: bool = False, + ) -> None: + super().__init__() + broadcastable_dims = (1, 1, 1) if not images else (1, 1) + shape = (dim, *broadcastable_dims) if channel_first else (dim,) + + self.channel_first = channel_first + self.scale = dim**0.5 + self.gamma = nn.Parameter(torch.ones(shape)) + self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0 + + def forward(self, x): + return ( + F.normalize(x, dim=(1 if self.channel_first else -1)) + * self.scale + * self.gamma + + self.bias + ) + + +class WanUpsample(nn.Upsample): + r""" + Perform upsampling while ensuring the output tensor has the same data type as the input. + """ + + def forward(self, x): + if current_platform.is_amp_supported(): + return super().forward(x) + return super().forward(x.float()).type_as(x) + + +def resample_forward(self, x): + b, c, t, h, w = x.size() + first_frame = is_first_frame.get() + if first_frame: + assert t == 1 + _feat_cache = feat_cache.get() + _feat_idx = feat_idx.get() + if self.mode == "upsample3d": + if _feat_cache is not None: + idx = _feat_idx + if _feat_cache[idx] is None: + _feat_cache[idx] = "Rep" + _feat_idx += 1 + else: + cache_x = x[:, :, -CACHE_T:, :, :].clone() + if ( + cache_x.shape[2] < 2 + and _feat_cache[idx] is not None + and _feat_cache[idx] != "Rep" + ): + # cache last frame of last two chunk + cache_x = torch.cat( + [ + _feat_cache[idx][:, :, -1, :, :] + .unsqueeze(2) + .to(cache_x.device), + cache_x, + ], + dim=2, + ) + if ( + cache_x.shape[2] < 2 + and _feat_cache[idx] is not None + and _feat_cache[idx] == "Rep" + ): + cache_x = torch.cat( + [torch.zeros_like(cache_x).to(cache_x.device), cache_x], + dim=2, + ) + if _feat_cache[idx] == "Rep": + x = self.time_conv(x) + else: + x = self.time_conv(x, _feat_cache[idx]) + _feat_cache[idx] = cache_x + _feat_idx += 1 + + x = x.reshape(b, 2, c, t, h, w) + x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3) + x = x.reshape(b, c, t * 2, h, w) + feat_cache.set(_feat_cache) + feat_idx.set(_feat_idx) + elif not first_frame and hasattr(self, "time_conv"): + x = self.time_conv(x) + x = x.reshape(b, 2, c, t, h, w) + x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3) + x = x.reshape(b, c, t * 2, h, w) + t = x.shape[2] + x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) + x = self.resample(x) + x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4) + + _feat_cache = feat_cache.get() + _feat_idx = feat_idx.get() + if self.mode == "downsample3d": + if _feat_cache is not None: + idx = _feat_idx + if _feat_cache[idx] is None: + _feat_cache[idx] = x.clone() + _feat_idx += 1 + else: + cache_x = x[:, :, -1:, :, :].clone() + x = self.time_conv(torch.cat([_feat_cache[idx][:, :, -1:, :, :], x], 2)) + _feat_cache[idx] = cache_x + _feat_idx += 1 + feat_cache.set(_feat_cache) + feat_idx.set(_feat_idx) + elif not first_frame and hasattr(self, "time_conv"): + x = self.time_conv(x) + return x + + +def residual_block_forward(self, x): + # Apply shortcut connection + h = self.conv_shortcut(x) + + # First normalization and activation + x = self.norm1(x) + x = self.nonlinearity(x) + + _feat_cache = feat_cache.get() + _feat_idx = feat_idx.get() + if _feat_cache is not None: + idx = _feat_idx + cache_x = x[:, :, -CACHE_T:, :, :].clone() + if cache_x.shape[2] < 2 and _feat_cache[idx] is not None: + cache_x = torch.cat( + [ + _feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), + cache_x, + ], + dim=2, + ) + + x = self.conv1(x, _feat_cache[idx]) + _feat_cache[idx] = cache_x + _feat_idx += 1 + feat_cache.set(_feat_cache) + feat_idx.set(_feat_idx) + else: + x = self.conv1(x) + + # Second normalization and activation + x = self.norm2(x) + x = self.nonlinearity(x) + + # Dropout + x = self.dropout(x) + + _feat_cache = feat_cache.get() + _feat_idx = feat_idx.get() + if _feat_cache is not None: + idx = _feat_idx + cache_x = x[:, :, -CACHE_T:, :, :].clone() + if cache_x.shape[2] < 2 and _feat_cache[idx] is not None: + cache_x = torch.cat( + [ + _feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), + cache_x, + ], + dim=2, + ) + + x = self.conv2(x, _feat_cache[idx]) + _feat_cache[idx] = cache_x + _feat_idx += 1 + feat_cache.set(_feat_cache) + feat_idx.set(_feat_idx) + else: + x = self.conv2(x) + + # Add residual connection + return x + h + + +def attention_block_forward(self, x): + identity = x + batch_size, channels, num_frames, height, width = x.size() + x = x.permute(0, 2, 1, 3, 4).reshape( + batch_size * num_frames, channels, height, width + ) + x = self.norm(x) + + # compute query, key, value + qkv = self.to_qkv(x) + qkv = qkv.reshape(batch_size * num_frames, 1, channels * 3, -1) + qkv = qkv.permute(0, 1, 3, 2).contiguous() + q, k, v = qkv.chunk(3, dim=-1) + + x = torch.nn.functional.scaled_dot_product_attention(q, k, v) + + x = ( + x.squeeze(1) + .permute(0, 2, 1) + .reshape(batch_size * num_frames, channels, height, width) + ) + + # output projection + x = self.proj(x) + + # Reshape back: [(b*t), c, h, w] -> [b, c, t, h, w] + x = x.view(batch_size, num_frames, channels, height, width) + x = x.permute(0, 2, 1, 3, 4) + + return x + identity + + +def mid_block_forward(self, x): + # First residual block + x = self.resnets[0](x) + + # Process through attention and residual blocks + for attn, resnet in zip(self.attentions, self.resnets[1:], strict=True): + if attn is not None: + x = attn(x) + + x = resnet(x) + + return x + + +def residual_down_block_forward(self, x): + x_copy = x + for resnet in self.resnets: + x = resnet(x) + if self.downsampler is not None: + x = self.downsampler(x) + + return x + self.avg_shortcut(x_copy) + + +def residual_up_block_forward(self, x): + if self.avg_shortcut is not None: + x_copy = x + + for resnet in self.resnets: + x = resnet(x) + + if self.upsampler is not None: + x = self.upsampler(x) + + if self.avg_shortcut is not None: + x = x + self.avg_shortcut(x_copy) + + return x + + +def up_block_forward(self, x): + for resnet in self.resnets: + x = resnet(x) + + if self.upsamplers is not None: + x = self.upsamplers[0](x) + return x + + +def split_for_parallel_encode( + x: torch.Tensor, downsample_count: int, world_size: int, rank: int +): + orig_height = x.shape[-2] + expected_height = orig_height // (2**downsample_count) + factor = world_size * (2**downsample_count) + pad_h = (factor - orig_height % factor) % factor + if pad_h: + x = F.pad(x, (0, 0, 0, pad_h, 0, 0)) + expected_local_height = (orig_height + pad_h) // (2**downsample_count) // world_size + x = torch.chunk(x, world_size, dim=-2)[rank] + return x, expected_height, expected_local_height + + +def ensure_local_height(x: torch.Tensor, expected_local_height: int | None): + if expected_local_height is None: + return x + if x.shape[-2] < expected_local_height: + pad = expected_local_height - x.shape[-2] + return F.pad(x, (0, 0, 0, pad, 0, 0)) + if x.shape[-2] > expected_local_height: + return x[..., :expected_local_height, :].contiguous() + return x @contextmanager @@ -107,7 +551,16 @@ class WanResample(nn.Module): - 'downsample3d': 3D downsampling with zero-padding, convolution, and causal 3D convolution. """ - def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None: + def __init__( + self, + dim: int, + mode: str, + upsample_out_dim: int = None, + *, + conv2d_cls=nn.Conv2d, + zero_pad2d_cls=nn.ZeroPad2d, + spatial_parallel: bool = False, + ) -> None: super().__init__() self.dim = dim self.mode = mode @@ -120,23 +573,37 @@ class WanResample(nn.Module): if mode == "upsample2d": self.resample = nn.Sequential( WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - nn.Conv2d(dim, upsample_out_dim, 3, padding=1), + conv2d_cls(dim, upsample_out_dim, 3, padding=1), ) elif mode == "upsample3d": self.resample = nn.Sequential( WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), - nn.Conv2d(dim, upsample_out_dim, 3, padding=1), + conv2d_cls(dim, upsample_out_dim, 3, padding=1), ) self.time_conv = WanCausalConv3d(dim, dim * 2, (3, 1, 1), padding=(1, 0, 0)) elif mode == "downsample2d": - self.resample = nn.Sequential( - nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)) - ) + if spatial_parallel: + self.resample = nn.Sequential( + zero_pad2d_cls((0, 1, 0, 0)), + conv2d_cls(dim, dim, 3, stride=(2, 2), height_padding=(0, 1)), + ) + else: + self.resample = nn.Sequential( + zero_pad2d_cls((0, 1, 0, 1)), + conv2d_cls(dim, dim, 3, stride=(2, 2)), + ) elif mode == "downsample3d": - self.resample = nn.Sequential( - nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)) - ) + if spatial_parallel: + self.resample = nn.Sequential( + zero_pad2d_cls((0, 1, 0, 0)), + conv2d_cls(dim, dim, 3, stride=(2, 2), height_padding=(0, 1)), + ) + else: + self.resample = nn.Sequential( + zero_pad2d_cls((0, 1, 0, 1)), + conv2d_cls(dim, dim, 3, stride=(2, 2)), + ) self.time_conv = WanCausalConv3d( dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0) ) @@ -165,6 +632,9 @@ class WanResidualBlock(nn.Module): out_dim: int, dropout: float = 0.0, non_linearity: str = "silu", + *, + causal_conv3d_cls=WanCausalConv3d, + shortcut_conv3d_cls=WanCausalConv3d, ) -> None: super().__init__() self.in_dim = in_dim @@ -173,12 +643,14 @@ class WanResidualBlock(nn.Module): # layers self.norm1 = WanRMS_norm(in_dim, images=False) - self.conv1 = WanCausalConv3d(in_dim, out_dim, 3, padding=1) + self.conv1 = causal_conv3d_cls(in_dim, out_dim, 3, padding=1) self.norm2 = WanRMS_norm(out_dim, images=False) self.dropout = nn.Dropout(dropout) - self.conv2 = WanCausalConv3d(out_dim, out_dim, 3, padding=1) + self.conv2 = causal_conv3d_cls(out_dim, out_dim, 3, padding=1) self.conv_shortcut = ( - WanCausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity() + shortcut_conv3d_cls(in_dim, out_dim, 1) + if in_dim != out_dim + else nn.Identity() ) def forward(self, x): @@ -193,9 +665,10 @@ class WanAttentionBlock(nn.Module): dim (int): The number of channels in the input tensor. """ - def __init__(self, dim) -> None: + def __init__(self, dim, *, spatial_parallel: bool = False) -> None: super().__init__() self.dim = dim + self.world_size = get_decode_parallel_world_size() if spatial_parallel else 1 # layers self.norm = WanRMS_norm(dim) @@ -203,7 +676,12 @@ class WanAttentionBlock(nn.Module): self.proj = nn.Conv2d(dim, dim, 1) def forward(self, x): - return attention_block_forward(self, x) + if self.world_size > 1: + x = gather_height_for_global_op(x).contiguous() + x = attention_block_forward(self, x) + if self.world_size > 1: + x = chunk_height_for_parallel_decode(x) + return x class WanMidBlock(nn.Module): @@ -222,16 +700,19 @@ class WanMidBlock(nn.Module): dropout: float = 0.0, non_linearity: str = "silu", num_layers: int = 1, + *, + residual_block_cls=WanResidualBlock, + attention_block_cls=WanAttentionBlock, ): super().__init__() self.dim = dim # Create the components - resnets = [WanResidualBlock(dim, dim, dropout, non_linearity)] + resnets = [residual_block_cls(dim, dim, dropout, non_linearity)] attentions = [] for _ in range(num_layers): - attentions.append(WanAttentionBlock(dim)) - resnets.append(WanResidualBlock(dim, dim, dropout, non_linearity)) + attentions.append(attention_block_cls(dim)) + resnets.append(residual_block_cls(dim, dim, dropout, non_linearity)) self.attentions = nn.ModuleList(attentions) self.resnets = nn.ModuleList(resnets) @@ -251,6 +732,9 @@ class WanResidualDownBlock(nn.Module): num_res_blocks, temperal_downsample=False, down_flag=False, + *, + residual_block_cls=WanResidualBlock, + resample_cls=WanResample, ): super().__init__() @@ -265,14 +749,14 @@ class WanResidualDownBlock(nn.Module): # Main path with residual blocks and downsample resnets = [] for _ in range(num_res_blocks): - resnets.append(WanResidualBlock(in_dim, out_dim, dropout)) + resnets.append(residual_block_cls(in_dim, out_dim, dropout)) in_dim = out_dim self.resnets = nn.ModuleList(resnets) # Add the final downsample block if down_flag: mode = "downsample3d" if temperal_downsample else "downsample2d" - self.downsampler = WanResample(out_dim, mode=mode) + self.downsampler = resample_cls(out_dim, mode=mode) else: self.downsampler = None @@ -280,6 +764,80 @@ class WanResidualDownBlock(nn.Module): return residual_down_block_forward(self, x) +class WanDistResample(WanResample): + def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None: + super().__init__( + dim, + mode, + upsample_out_dim=upsample_out_dim, + conv2d_cls=SpatialParallelConv2d, + zero_pad2d_cls=SpatialParallelZeroPad2d, + spatial_parallel=True, + ) + + +class WanDistResidualBlock(WanResidualBlock): + def __init__( + self, + in_dim: int, + out_dim: int, + dropout: float = 0.0, + non_linearity: str = "silu", + ) -> None: + super().__init__( + in_dim, + out_dim, + dropout, + non_linearity, + causal_conv3d_cls=SpatialParallelCausalConv3d, + ) + + +class WanDistAttentionBlock(WanAttentionBlock): + def __init__(self, dim) -> None: + super().__init__(dim, spatial_parallel=True) + + +class WanDistMidBlock(WanMidBlock): + def __init__( + self, + dim: int, + dropout: float = 0.0, + non_linearity: str = "silu", + num_layers: int = 1, + ): + super().__init__( + dim, + dropout, + non_linearity, + num_layers=num_layers, + residual_block_cls=WanDistResidualBlock, + attention_block_cls=WanDistAttentionBlock, + ) + + +class WanDistResidualDownBlock(WanResidualDownBlock): + def __init__( + self, + in_dim, + out_dim, + dropout, + num_res_blocks, + temperal_downsample=False, + down_flag=False, + ): + super().__init__( + in_dim, + out_dim, + dropout, + num_res_blocks, + temperal_downsample=temperal_downsample, + down_flag=down_flag, + residual_block_cls=WanDistResidualBlock, + resample_cls=WanDistResample, + ) + + class WanEncoder3d(nn.Module): r""" A 3D encoder module. @@ -330,7 +888,7 @@ class WanEncoder3d(nn.Module): world_size = get_sp_world_size() if use_parallel_encode and world_size > 1: - CausalConv3d = WanDistCausalConv3d + CausalConv3d = SpatialParallelCausalConv3d ResidualDownBlock = WanDistResidualDownBlock ResidualBlock = WanDistResidualBlock AttentionBlock = WanDistAttentionBlock @@ -488,6 +1046,9 @@ class WanResidualUpBlock(nn.Module): temperal_upsample: bool = False, up_flag: bool = False, non_linearity: str = "silu", + *, + residual_block_cls=WanResidualBlock, + resample_cls=WanResample, ): super().__init__() self.in_dim = in_dim @@ -508,7 +1069,7 @@ class WanResidualUpBlock(nn.Module): current_dim = in_dim for _ in range(num_res_blocks + 1): resnets.append( - WanResidualBlock(current_dim, out_dim, dropout, non_linearity) + residual_block_cls(current_dim, out_dim, dropout, non_linearity) ) current_dim = out_dim @@ -517,7 +1078,7 @@ class WanResidualUpBlock(nn.Module): # Add upsampling layer if needed if up_flag: upsample_mode = "upsample3d" if temperal_upsample else "upsample2d" - self.upsampler = WanResample( + self.upsampler = resample_cls( out_dim, mode=upsample_mode, upsample_out_dim=out_dim ) else: @@ -550,6 +1111,9 @@ class WanUpBlock(nn.Module): dropout: float = 0.0, upsample_mode: str | None = None, non_linearity: str = "silu", + *, + residual_block_cls=WanResidualBlock, + resample_cls=WanResample, ): super().__init__() self.in_dim = in_dim @@ -561,7 +1125,7 @@ class WanUpBlock(nn.Module): current_dim = in_dim for _ in range(num_res_blocks + 1): resnets.append( - WanResidualBlock(current_dim, out_dim, dropout, non_linearity) + residual_block_cls(current_dim, out_dim, dropout, non_linearity) ) current_dim = out_dim @@ -570,7 +1134,7 @@ class WanUpBlock(nn.Module): # Add upsampling layer if needed self.upsamplers = None if upsample_mode is not None: - self.upsamplers = nn.ModuleList([WanResample(out_dim, mode=upsample_mode)]) + self.upsamplers = nn.ModuleList([resample_cls(out_dim, mode=upsample_mode)]) self.gradient_checkpointing = False @@ -578,6 +1142,52 @@ class WanUpBlock(nn.Module): return up_block_forward(self, x) +class WanDistResidualUpBlock(WanResidualUpBlock): + def __init__( + self, + in_dim: int, + out_dim: int, + num_res_blocks: int, + dropout: float = 0.0, + temperal_upsample: bool = False, + up_flag: bool = False, + non_linearity: str = "silu", + ): + super().__init__( + in_dim, + out_dim, + num_res_blocks, + dropout=dropout, + temperal_upsample=temperal_upsample, + up_flag=up_flag, + non_linearity=non_linearity, + residual_block_cls=WanDistResidualBlock, + resample_cls=WanDistResample, + ) + + +class WanDistUpBlock(WanUpBlock): + def __init__( + self, + in_dim: int, + out_dim: int, + num_res_blocks: int, + dropout: float = 0.0, + upsample_mode: str | None = None, + non_linearity: str = "silu", + ): + super().__init__( + in_dim, + out_dim, + num_res_blocks, + dropout=dropout, + upsample_mode=upsample_mode, + non_linearity=non_linearity, + residual_block_cls=WanDistResidualBlock, + resample_cls=WanDistResample, + ) + + class WanDecoder3d(nn.Module): r""" A 3D decoder module. @@ -628,7 +1238,7 @@ class WanDecoder3d(nn.Module): world_size = get_decode_parallel_world_size() if use_parallel_decode and world_size > 1: - CausalConv3d = WanDistCausalConv3d + CausalConv3d = SpatialParallelCausalConv3d MidBlock = WanDistMidBlock ResidualUpBlock = WanDistResidualUpBlock UpBlock = WanDistUpBlock @@ -861,18 +1471,21 @@ class AutoencoderKLWan(ParallelTiledVAE): dropout=config.dropout, out_channels=config.out_channels, is_residual=config.is_residual, - use_parallel_decode=self.use_parallel_decode, + use_parallel_decode=should_use_spatial_shard_parallel_decode(config), ) self.use_feature_cache = config.use_feature_cache self._causal_decode_initialized = False + def _should_use_spatial_parallel_decode(self, z: torch.Tensor) -> bool: + return should_run_spatial_shard_parallel_decode(self.config, z) + def clear_cache(self) -> None: def _count_conv3d(model) -> int: count = 0 for m in model.modules(): - if isinstance(m, WanCausalConv3d) or isinstance(m, WanDistCausalConv3d): + if isinstance(m, (WanCausalConv3d, SpatialParallelCausalConv3d)): count += 1 return count @@ -904,13 +1517,19 @@ class AutoencoderKLWan(ParallelTiledVAE): iter_ = z.shape[2] x = self.post_quant_conv(z) outs = [] - with forward_context( - feat_cache_arg=self._feat_map, feat_idx_arg=self._conv_idx - ): - for i in range(iter_): - feat_idx.set(0) - first_chunk.set(is_first_chunk and i == 0) - outs.append(self.decoder(x[:, :, i : i + 1, :, :])) + spatial_context = ( + nullcontext() + if self._should_use_spatial_parallel_decode(z) + else disable_spatial_parallel_decode() + ) + with spatial_context: + with forward_context( + feat_cache_arg=self._feat_map, feat_idx_arg=self._conv_idx + ): + for i in range(iter_): + feat_idx.set(0) + first_chunk.set(is_first_chunk and i == 0) + outs.append(self.decoder(x[:, :, i : i + 1, :, :])) out = torch.cat(outs, 2) if self.config.patch_size is not None: @@ -984,16 +1603,25 @@ class AutoencoderKLWan(ParallelTiledVAE): self.clear_cache() iter_ = z.shape[2] x = self.post_quant_conv(z) - outs = [] - with forward_context( - feat_cache_arg=self._feat_map, feat_idx_arg=self._conv_idx - ): - out_chunks = [] - for i in range(iter_): - feat_idx.set(0) - first_chunk.set(i == 0) - out_chunks.append(self.decoder(x[:, :, i : i + 1, :, :])) - out = torch.cat(out_chunks, 2) if len(out_chunks) > 1 else out_chunks[0] + spatial_context = ( + nullcontext() + if self._should_use_spatial_parallel_decode(z) + else disable_spatial_parallel_decode() + ) + with spatial_context: + with forward_context( + feat_cache_arg=self._feat_map, feat_idx_arg=self._conv_idx + ): + out_chunks = [] + for i in range(iter_): + feat_idx.set(0) + first_chunk.set(i == 0) + out_chunks.append(self.decoder(x[:, :, i : i + 1, :, :])) + out = ( + torch.cat(out_chunks, 2) + if len(out_chunks) > 1 + else out_chunks[0] + ) if self.config.patch_size is not None: out = unpatchify(out, patch_size=self.config.patch_size) @@ -1008,8 +1636,14 @@ class AutoencoderKLWan(ParallelTiledVAE): def _decode(self, z: torch.Tensor, first_frame=False) -> torch.Tensor: x = self.post_quant_conv(z) - with forward_context(first_frame_arg=first_frame): - out = self.decoder(x) + spatial_context = ( + nullcontext() + if self._should_use_spatial_parallel_decode(z) + else disable_spatial_parallel_decode() + ) + with spatial_context: + with forward_context(first_frame_arg=first_frame): + out = self.decoder(x) out = torch.clamp(out, min=-1.0, max=1.0) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising_av.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising_av.py index dcae54d27..57e09ddef 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising_av.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising_av.py @@ -2,6 +2,10 @@ import torch from diffusers.utils.torch_utils import randn_tensor from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import is_ltx23_native_variant +from sglang.multimodal_gen.runtime.distributed import ( + get_decode_parallel_world_size, + model_parallel_is_initialized, +) from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( ComponentUse, ) @@ -146,8 +150,15 @@ class LTX2RefinementStage(LTX2AVDenoisingStage): @property def parallelism_type(self) -> StageParallelismType: # Stage 2 is distilled and always runs with CFG disabled, so non-main - # CFG ranks should wait at a barrier rather than run a redundant forward. + # CFG ranks only need the result when the decoder will use all ranks. if self.server_args.enable_cfg_parallel: + if ( + model_parallel_is_initialized() + and get_decode_parallel_world_size() > 1 + and self.vae is not None + and self.vae.use_parallel_decode + ): + return StageParallelismType.MAIN_RANK_ONLY_AND_SEND_TO_OTHERS return StageParallelismType.MAIN_RANK_ONLY return StageParallelismType.REPLICATED diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index 9bec42542..5db33bc58 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -34,7 +34,7 @@ if TYPE_CHECKING: logger = init_logger(__name__) -SGL_TEST_FILES_CI_DATA_REVISION = "caa56302ccf2d289e4488ed06d952edf5d2314cf" +SGL_TEST_FILES_CI_DATA_REVISION = "f0fd96eab85baed5256d142c7659e0634e7e4410" if current_platform.is_npu(): SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7" 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 index 9b308545b..8d679d6e8 100644 --- a/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py +++ b/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py @@ -6,6 +6,9 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( StageParallelismType, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import DecodingStage +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2.denoising_av import ( + LTX2RefinementStage, +) class TestDecodingStageParallelism(unittest.TestCase): @@ -95,6 +98,46 @@ class TestDecodingStageParallelism(unittest.TestCase): StageParallelismType.REPLICATED, ) + def test_ltx2_refinement_broadcasts_when_decode_uses_all_ranks(self): + stage = object.__new__(LTX2RefinementStage) + stage.server_args = SimpleNamespace(enable_cfg_parallel=True) + stage.vae = SimpleNamespace(use_parallel_decode=True) + + with ( + patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2.denoising_av.model_parallel_is_initialized", + return_value=True, + ), + patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2.denoising_av.get_decode_parallel_world_size", + return_value=2, + ), + ): + self.assertEqual( + stage.parallelism_type, + StageParallelismType.MAIN_RANK_ONLY_AND_SEND_TO_OTHERS, + ) + + def test_ltx2_refinement_stays_main_rank_only_without_parallel_decode(self): + stage = object.__new__(LTX2RefinementStage) + stage.server_args = SimpleNamespace(enable_cfg_parallel=True) + stage.vae = SimpleNamespace(use_parallel_decode=False) + + with ( + patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2.denoising_av.model_parallel_is_initialized", + return_value=True, + ), + patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2.denoising_av.get_decode_parallel_world_size", + return_value=2, + ), + ): + self.assertEqual( + stage.parallelism_type, + StageParallelismType.MAIN_RANK_ONLY, + ) + if __name__ == "__main__": unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py index 71dda79e5..4aa7ee1e3 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py @@ -17,7 +17,7 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import ( _backfill_ltx2_audio_vae_latent_stats, _should_use_channels_last_3d, ) -from sglang.multimodal_gen.runtime.models.vaes.parallel import wan_common_utils +from sglang.multimodal_gen.runtime.models.vaes import wanvae class _FakeServerArgs: @@ -176,14 +176,10 @@ class TestVAELoader(unittest.TestCase): ) with ( - patch.object( - wan_common_utils.current_platform, "is_cuda", return_value=False - ), - patch.object( - wan_common_utils.current_platform, "is_rocm", return_value=False - ), + patch.object(wanvae.current_platform, "is_cuda", return_value=False), + patch.object(wanvae.current_platform, "is_rocm", return_value=False), ): - out = wan_common_utils.match_conv3d_input_format(x, weight) + out = wanvae.match_conv3d_input_format(x, weight) self.assertIs(out, x) @@ -197,14 +193,10 @@ class TestVAELoader(unittest.TestCase): ) with ( - patch.object( - wan_common_utils.current_platform, "is_cuda", return_value=True - ), - patch.object( - wan_common_utils.current_platform, "is_rocm", return_value=False - ), + patch.object(wanvae.current_platform, "is_cuda", return_value=True), + patch.object(wanvae.current_platform, "is_rocm", return_value=False), ): - out = wan_common_utils.match_conv3d_input_format(x, weight) + out = wanvae.match_conv3d_input_format(x, weight) self.assertTrue(out.is_contiguous(memory_format=torch.channels_last_3d)) diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py b/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py new file mode 100644 index 000000000..bac9ff204 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py @@ -0,0 +1,512 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.multimodal_gen.configs.models.vaes.base import ( + VAEConfig, + is_spatial_shard_parallel_decode_mode, + should_use_spatial_shard_parallel_decode, +) +from sglang.multimodal_gen.configs.models.vaes.ernie_image import ErnieImageVAEConfig +from sglang.multimodal_gen.configs.models.vaes.flux import Flux2VAEConfig, FluxVAEConfig +from sglang.multimodal_gen.configs.models.vaes.glmimage import GlmImageVAEConfig +from sglang.multimodal_gen.configs.models.vaes.hunyuan3d import Hunyuan3DVAEConfig +from sglang.multimodal_gen.configs.models.vaes.hunyuanvae import HunyuanVAEConfig +from sglang.multimodal_gen.configs.models.vaes.ltx_audio import LTXAudioVAEConfig +from sglang.multimodal_gen.configs.models.vaes.ltx_video import LTXVideoVAEConfig +from sglang.multimodal_gen.configs.models.vaes.qwenimage import QwenImageVAEConfig +from sglang.multimodal_gen.configs.models.vaes.sana import SanaVAEConfig +from sglang.multimodal_gen.configs.models.vaes.stablediffusion3 import ( + StableDiffusion3VAEConfig, +) +from sglang.multimodal_gen.configs.models.vaes.wanvae import WanVAEConfig +from sglang.multimodal_gen.configs.utils import update_config_from_args +from sglang.multimodal_gen.runtime.distributed import parallel_state +from sglang.multimodal_gen.runtime.layers.parallel_conv import ( + SpatialParallelCausalConv3d, + SpatialParallelConv2d, + SpatialParallelConv3d, + chunk_height_by_sizes, + split_for_parallel_decode, +) +from sglang.multimodal_gen.runtime.models.vaes.autoencoder import AutoencoderKL +from sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage import ( + QwenImageAttentionBlock, + QwenImageDecoder3d, +) +from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE +from sglang.multimodal_gen.runtime.models.vaes.hunyuanvae import ( + HunyuanVideoDecoder3D, + HunyuanVideoMidBlock3D, + HunyuanVideoResnetBlockCausal3D, + _enable_hunyuan_decoder_spatial_parallel, +) +from sglang.multimodal_gen.runtime.models.vaes.ltx_2_vae import ( + LTX2VideoCausalConv3d, + LTX2VideoDecoder3d, + _enable_ltx_decoder_spatial_parallel, +) +from sglang.multimodal_gen.runtime.models.vaes.wanvae import ( + WanDecoder3d, + WanDistAttentionBlock, +) +from sglang.multimodal_gen.runtime.utils.distributed import RankGenerator +from sglang.multimodal_gen.utils import FlexibleArgumentParser + + +class _DispatchProbeVAE(ParallelTiledVAE): + def __init__(self, config): + super().__init__(config) + self.used_decode = False + self.used_parallel_tiled_decode = False + + def _encode(self, x): + return x + + def _decode(self, z): + self.used_decode = True + return z + + def parallel_tiled_decode(self, z): + self.used_parallel_tiled_decode = True + raise AssertionError( + "spatial_shard decode should not use parallel_tiled_decode" + ) + + +class TestVAESpatialParallelDecode(unittest.TestCase): + def test_base_vae_config_defaults_to_auto_parallel_decode(self): + config = VAEConfig() + + self.assertTrue(config.use_parallel_decode) + self.assertEqual(config.parallel_decode_mode, "auto") + + def test_image_video_vae_configs_default_to_auto_parallel_decode(self): + configs = ( + ErnieImageVAEConfig(), + FluxVAEConfig(), + Flux2VAEConfig(), + GlmImageVAEConfig(), + HunyuanVAEConfig(), + LTXVideoVAEConfig(), + QwenImageVAEConfig(), + SanaVAEConfig(), + StableDiffusion3VAEConfig(), + WanVAEConfig(), + ) + + for config in configs: + with self.subTest(config=type(config).__name__): + self.assertTrue(config.use_parallel_decode) + self.assertEqual(config.parallel_decode_mode, "auto") + + def test_auto_parallel_decode_policy_is_conservative(self): + self.assertFalse(is_spatial_shard_parallel_decode_mode("auto")) + self.assertFalse(should_use_spatial_shard_parallel_decode(VAEConfig())) + self.assertFalse(should_use_spatial_shard_parallel_decode(QwenImageVAEConfig())) + self.assertTrue(should_use_spatial_shard_parallel_decode(LTXVideoVAEConfig())) + self.assertTrue(should_use_spatial_shard_parallel_decode(WanVAEConfig())) + + ltx23_config = LTXVideoVAEConfig() + ltx23_config.arch_config.video_decoder_variant = "ltx_2_3" + self.assertTrue(should_use_spatial_shard_parallel_decode(ltx23_config)) + ltx23_config.parallel_decode_mode = "spatial_shard" + self.assertTrue(should_use_spatial_shard_parallel_decode(ltx23_config)) + + config = QwenImageVAEConfig() + self.assertFalse( + should_use_spatial_shard_parallel_decode( + config, torch.empty(1, 16, 1, 128, 128), 2 + ) + ) + config.parallel_decode_mode = "spatial_shard" + self.assertTrue(should_use_spatial_shard_parallel_decode(config)) + + hunyuan_config = HunyuanVAEConfig() + self.assertTrue(should_use_spatial_shard_parallel_decode(hunyuan_config)) + self.assertFalse( + should_use_spatial_shard_parallel_decode( + hunyuan_config, torch.empty(1, 16, 9, 16, 16), 2 + ) + ) + self.assertTrue( + should_use_spatial_shard_parallel_decode( + hunyuan_config, torch.empty(1, 16, 9, 96, 96), 2 + ) + ) + + def test_unsupported_vae_configs_opt_out_of_spatial_parallel_decode(self): + configs = (Hunyuan3DVAEConfig(), LTXAudioVAEConfig()) + + for config in configs: + with self.subTest(config=type(config).__name__): + self.assertFalse(config.use_parallel_decode) + self.assertEqual(config.parallel_decode_mode, "tiled") + + def test_vae_nested_cli_defaults_do_not_override_model_defaults(self): + parser = FlexibleArgumentParser() + VAEConfig.add_cli_args(parser) + parsed = vars(parser.parse_args([])) + config = FluxVAEConfig() + + update_config_from_args(config, parsed, "vae_config") + + self.assertTrue(config.use_parallel_decode) + self.assertEqual(config.parallel_decode_mode, "auto") + + def test_vae_nested_cli_explicit_args_override_model_defaults(self): + parser = FlexibleArgumentParser() + VAEConfig.add_cli_args(parser) + parsed = vars( + parser.parse_args( + [ + "--vae-config.use-parallel-decode", + "false", + "--vae-config.parallel-decode-mode", + "patch", + ] + ) + ) + config = FluxVAEConfig() + + update_config_from_args(config, parsed, "vae_config") + + self.assertFalse(config.use_parallel_decode) + self.assertEqual(config.parallel_decode_mode, "patch") + + def test_base_decode_prefers_spatial_parallel_dispatch(self): + config = VAEConfig() + config.arch_config.temporal_compression_ratio = 1 + config.arch_config.spatial_compression_ratio = 1 + config.use_parallel_decode = True + config.parallel_decode_mode = "spatial_shard" + vae = _DispatchProbeVAE(config) + z = torch.randn(1, 1, 1, 2, 2) + + with ( + patch( + "sglang.multimodal_gen.runtime.models.vaes.common.dist.is_initialized", + return_value=True, + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.common.model_parallel_is_initialized", + return_value=True, + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.common.get_decode_parallel_group_coordinator", + return_value=SimpleNamespace(), + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.common.get_decode_parallel_world_size", + return_value=2, + ), + ): + out = vae.decode(z) + + self.assertTrue(vae.used_decode) + self.assertFalse(vae.used_parallel_tiled_decode) + torch.testing.assert_close(out, z) + + def test_auto_decode_does_not_fall_back_to_parallel_tiling(self): + config = VAEConfig() + config.arch_config.temporal_compression_ratio = 1 + config.arch_config.spatial_compression_ratio = 1 + config.use_parallel_decode = True + config.parallel_decode_mode = "auto" + vae = _DispatchProbeVAE(config) + z = torch.randn(1, 1, 1, 2, 2) + + with patch( + "sglang.multimodal_gen.runtime.models.vaes.common.get_sp_world_size", + return_value=2, + ): + out = vae.decode(z) + + self.assertTrue(vae.used_decode) + self.assertFalse(vae.used_parallel_tiled_decode) + torch.testing.assert_close(out, z) + + def test_spatial_alias_still_uses_spatial_shard_dispatch(self): + self.assertTrue(is_spatial_shard_parallel_decode_mode("spatial")) + + def test_decode_parallel_group_uses_dedicated_group(self): + old_decode = parallel_state._VAE_DECODE + decode_group = SimpleNamespace(world_size=4, rank_in_group=2) + try: + parallel_state._VAE_DECODE = decode_group + self.assertIs( + parallel_state.get_decode_parallel_group_coordinator(), decode_group + ) + self.assertEqual(parallel_state.get_decode_parallel_world_size(), 4) + self.assertEqual(parallel_state.get_decode_parallel_rank(), 2) + finally: + parallel_state._VAE_DECODE = old_decode + + def test_decode_rank_groups_cover_non_dp_parallel_axes(self): + rank_generator = RankGenerator( + tp=2, + sp=2, + pp=1, + cfg=2, + dp=2, + order="tp-sp-pp-cfg-dp", + ) + + self.assertEqual( + rank_generator.get_ranks("tp-sp-pp-cfg"), + [list(range(0, 8)), list(range(8, 16))], + ) + + def test_spatial_split_keeps_uneven_height_lossless(self): + x = torch.arange(5).view(1, 1, 1, 5, 1) + + rank0, expected_height = split_for_parallel_decode( + x, upsample_count=1, world_size=2, rank=0 + ) + rank1, _ = split_for_parallel_decode(x, upsample_count=1, world_size=2, rank=1) + + self.assertEqual(expected_height, 10) + self.assertEqual(rank0.shape[-2], 3) + self.assertEqual(rank1.shape[-2], 2) + torch.testing.assert_close(torch.cat([rank0, rank1], dim=-2), x) + + def test_chunk_height_by_sizes_keeps_original_partition(self): + x = torch.arange(10).view(1, 1, 10, 1) + + with patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_rank", + return_value=0, + ): + rank0 = chunk_height_by_sizes(x, [6, 4]) + with patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_rank", + return_value=1, + ): + rank1 = chunk_height_by_sizes(x, [6, 4]) + + self.assertEqual(rank0.shape[-2], 6) + self.assertEqual(rank1.shape[-2], 4) + torch.testing.assert_close(torch.cat([rank0, rank1], dim=-2), x) + + def test_qwen_decoder_uses_spatial_parallel_components(self): + with ( + patch( + "sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage.dist.is_initialized", + return_value=True, + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage.model_parallel_is_initialized", + return_value=True, + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage.get_decode_parallel_world_size", + return_value=2, + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage.get_decode_parallel_rank", + return_value=0, + ), + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_world_size", + return_value=2, + ), + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_rank", + return_value=0, + ), + ): + decoder = QwenImageDecoder3d( + dim=4, + z_dim=2, + dim_mult=(1, 1), + num_res_blocks=1, + attn_scales=(), + temperal_upsample=(False,), + input_channels=3, + use_parallel_decode=True, + ) + + self.assertTrue( + any(isinstance(m, SpatialParallelCausalConv3d) for m in decoder.modules()) + ) + self.assertTrue( + any( + isinstance(m, QwenImageAttentionBlock) and m.spatial_parallel + for m in decoder.modules() + ) + ) + + def test_wan_decoder_uses_spatial_parallel_components(self): + with ( + patch( + "sglang.multimodal_gen.runtime.models.vaes.wanvae.dist.is_initialized", + return_value=True, + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.wanvae.get_decode_parallel_world_size", + return_value=2, + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.wanvae.get_decode_parallel_rank", + return_value=0, + ), + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_world_size", + return_value=2, + ), + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_rank", + return_value=0, + ), + ): + decoder = WanDecoder3d( + dim=4, + z_dim=2, + dim_mult=(1, 1), + num_res_blocks=1, + attn_scales=(), + temperal_upsample=(False,), + out_channels=3, + use_parallel_decode=True, + ) + + self.assertTrue( + any(isinstance(m, SpatialParallelCausalConv3d) for m in decoder.modules()) + ) + self.assertTrue( + any(isinstance(m, WanDistAttentionBlock) for m in decoder.modules()) + ) + + def test_diffusers_2d_decoder_uses_spatial_parallel_components(self): + config = StableDiffusion3VAEConfig() + config.use_parallel_decode = True + config.parallel_decode_mode = "spatial_shard" + config.arch_config.latent_channels = 2 + config.arch_config.block_out_channels = (4, 4) + config.arch_config.down_block_types = ( + "DownEncoderBlock2D", + "DownEncoderBlock2D", + ) + config.arch_config.up_block_types = ( + "UpDecoderBlock2D", + "UpDecoderBlock2D", + ) + config.arch_config.layers_per_block = 1 + config.arch_config.norm_num_groups = 1 + config.arch_config.sample_size = 8 + + with ( + patch( + "sglang.multimodal_gen.runtime.models.vaes.common.dist.is_initialized", + return_value=True, + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.common.model_parallel_is_initialized", + return_value=True, + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.common.get_decode_parallel_group_coordinator", + return_value=SimpleNamespace(), + ), + patch( + "sglang.multimodal_gen.runtime.models.vaes.common.get_decode_parallel_world_size", + return_value=2, + ), + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_world_size", + return_value=2, + ), + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_rank", + return_value=0, + ), + ): + vae = AutoencoderKL(config) + + self.assertTrue(vae._spatial_parallel_decode_enabled) + self.assertTrue( + any(isinstance(m, SpatialParallelConv2d) for m in vae.decoder.modules()) + ) + + def test_ltx_decoder_uses_spatial_parallel_conv3d(self): + with ( + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_world_size", + return_value=2, + ), + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_rank", + return_value=0, + ), + ): + decoder = LTX2VideoDecoder3d( + in_channels=4, + out_channels=3, + block_out_channels=(8,), + spatio_temporal_scaling=(False,), + layers_per_block=(1, 1), + patch_size=1, + upsample_residual=(False,), + upsample_factor=(1,), + spatial_padding_mode="reflect", + ) + _enable_ltx_decoder_spatial_parallel(decoder) + + causal_convs = [ + m for m in decoder.modules() if isinstance(m, LTX2VideoCausalConv3d) + ] + self.assertGreater(len(causal_convs), 0) + self.assertTrue( + all(isinstance(m.conv, SpatialParallelConv3d) for m in causal_convs) + ) + + def test_hunyuan_decoder_uses_spatial_parallel_components(self): + with ( + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_world_size", + return_value=2, + ), + patch( + "sglang.multimodal_gen.runtime.layers.parallel_conv.get_decode_parallel_rank", + return_value=0, + ), + ): + decoder = HunyuanVideoDecoder3D( + in_channels=4, + out_channels=3, + up_block_types=("HunyuanVideoUpBlock3D", "HunyuanVideoUpBlock3D"), + block_out_channels=(4, 8), + layers_per_block=1, + norm_num_groups=1, + mid_block_add_attention=True, + time_compression_ratio=4, + spatial_compression_ratio=2, + ) + _enable_hunyuan_decoder_spatial_parallel(decoder) + + self.assertTrue(decoder.spatial_parallel) + self.assertTrue( + any( + isinstance(m, HunyuanVideoMidBlock3D) and m.spatial_parallel + for m in decoder.modules() + ) + ) + self.assertTrue( + any(isinstance(m, SpatialParallelConv3d) for m in decoder.modules()) + ) + shortcut_convs = [ + m.conv_shortcut.conv + for m in decoder.modules() + if isinstance(m, HunyuanVideoResnetBlockCausal3D) + and m.conv_shortcut is not None + ] + self.assertTrue(shortcut_convs) + self.assertTrue( + all(not isinstance(conv, SpatialParallelConv3d) for conv in shortcut_convs) + ) + + +if __name__ == "__main__": + unittest.main()