Multi/InfiniteTalk sampling loop cleanup and optimizations, support FantasyPortrait within the loop
This commit is contained in:
@@ -1665,7 +1665,7 @@ class WanVideoSampler:
|
||||
#region Scheduler
|
||||
sample_scheduler = None
|
||||
if scheduler != "multitalk":
|
||||
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
|
||||
sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g)
|
||||
log.info(f"sigmas: {sample_scheduler.sigmas}")
|
||||
else:
|
||||
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
|
||||
@@ -1689,27 +1689,6 @@ class WanVideoSampler:
|
||||
steps = len(cfg)
|
||||
else:
|
||||
cfg = [cfg] * (steps + 1)
|
||||
|
||||
if end_step != -1:
|
||||
timesteps = timesteps[:end_step]
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
|
||||
log.info(f"Sampling until step {end_step}, timestep: {timesteps[-1]}")
|
||||
if start_step > 0:
|
||||
timesteps = timesteps[start_step:]
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[start_step:]
|
||||
log.info(f"Skipping first {start_step} steps, starting from timestep {timesteps[0]}")
|
||||
|
||||
log.info(f"timesteps: {timesteps}")
|
||||
|
||||
if sample_scheduler is not None:
|
||||
if hasattr(sample_scheduler, 'timesteps'):
|
||||
sample_scheduler.timesteps = timesteps
|
||||
|
||||
scheduler_step_args = {"generator": seed_g}
|
||||
step_sig = inspect.signature(sample_scheduler.step)
|
||||
for arg in list(scheduler_step_args.keys()):
|
||||
if arg not in step_sig.parameters:
|
||||
scheduler_step_args.pop(arg)
|
||||
|
||||
control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = None
|
||||
vace_data = vace_context = vace_scale = None
|
||||
@@ -2686,7 +2665,7 @@ class WanVideoSampler:
|
||||
# FreeInit noise reinitialization (after first iteration)
|
||||
if freeinit_args is not None and iter_idx > 0:
|
||||
# restart scheduler for each iteration
|
||||
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
|
||||
sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g)
|
||||
|
||||
# Re-apply start_step and end_step logic to timesteps and sigmas
|
||||
if end_step != -1:
|
||||
@@ -3066,9 +3045,10 @@ class WanVideoSampler:
|
||||
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[0])
|
||||
|
||||
audio_embedding = multitalk_audio_embedding
|
||||
human_num = len(audio_embedding)
|
||||
audio_embs = None
|
||||
|
||||
pcd_data = pcd_data_input = None
|
||||
if uni3c_embeds is not None:
|
||||
@@ -3082,13 +3062,10 @@ class WanVideoSampler:
|
||||
"end": uni3c_embeds["end"],
|
||||
}
|
||||
|
||||
total_frames = len(audio_embedding[0])
|
||||
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
|
||||
callback = prepare_callback(patcher, estimated_iterations)
|
||||
|
||||
audio_embedding = multitalk_audio_embedding
|
||||
human_num = len(audio_embedding)
|
||||
audio_embs = None
|
||||
|
||||
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
|
||||
|
||||
while True: # start video generation iteratively
|
||||
@@ -3105,29 +3082,13 @@ 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)
|
||||
# Pad original_images if needed
|
||||
num_frames = original_images.shape[2]
|
||||
required_frames = audio_end_idx - audio_start_idx
|
||||
if audio_end_idx > num_frames:
|
||||
pad_len = audio_end_idx - num_frames
|
||||
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
|
||||
padded_images = torch.cat([original_images, last_frame], dim=2)
|
||||
else:
|
||||
padded_images = original_images
|
||||
render_latent = vae.encode(
|
||||
padded_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]) if cond_image is not None else (target_h, target_w)
|
||||
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])
|
||||
latent_frame_num = (frame_num - 1) // 4 + 1
|
||||
|
||||
noise = torch.randn(
|
||||
16, (frame_num - 1) // 4 + 1,
|
||||
16, latent_frame_num,
|
||||
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
|
||||
|
||||
# Calculate the correct latent slice based on current iteration
|
||||
@@ -3180,52 +3141,27 @@ class WanVideoSampler:
|
||||
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
|
||||
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
|
||||
|
||||
window_vace_data = None
|
||||
if vace_data is not None:
|
||||
window_vace_data = []
|
||||
for vace_entry in vace_data:
|
||||
partial_context = vace_entry["context"][0][:, latent_start_idx:latent_end_idx]
|
||||
if has_ref:
|
||||
partial_context[:, 0] = vace_entry["context"][0][:, 0]
|
||||
|
||||
window_vace_data.append({
|
||||
"context": [partial_context],
|
||||
"scale": vace_entry["scale"],
|
||||
"start": vace_entry["start"],
|
||||
"end": vace_entry["end"],
|
||||
"seq_len": vace_entry["seq_len"]
|
||||
})
|
||||
|
||||
# get image cond mask
|
||||
msk = torch.ones(1, frame_num, lat_h, lat_w, device=device)
|
||||
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)
|
||||
msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w)
|
||||
msk = msk.transpose(1, 2).to(dtype) # B 4 T H W
|
||||
|
||||
mm.soft_empty_cache()
|
||||
|
||||
# zero padding and vae encode
|
||||
# zero padding and vae encode for img cond
|
||||
if cond_image is not None:
|
||||
video_frames = torch.zeros(1, 3, 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, pbar=False).to(dtype)
|
||||
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
|
||||
|
||||
if mode == "multitalk":
|
||||
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
|
||||
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
|
||||
else:
|
||||
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.model.clear_cache()
|
||||
vae.to(offload_device)
|
||||
y = torch.concat([msk, y], dim=1).squeeze(0) # 4+C T H W
|
||||
|
||||
motion_frame_index = cur_motion_frames_num if mode == "multitalk" else 1
|
||||
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
|
||||
msk[:, :motion_frame_index] = 1
|
||||
y = torch.cat([msk, y]) # 4+C T H W
|
||||
mm.soft_empty_cache()
|
||||
else:
|
||||
y = None
|
||||
@@ -3237,33 +3173,8 @@ class WanVideoSampler:
|
||||
timesteps = [torch.tensor([t], device=device) for t in timesteps]
|
||||
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
|
||||
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:
|
||||
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)
|
||||
transformed_timesteps.append(t_tensor)
|
||||
|
||||
transformed_timesteps.append(torch.tensor([0.], device=device))
|
||||
timesteps = transformed_timesteps
|
||||
sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g)
|
||||
timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)]
|
||||
|
||||
# sample videos
|
||||
latent = noise
|
||||
@@ -3316,6 +3227,43 @@ class WanVideoSampler:
|
||||
else:
|
||||
positive = text_embeds["prompt_embeds"]
|
||||
|
||||
window_vace_data = None
|
||||
# if vace_data is not None:
|
||||
# window_vace_data = []
|
||||
# for vace_entry in vace_data:
|
||||
# partial_context = vace_entry["context"][0][:, latent_start_idx:latent_end_idx]
|
||||
# if has_ref:
|
||||
# partial_context[:, 0] = vace_entry["context"][0][:, 0]
|
||||
|
||||
# window_vace_data.append({
|
||||
# "context": [partial_context],
|
||||
# "scale": vace_entry["scale"],
|
||||
# "start": vace_entry["start"],
|
||||
# "end": vace_entry["end"],
|
||||
# "seq_len": vace_entry["seq_len"]
|
||||
# })
|
||||
|
||||
# uni3c slices
|
||||
if uni3c_embeds is not None:
|
||||
vae.to(device)
|
||||
# Pad original_images if needed
|
||||
num_frames = original_images.shape[2]
|
||||
required_frames = audio_end_idx - audio_start_idx
|
||||
if audio_end_idx > num_frames:
|
||||
pad_len = audio_end_idx - num_frames
|
||||
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
|
||||
padded_images = torch.cat([original_images, last_frame], dim=2)
|
||||
else:
|
||||
padded_images = original_images
|
||||
render_latent = vae.encode(
|
||||
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
|
||||
device=device, tiled=tiled_vae
|
||||
).to(dtype)
|
||||
vae.model.clear_cache()
|
||||
vae.to(offload_device)
|
||||
pcd_data['render_latent'] = render_latent
|
||||
|
||||
# unianimate slices
|
||||
partial_unianim_data = None
|
||||
if unianim_data is not None:
|
||||
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
|
||||
@@ -3328,6 +3276,22 @@ class WanVideoSampler:
|
||||
"end_percent": unianimate_poses["end_percent"]
|
||||
}
|
||||
|
||||
# fantasy portrait slices
|
||||
partial_fantasy_portrait_input = None
|
||||
if fantasy_portrait_input is not None:
|
||||
adapter_proj = fantasy_portrait_input["adapter_proj"]
|
||||
if latent_end_idx > adapter_proj.shape[1]:
|
||||
pad_len = latent_end_idx - adapter_proj.shape[1]
|
||||
last_frame = adapter_proj[:, -1:, :, :].repeat(1, pad_len, 1, 1)
|
||||
padded_proj = torch.cat([adapter_proj, last_frame], dim=1)
|
||||
else:
|
||||
padded_proj = adapter_proj
|
||||
partial_fantasy_portrait_input = fantasy_portrait_input.copy()
|
||||
partial_fantasy_portrait_input["adapter_proj"] = padded_proj[:, latent_start_idx:latent_end_idx]
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
# sampling loop
|
||||
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]
|
||||
@@ -3338,7 +3302,7 @@ class WanVideoSampler:
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
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)
|
||||
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input)
|
||||
|
||||
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)
|
||||
@@ -3373,9 +3337,13 @@ class WanVideoSampler:
|
||||
else:
|
||||
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
del noise, y, msk, latent_motion_frames
|
||||
del noise, latent_motion_frames
|
||||
if offload:
|
||||
transformer.to(offload_device)
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
vae.to(device)
|
||||
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
|
||||
vae.model.clear_cache()
|
||||
@@ -3419,7 +3387,6 @@ class WanVideoSampler:
|
||||
cond_image = cond_
|
||||
|
||||
del videos, latent
|
||||
mm.soft_empty_cache()
|
||||
|
||||
# Repeat audio emb
|
||||
if multitalk_embeds is not None:
|
||||
|
||||
@@ -6,7 +6,7 @@ from .flowmatch_pusa import FlowMatchSchedulerPusa
|
||||
from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep
|
||||
from .scheduling_flow_match_lcm import FlowMatchLCMScheduler
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler
|
||||
import numpy as np
|
||||
import inspect
|
||||
from ...utils import log
|
||||
|
||||
scheduler_list = [
|
||||
@@ -23,7 +23,7 @@ scheduler_list = [
|
||||
"multitalk"
|
||||
]
|
||||
|
||||
def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_args, denoise_strength, sigmas=None):
|
||||
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim, flowedit_args, denoise_strength, sigmas=None, seed_g=None):
|
||||
timesteps = None
|
||||
if 'unipc' in scheduler:
|
||||
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
|
||||
@@ -99,4 +99,31 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg
|
||||
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
if timesteps is None:
|
||||
timesteps = sample_scheduler.timesteps
|
||||
return sample_scheduler, timesteps
|
||||
|
||||
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:
|
||||
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:]
|
||||
|
||||
log.info(f"timesteps: {timesteps}")
|
||||
|
||||
if hasattr(sample_scheduler, 'timesteps'):
|
||||
sample_scheduler.timesteps = timesteps
|
||||
|
||||
if seed_g is not None:
|
||||
scheduler_step_args = {"generator": seed_g}
|
||||
step_sig = inspect.signature(sample_scheduler.step)
|
||||
for arg in list(scheduler_step_args.keys()):
|
||||
if arg not in step_sig.parameters:
|
||||
scheduler_step_args.pop(arg)
|
||||
|
||||
return sample_scheduler, timesteps, scheduler_step_args
|
||||
Reference in New Issue
Block a user