Fix TTM for dual sampler setups

This commit is contained in:
kijai
2025-11-16 19:23:19 +02:00
parent b826642a83
commit f872460285
2 changed files with 18 additions and 15 deletions
+16 -15
View File
@@ -250,6 +250,7 @@ class WanVideoSampler:
if isinstance(scheduler, dict): if isinstance(scheduler, dict):
sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"]) sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"])
timesteps = scheduler["timesteps"] timesteps = scheduler["timesteps"]
start_step = scheduler.get("start_step", start_step)
elif scheduler != "multitalk": elif scheduler != "multitalk":
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, log_timesteps=True) sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, log_timesteps=True)
log.info(f"sigmas: {sample_scheduler.sigmas}") log.info(f"sigmas: {sample_scheduler.sigmas}")
@@ -1156,17 +1157,19 @@ class WanVideoSampler:
ttm_reference_latents = image_embeds.get("ttm_reference_latents", None) ttm_reference_latents = image_embeds.get("ttm_reference_latents", None)
if ttm_reference_latents is not None: if ttm_reference_latents is not None:
motion_mask = image_embeds["ttm_mask"].to(device, dtype) motion_mask = image_embeds["ttm_mask"].to(device, dtype)
background_mask = 1 - motion_mask ttm_start_step = max(image_embeds["ttm_start_step"] - start_step, 0)
ttm_start_step = image_embeds["ttm_start_step"] ttm_end_step = image_embeds["ttm_end_step"] - start_step
ttm_end_step = image_embeds["ttm_end_step"]
log.info("Using Time-to-move (TTM)") if ttm_start_step > steps:
log.info(f"TTM reference latents shape: {ttm_reference_latents.shape}") raise ValueError("TTM start step is beyond the total number of steps")
log.info(f"TTM motion mask shape: {motion_mask.shape}")
log.info(f"Applying TTM from step {ttm_start_step} to {ttm_end_step}")
tweak = torch.as_tensor(timesteps[ttm_start_step], device=device, dtype=torch.long).view(1) if ttm_end_step > ttm_start_step:
latent = sample_scheduler.add_noise(ttm_reference_latents, noise, tweak).to(latent) log.info("Using Time-to-move (TTM)")
log.info(f"TTM reference latents shape: {ttm_reference_latents.shape}")
log.info(f"TTM motion mask shape: {motion_mask.shape}")
log.info(f"Applying TTM from step {ttm_start_step} to {ttm_end_step}")
latent = add_noise(ttm_reference_latents, noise, timesteps[ttm_start_step].to(noise.device)).to(latent)
#region model pred #region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
@@ -3078,13 +3081,11 @@ class WanVideoSampler:
# TTM # TTM
if ttm_reference_latents is not None and (idx + ttm_start_step) < ttm_end_step: if ttm_reference_latents is not None and (idx + ttm_start_step) < ttm_end_step:
if idx + ttm_start_step + 1 < len(timesteps): if idx + ttm_start_step + 1 < len(sample_scheduler.all_timesteps):
prev_t = timesteps[idx + ttm_start_step + 1] noisy_latents = add_noise(ttm_reference_latents, noise, sample_scheduler.all_timesteps[idx + ttm_start_step + 1].to(noise.device)).to(latent)
prev_t = torch.as_tensor(prev_t, device=device, dtype=torch.long).view(1) latent = latent * (1 - motion_mask) + noisy_latents * motion_mask
noisy_latents = sample_scheduler.add_noise(ttm_reference_latents, noise, prev_t).to(latent)
latent = latent * background_mask + noisy_latents * motion_mask
else: else:
latent = latent * background_mask + ttm_reference_latents.to(latent) * motion_mask latent = latent * (1 - motion_mask) + ttm_reference_latents.to(latent) * motion_mask
if freeinit_args is not None: if freeinit_args is not None:
current_latent = latent.clone() current_latent = latent.clone()
+2
View File
@@ -156,6 +156,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
end_idx = end_step - 1 end_idx = end_step - 1
# Slice timesteps and sigmas once, based on indices # Slice timesteps and sigmas once, based on indices
all_timesteps = timesteps
timesteps = timesteps[start_idx:end_idx+1] timesteps = timesteps[start_idx:end_idx+1]
sample_scheduler.full_sigmas = sample_scheduler.sigmas.clone() sample_scheduler.full_sigmas = sample_scheduler.sigmas.clone()
sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer
@@ -167,5 +168,6 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
if hasattr(sample_scheduler, 'timesteps'): if hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps sample_scheduler.timesteps = timesteps
setattr(sample_scheduler, 'all_timesteps', all_timesteps)
return sample_scheduler, timesteps, start_idx, end_idx return sample_scheduler, timesteps, start_idx, end_idx