Cleanup Multi/InfiniteTalk sampling loop some
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user