diff --git a/python/sglang/multimodal_gen/configs/models/vaes/wanvae.py b/python/sglang/multimodal_gen/configs/models/vaes/wanvae.py index a1bd77ebf..f61f67dc9 100644 --- a/python/sglang/multimodal_gen/configs/models/vaes/wanvae.py +++ b/python/sglang/multimodal_gen/configs/models/vaes/wanvae.py @@ -82,6 +82,9 @@ class WanVAEConfig(VAEConfig): use_temporal_tiling: bool = False 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 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 new file mode 100644 index 000000000..25515f835 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_common_utils.py @@ -0,0 +1,457 @@ +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 + + +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.to(self.weight.dtype) if current_platform.is_mps() else x + ) # casting needed for mps since amp isn't supported + 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): + 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 new file mode 100644 index 000000000..f8aabc44f --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_dist_utils.py @@ -0,0 +1,677 @@ +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_sp_group, + get_sp_parallel_rank, + get_sp_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, + 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 gather_and_trim_height(x: torch.Tensor, expected_height: int | None): + if expected_height is None: + return x + x = get_sp_group().all_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: + if ( + recv_buf is None + or recv_buf.shape != reference.shape + or recv_buf.dtype != reference.dtype + or recv_buf.device != reference.device + ): + return torch.empty_like(reference) + 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 + + sp_group = get_sp_group() + rank = get_sp_parallel_rank() + world_size = get_sp_world_size() + group = sp_group.device_group + group_ranks = sp_group.ranks + + top_row = x[..., :height_halo_size, :].contiguous() + bottom_row = x[..., -height_halo_size:, :].contiguous() + + recv_top_buf = _ensure_recv_buf(recv_top_buf, top_row) + recv_bottom_buf = _ensure_recv_buf(recv_bottom_buf, bottom_row) + + reqs = [] + + 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] + reqs.append(dist.irecv(recv_top_buf, src=prev_rank, group=group)) + reqs.append(dist.isend(top_row, dst=prev_rank, group=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] + reqs.append(dist.isend(bottom_row, dst=next_rank, group=group)) + reqs.append(dist.irecv(recv_bottom_buf, src=next_rank, group=group)) + + if rank == 0: + recv_top_buf.zero_() + if rank == world_size - 1: + recv_bottom_buf.zero_() + + 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 = (self.padding[1], self.padding[1], 0, 0) + else: + self._padding = ( + self.padding[1], + self.padding[1], + self.padding[0], + self.padding[0], + ) + + self.padding = (0, 0) + self._halo_recv_top_buf: torch.Tensor | None = None + self._halo_recv_bottom_buf: torch.Tensor | None = None + self.rank = get_sp_parallel_rank() + self.world_size = get_sp_world_size() + + def forward(self, x): + 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_sp_parallel_rank() + self.world_size = get_sp_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.to(self.weight.dtype) if current_platform.is_mps() else x + ) # casting needed for mps since 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 + + 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_sp_parallel_rank() + self.world_size = get_sp_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_sp_parallel_rank() + self.world_size = get_sp_world_size() + self.sp_group = get_sp_group() + + def forward(self, x): + if self.world_size > 1: + x = self.sp_group.all_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 336c2fb5c..7279c2ed8 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py @@ -20,17 +20,49 @@ import contextvars from contextlib import contextmanager 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.runtime.distributed.parallel_state import ( + get_sp_parallel_rank, + get_sp_world_size, +) from sglang.multimodal_gen.runtime.layers.activation import get_act_fn from sglang.multimodal_gen.runtime.models.vaes.common import ( DiagonalGaussianDistribution, ParallelTiledVAE, ) -from sglang.multimodal_gen.runtime.platforms import current_platform +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, +) CACHE_T = 2 @@ -39,6 +71,8 @@ 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) + @contextmanager def forward_context( @@ -57,214 +91,6 @@ def forward_context( first_chunk.reset(first_chunk_token) -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: - 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. - - Args: - in_channels (int): Number of channels in the input image - out_channels (int): Number of channels produced by the convolution - kernel_size (int or tuple): Size of the convolving kernel - stride (int or tuple, optional): Stride of the convolution. Default: 1 - padding (int or tuple, optional): Zero-padding added to all three sides of the input. Default: 0 - """ - - 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.to(self.weight.dtype) if current_platform.is_mps() else x - ) # casting needed for mps since amp isn't supported - return super().forward(x) - - -class WanRMS_norm(nn.Module): - r""" - A custom RMS normalization layer. - - Args: - dim (int): The number of dimensions to normalize over. - channel_first (bool, optional): Whether the input tensor has channels as the first dimension. - Default is True. - images (bool, optional): Whether the input represents image data. Default is True. - bias (bool, optional): Whether to include a learnable bias term. Default is False. - """ - - 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. - - Args: - x (torch.Tensor): Input tensor to be upsampled. - - Returns: - torch.Tensor: Upsampled tensor with the same data type as the input. - """ - - def forward(self, x): - return super().forward(x.float()).type_as(x) - - class WanResample(nn.Module): r""" A custom resampling module for 2D and 3D data. @@ -317,86 +143,7 @@ class WanResample(nn.Module): self.resample = nn.Identity() def 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 + return resample_forward(self, x) class WanResidualBlock(nn.Module): @@ -433,70 +180,7 @@ class WanResidualBlock(nn.Module): ) def 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 + return residual_block_forward(self, x) class WanAttentionBlock(nn.Module): @@ -517,35 +201,7 @@ class WanAttentionBlock(nn.Module): self.proj = nn.Conv2d(dim, dim, 1) def forward(self, x): - identity = x - batch_size, channels, time, height, width = x.size() - - x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * time, channels, height, width) - x = self.norm(x) - - # compute query, key, value - qkv = self.to_qkv(x) - qkv = qkv.reshape(batch_size * time, 1, channels * 3, -1) - qkv = qkv.permute(0, 1, 3, 2).contiguous() - q, k, v = qkv.chunk(3, dim=-1) - - # apply attention - x = F.scaled_dot_product_attention(q, k, v) - - x = ( - x.squeeze(1) - .permute(0, 2, 1) - .reshape(batch_size * time, 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, time, channels, height, width) - x = x.permute(0, 2, 1, 3, 4) - - return x + identity + return attention_block_forward(self, x) class WanMidBlock(nn.Module): @@ -580,17 +236,7 @@ class WanMidBlock(nn.Module): self.gradient_checkpointing = False def 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 + return mid_block_forward(self, x) class WanResidualDownBlock(nn.Module): @@ -629,13 +275,7 @@ class WanResidualDownBlock(nn.Module): self.downsampler = None def forward(self, x): - x_copy = x.clone() - 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) + return residual_down_block_forward(self, x) class WanEncoder3d(nn.Module): @@ -665,6 +305,7 @@ class WanEncoder3d(nn.Module): dropout=0.0, non_linearity: str = "silu", is_residual: bool = False, # wan 2.2 vae use a residual downblock + use_parallel_encode: bool = False, ): super().__init__() self.dim = dim @@ -675,13 +316,34 @@ class WanEncoder3d(nn.Module): self.attn_scales = list(attn_scales) self.temperal_downsample = list(temperal_downsample) self.nonlinearity = get_act_fn(non_linearity) + self.use_parallel_encode = use_parallel_encode + self.downsample_count = max(len(dim_mult) - 1, 0) # dimensions dims = [dim * u for u in [1] + dim_mult] scale = 1.0 + world_size = 1 + if dist.is_initialized(): + world_size = get_sp_world_size() + + if use_parallel_encode and world_size > 1: + CausalConv3d = WanDistCausalConv3d + ResidualDownBlock = WanDistResidualDownBlock + ResidualBlock = WanDistResidualBlock + AttentionBlock = WanDistAttentionBlock + Resample = WanDistResample + MidBlock = WanDistMidBlock + else: + CausalConv3d = WanCausalConv3d + ResidualDownBlock = WanResidualDownBlock + ResidualBlock = WanResidualBlock + AttentionBlock = WanAttentionBlock + Resample = WanResample + MidBlock = WanMidBlock + # init block - self.conv_in = WanCausalConv3d(in_channels, dims[0], 3, padding=1) + self.conv_in = CausalConv3d(in_channels, dims[0], 3, padding=1) # downsample blocks self.down_blocks = nn.ModuleList([]) @@ -689,7 +351,7 @@ class WanEncoder3d(nn.Module): # residual (+attention) blocks if is_residual: self.down_blocks.append( - WanResidualDownBlock( + ResidualDownBlock( in_dim, out_dim, dropout, @@ -702,27 +364,39 @@ class WanEncoder3d(nn.Module): ) else: for _ in range(num_res_blocks): - self.down_blocks.append(WanResidualBlock(in_dim, out_dim, dropout)) + self.down_blocks.append(ResidualBlock(in_dim, out_dim, dropout)) if scale in attn_scales: - self.down_blocks.append(WanAttentionBlock(out_dim)) + self.down_blocks.append(AttentionBlock(out_dim)) in_dim = out_dim # downsample block if i != len(dim_mult) - 1: mode = "downsample3d" if temperal_downsample[i] else "downsample2d" - self.down_blocks.append(WanResample(out_dim, mode=mode)) + self.down_blocks.append(Resample(out_dim, mode=mode)) scale /= 2.0 # middle blocks - self.mid_block = WanMidBlock(out_dim, dropout, non_linearity, num_layers=1) + self.mid_block = MidBlock(out_dim, dropout, non_linearity, num_layers=1) # output blocks self.norm_out = WanRMS_norm(out_dim, images=False) - self.conv_out = WanCausalConv3d(out_dim, z_dim, 3, padding=1) + self.conv_out = CausalConv3d(out_dim, z_dim, 3, padding=1) self.gradient_checkpointing = False + self.world_size = 1 + self.rank = 0 + if dist.is_initialized(): + self.world_size = get_sp_world_size() + self.rank = get_sp_parallel_rank() def forward(self, x): + expected_local_height = None + expected_height = None + if self.use_parallel_encode and self.world_size > 1: + x, expected_height, expected_local_height = split_for_parallel_encode( + x, self.downsample_count, self.world_size, self.rank + ) + _feat_cache = feat_cache.get() _feat_idx = feat_idx.get() if _feat_cache is not None: @@ -752,6 +426,8 @@ class WanEncoder3d(nn.Module): x = layer(x) ## middle + if self.use_parallel_encode and self.world_size > 1: + x = ensure_local_height(x, expected_local_height) x = self.mid_block(x) ## head @@ -781,6 +457,9 @@ class WanEncoder3d(nn.Module): feat_idx.set(_feat_idx) else: x = self.conv_out(x) + + if self.use_parallel_encode and self.world_size > 1: + x = gather_and_trim_height(x, expected_height) return x @@ -845,28 +524,7 @@ class WanResidualUpBlock(nn.Module): self.gradient_checkpointing = False def forward(self, x): - """ - Forward pass through the upsampling block. - Args: - x (torch.Tensor): Input tensor - feat_cache (list, optional): Feature cache for causal convolutions - feat_idx (list, optional): Feature index for cache management - Returns: - torch.Tensor: Output tensor - """ - if self.avg_shortcut is not None: - x_copy = x.clone() - - 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 + return residual_up_block_forward(self, x) class WanUpBlock(nn.Module): @@ -915,23 +573,7 @@ class WanUpBlock(nn.Module): self.gradient_checkpointing = False def forward(self, x): - """ - Forward pass through the upsampling block. - - Args: - x (torch.Tensor): Input tensor - feat_cache (list, optional): Feature cache for causal convolutions - feat_idx (list, optional): Feature index for cache management - - Returns: - torch.Tensor: Output tensor - """ - for resnet in self.resnets: - x = resnet(x) - - if self.upsamplers is not None: - x = self.upsamplers[0](x) - return x + return up_block_forward(self, x) class WanDecoder3d(nn.Module): @@ -961,6 +603,7 @@ class WanDecoder3d(nn.Module): non_linearity: str = "silu", out_channels: int = 3, is_residual: bool = False, + use_parallel_decode: bool = False, ): super().__init__() self.dim = dim @@ -972,17 +615,35 @@ class WanDecoder3d(nn.Module): self.temperal_upsample = list(temperal_upsample) self.nonlinearity = get_act_fn(non_linearity) + self.use_parallel_decode = use_parallel_decode + self.upsample_count = 0 # dimensions dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]] + world_size = 1 + if dist.is_initialized(): + world_size = get_sp_world_size() + + if use_parallel_decode and world_size > 1: + CausalConv3d = WanDistCausalConv3d + MidBlock = WanDistMidBlock + ResidualUpBlock = WanDistResidualUpBlock + UpBlock = WanDistUpBlock + else: + CausalConv3d = WanCausalConv3d + MidBlock = WanMidBlock + ResidualUpBlock = WanResidualUpBlock + UpBlock = WanUpBlock + # init block - self.conv_in = WanCausalConv3d(z_dim, dims[0], 3, padding=1) + self.conv_in = CausalConv3d(z_dim, dims[0], 3, padding=1) # middle blocks - self.mid_block = WanMidBlock(dims[0], dropout, non_linearity, num_layers=1) + self.mid_block = MidBlock(dims[0], dropout, non_linearity, num_layers=1) # upsample blocks + self.upsample_count = 0 self.up_blocks = nn.ModuleList([]) for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=True)): # residual (+attention) blocks @@ -1001,7 +662,7 @@ class WanDecoder3d(nn.Module): # Create and add the upsampling block if is_residual: - up_block = WanResidualUpBlock( + up_block = ResidualUpBlock( in_dim=in_dim, out_dim=out_dim, num_res_blocks=num_res_blocks, @@ -1011,7 +672,7 @@ class WanDecoder3d(nn.Module): non_linearity=non_linearity, ) else: - up_block = WanUpBlock( + up_block = UpBlock( in_dim=in_dim, out_dim=out_dim, num_res_blocks=num_res_blocks, @@ -1020,14 +681,27 @@ class WanDecoder3d(nn.Module): non_linearity=non_linearity, ) self.up_blocks.append(up_block) + if up_flag: + self.upsample_count += 1 # output blocks self.norm_out = WanRMS_norm(out_dim, images=False) - self.conv_out = WanCausalConv3d(out_dim, out_channels, 3, padding=1) + self.conv_out = CausalConv3d(out_dim, out_channels, 3, padding=1) self.gradient_checkpointing = False + self.world_size = 1 + self.rank = 0 + if dist.is_initialized(): + self.world_size = get_sp_world_size() + self.rank = get_sp_parallel_rank() def forward(self, x): + 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 _feat_cache = feat_cache.get() _feat_idx = feat_idx.get() @@ -1086,6 +760,9 @@ class WanDecoder3d(nn.Module): feat_idx.set(_feat_idx) 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 @@ -1152,6 +829,8 @@ class AutoencoderKLWan(ParallelTiledVAE): self.latents_mean = list(config.latents_mean) self.latents_std = list(config.latents_std) self.shift_factor = config.shift_factor + self.use_parallel_encode = getattr(config, "use_parallel_encode", False) + self.use_parallel_decode = getattr(config, "use_parallel_decode", False) if config.load_encoder: self.encoder = WanEncoder3d( @@ -1164,6 +843,7 @@ class AutoencoderKLWan(ParallelTiledVAE): temperal_downsample=self.temperal_downsample, dropout=config.dropout, is_residual=config.is_residual, + use_parallel_encode=self.use_parallel_encode, ) self.quant_conv = WanCausalConv3d(self.z_dim * 2, self.z_dim * 2, 1) self.post_quant_conv = WanCausalConv3d(self.z_dim, self.z_dim, 1) @@ -1179,6 +859,7 @@ class AutoencoderKLWan(ParallelTiledVAE): dropout=config.dropout, out_channels=config.out_channels, is_residual=config.is_residual, + use_parallel_decode=self.use_parallel_decode, ) self.use_feature_cache = config.use_feature_cache @@ -1188,7 +869,7 @@ class AutoencoderKLWan(ParallelTiledVAE): def _count_conv3d(model) -> int: count = 0 for m in model.modules(): - if isinstance(m, WanCausalConv3d): + if isinstance(m, WanCausalConv3d) or isinstance(m, WanDistCausalConv3d): count += 1 return count