Add TTM support (Time To Move)

https://github.com/time-to-move/TTM
This commit is contained in:
kijai
2025-11-16 17:52:49 +02:00
parent e3c2a1431b
commit b826642a83
2 changed files with 77 additions and 3 deletions
+47
View File
@@ -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",
}
+30 -3
View File
@@ -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()