Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bb06c1d634 | ||
|
|
db59678e7e |
@@ -19,14 +19,34 @@ from typing import Optional, Tuple, Union
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
from contextlib import contextmanager
|
||||
import contextvars
|
||||
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.models.utils import auto_attributes
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE, DiagonalGaussianDistribution
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
is_first_frame = contextvars.ContextVar("is_first_frame", default=False)
|
||||
feat_cache = contextvars.ContextVar("feat_cache", default=None)
|
||||
feat_idx = contextvars.ContextVar("feat_idx", default=0)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def forward_context(first_frame_arg=False,
|
||||
feat_cache_arg=None,
|
||||
feat_idx_arg=None):
|
||||
is_first_frame_token = is_first_frame.set(first_frame_arg)
|
||||
feat_cache_token = feat_cache.set(feat_cache_arg)
|
||||
feat_idx_token = feat_idx.set(feat_idx_arg)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
is_first_frame.reset(is_first_frame_token)
|
||||
feat_cache.reset(feat_cache_token)
|
||||
feat_idx.reset(feat_idx_token)
|
||||
|
||||
|
||||
class WanCausalConv3d(nn.Conv3d):
|
||||
r"""
|
||||
@@ -60,12 +80,17 @@ class WanCausalConv3d(nn.Conv3d):
|
||||
)
|
||||
self.padding: Tuple[int, int, int]
|
||||
# Set up causal padding
|
||||
self._padding = (self.padding[2], self.padding[2], self.padding[1],
|
||||
self.padding[1], 2 * self.padding[0], 0)
|
||||
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):
|
||||
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)
|
||||
return super().forward(x)
|
||||
|
||||
@@ -157,28 +182,82 @@ class WanResample(nn.Module):
|
||||
self.time_conv = WanCausalConv3d(dim,
|
||||
dim, (3, 1, 1),
|
||||
stride=(2, 1, 1),
|
||||
padding=(1, 0, 0))
|
||||
padding=(0, 0, 0))
|
||||
|
||||
else:
|
||||
self.resample = nn.Identity()
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
def forward(self, x):
|
||||
b, c, t, h, w = x.size()
|
||||
first_frame = is_first_frame.get()
|
||||
if first_frame:
|
||||
assert t == 1
|
||||
if self.mode == "upsample3d" and 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)
|
||||
_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)
|
||||
if self.mode == "downsample3d" and not first_frame and hasattr(
|
||||
self, "time_conv"):
|
||||
x = self.time_conv(x)
|
||||
|
||||
_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
|
||||
|
||||
|
||||
@@ -222,7 +301,25 @@ class WanResidualBlock(nn.Module):
|
||||
x = self.norm1(x)
|
||||
x = self.nonlinearity(x)
|
||||
|
||||
x = self.conv1(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)
|
||||
@@ -231,7 +328,25 @@ class WanResidualBlock(nn.Module):
|
||||
# Dropout
|
||||
x = self.dropout(x)
|
||||
|
||||
x = self.conv2(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
|
||||
@@ -400,15 +515,30 @@ class WanEncoder3d(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
x = self.conv_in(x)
|
||||
def forward(self, 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 last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_in(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.conv_in(x)
|
||||
|
||||
## downsamples
|
||||
for layer in self.down_blocks:
|
||||
if isinstance(layer, WanResample):
|
||||
x = layer(x, first_frame=first_frame)
|
||||
else:
|
||||
x = layer(x)
|
||||
x = layer(x)
|
||||
|
||||
## middle
|
||||
x = self.mid_block(x)
|
||||
@@ -416,7 +546,26 @@ class WanEncoder3d(nn.Module):
|
||||
## head
|
||||
x = self.norm_out(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.conv_out(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 last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_out(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.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -465,7 +614,7 @@ class WanUpBlock(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
def forward(self, x):
|
||||
"""
|
||||
Forward pass through the upsampling block.
|
||||
|
||||
@@ -481,7 +630,7 @@ class WanUpBlock(nn.Module):
|
||||
x = resnet(x)
|
||||
|
||||
if self.upsamplers is not None:
|
||||
x = self.upsamplers[0](x, first_frame=first_frame)
|
||||
x = self.upsamplers[0](x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -569,21 +718,57 @@ class WanDecoder3d(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
def forward(self, x):
|
||||
## conv1
|
||||
x = self.conv_in(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 last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_in(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.conv_in(x)
|
||||
|
||||
## middle
|
||||
x = self.mid_block(x)
|
||||
|
||||
## upsamples
|
||||
for up_block in self.up_blocks:
|
||||
x = up_block(x, first_frame=first_frame)
|
||||
x = up_block(x)
|
||||
|
||||
## head
|
||||
x = self.norm_out(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.conv_out(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 last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_out(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.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -681,10 +866,63 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
self.tile_sample_stride_height = 192
|
||||
self.tile_sample_stride_width = 192
|
||||
self.tile_sample_stride_num_frames = 12
|
||||
|
||||
# Whether to use the feature cache algorithm used by diffusers and Wan2.1
|
||||
self.use_feature_cache = True # default to True for best performance
|
||||
ParallelTiledVAE.__init__(self)
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
|
||||
def _count_conv3d(model) -> int:
|
||||
count = 0
|
||||
for m in model.modules():
|
||||
if isinstance(m, WanCausalConv3d):
|
||||
count += 1
|
||||
return count
|
||||
|
||||
self._conv_num = _count_conv3d(self.decoder)
|
||||
self._conv_idx = 0
|
||||
self._feat_map = [None] * self._conv_num
|
||||
# cache encode
|
||||
self._enc_conv_num = _count_conv3d(self.encoder)
|
||||
self._enc_conv_idx = 0
|
||||
self._enc_feat_map = [None] * self._enc_conv_num
|
||||
|
||||
def encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.use_feature_cache:
|
||||
self.clear_cache()
|
||||
with forward_context(feat_cache_arg=self._enc_feat_map,
|
||||
feat_idx_arg=self._enc_conv_idx):
|
||||
t = x.shape[2]
|
||||
iter_ = 1 + (t - 1) // 4
|
||||
for i in range(iter_):
|
||||
feat_idx.set(0)
|
||||
if i == 0:
|
||||
out = self.encoder(x[:, :, :1, :, :])
|
||||
else:
|
||||
out_ = self.encoder(x[:, :,
|
||||
1 + 4 * (i - 1):1 + 4 * i, :, :])
|
||||
out = torch.cat([out, out_], 2)
|
||||
enc = self.quant_conv(out)
|
||||
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:,
|
||||
self.z_dim:, :, :, :]
|
||||
enc = torch.cat([mu, logvar], dim=1)
|
||||
enc = DiagonalGaussianDistribution(enc)
|
||||
self.clear_cache()
|
||||
else:
|
||||
for block in self.encoder.down_blocks:
|
||||
if isinstance(block,
|
||||
WanResample) and block.mode == "downsample3d":
|
||||
_padding = list(block.time_conv._padding)
|
||||
_padding[4] = 2
|
||||
block.time_conv._padding = tuple(_padding)
|
||||
enc = ParallelTiledVAE.encode(self, x)
|
||||
|
||||
return enc
|
||||
|
||||
def _encode(self, x: torch.Tensor, first_frame=False) -> torch.Tensor:
|
||||
out = self.encoder(x, first_frame=first_frame)
|
||||
with forward_context(first_frame_arg=first_frame):
|
||||
out = self.encoder(x)
|
||||
enc = self.quant_conv(out)
|
||||
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:, self.z_dim:, :, :, :]
|
||||
enc = torch.cat([mu, logvar], dim=1)
|
||||
@@ -708,9 +946,32 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
enc = torch.cat([first_frame, enc], dim=2)
|
||||
return enc
|
||||
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
if self.use_feature_cache:
|
||||
self.clear_cache()
|
||||
iter_ = z.shape[2]
|
||||
x = self.post_quant_conv(z)
|
||||
with forward_context(feat_cache_arg=self._feat_map,
|
||||
feat_idx_arg=self._conv_idx):
|
||||
for i in range(iter_):
|
||||
feat_idx.set(0)
|
||||
if i == 0:
|
||||
out = self.decoder(x[:, :, i:i + 1, :, :])
|
||||
else:
|
||||
out_ = self.decoder(x[:, :, i:i + 1, :, :])
|
||||
out = torch.cat([out, out_], 2)
|
||||
|
||||
out = torch.clamp(out, min=-1.0, max=1.0)
|
||||
self.clear_cache()
|
||||
else:
|
||||
out = ParallelTiledVAE.decode(self, z)
|
||||
|
||||
return out
|
||||
|
||||
def _decode(self, z: torch.Tensor, first_frame=False) -> torch.Tensor:
|
||||
x = self.post_quant_conv(z)
|
||||
out = self.decoder(x, first_frame=first_frame)
|
||||
with forward_context(first_frame_arg=first_frame):
|
||||
out = self.decoder(x)
|
||||
|
||||
out = torch.clamp(out, min=-1.0, max=1.0)
|
||||
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -33,6 +33,7 @@ def test_wan_vae():
|
||||
|
||||
loader = VAELoader()
|
||||
model2 = loader.load(VAE_PATH, "", args)
|
||||
assert model2.use_feature_cache # Default to use the original WanVAE algorithm
|
||||
|
||||
model1 = AutoencoderKLWan.from_pretrained(
|
||||
VAE_PATH, torch_dtype=precision).to(device).eval()
|
||||
@@ -48,43 +49,52 @@ def test_wan_vae():
|
||||
32,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
latent_tensor = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
32,
|
||||
32,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
# latent_tensor = torch.randn(batch_size,
|
||||
# 16,
|
||||
# 21,
|
||||
# 32,
|
||||
# 32,
|
||||
# device=device,
|
||||
# dtype=precision)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
# Test encoding
|
||||
logger.info("Testing encoding...")
|
||||
latent1 = model1.encode(input_tensor).latent_dist.mean
|
||||
latent1 = model1.encode(input_tensor).latent_dist
|
||||
print("--------------------------------")
|
||||
latent2 = model2.encode(input_tensor).mean
|
||||
latent2 = model2.encode(input_tensor)
|
||||
# Check if latents have the same shape
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
assert latent1.mean.shape == latent2.mean.shape, f"Latent shapes don't match: {latent1.mean.shape} vs {latent2.mean.shape}"
|
||||
# Check if latents are similar
|
||||
max_diff_encode = torch.max(torch.abs(latent1 - latent2))
|
||||
mean_diff_encode = torch.mean(torch.abs(latent1 - latent2))
|
||||
max_diff_encode = torch.max(torch.abs(latent1.mean - latent2.mean))
|
||||
mean_diff_encode = torch.mean(torch.abs(latent1.mean - latent2.mean))
|
||||
logger.info("Maximum difference between encoded latents: %s",
|
||||
max_diff_encode.item())
|
||||
logger.info("Mean difference between encoded latents: %s",
|
||||
mean_diff_encode.item())
|
||||
assert mean_diff_encode < 5e-1, f"Encoded latents differ significantly: mean diff = {mean_diff_encode.item()}"
|
||||
assert max_diff_encode < 1e-5, f"Encoded latents differ significantly: max diff = {mean_diff_encode.item()}"
|
||||
# Test decoding
|
||||
logger.info("Testing decoding...")
|
||||
latent1_tensor = latent1.mode()
|
||||
latents_mean = (torch.tensor(model1.config.latents_mean).view(
|
||||
1, model1.config.z_dim, 1, 1, 1).to(latent_tensor.device,
|
||||
latent_tensor.dtype))
|
||||
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
|
||||
input_tensor.dtype))
|
||||
latents_std = 1.0 / torch.tensor(model1.config.latents_std).view(
|
||||
1, model1.config.z_dim, 1, 1, 1).to(latent_tensor.device,
|
||||
latent_tensor.dtype)
|
||||
latent_tensor = latent_tensor / latents_std + latents_mean
|
||||
output2 = model2.decode(latent_tensor)
|
||||
output1 = model1.decode(latent_tensor).sample
|
||||
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
|
||||
input_tensor.dtype)
|
||||
latent1_tensor = latent1_tensor / latents_std + latents_mean
|
||||
output1 = model1.decode(latent1_tensor).sample
|
||||
|
||||
latent2_tensor = latent2.mode()
|
||||
latents_mean = (torch.tensor(model2.config.latents_mean).view(
|
||||
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
|
||||
input_tensor.dtype))
|
||||
latents_std = 1.0 / torch.tensor(model2.config.latents_std).view(
|
||||
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
|
||||
input_tensor.dtype)
|
||||
latent2_tensor = latent2_tensor / latents_std + latents_mean
|
||||
output2 = model2.decode(latent2_tensor)
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
@@ -95,4 +105,4 @@ def test_wan_vae():
|
||||
max_diff_decode.item())
|
||||
logger.info("Mean difference between decoded outputs: %s",
|
||||
mean_diff_decode.item())
|
||||
assert mean_diff_decode < 1e-1, f"Decoded outputs differ significantly: mean diff = {mean_diff_decode.item()}"
|
||||
assert max_diff_decode < 1e-5, f"Decoded outputs differ significantly: max diff = {mean_diff_decode.item()}"
|
||||
Reference in New Issue
Block a user