From b826642a83cb0e752f18949c3e858eabb8f49646 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 16 Nov 2025 17:52:49 +0200 Subject: [PATCH] Add TTM support (Time To Move) https://github.com/time-to-move/TTM --- nodes.py | 47 +++++++++++++++++++++++++++++++++++++++++++++++ nodes_sampler.py | 33 ++++++++++++++++++++++++++++++--- 2 files changed, 77 insertions(+), 3 deletions(-) diff --git a/nodes.py b/nodes.py index 0a611e9..d76a61c 100644 --- a/nodes.py +++ b/nodes.py @@ -2083,6 +2083,49 @@ class WanVideoRoPEFunction: return (rope_func_dict,) return (rope_function,) +#region TTM +class WanVideoAddTTMLatents: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "embeds": ("WANVIDIMAGE_EMBEDS",), + "reference_latents": ("LATENT", {"tooltip": "Latents used as reference for TTM"}), + "mask": ("MASK", {"tooltip": "Mask used for TTM"}), + "start_step": ("INT", {"default": 0, "min": -1, "max": 1000, "step": 1, "tooltip": "Start step for whole denoising process"}), + "end_step": ("INT", {"default": 1, "min": 1, "max": 1000, "step": 1, "tooltip": "The step to stop applying TTM"}), + }, + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", ) + RETURN_NAMES = ("image_embeds", ) + FUNCTION = "add" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "https://github.com/time-to-move/TTM" + + def add(self, embeds, reference_latents, mask, start_step, end_step): + + if end_step < max(0, start_step): + raise ValueError(f"`end_step` ({end_step}) must be >= `start_step` ({start_step}).") + + mask_sampled = mask[::VAE_STRIDE[0]] + mask_sampled = mask_sampled.unsqueeze(1).unsqueeze(0) # [1, T, 1, H, W] + + # Upsample spatially to latent resolution + H_latent = mask_sampled.shape[-2] // VAE_STRIDE[1] + W_latent = mask_sampled.shape[-1] // VAE_STRIDE[1] + mask_latent = F.interpolate( + mask_sampled.float(), + size=(mask_sampled.shape[2], H_latent, W_latent), + mode="nearest" + ) + + updated = dict(embeds) + updated["ttm_reference_latents"] = reference_latents["samples"].squeeze(0) + updated["ttm_mask"] = mask_latent.squeeze(0).movedim(1, 0) # [T, 1, H, W] + updated["ttm_start_step"] = start_step + updated["ttm_end_step"] = end_step + + return (updated,) #region VideoDecode class WanVideoDecode: @@ -2292,6 +2335,8 @@ class WanVideoEncode: if latent_strength != 1.0: latents *= latent_strength + latents = latents.cpu() + log.info(f"WanVideoEncode: Encoded latents shape {latents.shape}") mm.soft_empty_cache() @@ -2337,6 +2382,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds, "TextImageEncodeQwenVL": TextImageEncodeQwenVL, "WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds, + "WanVideoAddTTMLatents": WanVideoAddTTMLatents, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -2378,4 +2424,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSchedulerSA_ODE": "WanVideo Scheduler SA-ODE", "WanVideoAddBindweaveEmbeds": "WanVideo Add Bindweave Embeds", "WanVideoUniLumosEmbeds": "WanVideo UniLumos Embeds", + "WanVideoAddTTMLatents": "WanVideo Add TTMLatents", } diff --git a/nodes_sampler.py b/nodes_sampler.py index b8cd80a..7263fe0 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1151,6 +1151,23 @@ class WanVideoSampler: log.info(f"UniLumos background latent input shape: {background_latents.shape}") background_latents = background_latents.to(device, dtype) + #Time-to-move (TTM) + ttm_start_step = 0 + ttm_reference_latents = image_embeds.get("ttm_reference_latents", None) + if ttm_reference_latents is not None: + motion_mask = image_embeds["ttm_mask"].to(device, dtype) + background_mask = 1 - motion_mask + ttm_start_step = image_embeds["ttm_start_step"] + ttm_end_step = image_embeds["ttm_end_step"] + + 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}") + + tweak = torch.as_tensor(timesteps[ttm_start_step], device=device, dtype=torch.long).view(1) + latent = sample_scheduler.add_noise(ttm_reference_latents, noise, tweak).to(latent) + #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, @@ -1621,7 +1638,7 @@ class WanVideoSampler: if not multitalk_sampling and not framepack and not wananimate_loop: log.info(f"Input sequence length: {seq_len}") - log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps") + log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps-ttm_start_step} steps") intermediate_device = device @@ -1705,9 +1722,9 @@ class WanVideoSampler: if pusa_noisy_steps == -1: pusa_noisy_steps = len(timesteps) try: - pbar = ProgressBar(len(timesteps)) + pbar = ProgressBar(len(timesteps) - ttm_start_step) #region main loop start - for idx, t in enumerate(tqdm(timesteps, disable=multitalk_sampling or wananimate_loop)): + for idx, t in enumerate(tqdm(timesteps[ttm_start_step:], disable=multitalk_sampling or wananimate_loop)): if flowedit_args is not None: if idx < skip_steps: continue @@ -3058,6 +3075,16 @@ class WanVideoSampler: ) mask = masks[idx].to(latent) latent = image_latent * mask + latent * (1-mask) + + # TTM + if ttm_reference_latents is not None and (idx + ttm_start_step) < ttm_end_step: + if idx + ttm_start_step + 1 < len(timesteps): + prev_t = timesteps[idx + ttm_start_step + 1] + prev_t = torch.as_tensor(prev_t, device=device, dtype=torch.long).view(1) + noisy_latents = sample_scheduler.add_noise(ttm_reference_latents, noise, prev_t).to(latent) + latent = latent * background_mask + noisy_latents * motion_mask + else: + latent = latent * background_mask + ttm_reference_latents.to(latent) * motion_mask if freeinit_args is not None: current_latent = latent.clone()