From a5621b87391013155b4f688fbe01dba10a8104aa Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 22:43:28 +0300 Subject: [PATCH] Add pose input --- nodes.py | 96 +++++++++++++++++++++++++++++++++++---- nodes_model_loading.py | 2 +- s2v/nodes.py | 53 +++++++++++---------- wanvideo/modules/model.py | 13 ++++-- 4 files changed, 126 insertions(+), 38 deletions(-) diff --git a/nodes.py b/nodes.py index dd2810b..50a570e 100644 --- a/nodes.py +++ b/nodes.py @@ -2228,17 +2228,23 @@ class WanVideoSampler: s2v_audio_embeds = image_embeds.get("audio_embeds", None) if s2v_audio_embeds is not None: log.info(f"Using S2V audio embeddings") - s2v_audio_input = s2v_audio_embeds["audio_embed_bucket"].to(device, dtype) + s2v_audio_input = s2v_audio_embeds.get("audio_embed_bucket", None) + if s2v_audio_input is not None: + s2v_audio_input = s2v_audio_input[..., 0:image_embeds["num_frames"]].to(device, dtype) s2v_audio_scale = s2v_audio_embeds["audio_scale"] - s2v_ref_latent = s2v_audio_embeds["ref_latent"].to(device, dtype) if "ref_latent" in s2v_audio_embeds else None - s2v_ref_motion = s2v_audio_embeds["ref_motion"].to(device, dtype) if "ref_motion" in s2v_audio_embeds else None - s2v_audio_input = s2v_audio_input[..., 0:image_embeds["num_frames"]] + s2v_ref_latent = s2v_audio_embeds.get("ref_latent", None) + if s2v_ref_latent is not None: + s2v_ref_latent = s2v_ref_latent.to(device, dtype) + s2v_ref_motion = s2v_audio_embeds.get("ref_motion", None) + if s2v_ref_motion is not None: + s2v_ref_motion = s2v_ref_motion.to(device, dtype) + s2v_pose = s2v_audio_embeds.get("pose_latent", None) + if s2v_pose is not None: + s2v_pose = s2v_pose.to(device, dtype) + s2v_num_repeat = s2v_audio_embeds.get("num_repeat", 1) vae = image_embeds.get("vae", None) framepack = False - #s2v_audio_input_all_layers = s2v_audio_embeds["audio_encoder_output"]["encoded_audio_all_layers"] - print(s2v_audio_input.shape) - ##print(s2v_audio_input_all_layers[0].shape) # vid2vid noise_mask=original_image=None @@ -2700,7 +2706,8 @@ class WanVideoSampler: "s2v_audio_input": s2v_audio_input, # official speech-to-video audio input "s2v_ref_latent": s2v_ref_latent, # speech-to-video reference latent "s2v_ref_motion": s2v_ref_motion, # speech-to-video reference motion latent - "s2v_audio_scale": s2v_audio_scale if s2v_audio_input is not None else 1.0 # speech-to-video audio scale + "s2v_audio_scale": s2v_audio_scale if s2v_audio_input is not None else 1.0, # speech-to-video audio scale + "s2v_pose": s2v_pose if s2v_pose is not None else None # speech-to-video pose control } batch_size = 1 @@ -3664,7 +3671,78 @@ class WanVideoSampler: except: pass return {"video": gen_video_samples.permute(1, 2, 3, 0)}, - + elif framepack: + framepack_out = [] + ref_motion_image = None + motion_frames = 5 + infer_frames = image_embeds["num_frames"] + + for r in range(s2v_num_repeat): + if ref_motion_image is not None: + if ref_motion_image.shape[0] > 73: + ref_motion_image = ref_motion_image[-73:] + + if ref_motion_image.shape[0] < 73: + ref = torch.ones([73, ref_motion_image.shape[1], ref_motion_image.shape[2], 3]) * 0.5 + ref[-ref_motion_image.shape[0]:] = ref_motion_image + ref_motion_image = ref + + vae.to(device) + ref_motion = vae.encode(ref_motion_image[:, :, :, :3], device=device, pbar=False)[0].to(dtype) + vae.to(offload_device) + + left_idx = r * infer_frames + right_idx = r * infer_frames + infer_frames + #cond_latents = COND[r] if pose_video else COND[0] * 0 + #cond_latents = cond_latents.to(dtype=self.param_dtype, device=self.device) + s2v_audio_input = s2v_audio_embeds[..., left_idx:right_idx] + input_motion_latents = ref_motion.clone() + + noise_pred, self.cache_state = predict_with_cfg( + latent_model_input, + cfg[idx], + text_embeds["prompt_embeds"], + text_embeds["negative_prompt_embeds"], + 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, mtv_motion_tokens=mtv_motion_tokens, + s2v_audio_input=s2v_audio_input, s2v_ref_motion=input_motion_latents) + + latent = sample_scheduler.step( + noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0), + **scheduler_step_args)[0].squeeze(0) + + latents = torch.stack(latent) + #if not (drop_first_motion and r == 0): + # decode_latents = torch.cat([motion_latents, latents], dim=2) + #else: + decode_latents = torch.cat([s2v_ref_latent, latents], dim=2) + image = torch.stack(vae.decode(decode_latents), device=device) + image = image[:, :, -(infer_frames):] + #if (drop_first_motion and r == 0): + # image = image[:, :, 3:] + + overlap_frames_num = min(motion_frames, image.shape[2]) + videos_last_frames = torch.cat([ + videos_last_frames[:, :, overlap_frames_num:], + image[:, :, -overlap_frames_num:]], dim=2).to(vae.device, vae.dtype) + + vae.to(device) + ref_motion_image = torch.stack(vae.encode(videos_last_frames, device=device, pbar=False)[0]) + vae.to(device) + framepack_out.append(image.cpu()) + + gen_video_samples = torch.cat(framepack_out, dim=1) + + if force_offload: + if not model["auto_cpu_offload"]: + offload_transformer(transformer) + try: + print_memory(device) + torch.cuda.reset_peak_memory_stats(device) + except: + pass + return {"video": gen_video_samples.permute(1, 2, 3, 0)}, + #region normal inference else: noise_pred, self.cache_state = predict_with_cfg( diff --git a/nodes_model_loading.py b/nodes_model_loading.py index f053a8c..c5f9002 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -731,7 +731,7 @@ class WanVideoSetLoRAs: def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None): - params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio"} + params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio", "cond_encoder"} param_count = sum(1 for _ in transformer.named_parameters()) pbar = ProgressBar(param_count) cnt = 0 diff --git a/s2v/nodes.py b/s2v/nodes.py index 9bd8ab1..f854e16 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -51,12 +51,13 @@ class WanVideoAddAudioEmbeds: def INPUT_TYPES(s): return {"required": { "embeds": ("WANVIDIMAGE_EMBEDS",), - "audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",), "frames": ("INT", {"default": 81, "min": 1, "max": 100000, "step": 1, "tooltip": "Number of frames to process"}), "audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for audio embeddings"}) }, "optional": { - "ref_latent": ("LATENT",) + "audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",), + "ref_latent": ("LATENT",), + "pose_latent": ("LATENT",) } } @@ -66,38 +67,40 @@ class WanVideoAddAudioEmbeds: FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, embeds, frames, audio_encoder_output, audio_scale, ref_latent=None): - all_layers = audio_encoder_output["encoded_audio_all_layers"] - audio_feat = torch.stack(all_layers, dim=0).squeeze(1) # shape: [num_layers, T, 512] + def add(self, embeds, frames, audio_encoder_output=None, audio_scale=1.0, ref_latent=None, pose_latent=None): + if audio_encoder_output is not None: + all_layers = audio_encoder_output["encoded_audio_all_layers"] + audio_feat = torch.stack(all_layers, dim=0).squeeze(1) # shape: [num_layers, T, 512] - print("audio_feat", audio_feat.shape) - input_fps = 50 - output_fps = 30 - bucket_fps = 16 + print("audio_feat", audio_feat.shape) + input_fps = 50 + output_fps = 30 + bucket_fps = 16 - if input_fps != output_fps: - audio_feat = linear_interpolation(audio_feat, input_fps=input_fps, output_fps=output_fps) + if input_fps != output_fps: + audio_feat = linear_interpolation(audio_feat, input_fps=input_fps, output_fps=output_fps) - self.video_rate = output_fps + self.video_rate = output_fps - audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps( - audio_feat, - fps=bucket_fps, - batch_frames=frames-1 - ) + audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps( + audio_feat, + fps=bucket_fps, + batch_frames=frames-1 + ) - audio_embed_bucket = audio_embed_bucket.unsqueeze(0) - if len(audio_embed_bucket.shape) == 3: - audio_embed_bucket = audio_embed_bucket.permute(0, 2, 1) - elif len(audio_embed_bucket.shape) == 4: - audio_embed_bucket = audio_embed_bucket.permute(0, 2, 3, 1) + audio_embed_bucket = audio_embed_bucket.unsqueeze(0) + if len(audio_embed_bucket.shape) == 3: + audio_embed_bucket = audio_embed_bucket.permute(0, 2, 1) + elif len(audio_embed_bucket.shape) == 4: + audio_embed_bucket = audio_embed_bucket.permute(0, 2, 3, 1) - print("audio_embed_bucket", audio_embed_bucket.shape) + print("audio_embed_bucket", audio_embed_bucket.shape) new_entry = { - "audio_embed_bucket": audio_embed_bucket, - "num_repeat": num_repeat, + "audio_embed_bucket": audio_embed_bucket if audio_encoder_output is not None else None, + "num_repeat": num_repeat if audio_encoder_output is not None else None, "ref_latent": ref_latent["samples"] if ref_latent is not None else None, + "pose_latent": pose_latent["samples"] if pose_latent is not None else None, "audio_scale": audio_scale } updated = dict(embeds) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 21f2a91..ce9cc2f 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2138,7 +2138,8 @@ class WanModel(torch.nn.Module): s2v_audio_input=None, s2v_ref_latent=None, s2v_audio_scale=1.0, - s2v_ref_motion=None + s2v_ref_motion=None, + s2v_pose=None ): r""" @@ -2243,6 +2244,12 @@ class WanModel(torch.nn.Module): for u in x ] + if s2v_pose is not None: + print("s2v_pose.shape:", s2v_pose.shape) + print("x[0].shape:", x[0].shape) + x[0] = x[0] + self.cond_encoder(s2v_pose.to(self.cond_encoder.weight.dtype)).to(x[0].dtype) + + if self.control_adapter is not None and fun_camera is not None: fun_camera = self.control_adapter(fun_camera) x = [u + v for u, v in zip(x, fun_camera)] @@ -2371,8 +2378,8 @@ class WanModel(torch.nn.Module): x = torch.cat([x, motion_encoded], dim=1) freqs = torch.cat([freqs, freqs_motion], dim=1) - t = torch.repeat_interleave(t, 2, dim=1) - t = torch.cat([t, torch.zeros((t.shape[0], 3), device=t.device, dtype=t.dtype)], dim=1) + #t = torch.repeat_interleave(t, 2, dim=1) + #t = torch.cat([t, torch.zeros((t.shape[0], 3), device=t.device, dtype=t.dtype)], dim=1) # time embeddings if t.dim() == 2: