Multi/InfiniteTalk v2v

This commit is contained in:
kijai
2025-08-20 16:40:13 +03:00
parent edf2a24519
commit 053f26ce82
3 changed files with 96 additions and 38 deletions
+59 -13
View File
@@ -2033,7 +2033,7 @@ class WanVideoSampler:
context = get_context_scheduler(context_schedule)
# vid2vid
if samples is not None:
if samples is not None and not multitalk_sampling:
saved_generator_state = samples.get("generator_state", None)
if saved_generator_state is not None:
seed_g.set_state(saved_generator_state)
@@ -2623,7 +2623,7 @@ class WanVideoSampler:
# diff diff prep
masks = None
if samples is not None and mask is not None:
if not multitalk_sampling and samples is not None and mask is not None:
mask = 1 - mask
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device)
@@ -2704,7 +2704,7 @@ class WanVideoSampler:
try:
pbar = ProgressBar(len(timesteps))
#region main loop start
for idx, t in enumerate(tqdm(timesteps)):
for idx, t in enumerate(tqdm(timesteps, disable=multitalk_sampling)):
if flowedit_args is not None:
if idx < skip_steps:
continue
@@ -3071,17 +3071,16 @@ class WanVideoSampler:
}
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
loop_pbar = tqdm(total=estimated_iterations, desc="Generating video clips")
loop_pbar = tqdm(total=estimated_iterations, desc="Total progress", position=1, leave=True)
callback = prepare_callback(patcher, estimated_iterations)
audio_embedding = multitalk_audio_embedding
human_num = len(audio_embedding)
audio_embs = None
while True: # start video generation iteratively
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
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_indices: {audio_start_idx}-{audio_end_idx}|{clip_length}")
if multitalk_embeds is not None:
audio_embs = []
# split audio with window size
@@ -3094,7 +3093,6 @@ class WanVideoSampler:
if uni3c_embeds is not None:
vae.to(device)
print("original_images", original_images.shape)
# Pad original_images if needed
num_frames = original_images.shape[2]
required_frames = audio_end_idx - audio_start_idx
@@ -3117,6 +3115,34 @@ class WanVideoSampler:
noise = torch.randn(
16, (frame_num - 1) // 4 + 1,
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
if samples is not None:
input_samples = samples["samples"].squeeze(0).to(noise)
# Calculate the correct slice based on current iteration
if is_first_clip:
latent_start_idx = 0
latent_end_idx = noise.shape[1]
else:
new_frames_per_iteration = frame_num - motion_frame
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
latent_start_idx = iteration_count * new_latent_frames_per_iteration
latent_end_idx = latent_start_idx + noise.shape[1]
# Check if we have enough frames in input_samples
if latent_end_idx > input_samples.shape[1]:
# We need more frames than available - pad the input_samples at the end
pad_length = latent_end_idx - input_samples.shape[1]
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
input_samples = torch.cat([input_samples, last_frame], dim=1)
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
if add_noise_to_samples:
latent_timestep = timesteps[0]
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
else:
noise = input_samples
# get mask
msk = torch.ones(1, frame_num, lat_h, lat_w, device=device)
@@ -3138,8 +3164,7 @@ class WanVideoSampler:
# 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)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)
if mode == "multitalk":
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
@@ -3162,6 +3187,24 @@ class WanVideoSampler:
else:
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
steps = len(timesteps)
if end_step != -1 and start_step >= end_step:
raise ValueError("start_step must be less than end_step")
if denoise_strength < 1.0:
if start_step != 0:
raise ValueError("start_step must be 0 when denoise_strength is used")
start_step = steps - int(steps * denoise_strength) - 1
if (end_step != -1 or end_step >= steps):
timesteps = timesteps[:end_step]
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
if start_step > 0:
timesteps = timesteps[start_step:]
sample_scheduler.sigmas = sample_scheduler.sigmas[start_step:]
if sample_scheduler is not None:
if hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps
transformed_timesteps = []
for t in timesteps:
t_tensor = torch.tensor([t.item()], device=device)
@@ -3214,8 +3257,8 @@ class WanVideoSampler:
elif model["manual_offloading"]:
transformer.to(device)
comfy_pbar = ProgressBar(len(timesteps)-1)
for i in tqdm(range(len(timesteps)-1)):
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
for i in range(len(timesteps)-1):
timestep = timesteps[i]
latent_model_input = latent.to(device)
if mode == "infinitetalk":
@@ -3229,6 +3272,8 @@ class WanVideoSampler:
timestep, idx, y.squeeze(0), clip_embeds, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs)
sampling_pbar.update(1)
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)
callback(iteration_count, callback_latent, None, estimated_iterations)
@@ -3261,13 +3306,14 @@ class WanVideoSampler:
x0 = latent.to(device)
del latent_model_input, timestep
comfy_pbar.update(1)
if offload:
transformer.to(offload_device)
vae.to(device)
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae)
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae, pbar=False)
vae.to(offload_device)
sampling_pbar.close()
# cache generated samples
videos = torch.stack(videos).cpu() # B C T H W
+1
View File
@@ -1098,6 +1098,7 @@ class WanVideoModelLoader:
#for name, param in transformer.named_parameters():
# print(name, param.dtype, param.device, param.shape)
pbar.update_absolute(param_count)
pbar.update_absolute(0)
comfy_model.diffusion_model = transformer
comfy_model.load_device = transformer_load_device
+36 -25
View File
@@ -1019,10 +1019,11 @@ class VideoVAE_(nn.Module):
return mu
def encode(self, x):
def encode(self, x, pbar=True):
self.clear_cache()
## cache
pbar = ProgressBar(x.shape[2])
if pbar:
pbar = ProgressBar(x.shape[2])
t = x.shape[2]
iter_ = 1 + (t - 1) // 4
@@ -1037,10 +1038,13 @@ class VideoVAE_(nn.Module):
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
out = torch.cat([out, out_], 2)
pbar.update(iter_)
if pbar:
pbar.update(iter_)
mu = self.conv1(out).chunk(2, dim=1)[0]
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
if pbar:
pbar.update_absolute(0)
return mu
@@ -1076,11 +1080,12 @@ class VideoVAE_(nn.Module):
def decode(self, z):
def decode(self, z, pbar=True):
self.clear_cache()
# z: [b,c,t,h,w]
pbar = ProgressBar(z.shape[2])
if pbar:
pbar = ProgressBar(z.shape[2])
z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2]
@@ -1096,7 +1101,11 @@ class VideoVAE_(nn.Module):
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2) # may add tensor offload
pbar.update(1)
if pbar:
pbar.update(1)
if pbar:
pbar.update_absolute(0)
return out
def reparameterize(self, mu, log_var):
@@ -1167,7 +1176,7 @@ class WanVideoVAE(nn.Module):
return mask
def tiled_decode(self, hidden_states, device, tile_size, tile_stride):
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, pbar=True):
_, _, T, H, W = hidden_states.shape
size_h, size_w = tile_size
stride_h, stride_w = tile_stride
@@ -1187,8 +1196,8 @@ class WanVideoVAE(nn.Module):
out_T = T * 4 - 3
weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
pbar = ProgressBar(len(tasks))
if pbar:
pbar = ProgressBar(len(tasks))
for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
@@ -1215,13 +1224,14 @@ class WanVideoVAE(nn.Module):
target_h: target_h + hidden_states_batch.shape[3],
target_w: target_w + hidden_states_batch.shape[4],
] += mask
pbar.update(1)
if pbar:
pbar.update(1)
values = values / weight
values = values.float().clamp_(-1, 1)
return values
def tiled_encode(self, video, device, tile_size, tile_stride, end_=False):
def tiled_encode(self, video, device, tile_size, tile_stride, end_=False, pbar=True):
_, _, T, H, W = video.shape
if tile_size is None and tile_stride is None:
@@ -1248,8 +1258,8 @@ class WanVideoVAE(nn.Module):
out_T += 1
weight = torch.zeros((1, 1, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device)
values = torch.zeros((1, self.z_dim, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device)
pbar = ProgressBar(len(tasks))
if pbar:
pbar = ProgressBar(len(tasks))
for h, h_, w, w_ in tqdm(tasks, desc="VAE encoding"):
hidden_states_batch = video[:, :, :, h:h_, w:w_].to(computation_device)
if end_:
@@ -1279,21 +1289,22 @@ class WanVideoVAE(nn.Module):
target_h: target_h + hidden_states_batch.shape[3],
target_w: target_w + hidden_states_batch.shape[4],
] += mask
pbar.update(1)
if pbar:
pbar.update(1)
values = values / weight
values = values.float()
return values
def single_encode(self, video, device):
def single_encode(self, video, device, pbar=True):
video = video.to(device)
x = self.model.encode(video)
x = self.model.encode(video, pbar=pbar)
return x.float()
def single_decode(self, hidden_state, device):
def single_decode(self, hidden_state, device, pbar=True):
hidden_state = hidden_state.to(device)
video = self.model.decode(hidden_state)
video = self.model.decode(hidden_state, pbar=pbar)
return video
def double_encode(self, video, device):
@@ -1308,36 +1319,36 @@ class WanVideoVAE(nn.Module):
video = self.model.decode_2(hidden_state)
return video
def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None):
def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None, pbar=True):
videos = [video.to("cpu") for video in videos]
hidden_states = []
for video in videos:
video = video.unsqueeze(0)
if tiled:
hidden_state = self.tiled_encode(video, device, tile_size, tile_stride, end_=end_)
hidden_state = self.tiled_encode(video, device, tile_size, tile_stride, end_=end_, pbar=pbar)
else:
if end_:
hidden_state = self.double_encode(video, device)
else:
hidden_state = self.single_encode(video, device)
hidden_state = self.single_encode(video, device, pbar=pbar)
hidden_state = hidden_state.squeeze(0)
hidden_states.append(hidden_state)
hidden_states = torch.stack(hidden_states)
return hidden_states
def decode(self, hidden_states, device, tiled=False, end_=False, tile_size=(34, 34), tile_stride=(18, 16)):
def decode(self, hidden_states, device, tiled=False, end_=False, tile_size=(34, 34), tile_stride=(18, 16), pbar=True):
hidden_states = [hidden_state.to("cpu") for hidden_state in hidden_states]
videos = []
for hidden_state in hidden_states:
hidden_state = hidden_state.unsqueeze(0)
if tiled:
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride)
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, pbar=pbar)
else:
if end_:
video = self.double_decode(hidden_state, device)
else:
video = self.single_decode(hidden_state, device)
video = self.single_decode(hidden_state, device, pbar=pbar)
video = video.squeeze(0)
videos.append(video)
return videos