Add pose input
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user