Merge branch 'main' into dev
This commit is contained in:
@@ -2,7 +2,6 @@ from diffusers import ModelMixin, ConfigMixin
|
||||
from einops import rearrange, repeat
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from functools import lru_cache
|
||||
from ..wanvideo.modules.attention import attention
|
||||
|
||||
from comfy import model_management as mm
|
||||
@@ -118,8 +117,6 @@ class RotaryPositionalEmbedding1D(nn.Module):
|
||||
self.head_dim = head_dim
|
||||
self.base = 10000
|
||||
|
||||
|
||||
#@lru_cache(maxsize=32)
|
||||
def precompute_freqs_cis_1d(self, pos_indices):
|
||||
|
||||
freqs = 1.0 / (self.base ** (torch.arange(0, self.head_dim, 2)[: (self.head_dim // 2)].float() / self.head_dim))
|
||||
|
||||
+7
-2
@@ -273,6 +273,10 @@ class WanVideoImageToVideoMultiTalk:
|
||||
"start_image": ("IMAGE", {"tooltip": "Image to encode"}),
|
||||
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
|
||||
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
|
||||
"mode": ([
|
||||
"multitalk",
|
||||
"infinitetalk"
|
||||
], {"default": "multitalk", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -281,7 +285,7 @@ class WanVideoImageToVideoMultiTalk:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, vae, width, height, frame_window_size, motion_frame, force_offload, colormatch, start_image=None, tiled_vae=False, clip_embeds=None):
|
||||
def process(self, vae, width, height, frame_window_size, motion_frame, force_offload, colormatch, start_image=None, tiled_vae=False, clip_embeds=None, mode="multitalk"):
|
||||
|
||||
H = height
|
||||
W = width
|
||||
@@ -311,7 +315,8 @@ class WanVideoImageToVideoMultiTalk:
|
||||
"vae": vae,
|
||||
"target_shape": target_shape,
|
||||
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
|
||||
"colormatch": colormatch
|
||||
"colormatch": colormatch,
|
||||
"multitalk_mode": mode
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
|
||||
@@ -3096,7 +3096,9 @@ class WanVideoSampler:
|
||||
noise_pred /= counter
|
||||
#region multitalk
|
||||
elif multitalk_sampling:
|
||||
original_image = cond_image = image_embeds.get("multitalk_start_image", None)
|
||||
mode = image_embeds.get("multitalk_mode", "multitalk")
|
||||
log.info(f"Multitalk mode: {mode}")
|
||||
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
|
||||
offload = image_embeds.get("force_offload", False)
|
||||
tiled_vae = image_embeds.get("tiled_vae", False)
|
||||
frame_num = clip_length = image_embeds.get("num_frames", 81)
|
||||
@@ -3107,6 +3109,25 @@ class WanVideoSampler:
|
||||
target_w = image_embeds.get("target_w", None)
|
||||
target_h = image_embeds.get("target_h", None)
|
||||
|
||||
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
|
||||
face_scale = 0.1
|
||||
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
|
||||
background_mask = torch.zeros([target_h, target_w])
|
||||
background_mask = torch.zeros([target_h, target_w])
|
||||
human_mask1 = torch.zeros([target_h, target_w])
|
||||
human_mask2 = torch.zeros([target_h, target_w])
|
||||
lefty_min, lefty_max = int((target_w//2) * face_scale), int((target_w//2) * (1 - face_scale))
|
||||
righty_min, righty_max = int((target_w//2) * face_scale + (target_w//2)), int((target_w//2) * (1 - face_scale) + (target_w//2))
|
||||
human_mask1[x_min:x_max, lefty_min:lefty_max] = 1
|
||||
human_mask2[x_min:x_max, righty_min:righty_max] = 1
|
||||
background_mask += human_mask1
|
||||
background_mask += human_mask2
|
||||
human_masks = [human_mask1, human_mask2]
|
||||
background_mask = torch.where(background_mask > 0, torch.tensor(0), torch.tensor(1))
|
||||
human_masks.append(background_mask)
|
||||
ref_target_masks = torch.stack(human_masks, dim=0)
|
||||
multitalk_embeds['ref_target_masks'] = ref_target_masks
|
||||
|
||||
gen_video_list = []
|
||||
is_first_clip = True
|
||||
arrive_last_frame = False
|
||||
@@ -3114,9 +3135,22 @@ class WanVideoSampler:
|
||||
audio_start_idx = iteration_count = 0
|
||||
audio_end_idx = audio_start_idx + clip_length
|
||||
indices = (torch.arange(4 + 1) - 2) * 1
|
||||
current_condframe_index = 0
|
||||
|
||||
if multitalk_embeds is not None:
|
||||
total_frames = len(multitalk_audio_embedding)
|
||||
|
||||
pcd_data = pcd_data_input = None
|
||||
if uni3c_embeds is not None:
|
||||
transformer.controlnet = uni3c_embeds["controlnet"]
|
||||
pcd_data = {
|
||||
"render_latent": uni3c_embeds["render_latent"].to(dtype),
|
||||
"render_mask": uni3c_embeds["render_mask"],
|
||||
"camera_embedding": uni3c_embeds["camera_embedding"],
|
||||
"controlnet_weight": uni3c_embeds["controlnet_weight"],
|
||||
"start": uni3c_embeds["start"],
|
||||
"end": uni3c_embeds["end"],
|
||||
}
|
||||
|
||||
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
|
||||
loop_pbar = tqdm(total=estimated_iterations, desc="Generating video clips")
|
||||
@@ -3126,6 +3160,10 @@ class WanVideoSampler:
|
||||
human_num = len(audio_embedding)
|
||||
audio_embs = None
|
||||
while True: # start video generation iteratively
|
||||
if mode == "infinitetalk":
|
||||
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1]
|
||||
log.info(f"current_condframe_index: {current_condframe_index}")
|
||||
log.info(f"audio_start_idx: {audio_start_idx}")
|
||||
if multitalk_embeds is not None:
|
||||
audio_embs = []
|
||||
# split audio with window size
|
||||
@@ -3136,6 +3174,11 @@ class WanVideoSampler:
|
||||
audio_embs.append(audio_emb)
|
||||
audio_embs = torch.concat(audio_embs, dim=0).to(dtype)
|
||||
|
||||
if uni3c_embeds is not None:
|
||||
vae.to(device)
|
||||
render_latent = vae.encode(original_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype), device=device, tiled=tiled_vae).to(dtype)
|
||||
pcd_data['render_latent'] = render_latent
|
||||
|
||||
h, w = cond_image.shape[-2], cond_image.shape[-1]
|
||||
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
|
||||
seq_len = ((frame_num - 1) // VAE_STRIDE[0] + 1) * lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
|
||||
@@ -3146,7 +3189,10 @@ class WanVideoSampler:
|
||||
|
||||
# get mask
|
||||
msk = torch.ones(1, frame_num, lat_h, lat_w, device=device)
|
||||
msk[:, cur_motion_frames_num:] = 0
|
||||
if mode == "multitalk":
|
||||
msk[:, cur_motion_frames_num:] = 0
|
||||
else:
|
||||
msk[:, 1:] = 0
|
||||
msk = torch.concat([
|
||||
torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]
|
||||
], dim=1)
|
||||
@@ -3159,12 +3205,21 @@ class WanVideoSampler:
|
||||
video_frames = torch.zeros(1, cond_image.shape[1], frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype)
|
||||
padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2)
|
||||
|
||||
# encode
|
||||
vae.to(device)
|
||||
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae).to(dtype)
|
||||
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
|
||||
|
||||
if mode == "multitalk":
|
||||
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
|
||||
else:
|
||||
if is_first_clip:
|
||||
latent_motion_frames = vae.encode(cond_image.to(device, vae.dtype), device=device, tiled=tiled_vae).to(dtype)
|
||||
else:
|
||||
latent_motion_frames = vae.encode(cond_frame.to(device, vae.dtype), device=device, tiled=tiled_vae).to(dtype)
|
||||
latent_motion_frames = latent_motion_frames[0]
|
||||
vae.to(offload_device)
|
||||
|
||||
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
|
||||
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
|
||||
y = torch.concat([msk, y], dim=1) # B 4+C T H W
|
||||
mm.soft_empty_cache()
|
||||
|
||||
@@ -3188,7 +3243,7 @@ class WanVideoSampler:
|
||||
latent = noise
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip:
|
||||
if not is_first_clip and mode == "multitalk":
|
||||
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
|
||||
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
|
||||
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
|
||||
@@ -3232,6 +3287,8 @@ class WanVideoSampler:
|
||||
for i in tqdm(range(len(timesteps)-1)):
|
||||
timestep = timesteps[i]
|
||||
latent_model_input = latent.to(device)
|
||||
if mode == "infinitetalk":
|
||||
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
latent_model_input,
|
||||
@@ -3262,12 +3319,14 @@ class WanVideoSampler:
|
||||
latent = temp_x0.squeeze(0)
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip:
|
||||
if not is_first_clip and mode == "multitalk":
|
||||
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
|
||||
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
|
||||
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
|
||||
_, T_m, _, _ = add_latent.shape
|
||||
latent[:, :T_m] = add_latent
|
||||
else:
|
||||
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
x0 = latent.to(device)
|
||||
del latent_model_input, timestep
|
||||
@@ -3287,7 +3346,10 @@ class WanVideoSampler:
|
||||
cm = ColorMatcher()
|
||||
cm_result_list = []
|
||||
for img in videos:
|
||||
cm_result = cm.transfer(src=img, ref=original_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch)
|
||||
if mode == "multitalk":
|
||||
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch)
|
||||
else:
|
||||
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch)
|
||||
cm_result_list.append(torch.from_numpy(cm_result))
|
||||
|
||||
videos = torch.stack(cm_result_list, dim=0).to(torch.float32).permute(3, 0, 1, 2).unsqueeze(0)
|
||||
@@ -3296,6 +3358,7 @@ class WanVideoSampler:
|
||||
gen_video_list.append(videos)
|
||||
else:
|
||||
gen_video_list.append(videos[:, :, cur_motion_frames_num:])
|
||||
current_condframe_index += 1
|
||||
|
||||
# decide whether is done
|
||||
if arrive_last_frame:
|
||||
@@ -3307,7 +3370,10 @@ class WanVideoSampler:
|
||||
is_first_clip = False
|
||||
cur_motion_frames_num = motion_frame
|
||||
|
||||
cond_image = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device)
|
||||
if mode == "infinitetalk":
|
||||
cond_frame = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device)
|
||||
else:
|
||||
cond_image = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device)
|
||||
|
||||
# Update progress bar
|
||||
iteration_count += 1
|
||||
@@ -3331,7 +3397,11 @@ class WanVideoSampler:
|
||||
miss_lengths.append(miss_length)
|
||||
else:
|
||||
miss_lengths.append(0)
|
||||
|
||||
if mode == "infinitetalk" and current_condframe_index >= original_images.shape[2]:
|
||||
last_frame = original_images[:, :, -1:, :, :]
|
||||
miss_length = 1
|
||||
original_images = torch.cat([original_images, last_frame.repeat(1, 1, miss_length, 1, 1)], dim=2)
|
||||
|
||||
gen_video_samples = torch.cat(gen_video_list, dim=2).to(torch.float32)
|
||||
|
||||
del noise, latent
|
||||
|
||||
Reference in New Issue
Block a user