diff --git a/nodes.py b/nodes.py index c18dc75..923656c 100644 --- a/nodes.py +++ b/nodes.py @@ -2661,13 +2661,13 @@ class WanVideoSampler: controlnet["controlnet_stride"] = controlnet["control_stride"] #uni3c - pcd_data = pcd_data_input = None + uni3c_data = uni3c_data_input = None if uni3c_embeds is not None: transformer.controlnet = uni3c_embeds["controlnet"] render_latent = uni3c_embeds["render_latent"].to(device) if render_latent.shape != noise.shape: render_latent = torch.nn.functional.interpolate(render_latent, size=(noise.shape[1], noise.shape[2], noise.shape[3]), mode='trilinear', align_corners=False) - pcd_data = { + uni3c_data = { "render_latent": render_latent, "render_mask": uni3c_embeds["render_mask"], "camera_embedding": uni3c_embeds["camera_embedding"], @@ -2852,7 +2852,7 @@ class WanVideoSampler: add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False, mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None, s2v_motion_frames=[1, 0], s2v_pose=None, humo_image_cond=None, humo_image_cond_neg=None, humo_audio=None, humo_audio_neg=None, wananim_pose_latents=None, - wananim_face_pixels=None): + wananim_face_pixels=None, uni3c_data=None,): nonlocal transformer z = z.to(dtype) autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear) @@ -2999,13 +2999,13 @@ class WanVideoSampler: elif multitalk_sampling and multitalk_audio_embeds is not None: multitalk_audio_input = multitalk_audio_embeds - if context_window is not None and pcd_data is not None and pcd_data["render_latent"].shape[2] != context_frames: - pcd_data_input = {"render_latent": pcd_data["render_latent"][:, :, context_window]} - for k in pcd_data: + if context_window is not None and uni3c_data is not None and uni3c_data["render_latent"].shape[2] != context_frames: + uni3c_data_input = {"render_latent": uni3c_data["render_latent"][:, :, context_window]} + for k in uni3c_data: if k != "render_latent": - pcd_data_input[k] = pcd_data[k] + uni3c_data_input[k] = uni3c_data[k] else: - pcd_data_input = pcd_data + uni3c_data_input = uni3c_data if s2v_pose is not None: if not ((s2v_pose_start_percent <= current_step_percentage <= s2v_pose_end_percent) or \ @@ -3046,7 +3046,7 @@ class WanVideoSampler: 'fun_camera': control_camera_input if control_camera_latents is not None else None, # Fun model camera embed 'audio_proj': audio_proj if fantasytalking_embeds is not None else None, # FantasyTalking audio projection 'audio_scale': audio_scale, # FantasyTalking audio scale - "pcd_data": pcd_data_input, # Uni3C input + "uni3c_data": uni3c_data_input, # Uni3C input "controlnet": controlnet, # TheDenk's controlnet input "add_cond": add_cond_input, # additional conditioning input "nag_params": text_embeds.get("nag_params", {}), # normalized attention guidance @@ -3725,10 +3725,10 @@ class WanVideoSampler: audio_embs = None cond_frame = None - pcd_data = pcd_data_input = None + uni3c_data = uni3c_data_input = None if uni3c_embeds is not None: transformer.controlnet = uni3c_embeds["controlnet"] - pcd_data = { + uni3c_data = { "render_latent": uni3c_embeds["render_latent"], "render_mask": uni3c_embeds["render_mask"], "camera_embedding": uni3c_embeds["camera_embedding"], @@ -3947,7 +3947,7 @@ class WanVideoSampler: ).to(dtype) vae.model.clear_cache() vae.to(offload_device) - pcd_data['render_latent'] = render_latent + uni3c_data['render_latent'] = render_latent # unianimate slices partial_unianim_data = None @@ -3989,7 +3989,8 @@ class WanVideoSampler: latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"], timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond, cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input, - humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg) + humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg, + uni3c_data = uni3c_data) if callback is not None: callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3) @@ -4442,24 +4443,15 @@ class WanVideoSampler: positive = text_embeds["prompt_embeds"] # uni3c slices + uni3c_data_input = None if uni3c_embeds is not None: - vae.to(device) - # Pad original_images if needed - num_frames = original_images.shape[2] - required_frames = audio_end_idx - audio_start_idx - if audio_end_idx > num_frames: - pad_len = audio_end_idx - num_frames - last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1) - padded_images = torch.cat([original_images, last_frame], dim=2) - else: - padded_images = original_images - render_latent = vae.encode( - 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) - pcd_data['render_latent'] = render_latent + render_latent = uni3c_embeds["render_latent"][:,:,start_latent:end_latent].to(device) + if render_latent.shape[2] < noise.shape[1]: + render_latent = torch.nn.functional.interpolate(render_latent, size=(noise.shape[1], noise.shape[2], noise.shape[3]), mode='trilinear', align_corners=False) + uni3c_data_input = {"render_latent": render_latent} + for k in uni3c_data: + if k != "render_latent": + uni3c_data_input[k] = uni3c_data[k] mm.soft_empty_cache() gc.collect() @@ -4474,7 +4466,7 @@ class WanVideoSampler: timestep, i, cache_state=self.cache_state, image_cond = image_cond_in, wananim_face_pixels=face_images[:, :, start:end].to(device, torch.float32) if face_images is not None else None, - wananim_pose_latents=pose_input_slice + wananim_pose_latents=pose_input_slice, uni3c_data = uni3c_data_input, ) if callback is not None: callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3) @@ -4565,7 +4557,7 @@ class WanVideoSampler: timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, multitalk_audio_embeds=multitalk_audio_embeds, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input, humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg, - wananim_face_pixels=wananim_face_pixels, wananim_pose_latents=wananim_pose_latents, + wananim_face_pixels=wananim_face_pixels, wananim_pose_latents=wananim_pose_latents, uni3c_data = uni3c_data, ) if bidirectional_sampling: noise_pred_flipped, self.cache_state = predict_with_cfg( diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 1150ddf..e8ff190 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1906,9 +1906,7 @@ class WanModel(torch.nn.Module): def wananimate_forward(self, block_idx, x, motion_vec, strength=1.0, motion_masks=None): if block_idx % 5 == 0: adapter_args = [x, motion_vec, motion_masks] - self.face_adapter.to(self.main_device) residual_out = self.face_adapter.fuser_blocks[block_idx // 5](*adapter_args) - self.face_adapter.to(self.offload_device) return x.add(residual_out, alpha=strength) return x @@ -1986,7 +1984,7 @@ class WanModel(torch.nn.Module): fun_camera=None, audio_proj=None, audio_scale=1.0, - pcd_data=None, + uni3c_data=None, controlnet=None, add_cond=None, attn_cond=None, @@ -2092,8 +2090,8 @@ class WanModel(torch.nn.Module): x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)] #uni3c controlnet - if pcd_data is not None: - render_latent = pcd_data["render_latent"].to(x[0].dtype) + if uni3c_data is not None: + render_latent = uni3c_data["render_latent"].to(x[0].dtype) hidden_states = x[0].unsqueeze(0).clone().float() if hidden_states.shape[1] == 16: #T2V work around hidden_states = torch.cat([hidden_states, torch.zeros_like(hidden_states[:, :4])], dim=1) @@ -2572,16 +2570,16 @@ class WanModel(torch.nn.Module): kwargs['vace_context_scale'] = vace_scale_list #uni3c controlnet - pdc_controlnet_states = None - if pcd_data is not None: - if (pcd_data["start"] <= current_step_percentage <= pcd_data["end"]) or \ - (pcd_data["end"] > 0 and current_step == 0 and current_step_percentage >= pcd_data["start"]): + uni3c_controlnet_states = None + if uni3c_data is not None: + if (uni3c_data["start"] <= current_step_percentage <= uni3c_data["end"]) or \ + (uni3c_data["end"] > 0 and current_step == 0 and current_step_percentage >= uni3c_data["start"]): self.controlnet.to(self.main_device) with torch.autocast(device_type=mm.get_autocast_device(device), dtype=x.dtype, enabled=True): - pdc_controlnet_states = self.controlnet( + uni3c_controlnet_states = self.controlnet( render_latent=render_latent.to(self.main_device, self.controlnet.dtype), - render_mask=pcd_data["render_mask"], - camera_embedding=pcd_data["camera_embedding"], + render_mask=uni3c_data["render_mask"], + camera_embedding=uni3c_data["camera_embedding"], temb=e.to(self.main_device), device=self.offload_device) self.controlnet.to(self.offload_device) @@ -2641,8 +2639,8 @@ class WanModel(torch.nn.Module): log.info(f"Block {b}: transfer_time={transfer_time:.4f}s, compute_time={compute_time:.4f}s, to_cpu_transfer_time={to_cpu_transfer_time:.4f}s") #uni3c controlnet - if pdc_controlnet_states is not None and b < len(pdc_controlnet_states): - x[:, :self.original_seq_len] += pdc_controlnet_states[b].to(x) * pcd_data["controlnet_weight"] + if uni3c_controlnet_states is not None and b < len(uni3c_controlnet_states): + x[:, :self.original_seq_len] += uni3c_controlnet_states[b].to(x) * uni3c_data["controlnet_weight"] #controlnet if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])): x[:, :self.original_seq_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]