From 6a37c0b2d6deefa801cb204ed79df1492aa6d31d Mon Sep 17 00:00:00 2001 From: wzxysf Date: Fri, 24 Oct 2025 22:20:56 +0800 Subject: [PATCH 1/2] Enable vae tiling with end frame --- nodes.py | 2 -- wanvideo/wan_video_vae.py | 18 ++++++++++++------ 2 files changed, 12 insertions(+), 8 deletions(-) diff --git a/nodes.py b/nodes.py index 9eb1d64..b332a1c 100644 --- a/nodes.py +++ b/nodes.py @@ -2023,8 +2023,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] diff --git a/wanvideo/wan_video_vae.py b/wanvideo/wan_video_vae.py index cb2fd58..97ab234 100644 --- a/wanvideo/wan_video_vae.py +++ b/wanvideo/wan_video_vae.py @@ -1193,7 +1193,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 @@ -1210,14 +1210,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, @@ -1361,7 +1367,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) From aa9f4749587c0f8a5041a56bcc4e4a07ca76c4f0 Mon Sep 17 00:00:00 2001 From: jamesjjcondon <33615047+jamesjjcondon@users.noreply.github.com> Date: Fri, 14 Nov 2025 11:57:26 +1030 Subject: [PATCH 2/2] Refactor data_mean and data_std initialization Refactor data_mean and data_std to use register_buffer and remove nn.Buffer. Tested with pytorch version: 2.4.0+cu121 xformers version: 0.0.27.post2 Set vram state to: LOW_VRAM Device: cuda:0 NVIDIA GeForce RTX 3090 : cudaMallocAsync Enabled pinned memory 122281.0 Using xformers attention Python version: 3.10.12 (main, Aug 15 2025, 14:32:43) [GCC 11.4.0] ComfyUI version: 0.3.68 ComfyUI frontend version: 1.28.8 --- Ovi/vae/vae.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/Ovi/vae/vae.py b/Ovi/vae/vae.py index d1b9254..0e0d207 100644 --- a/Ovi/vae/vae.py +++ b/Ovi/vae/vae.py @@ -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,