Add pose input

This commit is contained in:
kijai
2025-08-27 22:43:28 +03:00
parent 5266959a93
commit a5621b8739
4 changed files with 126 additions and 38 deletions
+87 -9
View File
@@ -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(
+1 -1
View File
@@ -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
+28 -25
View File
@@ -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)
+10 -3
View File
@@ -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: