This commit is contained in:
kijai
2025-12-04 16:44:29 +02:00
3 changed files with 25 additions and 14 deletions
+13 -6
View File
@@ -75,14 +75,21 @@ class VAE(nn.Module):
super().__init__()
if data_dim == 80:
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_80D, dtype=torch.float32))
self.data_std = nn.Buffer(torch.tensor(DATA_STD_80D, dtype=torch.float32))
data_mean = torch.tensor(DATA_MEAN_80D, dtype=torch.float32)
data_std = torch.tensor(DATA_STD_80D, dtype=torch.float32)
elif data_dim == 128:
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_128D, dtype=torch.float32))
self.data_std = nn.Buffer(torch.tensor(DATA_STD_128D, dtype=torch.float32))
data_mean = torch.tensor(DATA_MEAN_128D, dtype=torch.float32)
data_std = torch.tensor(DATA_STD_128D, dtype=torch.float32)
else:
raise ValueError(f"Unsupported data_dim={data_dim}, expected 80 or 128")
self.data_mean = self.data_mean.view(1, -1, 1)
self.data_std = self.data_std.view(1, -1, 1)
# match old shape: (1, channels, 1)
data_mean = data_mean.view(1, -1, 1)
data_std = data_std.view(1, -1, 1)
# register as buffers so they move with .to(device) / .cuda()
self.register_buffer("data_mean", data_mean)
self.register_buffer("data_std", data_std)
self.encoder = Encoder1D(
dim=hidden_dim,
-2
View File
@@ -2204,8 +2204,6 @@ class WanVideoDecode:
images = images.permute(1, 2, 3, 0).cpu().float()
return (images,)
else:
if end_image is not None:
enable_vae_tiling = False
images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0]
+12 -6
View File
@@ -1231,7 +1231,7 @@ class WanVideoVAE(nn.Module):
return mask
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, pbar=True):
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, end_=False, pbar=True):
_, _, T, H, W = hidden_states.shape
size_h, size_w = tile_size
stride_h, stride_w = tile_stride
@@ -1248,14 +1248,20 @@ class WanVideoVAE(nn.Module):
data_device = "cpu"
computation_device = device
out_T = T * 4 - 3
weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
weight, values = None, None
if pbar:
pbar = ProgressBar(len(tasks))
for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
if end_:
hidden_states_batch = self.model.decode_2(hidden_states_batch).to(data_device)
else:
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
if weight is None:
weight = torch.zeros((1, 1, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
if values is None:
values = torch.zeros((1, 3, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
mask = self.build_mask(
hidden_states_batch,
@@ -1399,7 +1405,7 @@ class WanVideoVAE(nn.Module):
for hidden_state in hidden_states:
hidden_state = hidden_state.unsqueeze(0)
if tiled:
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, pbar=pbar)
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, end_=end_, pbar=pbar)
else:
if end_:
video = self.double_decode(hidden_state, device)