Cleanup Multi/InfiniteTalk sampling loop some

This commit is contained in:
kijai
2025-08-23 16:27:10 +03:00
parent 5f521dc169
commit 73cff6ebab
+32 -53
View File
@@ -3085,7 +3085,6 @@ class WanVideoSampler:
}
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
loop_pbar = tqdm(total=estimated_iterations, desc="Total progress", position=1, leave=True)
callback = prepare_callback(patcher, estimated_iterations)
audio_embedding = multitalk_audio_embedding
@@ -3221,15 +3220,11 @@ class WanVideoSampler:
# encode
vae.to(device)
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
else:
if is_first_clip:
latent_motion_frames = vae.encode(cond_image.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)
else:
latent_motion_frames = vae.encode(cond_frame.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)
latent_motion_frames = latent_motion_frames[0]
cond_ = cond_image if is_first_clip else cond_frame
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
vae.to(offload_device)
y = torch.concat([msk, y], dim=1).squeeze(0) # 4+C T H W
mm.soft_empty_cache()
@@ -3279,8 +3274,7 @@ class WanVideoSampler:
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])
_, T_m, _, _ = add_latent.shape
latent[:, :T_m] = add_latent
latent[:, :add_latent.shape[1]] = add_latent
if offload:
#blockswap init
@@ -3343,19 +3337,16 @@ class WanVideoSampler:
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
noise_pred, self.cache_state = predict_with_cfg(
latent_model_input,
cfg[i],
positive,
text_embeds["negative_prompt_embeds"],
latent_model_input, cfg[i], positive, text_embeds["negative_prompt_embeds"],
timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_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(step_iteration_count, callback_latent, None, estimated_iterations*(len(timesteps)-1))
del callback_latent
sampling_pbar.update(1)
step_iteration_count += 1
# update latent
@@ -3364,13 +3355,8 @@ class WanVideoSampler:
dt = (timesteps[i] - timesteps[i + 1]) / 1000
latent = latent + noise_pred * dt[:, None, None, None]
else:
latent = latent.to(intermediate_device)
temp_x0 = sample_scheduler.step(
noise_pred.unsqueeze(0),
timestep,
latent.unsqueeze(0),
**scheduler_step_args)[0]
latent = temp_x0.squeeze(0)
latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0)
del noise_pred, latent_model_input, timestep
# differential diffusion inpaint
if masks is not None:
@@ -3384,62 +3370,56 @@ class WanVideoSampler:
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
latent[:, :add_latent.shape[1]] = add_latent
else:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
x0 = latent.to(device)
del latent_model_input, timestep
del noise, y, msk, latent_motion_frames
if offload:
transformer.to(offload_device)
vae.to(device)
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae, pbar=False)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.to(offload_device)
sampling_pbar.close()
# cache generated samples
videos = torch.stack(videos).cpu() # B C T H W
# optional color correction (less relevant for InfiniteTalk)
if colormatch != "disabled":
videos = videos[0].permute(1, 2, 3, 0).cpu().float().numpy()
videos = videos.permute(1, 2, 3, 0).float().numpy()
from color_matcher import ColorMatcher
cm = ColorMatcher()
cm_result_list = []
for img in videos:
if mode == "multitalk":
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch)
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().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))
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
videos = torch.stack(cm_result_list, dim=0).to(torch.float32).permute(3, 0, 1, 2).unsqueeze(0)
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
# cache generated samples
gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:])
if is_first_clip:
gen_video_list.append(videos)
else:
gen_video_list.append(videos[:, :, cur_motion_frames_num:])
current_condframe_index += 1
iteration_count += 1
# decide whether is done
if arrive_last_frame:
loop_pbar.update(estimated_iterations - iteration_count)
loop_pbar.close()
break
# update next condition frames
is_first_clip = False
cur_motion_frames_num = motion_frame
cond_ = videos[:, -cur_motion_frames_num:].unsqueeze(0)
if mode == "infinitetalk":
cond_frame = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device)
cond_frame = cond_
else:
cond_image = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device)
cond_image = cond_
# Update progress bar
iteration_count += 1
loop_pbar.update(1)
del videos, latent
mm.soft_empty_cache()
# Repeat audio emb
if multitalk_embeds is not None:
@@ -3464,9 +3444,8 @@ class WanVideoSampler:
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
gen_video_samples = torch.cat(gen_video_list, dim=1)
if force_offload:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
@@ -3475,7 +3454,7 @@ class WanVideoSampler:
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return {"video": gen_video_samples[0].permute(1, 2, 3, 0).cpu()},
return {"video": gen_video_samples.permute(1, 2, 3, 0)},
#region normal inference
else:
@@ -3657,9 +3636,9 @@ class WanVideoDecode:
mm.soft_empty_cache()
video = samples.get("video", None)
if video is not None:
video = torch.clamp(video, -1.0, 1.0)
video = (video + 1.0) / 2.0
return video.cpu(),
video.clamp_(-1.0, 1.0)
video.add_(1.0).div_(2.0)
return video.cpu().float(),
latents = samples["samples"]
end_image = samples.get("end_image", None)
has_ref = samples.get("has_ref", False)