diff --git a/HuMo/nodes.py b/HuMo/nodes.py index 9659d63..3a98ba7 100644 --- a/HuMo/nodes.py +++ b/HuMo/nodes.py @@ -219,7 +219,7 @@ class HuMoEmbeds: vae.to(device) zero_frames = torch.zeros(1, 3, pixel_frame_num + 4*num_refs, height, width, device=device, dtype=vae.dtype) zero_latents = vae.encode(zero_frames, device=device, tiled=tiled_vae)[0].to(offload_device) - vae.model.clear_cache() + vae.to(offload_device) mm.soft_empty_cache() diff --git a/nodes.py b/nodes.py index ddd953f..4688976 100644 --- a/nodes.py +++ b/nodes.py @@ -1018,7 +1018,6 @@ class WanVideoImageToVideoEncode: vae.to(device) y = vae.encode([concatenated], device, end_=(end_image is not None and not fun_or_fl2v_model),tiled=tiled_vae)[0] - vae.model.clear_cache() del concatenated has_ref = False @@ -1141,7 +1140,7 @@ class WanVideoAnimateEmbeds: resized_pose_images = resized_pose_images * 2 - 1 pose_latents = vae.encode([resized_pose_images.to(device, vae.dtype)], device,tiled=tiled_vae) pose_latents = pose_latents.to(offload_device) - vae.model.clear_cache() + if not looping and pose_latents.shape[2] < latent_window_size: log.info(f"WanAnimate: Padding pose latents from {pose_latents.shape} to length {latent_window_size}") pad_len = latent_window_size - pose_latents.shape[2] @@ -1156,15 +1155,15 @@ class WanVideoAnimateEmbeds: resized_bg_images = common_upscale(bg_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1) else: resized_bg_images = bg_images.permute(3, 0, 1, 2) # C, T, H, W - resized_bg_images = resized_bg_images[:3] * 2 - 1 - if not looping: - bg_latents = vae.encode([resized_bg_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0] - bg_latents = bg_latents.to(offload_device) - vae.model.clear_cache() - print("bg_latents", bg_latents.shape) - del resized_bg_images - else: - resized_bg_images = resized_bg_images.to(offload_device, dtype=vae.dtype) + resized_bg_images = (resized_bg_images[:3] * 2 - 1) + + if not looping: + if bg_images is None: + resized_bg_images = torch.zeros(3, num_frames - num_refs, H, W, device=device, dtype=vae.dtype) + bg_latents = vae.encode([resized_bg_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0].to(offload_device) + del resized_bg_images + elif bg_images is not None: + resized_bg_images = resized_bg_images.to(offload_device, dtype=vae.dtype) if ref_images is not None: if ref_images.shape[1] != H or ref_images.shape[2] != W: @@ -1173,46 +1172,31 @@ class WanVideoAnimateEmbeds: resized_ref_images = ref_images.permute(3, 0, 1, 2) # C, T, H, W resized_ref_images = resized_ref_images[:3] * 2 - 1 - if looping or bg_images is not None: # looping or when using background, encode refs separately - ref_latent = vae.encode([resized_ref_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0] - msk = torch.zeros(4, 1, lat_h, lat_w, device=device, dtype=vae.dtype) - msk[:, :1] = 1 - ref_latent_masked = torch.cat([msk, ref_latent], dim=0) # 4+C 1 H W - ref_latent_masked = ref_latent_masked.to(offload_device) - - if bg_images is None: - zero_frames = torch.zeros(3, num_frames - num_refs, H, W, device=device, dtype=vae.dtype) - concatenated = torch.cat([resized_ref_images.to(device, dtype=vae.dtype), zero_frames], dim=1) - del zero_frames - ref_latent = vae.encode([concatenated.to(device, vae.dtype)], device,tiled=tiled_vae)[0] - ref_latent = ref_latent.to(offload_device) - del concatenated - - vae.model.clear_cache() + ref_latent = vae.encode([resized_ref_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0] + msk = torch.zeros(4, 1, lat_h, lat_w, device=device, dtype=vae.dtype) + msk[:, :num_refs] = 1 + ref_latent_masked = torch.cat([msk, ref_latent], dim=0).to(offload_device) # 4+C 1 H W if mask is None: - ref_mask = torch.zeros(1, num_frames, lat_h, lat_w, device=offload_device, dtype=vae.dtype) + bg_mask = torch.zeros(1, num_frames, lat_h, lat_w, device=offload_device, dtype=vae.dtype) else: - ref_mask = 1 - mask[:num_frames] - if ref_mask.shape[0] < num_frames and not looping: - ref_mask = torch.cat([ref_mask, ref_mask[-1:].repeat(num_frames - ref_mask.shape[0], 1, 1)], dim=0) - ref_mask = common_upscale(ref_mask.unsqueeze(1), lat_w, lat_h, "nearest", "disabled").squeeze(1) - ref_mask = ref_mask.to(vae.dtype).to(offload_device) - ref_mask = ref_mask.unsqueeze(-1).permute(3, 0, 1, 2) # C, T, H, W + bg_mask = 1 - mask[:num_frames] + if bg_mask.shape[0] < num_frames and not looping: + bg_mask = torch.cat([bg_mask, bg_mask[-1:].repeat(num_frames - bg_mask.shape[0], 1, 1)], dim=0) + bg_mask = common_upscale(bg_mask.unsqueeze(1), lat_w, lat_h, "nearest", "disabled").squeeze(1) + bg_mask = bg_mask.unsqueeze(-1).permute(3, 0, 1, 2).to(offload_device, vae.dtype) # C, T, H, W - if bg_images is None: - ref_mask[:, :num_refs] = 1 - ref_mask_mask_repeated = torch.repeat_interleave(ref_mask[:, 0:1], repeats=4, dim=1) # T, C, H, W - ref_mask = torch.cat([ref_mask_mask_repeated, ref_mask[:, 1:]], dim=1) - ref_mask = ref_mask.view(1, ref_mask.shape[1] // 4, 4, lat_h, lat_w) # 1, T, C, H, W - ref_mask = ref_mask.movedim(1, 2)[0]# C, T, H, W + if bg_images is None and looping: + bg_mask[:, :num_refs] = 1 + bg_mask_mask_repeated = torch.repeat_interleave(bg_mask[:, 0:1], repeats=4, dim=1) # T, C, H, W + bg_mask = torch.cat([bg_mask_mask_repeated, bg_mask[:, 1:]], dim=1) + bg_mask = bg_mask.view(1, bg_mask.shape[1] // 4, 4, lat_h, lat_w) # 1, T, C, H, W + bg_mask = bg_mask.movedim(1, 2)[0]# C, T, H, W if not looping: - if bg_images is not None: - bg_latents_masked = torch.cat([ref_mask[:, :bg_latents.shape[1]], bg_latents], dim=0) - ref_latent = torch.cat([ref_latent_masked, bg_latents_masked], dim=1) - else: - ref_latent = torch.cat([ref_mask, ref_latent], dim=0) + bg_latents_masked = torch.cat([bg_mask[:, :bg_latents.shape[1]], bg_latents], dim=0) + del bg_mask, bg_latents + ref_latent = torch.cat([ref_latent_masked, bg_latents_masked], dim=1) else: ref_latent = ref_latent_masked @@ -1225,7 +1209,6 @@ class WanVideoAnimateEmbeds: resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0) resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype) - vae.model.clear_cache() seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1]) @@ -1240,7 +1223,7 @@ class WanVideoAnimateEmbeds: "max_seq_len": seq_len, "pose_latents": pose_latents, "bg_images": resized_bg_images if bg_images is not None and looping else None, - "ref_masks": ref_mask if mask is not None and looping else None, + "ref_masks": bg_mask if mask is not None and looping else None, "ref_latent": ref_latent, "ref_image": resized_ref_images if ref_images is not None else None, "face_pixels": resized_face_images if face_images is not None else None, @@ -1625,7 +1608,7 @@ class WanVideoVACEEncode: vae = vae.to(device) z0 = self.vace_encode_frames(vae, input_frames, ref_images, masks=input_masks, tiled_vae=tiled_vae) - vae.model.clear_cache() + m0 = self.vace_encode_masks(input_masks, ref_images) z = self.vace_latent(z0, m0) vae.to(offload_device) @@ -1666,7 +1649,7 @@ class WanVideoVACEEncode: reactive = vae.encode(reactive, device=device, tiled=tiled_vae) latents = [torch.cat((u, c), dim=0) for u, c in zip(inactive, reactive)] del inactive, reactive - vae.model.clear_cache() + cat_latents = [] for latent, refs in zip(latents, ref_images): @@ -3858,7 +3841,7 @@ class WanVideoSampler: else: cond_ = cond_image if is_first_clip else cond_frame latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0] - vae.model.clear_cache() + vae.to(offload_device) #motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1 @@ -3957,7 +3940,7 @@ class WanVideoSampler: padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype), device=device, tiled=tiled_vae ).to(dtype) - vae.model.clear_cache() + vae.to(offload_device) uni3c_data['render_latent'] = render_latent @@ -4046,7 +4029,7 @@ class WanVideoSampler: latent = latent[:,:-humo_reference_count] vae.to(device) videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu() - vae.model.clear_cache() + vae.to(offload_device) sampling_pbar.close() @@ -4190,13 +4173,13 @@ class WanVideoSampler: log.info(f"Sampling {total_frames} frames in {s2v_num_repeat} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps") # sample for r in range(s2v_num_repeat): - vae.model.clear_cache() + mm.soft_empty_cache() gc.collect() if ref_motion_image is not None: vae.to(device) ref_motion = vae.encode(ref_motion_image.to(vae.dtype), device=device, pbar=False).to(dtype)[0] - vae.model.clear_cache() + vae.to(offload_device) left_idx = r * infer_frames @@ -4265,7 +4248,7 @@ class WanVideoSampler: ref_motion_image = videos_last_frames vae.to(offload_device) - vae.model.clear_cache() + mm.soft_empty_cache() gen_video_samples = torch.cat(framepack_out, dim=2).squeeze(0).permute(1, 2, 3, 0) @@ -4300,7 +4283,7 @@ class WanVideoSampler: ref_masks = image_embeds.get("ref_masks", None) bg_images = image_embeds.get("bg_images", None) - pose_input_latents = current_ref_images = face_images = None + current_ref_images = face_images = None if wananim_face_pixels is not None: face_images = tensor_pingpong_pad(wananim_face_pixels, target_len) @@ -4365,7 +4348,7 @@ class WanVideoSampler: else: temporal_ref_latents = temporal_ref_latents[:, :msk.shape[1]] - vae.model.clear_cache() + vae.to(offload_device) temporal_ref_latents = torch.cat([msk, temporal_ref_latents], dim=0) # 4+C T H W @@ -4501,7 +4484,7 @@ class WanVideoSampler: vae.to(device) videos = vae.decode(latent[:, 1:].unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu() del latent - vae.model.clear_cache() + sampling_pbar.close() @@ -4778,7 +4761,7 @@ class WanVideoDecode: 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] - vae.model.clear_cache() + images = images.cpu().float() @@ -4798,7 +4781,7 @@ class WanVideoDecode: if end_image is not None: images = images[:, 0:-1] - vae.model.clear_cache() + vae.to(offload_device) mm.soft_empty_cache() @@ -4849,7 +4832,7 @@ class WanVideoEncodeLatentBatch: latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor)) else: latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling) - vae.model.clear_cache() + if latent_strength != 1.0: latent *= latent_strength latent_list.append(latent.squeeze(0).cpu()) @@ -4909,7 +4892,7 @@ class WanVideoEncode: latents = latents.permute(0, 2, 1, 3, 4) else: latents = vae.encode(image * 2.0 - 1.0, device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor)) - vae.model.clear_cache() + vae.to(offload_device) if latent_strength != 1.0: latents *= latent_strength diff --git a/nodes_deprecated.py b/nodes_deprecated.py index cfad1ef..df8769d 100644 --- a/nodes_deprecated.py +++ b/nodes_deprecated.py @@ -124,7 +124,6 @@ class WanVideoImageClipEncode: y = torch.concat([mask, y]) - vae.model.clear_cache() vae.to(offload_device) image_embeds = { diff --git a/wanvideo/wan_video_vae.py b/wanvideo/wan_video_vae.py index 2f55bc4..a4df70c 100644 --- a/wanvideo/wan_video_vae.py +++ b/wanvideo/wan_video_vae.py @@ -17,20 +17,6 @@ def check_is_instance(model, module_class): return False -def block_causal_mask(x, block_size): - # params - b, n, s, _, device = *x.size(), x.device - assert s % block_size == 0 - num_blocks = s // block_size - - # build mask - mask = torch.zeros(b, n, s, s, dtype=torch.bool, device=device) - for i in range(num_blocks): - mask[:, :, - i * block_size:(i + 1) * block_size, :(i + 1) * block_size] = 1 - return mask - - class CausalConv3d(nn.Conv3d): """ Causal 3d convolusion. @@ -329,12 +315,7 @@ class AttentionBlock(nn.Module): 0, 1, 3, 2).contiguous().chunk(3, dim=-1) # apply attention - x = F.scaled_dot_product_attention( - q, - k, - v, - #attn_mask=block_causal_mask(q, block_size=h * w) - ) + x = F.scaled_dot_product_attention(q, k, v) x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w) # output @@ -860,7 +841,7 @@ class Decoder3d_38(nn.Module): # dimensions dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]] - scale = 1.0 / 2 ** (len(dim_mult) - 2) + # init block self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1) @@ -989,8 +970,6 @@ class VideoVAE_(nn.Module): #modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames def encode_2(self, x): - self.clear_cache() - ## cache t = x.shape[2] iter_ = 2 + (t - 2) // 4 @@ -1010,18 +989,16 @@ class VideoVAE_(nn.Module): feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx) out = torch.cat([out, out_], 2) + self.clear_cache() out_head = out[:, :, :iter_ - 1, :, :] out_tail = out[:, :, -1, :, :].unsqueeze(2) mu = torch.cat([self.conv1(out_head), self.conv1(out_tail)], dim=2).chunk(2, dim=1)[0] - mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu) return mu def encode(self, x, pbar=True): - self.clear_cache() - ## cache t = x.shape[2] iter_ = 1 + (t - 1) // 4 if pbar: @@ -1040,18 +1017,17 @@ class VideoVAE_(nn.Module): out = torch.cat([out, out_], 2) if pbar: pbar.update(1) + self.clear_cache() mu = self.conv1(out).chunk(2, dim=1)[0] - mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu) if pbar: pbar.update_absolute(0) - + return mu #modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames def decode_2(self, z): - self.clear_cache() # z: [b,c,t,h,w] z = z / self.inv_std.to(z) + self.mean.to(z) @@ -1076,12 +1052,12 @@ class VideoVAE_(nn.Module): feat_cache=self._feat_map, feat_idx=self._conv_idx) out = torch.cat([out, out_], 2) # may add tensor offload + self.clear_cache() return out def decode(self, z, pbar=True): - self.clear_cache() # z: [b,c,t,h,w] z = z / self.inv_std.to(z) + self.mean.to(z) iter_ = z.shape[2] @@ -1104,6 +1080,7 @@ class VideoVAE_(nn.Module): pbar.update(1) if pbar: pbar.update_absolute(0) + self.clear_cache() return out def reparameterize(self, mu, log_var): @@ -1318,6 +1295,7 @@ class WanVideoVAE(nn.Module): return video def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None, pbar=True): + self.model.clear_cache() videos = [video.to("cpu") for video in videos] hidden_states = [] for video in videos: @@ -1336,6 +1314,7 @@ class WanVideoVAE(nn.Module): def decode(self, hidden_states, device, tiled=False, end_=False, tile_size=(34, 34), tile_stride=(18, 16), pbar=True): + self.model.clear_cache() hidden_states = [hidden_state.to("cpu") for hidden_state in hidden_states] videos = [] for hidden_state in hidden_states: @@ -1352,25 +1331,6 @@ class WanVideoVAE(nn.Module): return videos - @staticmethod - def state_dict_converter(): - return WanVideoVAEStateDictConverter() - - -class WanVideoVAEStateDictConverter: - - def __init__(self): - pass - - def from_civitai(self, state_dict): - state_dict_ = {} - if 'model_state' in state_dict: - state_dict = state_dict['model_state'] - for name in state_dict: - state_dict_['model.' + name] = state_dict[name] - return state_dict_ - - class VideoVAE38_(VideoVAE_): def __init__(self, @@ -1427,9 +1387,7 @@ class VideoVAE38_(VideoVAE_): if pbar: pbar.update(1) mu = self.conv1(out).chunk(2, dim=1)[0] - mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu) - self.clear_cache() return mu