[diffusion] feat: support parallel wan-vae decode (#18179)
This commit is contained in:
@@ -82,6 +82,9 @@ class WanVAEConfig(VAEConfig):
|
|||||||
use_temporal_tiling: bool = False
|
use_temporal_tiling: bool = False
|
||||||
use_parallel_tiling: bool = False
|
use_parallel_tiling: bool = False
|
||||||
|
|
||||||
|
use_parallel_encode: bool = True
|
||||||
|
use_parallel_decode: bool = True
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self.blend_num_frames = (
|
self.blend_num_frames = (
|
||||||
self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
|
self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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)
|
||||||
@@ -20,17 +20,49 @@ import contextvars
|
|||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.vaes import WanVAEConfig
|
from sglang.multimodal_gen.configs.models.vaes import WanVAEConfig
|
||||||
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
|
get_sp_parallel_rank,
|
||||||
|
get_sp_world_size,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
|
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
|
||||||
from sglang.multimodal_gen.runtime.models.vaes.common import (
|
from sglang.multimodal_gen.runtime.models.vaes.common import (
|
||||||
DiagonalGaussianDistribution,
|
DiagonalGaussianDistribution,
|
||||||
ParallelTiledVAE,
|
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
|
CACHE_T = 2
|
||||||
|
|
||||||
@@ -39,6 +71,8 @@ feat_cache = contextvars.ContextVar("feat_cache", default=None)
|
|||||||
feat_idx = contextvars.ContextVar("feat_idx", default=0)
|
feat_idx = contextvars.ContextVar("feat_idx", default=0)
|
||||||
first_chunk = contextvars.ContextVar("first_chunk", default=None)
|
first_chunk = contextvars.ContextVar("first_chunk", default=None)
|
||||||
|
|
||||||
|
bind_context(is_first_frame, feat_cache, feat_idx, CACHE_T, first_chunk)
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def forward_context(
|
def forward_context(
|
||||||
@@ -57,214 +91,6 @@ def forward_context(
|
|||||||
first_chunk.reset(first_chunk_token)
|
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):
|
class WanResample(nn.Module):
|
||||||
r"""
|
r"""
|
||||||
A custom resampling module for 2D and 3D data.
|
A custom resampling module for 2D and 3D data.
|
||||||
@@ -317,86 +143,7 @@ class WanResample(nn.Module):
|
|||||||
self.resample = nn.Identity()
|
self.resample = nn.Identity()
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
b, c, t, h, w = x.size()
|
return resample_forward(self, x)
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
class WanResidualBlock(nn.Module):
|
class WanResidualBlock(nn.Module):
|
||||||
@@ -433,70 +180,7 @@ class WanResidualBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
# Apply shortcut connection
|
return residual_block_forward(self, x)
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
class WanAttentionBlock(nn.Module):
|
class WanAttentionBlock(nn.Module):
|
||||||
@@ -517,35 +201,7 @@ class WanAttentionBlock(nn.Module):
|
|||||||
self.proj = nn.Conv2d(dim, dim, 1)
|
self.proj = nn.Conv2d(dim, dim, 1)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
identity = x
|
return attention_block_forward(self, 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
|
|
||||||
|
|
||||||
|
|
||||||
class WanMidBlock(nn.Module):
|
class WanMidBlock(nn.Module):
|
||||||
@@ -580,17 +236,7 @@ class WanMidBlock(nn.Module):
|
|||||||
self.gradient_checkpointing = False
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
# First residual block
|
return mid_block_forward(self, x)
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
class WanResidualDownBlock(nn.Module):
|
class WanResidualDownBlock(nn.Module):
|
||||||
@@ -629,13 +275,7 @@ class WanResidualDownBlock(nn.Module):
|
|||||||
self.downsampler = None
|
self.downsampler = None
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
x_copy = x.clone()
|
return residual_down_block_forward(self, 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)
|
|
||||||
|
|
||||||
|
|
||||||
class WanEncoder3d(nn.Module):
|
class WanEncoder3d(nn.Module):
|
||||||
@@ -665,6 +305,7 @@ class WanEncoder3d(nn.Module):
|
|||||||
dropout=0.0,
|
dropout=0.0,
|
||||||
non_linearity: str = "silu",
|
non_linearity: str = "silu",
|
||||||
is_residual: bool = False, # wan 2.2 vae use a residual downblock
|
is_residual: bool = False, # wan 2.2 vae use a residual downblock
|
||||||
|
use_parallel_encode: bool = False,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dim = dim
|
self.dim = dim
|
||||||
@@ -675,13 +316,34 @@ class WanEncoder3d(nn.Module):
|
|||||||
self.attn_scales = list(attn_scales)
|
self.attn_scales = list(attn_scales)
|
||||||
self.temperal_downsample = list(temperal_downsample)
|
self.temperal_downsample = list(temperal_downsample)
|
||||||
self.nonlinearity = get_act_fn(non_linearity)
|
self.nonlinearity = get_act_fn(non_linearity)
|
||||||
|
self.use_parallel_encode = use_parallel_encode
|
||||||
|
self.downsample_count = max(len(dim_mult) - 1, 0)
|
||||||
|
|
||||||
# dimensions
|
# dimensions
|
||||||
dims = [dim * u for u in [1] + dim_mult]
|
dims = [dim * u for u in [1] + dim_mult]
|
||||||
scale = 1.0
|
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
|
# 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
|
# downsample blocks
|
||||||
self.down_blocks = nn.ModuleList([])
|
self.down_blocks = nn.ModuleList([])
|
||||||
@@ -689,7 +351,7 @@ class WanEncoder3d(nn.Module):
|
|||||||
# residual (+attention) blocks
|
# residual (+attention) blocks
|
||||||
if is_residual:
|
if is_residual:
|
||||||
self.down_blocks.append(
|
self.down_blocks.append(
|
||||||
WanResidualDownBlock(
|
ResidualDownBlock(
|
||||||
in_dim,
|
in_dim,
|
||||||
out_dim,
|
out_dim,
|
||||||
dropout,
|
dropout,
|
||||||
@@ -702,27 +364,39 @@ class WanEncoder3d(nn.Module):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
for _ in range(num_res_blocks):
|
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:
|
if scale in attn_scales:
|
||||||
self.down_blocks.append(WanAttentionBlock(out_dim))
|
self.down_blocks.append(AttentionBlock(out_dim))
|
||||||
in_dim = out_dim
|
in_dim = out_dim
|
||||||
|
|
||||||
# downsample block
|
# downsample block
|
||||||
if i != len(dim_mult) - 1:
|
if i != len(dim_mult) - 1:
|
||||||
mode = "downsample3d" if temperal_downsample[i] else "downsample2d"
|
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
|
scale /= 2.0
|
||||||
|
|
||||||
# middle blocks
|
# 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
|
# output blocks
|
||||||
self.norm_out = WanRMS_norm(out_dim, images=False)
|
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.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):
|
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_cache = feat_cache.get()
|
||||||
_feat_idx = feat_idx.get()
|
_feat_idx = feat_idx.get()
|
||||||
if _feat_cache is not None:
|
if _feat_cache is not None:
|
||||||
@@ -752,6 +426,8 @@ class WanEncoder3d(nn.Module):
|
|||||||
x = layer(x)
|
x = layer(x)
|
||||||
|
|
||||||
## middle
|
## middle
|
||||||
|
if self.use_parallel_encode and self.world_size > 1:
|
||||||
|
x = ensure_local_height(x, expected_local_height)
|
||||||
x = self.mid_block(x)
|
x = self.mid_block(x)
|
||||||
|
|
||||||
## head
|
## head
|
||||||
@@ -781,6 +457,9 @@ class WanEncoder3d(nn.Module):
|
|||||||
feat_idx.set(_feat_idx)
|
feat_idx.set(_feat_idx)
|
||||||
else:
|
else:
|
||||||
x = self.conv_out(x)
|
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
|
return x
|
||||||
|
|
||||||
|
|
||||||
@@ -845,28 +524,7 @@ class WanResidualUpBlock(nn.Module):
|
|||||||
self.gradient_checkpointing = False
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
"""
|
return residual_up_block_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
|
|
||||||
|
|
||||||
|
|
||||||
class WanUpBlock(nn.Module):
|
class WanUpBlock(nn.Module):
|
||||||
@@ -915,23 +573,7 @@ class WanUpBlock(nn.Module):
|
|||||||
self.gradient_checkpointing = False
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
"""
|
return up_block_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
|
|
||||||
|
|
||||||
|
|
||||||
class WanDecoder3d(nn.Module):
|
class WanDecoder3d(nn.Module):
|
||||||
@@ -961,6 +603,7 @@ class WanDecoder3d(nn.Module):
|
|||||||
non_linearity: str = "silu",
|
non_linearity: str = "silu",
|
||||||
out_channels: int = 3,
|
out_channels: int = 3,
|
||||||
is_residual: bool = False,
|
is_residual: bool = False,
|
||||||
|
use_parallel_decode: bool = False,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dim = dim
|
self.dim = dim
|
||||||
@@ -972,17 +615,35 @@ class WanDecoder3d(nn.Module):
|
|||||||
self.temperal_upsample = list(temperal_upsample)
|
self.temperal_upsample = list(temperal_upsample)
|
||||||
|
|
||||||
self.nonlinearity = get_act_fn(non_linearity)
|
self.nonlinearity = get_act_fn(non_linearity)
|
||||||
|
self.use_parallel_decode = use_parallel_decode
|
||||||
|
self.upsample_count = 0
|
||||||
|
|
||||||
# dimensions
|
# dimensions
|
||||||
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
|
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
|
# 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
|
# 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
|
# upsample blocks
|
||||||
|
self.upsample_count = 0
|
||||||
self.up_blocks = nn.ModuleList([])
|
self.up_blocks = nn.ModuleList([])
|
||||||
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=True)):
|
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=True)):
|
||||||
# residual (+attention) blocks
|
# residual (+attention) blocks
|
||||||
@@ -1001,7 +662,7 @@ class WanDecoder3d(nn.Module):
|
|||||||
|
|
||||||
# Create and add the upsampling block
|
# Create and add the upsampling block
|
||||||
if is_residual:
|
if is_residual:
|
||||||
up_block = WanResidualUpBlock(
|
up_block = ResidualUpBlock(
|
||||||
in_dim=in_dim,
|
in_dim=in_dim,
|
||||||
out_dim=out_dim,
|
out_dim=out_dim,
|
||||||
num_res_blocks=num_res_blocks,
|
num_res_blocks=num_res_blocks,
|
||||||
@@ -1011,7 +672,7 @@ class WanDecoder3d(nn.Module):
|
|||||||
non_linearity=non_linearity,
|
non_linearity=non_linearity,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
up_block = WanUpBlock(
|
up_block = UpBlock(
|
||||||
in_dim=in_dim,
|
in_dim=in_dim,
|
||||||
out_dim=out_dim,
|
out_dim=out_dim,
|
||||||
num_res_blocks=num_res_blocks,
|
num_res_blocks=num_res_blocks,
|
||||||
@@ -1020,14 +681,27 @@ class WanDecoder3d(nn.Module):
|
|||||||
non_linearity=non_linearity,
|
non_linearity=non_linearity,
|
||||||
)
|
)
|
||||||
self.up_blocks.append(up_block)
|
self.up_blocks.append(up_block)
|
||||||
|
if up_flag:
|
||||||
|
self.upsample_count += 1
|
||||||
|
|
||||||
# output blocks
|
# output blocks
|
||||||
self.norm_out = WanRMS_norm(out_dim, images=False)
|
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.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):
|
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
|
## conv1
|
||||||
_feat_cache = feat_cache.get()
|
_feat_cache = feat_cache.get()
|
||||||
_feat_idx = feat_idx.get()
|
_feat_idx = feat_idx.get()
|
||||||
@@ -1086,6 +760,9 @@ class WanDecoder3d(nn.Module):
|
|||||||
feat_idx.set(_feat_idx)
|
feat_idx.set(_feat_idx)
|
||||||
else:
|
else:
|
||||||
x = self.conv_out(x)
|
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
|
return x
|
||||||
|
|
||||||
|
|
||||||
@@ -1152,6 +829,8 @@ class AutoencoderKLWan(ParallelTiledVAE):
|
|||||||
self.latents_mean = list(config.latents_mean)
|
self.latents_mean = list(config.latents_mean)
|
||||||
self.latents_std = list(config.latents_std)
|
self.latents_std = list(config.latents_std)
|
||||||
self.shift_factor = config.shift_factor
|
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:
|
if config.load_encoder:
|
||||||
self.encoder = WanEncoder3d(
|
self.encoder = WanEncoder3d(
|
||||||
@@ -1164,6 +843,7 @@ class AutoencoderKLWan(ParallelTiledVAE):
|
|||||||
temperal_downsample=self.temperal_downsample,
|
temperal_downsample=self.temperal_downsample,
|
||||||
dropout=config.dropout,
|
dropout=config.dropout,
|
||||||
is_residual=config.is_residual,
|
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.quant_conv = WanCausalConv3d(self.z_dim * 2, self.z_dim * 2, 1)
|
||||||
self.post_quant_conv = WanCausalConv3d(self.z_dim, self.z_dim, 1)
|
self.post_quant_conv = WanCausalConv3d(self.z_dim, self.z_dim, 1)
|
||||||
@@ -1179,6 +859,7 @@ class AutoencoderKLWan(ParallelTiledVAE):
|
|||||||
dropout=config.dropout,
|
dropout=config.dropout,
|
||||||
out_channels=config.out_channels,
|
out_channels=config.out_channels,
|
||||||
is_residual=config.is_residual,
|
is_residual=config.is_residual,
|
||||||
|
use_parallel_decode=self.use_parallel_decode,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.use_feature_cache = config.use_feature_cache
|
self.use_feature_cache = config.use_feature_cache
|
||||||
@@ -1188,7 +869,7 @@ class AutoencoderKLWan(ParallelTiledVAE):
|
|||||||
def _count_conv3d(model) -> int:
|
def _count_conv3d(model) -> int:
|
||||||
count = 0
|
count = 0
|
||||||
for m in model.modules():
|
for m in model.modules():
|
||||||
if isinstance(m, WanCausalConv3d):
|
if isinstance(m, WanCausalConv3d) or isinstance(m, WanDistCausalConv3d):
|
||||||
count += 1
|
count += 1
|
||||||
return count
|
return count
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user