[diffusion] feat: support parallel wan-vae decode (#18179)

This commit is contained in:
wxy
2026-02-10 18:32:00 +08:00
committed by GitHub
parent 26f2b3798d
commit 47978ee858
4 changed files with 1271 additions and 453 deletions

View File

@@ -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

View File

@@ -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

View File

@@ -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)

View File

@@ -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