Merge pull request #3 from GiusTex/TTM-v2

TTM v2. Updated code to fix color  and some other bugs. Updated readme to show updates and some more details.
This commit is contained in:
Gius
2026-01-29 16:36:55 +01:00
committed by GitHub
8 changed files with 5118 additions and 4859 deletions
+16 -5
View File
@@ -1,13 +1,24 @@
# ComfyUI-Wan-TimeToMove
A native comfyui port of kijai's WanVideo-Wrapper TimeToMove
A native comfyui port of Kijai's WanVideo-Wrapper TimeToMove
<img width="1577" height="691" alt="WanVideo TTM nodes image" src="https://github.com/user-attachments/assets/0a7b34a6-805d-4c4a-80f3-ed2143907068" />
<img width="763" height="449" alt="ComfyUI-TTM-nodes" src="https://github.com/user-attachments/assets/af01cab0-5ccf-435c-bab0-f138cb227f1c" />
https://github.com/user-attachments/assets/551eac0d-c5fe-49a8-b1a2-3884d0ece746
https://github.com/user-attachments/assets/0b201e7a-d3c6-417f-8293-e20e8d2872fb
**This node is still WIP.** For now only the [lcm sampler](https://github.com/GiusTex/ComfyUI-Wan-TimeToMove/blob/main/k_diffusion/sampling.py#L1020) supports TimeToMove, and the generated frames are a bit dark (this color difference is seen especially when a first frame is passed).
### Updates:
- Solved color issue.
- Fixed other bugs.
The second sampler can be found here: `https://github.com/GiusTex/ComfyUI-MoreEfficientSamplers` but you can change it, and the scheduler used is this: `https://github.com/BigStationW/flowmatch_scheduler-comfyui`, useful when you use lightx loras.
### Nodes
The custom node contains 4 new nodes:
- `Encode WanVideo`: taken from wanvideo-wrapper, it encodes the reference video.
- `TTM Latent Add`: taken from wanvideo-wrapper, it embeds in the latent the reference to the driving video.
- `Timove To Move Guider`: this node adds the ttm latent to the latent noise before passing it to the sampling function. This node removes the necessity of a dedicated sampler.
- `CFG Float List Scheduler`: taken from wanvideo-wrapper, it creates a list of cfg values, and submits them step by step, making possible using different cfg values at different steps.
### Other custom nodes used:
- The advanced sampler used in the [second workflow](https://github.com/GiusTex/ComfyUI-Wan-TimeToMove/blob/TTM-v2/wanvideo_2_2_I2V_A14B_TimeToMove_workflow2.json) can be found [here](https://github.com/GiusTex/ComfyUI-MoreEfficientSamplers). You can still use the native comfyui `sampler custom advanced` using [this](https://github.com/GiusTex/ComfyUI-Wan-TimeToMove/blob/TTM-v2/wanvideo_2_2_I2V_A14B_TimeToMove_workflow1.json) workflow.
- The scheduler used is [this](https://github.com/BigStationW/flowmatch_scheduler-comfyui), useful for models using lightx loras. You can still use other samplers/schedulers.
### Download
To install ComfyUI-Wan-TimeToMove, follow these steps:
+13 -13
View File
@@ -1,20 +1,20 @@
from .nodes import (WanVideoEncode,
TTMKSamplerSelect,
AddTTMLatent,
WanVideoSamplerCustomUltraAdvancedEfficient)
from .nodes import (EncodeWanVideo,
TTMLatentAdd,
TimeToMoveGuider,
CFGFloatListScheduler)
NODE_CLASS_MAPPINGS = {
"WanVideoEncode": WanVideoEncode,
"TTMKSamplerSelect": TTMKSamplerSelect,
"AddTTMLatent": AddTTMLatent,
"WanVideoSamplerCustomUltraAdvancedEfficient": WanVideoSamplerCustomUltraAdvancedEfficient,
"EncodeWanVideo": EncodeWanVideo,
"TTMLatentAdd": TTMLatentAdd,
"TimeToMoveGuider": TimeToMoveGuider,
"CFGFloatListScheduler": CFGFloatListScheduler,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoEncode": "WanVideo Encode",
"TTMKSamplerSelect": "TimeToMove KSampler Select",
"AddTTMLatent": "Add TTM Latent",
"WanVideoSamplerCustomUltraAdvancedEfficient": "WanVideoSampler Custom Ultra Advanced Efficient",
"EncodeWanVideo": "Encode WanVideo",
"TTMLatentAdd": "TTM Latent Add",
"TimeToMoveGuider": "TimeToMove Guider",
"CFGFloatListScheduler": "CFGFloatListScheduler",
}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
File diff suppressed because it is too large Load Diff
+83 -223
View File
@@ -1,26 +1,10 @@
import torch
import logging
from comfy_api.latest import io
from comfy.utils import PROGRESS_BAR_ENABLED
import torch.nn.functional as F
import latent_preview
import comfy
from nodes import VAEDecodeTiled, PreviewImage, VAEDecode
from comfy_extras.nodes_custom_sampler import Noise_EmptyNoise, Noise_RandomNoise
from comfy.samplers import SAMPLER_NAMES
from PIL import Image
from .utils import (pil2tensor, warning, set_preview_method, sample_custom_ultra,
global_preview_method, store_ksampler_results, globals_cleanup,
add_noise_at_step, add_noise_to_reference_video)
from .samplers import sampler_object
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
log = logging.getLogger(__name__)
from .samplers import TTMGuider
from .utils import add_noise_to_reference_video
# Copied from ComfyUI Wanvideo Wrapper
class WanVideoEncode:
class EncodeWanVideo:
@classmethod
def INPUT_TYPES(s):
return {"required": {
@@ -60,21 +44,21 @@ class WanVideoEncode:
if latent_strength != 1.0:
latents *= latent_strength
log.info(f"WanVideo Encode: Encoded latents shape {latents.shape}")
print(f"WanVideo Encode: Encoded latents shape {latents.shape}")
return ({"samples": latents, "noise_mask": mask},)
# Copied from ComfyUI Wanvideo Wrapper
class AddTTMLatent:
class TTMLatentAdd:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"latent": ("LATENT", {"tooltip": "wanvideo latent"}),
"reference_latents": ("LATENT", {"tooltip": "Reference image to encode"}),
"start_step": ("INT", {"default": 0, "min": -1, "max": 1000, "step": 1, "tooltip": "Start step for whole denoising process"}),
"end_step": ("INT", {"default": 2, "min": 1, "max": 1000, "step": 1, "tooltip": "The step to stop applying TTM"}),
"ttm_start_step": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Start step to apply TTM latent guide"}),
"ttm_end_step": ("INT", {"default": 3, "min": 1, "max": 1000, "step": 1, "tooltip": "The step to stop applying TTM"}),
"ref_masks": ("MASK", {"tooltip": "Reference mask to encode"}),
}
}
@@ -84,9 +68,10 @@ class AddTTMLatent:
FUNCTION = "add"
CATEGORY = "Wan22 TimeToMove"
def add(self, latent, reference_latents, start_step, end_step, ref_masks):
if end_step < max(0, start_step):
raise ValueError(f"`end_step` ({end_step}) must be >= `start_step` ({start_step}).")
def add(self, latent, reference_latents, ttm_start_step, ttm_end_step, ref_masks):
if ttm_end_step < max(0, ttm_start_step):
raise ValueError(f"`ttm_end_step` ({ttm_end_step}) must be >= `ttm_start_step` ({ttm_start_step}).")
mask_sampled = ref_masks[::4]
mask_sampled = mask_sampled.unsqueeze(1).unsqueeze(0) # [1, T, 1, H, W]
@@ -104,224 +89,99 @@ class AddTTMLatent:
mode="nearest"
)
latent["ttm_reference_latents"] = reference_latents["samples"].squeeze(0) # [16, T, H, W]
latent["ttm_mask"] = mask_latent.squeeze(0).movedim(1, 0) # [1, T, H, W]
latent["ttm_start_step"] = start_step
latent["ttm_end_step"] = end_step
latent["ttm_reference_latents"] = reference_latents["samples"]
latent["ttm_mask"] = mask_latent.movedim(2, 1)
latent["ttm_start_step"] = ttm_start_step
latent["ttm_end_step"] = ttm_end_step
return (latent,)
class TTMKSamplerSelect(io.ComfyNode):
class TimeToMoveGuider:
@classmethod
def define_schema(cls):
return io.Schema(
node_id="TTMKSamplerSelect",
category="Wan Animate End Reference",
inputs=[
io.Combo.Input("sampler_name", options=SAMPLER_NAMES, default="lcm"),
io.Latent.Input("latent"),
],
outputs=[
io.Sampler.Output(),
]
)
def INPUT_TYPES(s):
return {"required":
{"model": ("MODEL", ),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "Works with a list of floats too (one cfg float per step)"}),
"latent": ("LATENT", {"tooltip": "You can connect here the latent from TTM Latent Add, to pass reference video and ttm options"}),
"start_sampler_step": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Start step of the whole sampling process. It will automatically skip the selected number of sigmas (starting from the first ones); if the sampler has a start_step option and you changed its value, set the same here"}),
},
}
RETURN_TYPES = ("GUIDER",)
RETURN_NAMES = ("guider",)
FUNCTION = "guide"
CATEGORY = "Wan22 TimeToMove"
def guide(cls, model, positive, negative, cfg, latent, start_sampler_step):
guider = TTMGuider(model)
guider.set_conds(positive, negative)
guider.set_cfg(cfg)
@classmethod
def execute(cls, sampler_name, latent) -> io.NodeOutput:
ttm_options = {}
ttm_options["ttm_reference_latents"] = latent.get("ttm_reference_latents", None)
ttm_options["ttm_start_step"] = latent["ttm_start_step"]
ttm_options["ttm_end_step"] = latent["ttm_end_step"]
ttm_options["latent_image"] = latent["samples"]
ttm_options["motion_mask"] = latent["ttm_mask"]
ttm_options["start_sampler_step"] = start_sampler_step
guider.set_ttm_options(ttm_options)
sampler = sampler_object(sampler_name, ttm_options)
return io.NodeOutput(sampler)
get_sampler = execute
return (guider,)
class WanVideoSamplerCustomUltraAdvancedEfficient:
# Image Preview code taken from jags111's efficiency-nodes (TSC_KSampler)
empty_image = pil2tensor(Image.new('RGBA', (1, 1), (0, 0, 0, 0)))
# Taken from kijai WanVideo-Wrapper
class CFGFloatListScheduler:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"model": ("MODEL",),
"add_noise": ("BOOLEAN", {"default": True}),
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "control_after_generate": True}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"sampler": ("SAMPLER", ),
"sigmas": ("SIGMAS", ),
"latent": ("LATENT", ),
"start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}),
"end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}),
"return_with_leftover_noise": ("BOOLEAN", {"default": False}),
"preview_method": (["auto", "latent2rgb", "taesd", "vae_decoded_only", "none"],),
"vae_decode": (["true", "true (tiled)", "false"],),
},
"optional": {
"optional_vae": ("VAE",),
},
"hidden": {
"prompt": "PROMPT",
"extra_pnginfo": "EXTRA_PNGINFO",
"my_unique_id": "UNIQUE_ID",
},
}
return {"required": {
"steps": ("INT", {"default": 30, "min": 2, "max": 1000, "step": 1, "tooltip": "Number of steps to schedule cfg for"} ),
"cfg_scale_start": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 30.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}),
"cfg_scale_end": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 30.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}),
"interpolation": (["linear", "ease_in", "ease_out"], {"default": "linear", "tooltip": "Interpolation method to use for the cfg scale"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "Start percent of the steps to apply cfg"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "End percent of the steps to apply cfg"}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
},
}
RETURN_TYPES = ("MODEL", "CONDITIONING", "CONDITIONING", "SAMPLER", "SIGMAS", "LATENT","LATENT", "IMAGE", "VAE",)
RETURN_NAMES = ("model", "positive", "negative", "sampler", "sigmas", "output", "denoised_output", "image", "vae", )
FUNCTION = "sample"
RETURN_TYPES = ("FLOAT", )
RETURN_NAMES = ("float_list",)
FUNCTION = "process"
CATEGORY = "Wan22 TimeToMove"
DESCRIPTION = "Helper node to generate a list of floats that can be used to schedule cfg scale for the steps, outside the set range cfg is set to 1.0. Taken from Kijai WanVideo-Wrapper"
def sample(self, model, add_noise, noise_seed, cfg, positive, negative, sampler, sigmas, latent, start_at_step, end_at_step, return_with_leftover_noise, preview_method, vae_decode, optional_vae=(None,), prompt=None, extra_pnginfo=None, my_unique_id=None):
latent_image = latent["samples"]
latent_image = comfy.sample.fix_empty_latent_channels(model, latent_image)
latent["samples"] = latent_image
# Rename the vae variable
vae = optional_vae
# If vae is not connected, disable vae decoding
if vae == (None,) and vae_decode != "false":
print(f"{warning('Sampler Custom Ultra Advanced Warning:')} No vae input detected, proceeding as if vae_decode was false.\n")
vae_decode = "false"
# ------------------------------------------------------------------------------------------------------
def vae_decode_latent(vae, out, vae_decode):
return VAEDecodeTiled().decode(vae,out,320)[0] if "tiled" in vae_decode else VAEDecode().decode(vae,out)[0]
# ---------------------------------------------------------------------------------------------------------------
def process(self, steps, cfg_scale_start, cfg_scale_end, interpolation, start_percent, end_percent, unique_id):
noise_mask = None
if "noise_mask" in latent:
noise_mask = latent["noise_mask"]
def process_latents():
x0_output = {}
# Initialize output variables
out = out_denoised = images = preview = previous_preview_method = None
# Create a list of floats for the cfg schedule
cfg_list = [1.0] * steps
start_idx = min(int(steps * start_percent), steps - 1)
end_idx = min(int(steps * end_percent), steps - 1)
if not add_noise:
noise = Noise_EmptyNoise().generate_noise(latent)
for i in range(start_idx, end_idx + 1):
if i >= steps:
break
if end_idx == start_idx:
t = 0
else:
noise = Noise_RandomNoise(noise_seed).generate_noise(latent)
#Time-to-move (TTM)
ttm_start_step = 0
ttm_reference_latents = latent.get("ttm_reference_latents", None)
if ttm_reference_latents is not None:
motion_mask = latent["ttm_mask"].to(latent_image.device, latent_image.dtype)
ttm_start_step = max(latent["ttm_start_step"] - start_at_step, 0)
ttm_end_step = latent["ttm_end_step"] - start_at_step
if ttm_start_step > end_at_step:
raise ValueError("TTM start step is beyond the total number of steps")
sigma = sigmas[ttm_start_step]
t = (i - start_idx) / (end_idx - start_idx)
if ttm_end_step > ttm_start_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}")
if interpolation == "linear":
factor = t
elif interpolation == "ease_in":
factor = t * t
elif interpolation == "ease_out":
factor = t * (2 - t)
noise = add_noise_at_step(ttm_reference_latents,
noise,
sigma
).to(latent_image.device, latent_image.dtype)
#--------------------------------------------------------------
try:
# Change the global preview method (temporarily)
set_preview_method(preview_method)
x0_output = {}
callback = latent_preview.prepare_callback(model, sigmas.shape[-1] - 1, x0_output)
cfg_list[i] = round(cfg_scale_start + factor * (cfg_scale_end - cfg_scale_start), 2)
disable_pbar = not PROGRESS_BAR_ENABLED
disable_noise = False
if not add_noise:
disable_noise = True
# Prepare noise for img specified by batch_inds
if disable_noise:
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
else:
batch_inds = latent["batch_index"] if "batch_index" in latent else None
noise = comfy.sample.prepare_noise(latent_image, noise_seed, batch_inds)
force_full_denoise = True
if return_with_leftover_noise:
force_full_denoise = False
device = comfy.model_management.intermediate_device()
model_options = model.model_options
start_step = start_at_step
last_step = end_at_step
denoise_mask = noise_mask
# If start_percent > 0, always include the first step
if start_percent > 0:
cfg_list[0] = 1.0
samples = sample_custom_ultra(model, device,
noise,
sampler,
positive, negative,
cfg, model_options,
latent_image,
start_step, last_step,
force_full_denoise, denoise_mask,
sigmas,
callback, disable_pbar, noise_seed)
samples = samples.to(comfy.model_management.intermediate_device())
out = latent.copy()
out["samples"] = samples
if "x0" in x0_output:
out_denoised = latent.copy()
out_denoised["samples"] = model.model.process_latent_out(x0_output["x0"].cpu())
else:
out_denoised = out
previous_preview_method = global_preview_method()
# ---------------------------------------------------------------------------------------------------------------
# Decode image if not yet decoded
if "true" in vae_decode:
if images is None:
images = vae_decode_latent(vae, out, vae_decode)
# Store decoded image as base image of no script is detected
store_ksampler_results("image", my_unique_id, images)
# Define preview images
if preview_method == "none" or (preview_method == "vae_decoded_only" and vae_decode == "false"):
preview = {"images": list()}
elif images is not None:
preview = PreviewImage().save_images(images, prompt=prompt, extra_pnginfo=extra_pnginfo)["ui"]
# Define a dummy output image
if images is None and vae_decode == "false":
images = WanVideoSamplerCustomUltraAdvancedEfficient.empty_image
finally:
# Restore global changes
set_preview_method(previous_preview_method)
return out, out_denoised, preview, images
# ---------------------------------------------------------------------------------------------------------------
# Clean globally stored objects of non-existant nodes
globals_cleanup(prompt)
# ---------------------------------------------------------------------------------------------------------------
out, out_denoised, preview, images = process_latents()
result = (model, positive, negative, sampler, sigmas,
out, out_denoised, images, vae,)
if preview is None:
return {"result": result}
else:
return {"ui": preview, "result": result}
return (cfg_list,)
+201 -94
View File
@@ -1,109 +1,216 @@
import torch
from comfy.samplers import Sampler
from comfy.extra_samplers import uni_pc
from .k_diffusion import sampling as k_diffusion_sampling
import comfy
from comfy.model_patcher import ModelPatcher
from comfy.samplers import (sampling_function, process_conds, cast_to_load_options,
preprocess_conds_hooks, get_total_hook_groups_in_conds,
filter_registered_hooks_on_conds)
from .utils import add_noise_at_step
class KSamplerX0Inpaint:
def __init__(self, model, sigmas):
self.inner_model = model
self.sigmas = sigmas
# Add ttm_options to extra_args
def __call__(self, x, sigma, denoise_mask, model_options={}, seed=None,
ttm_reference_latents=None, ttm_start_step=None,
ttm_end_step=None, latent_image=None, motion_mask=None):
if denoise_mask is not None:
if "denoise_mask_function" in model_options:
denoise_mask = model_options["denoise_mask_function"](sigma, denoise_mask, extra_options={"model": self.inner_model, "sigmas": self.sigmas})
latent_mask = 1. - denoise_mask
x = x * denoise_mask + self.inner_model.inner_model.scale_latent_inpaint(x=x, sigma=sigma, noise=self.noise, latent_image=self.latent_image) * latent_mask
model_options["ttm_reference_latents"] = ttm_reference_latents
model_options["ttm_start_step"] = ttm_start_step
model_options["ttm_end_step"] = ttm_end_step
model_options["latent_image"] = latent_image
model_options["motion_mask"] = motion_mask
out = self.inner_model(x, sigma, model_options=model_options, seed=seed)
if denoise_mask is not None:
out = out * denoise_mask + self.latent_image * latent_mask
return out
class TTMGuider:
def __init__(self, model_patcher: ModelPatcher):
self.model_patcher = model_patcher
self.model_options = model_patcher.model_options
self.original_conds = {}
self.cfg = 1.0
def set_conds(self, positive, negative):
self.inner_set_conds({"positive": positive, "negative": negative})
def set_cfg(self, cfg):
self.cfg = cfg
def set_ttm_options(self, ttm_options):
self.ttm_reference_latents = ttm_options["ttm_reference_latents"]
self.ttm_start_step = ttm_options["ttm_start_step"]
self.ttm_end_step = ttm_options["ttm_end_step"]
self.latent_image = ttm_options["latent_image"]
self.motion_mask = ttm_options["motion_mask"]
self.start_sampler_step = ttm_options["start_sampler_step"]
class KSAMPLER(Sampler):
def __init__(self, sampler_function, extra_options={}, inpaint_options={}):
self.sampler_function = sampler_function
self.extra_options = extra_options
self.inpaint_options = inpaint_options
def inner_set_conds(self, conds):
for k in conds:
self.original_conds[k] = comfy.sampler_helpers.convert_cond(conds[k])
def sample(self, model_wrap, sigmas, extra_args, callback, noise, latent_image=None, denoise_mask=None, disable_pbar=False):
extra_args["denoise_mask"] = denoise_mask
model_k = KSamplerX0Inpaint(model_wrap, sigmas)
model_k.latent_image = latent_image
if self.inpaint_options.get("random", False): #TODO: Should this be the default?
generator = torch.manual_seed(extra_args.get("seed", 41) + 1)
model_k.noise = torch.randn(noise.shape, generator=generator, device="cpu").to(noise.dtype).to(noise.device)
def __call__(self, *args, **kwargs):
return self.outer_predict_noise(*args, **kwargs)
def outer_predict_noise(self, x, timestep, model_options={}, seed=None):
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
self.predict_noise,
self,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.PREDICT_NOISE, self.model_options, is_model_options=True)
).execute(x, timestep, model_options, seed)
def predict_noise(self, x, timestep, model_options={}, seed=None):
#---------------------------------------------------------
sigmas = model_options["sigmas"]
noise = model_options["noise"]
i = torch.argmin(torch.abs(sigmas - timestep)).item()
ttm_ref_latent = model_options["ttm_reference_latents"]
ttm_start_step = model_options["ttm_start_step"]
ttm_end_step = model_options["ttm_end_step"]
ttm_mask = model_options["motion_mask"]
# Time-to-move (TTM)
if (i + ttm_start_step) < ttm_end_step:
if i + ttm_start_step < len(sigmas):
sigma_next = sigmas[i + ttm_start_step]
noisy_latents = add_noise_at_step(ttm_ref_latent,
noise,
sigma_next.to(x.device)
).to(x)
x = x * (1 - ttm_mask) + noisy_latents * ttm_mask
else:
x = x * (1 - ttm_mask) + ttm_ref_latent * ttm_mask
#---------------------------------------------------------
return sampling_function(self.inner_model, x, timestep,
self.conds.get("negative", None),
self.conds.get("positive", None),
self.cfg[i],
model_options=model_options, seed=seed)
def inner_sample(self, noise, latent_image, device, sampler, sigmas, denoise_mask, callback, disable_pbar, seed, latent_shapes=None):
if latent_image is not None and torch.count_nonzero(latent_image) > 0: #Don't shift the empty latent image.
latent_image = self.inner_model.process_latent_in(latent_image)
self.conds = process_conds(self.inner_model, noise, self.conds, device, latent_image, denoise_mask, seed, latent_shapes=latent_shapes)
extra_model_options = comfy.model_patcher.create_model_options_clone(self.model_options)
extra_model_options.setdefault("transformer_options", {})["sample_sigmas"] = sigmas
extra_args = {"model_options": extra_model_options, "seed": seed}
#---------------------------------------------------------
skipped_sigmas = sigmas[self.start_sampler_step:]
# 4 < 5
if len(skipped_sigmas) < len(sigmas): # sampler doesn't have start_step option
sigmas = skipped_sigmas
# 4 == 4
elif len(skipped_sigmas) == len(sigmas): # sampler already has option
pass # we don't want another sigma less
steps = len(sigmas)-1
extra_args["model_options"]["steps"] = steps
#---------------------------------------------------------
# Pass ttm options to KSAMPLER.sample
ttm_start_step = max(self.ttm_start_step - self.start_sampler_step, 0)
ttm_end_step = self.ttm_end_step - self.start_sampler_step
extra_args["model_options"]["ttm_reference_latents"] = self.ttm_reference_latents.to(noise.device)
extra_args["model_options"]["ttm_start_step"] = ttm_start_step
extra_args["model_options"]["ttm_end_step"] = ttm_end_step
extra_args["model_options"]["motion_mask"] = self.motion_mask.to(noise.device)
extra_args["model_options"]["sigmas"] = sigmas
extra_args["model_options"]["noise"] = noise
if ttm_start_step > steps:
raise ValueError("TTM start step is beyond the total number of steps")
if ttm_end_step > ttm_start_step:
print("Using Time-to-move (TTM)")
print(f"TTM reference latents shape: {self.ttm_reference_latents.shape}")
print(f"TTM motion mask shape: {self.motion_mask.shape}")
print(f"Applying TTM from step {ttm_start_step} to {ttm_end_step}")
#---------------------------------------------------------
# Cfg schedule taken from Kijai WanVideo-Wrapper
if isinstance(self.cfg, list):
if steps < len(self.cfg):
print(f"Received {len(self.cfg)} cfg values, but only {steps} steps. Slicing cfg list to match steps.")
self.cfg = self.cfg[:steps]
elif steps > len(self.cfg):
print(f"Received only {len(self.cfg)} cfg values, but {steps} steps. Extending cfg list to match steps.")
self.cfg.extend([self.cfg[-1]] * (steps - len(self.cfg)))
print(f"Using per-step cfg list: {self.cfg}")
else:
model_k.noise = noise
self.cfg = [self.cfg] * (steps + 1)
#---------------------------------------------------------
noise = model_wrap.inner_model.model_sampling.noise_scaling(sigmas[0], noise, latent_image, self.max_denoise(model_wrap, sigmas))
k_callback = None
total_steps = len(sigmas) - 1
if callback is not None:
k_callback = lambda x: callback(x["i"], x["denoised"], x["x"], total_steps)
executor = comfy.patcher_extension.WrapperExecutor.new_class_executor(
sampler.sample,
sampler,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE, extra_args["model_options"], is_model_options=True)
)
samples = self.sampler_function(model_k, noise, sigmas, extra_args=extra_args, callback=k_callback, disable=disable_pbar, **self.extra_options)
samples = model_wrap.inner_model.model_sampling.inverse_noise_scaling(sigmas[-1], samples)
return samples
# run steps and get final samples
samples = executor.execute(self, sigmas, extra_args, callback, noise, latent_image, denoise_mask, disable_pbar)
return self.inner_model.process_latent_out(samples.to(torch.float32))
def ksampler(sampler_name, ttm_options, extra_options={}, inpaint_options={}):
if sampler_name == "dpm_fast":
def dpm_fast_function(model, noise, sigmas, extra_args, callback, disable):
if len(sigmas) <= 1:
return noise
def outer_sample(self, noise, latent_image, sampler, sigmas, denoise_mask=None, callback=None, disable_pbar=False, seed=None, latent_shapes=None):
self.inner_model, self.conds, self.loaded_models = comfy.sampler_helpers.prepare_sampling(self.model_patcher, noise.shape, self.conds, self.model_options)
device = self.model_patcher.load_device
sigma_min = sigmas[-1]
if sigma_min == 0:
sigma_min = sigmas[-2]
total_steps = len(sigmas) - 1
return k_diffusion_sampling.sample_dpm_fast(model, noise, sigma_min, sigmas[0], total_steps, extra_args=extra_args, callback=callback, disable=disable)
sampler_function = dpm_fast_function
elif sampler_name == "dpm_adaptive":
def dpm_adaptive_function(model, noise, sigmas, extra_args, callback, disable, **extra_options):
if len(sigmas) <= 1:
return noise
noise = noise.to(device)
latent_image = latent_image.to(device)
sigmas = sigmas.to(device)
cast_to_load_options(self.model_options, device=device, dtype=self.model_patcher.model_dtype())
sigma_min = sigmas[-1]
if sigma_min == 0:
sigma_min = sigmas[-2]
return k_diffusion_sampling.sample_dpm_adaptive(model, noise, sigma_min, sigmas[0], extra_args=extra_args, callback=callback, disable=disable, **extra_options)
sampler_function = dpm_adaptive_function
elif sampler_name == "lcm":
def lcm_function(model, noise, sigmas, extra_args, callback, disable, **extra_options):
extra_args["ttm_reference_latents"] = ttm_options["ttm_reference_latents"]
extra_args["ttm_start_step"] = ttm_options["ttm_start_step"]
extra_args["ttm_end_step"] = ttm_options["ttm_end_step"]
extra_args["latent_image"] = ttm_options["latent_image"]
extra_args["motion_mask"] = ttm_options["motion_mask"]
return k_diffusion_sampling.sample_lcm(model, noise, sigmas, extra_args=extra_args, callback=callback, disable=disable, **extra_options)
sampler_function = lcm_function
else:
sampler_function = getattr(k_diffusion_sampling, "sample_{}".format(sampler_name))
try:
self.model_patcher.pre_run()
output = self.inner_sample(noise, latent_image, device, sampler, sigmas, denoise_mask, callback, disable_pbar, seed, latent_shapes=latent_shapes)
finally:
self.model_patcher.cleanup()
return KSAMPLER(sampler_function, extra_options, inpaint_options)
comfy.sampler_helpers.cleanup_models(self.conds, self.loaded_models)
del self.inner_model
del self.loaded_models
return output
def sample(self, noise, latent_image, sampler, sigmas, denoise_mask=None, callback=None, disable_pbar=False, seed=None):
if sigmas.shape[-1] == 0:
return latent_image
def sampler_object(name, ttm_options):
if name == "uni_pc":
sampler = KSAMPLER(uni_pc.sample_unipc)
elif name == "uni_pc_bh2":
sampler = KSAMPLER(uni_pc.sample_unipc_bh2)
elif name == "ddim":
sampler = ksampler("euler", inpaint_options={"random": True})
elif name == "lcm":
sampler = ksampler(name, ttm_options)
else:
sampler = ksampler(name)
return sampler
if latent_image.is_nested:
latent_image, latent_shapes = comfy.utils.pack_latents(latent_image.unbind())
noise, _ = comfy.utils.pack_latents(noise.unbind())
else:
latent_shapes = [latent_image.shape]
if denoise_mask is not None:
if denoise_mask.is_nested:
denoise_masks = denoise_mask.unbind()
denoise_masks = denoise_masks[:len(latent_shapes)]
else:
denoise_masks = [denoise_mask]
for i in range(len(denoise_masks), len(latent_shapes)):
denoise_masks.append(torch.ones(latent_shapes[i]))
for i in range(len(denoise_masks)):
denoise_masks[i] = comfy.sampler_helpers.prepare_mask(denoise_masks[i], latent_shapes[i], self.model_patcher.load_device)
if len(denoise_masks) > 1:
denoise_mask, _ = comfy.utils.pack_latents(denoise_masks)
else:
denoise_mask = denoise_masks[0]
self.conds = {}
for k in self.original_conds:
self.conds[k] = list(map(lambda a: a.copy(), self.original_conds[k]))
preprocess_conds_hooks(self.conds)
try:
orig_model_options = self.model_options
self.model_options = comfy.model_patcher.create_model_options_clone(self.model_options)
# if one hook type (or just None), then don't bother caching weights for hooks (will never change after first step)
orig_hook_mode = self.model_patcher.hook_mode
if get_total_hook_groups_in_conds(self.conds) <= 1:
self.model_patcher.hook_mode = comfy.hooks.EnumHookMode.MinVram
comfy.sampler_helpers.prepare_model_patcher(self.model_patcher, self.conds, self.model_options)
filter_registered_hooks_on_conds(self.conds, self.model_options)
executor = comfy.patcher_extension.WrapperExecutor.new_class_executor(
self.outer_sample,
self,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, self.model_options, is_model_options=True)
)
output = executor.execute(noise, latent_image, sampler, sigmas, denoise_mask, callback, disable_pbar, seed, latent_shapes=latent_shapes)
finally:
cast_to_load_options(self.model_options, device=self.model_patcher.offload_device)
self.model_options = orig_model_options
self.model_patcher.hook_mode = orig_hook_mode
self.model_patcher.restore_hook_patches()
del self.conds
if len(latent_shapes) > 1:
output = comfy.nested_tensor.NestedTensor(comfy.utils.unpack_latents(output, latent_shapes))
return output
-117
View File
@@ -1,121 +1,4 @@
import torch
from PIL import Image
import numpy as np
import latent_preview
from comfy.cli_args import args
from comfy.samplers import sample
# Convert PIL to Tensor (grabbed from WAS Suite)
def pil2tensor(image: Image.Image) -> torch.Tensor:
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
def format_message(text, color_code):
RESET_COLOR = "\033[0m"
return f"{color_code}{text}{RESET_COLOR}"
WARNING_COLOR = "\033[93m" # Yellow
def warning(text):
return format_message(text, WARNING_COLOR)
# Set global preview_method
def set_preview_method(method):
if method == 'auto' or method == 'LatentPreviewMethod.Auto':
args.preview_method = latent_preview.LatentPreviewMethod.Auto
elif method == 'latent2rgb' or method == 'LatentPreviewMethod.Latent2RGB':
args.preview_method = latent_preview.LatentPreviewMethod.Latent2RGB
elif method == 'taesd' or method == 'LatentPreviewMethod.TAESD':
args.preview_method = latent_preview.LatentPreviewMethod.TAESD
else:
args.preview_method = latent_preview.LatentPreviewMethod.NoPreviews
def sample_custom_ultra(model, device, noise, sampler, positive, negative, cfg, model_options={}, latent_image=None, start_step=None, last_step=None, force_full_denoise=False, denoise_mask=None, sigmas=None, callback=None, disable_pbar=False, seed=None):
if last_step is not None and last_step < (len(sigmas) - 1):
sigmas = sigmas[:last_step + 1]
if force_full_denoise:
sigmas[-1] = 0
if start_step is not None:
if start_step < (len(sigmas) - 1):
sigmas = sigmas[start_step:]
else:
if latent_image is not None:
return latent_image
else:
return torch.zeros_like(noise)
return sample(model, noise, positive, negative, cfg, device, sampler, sigmas, model_options, latent_image=latent_image, denoise_mask=denoise_mask, callback=callback, disable_pbar=disable_pbar, seed=seed)
# Extract global preview_method
def global_preview_method():
return args.preview_method
# Cache for Efficiency Node models
loaded_objects = {
"ckpt": [], # (ckpt_name, ckpt_model, clip, bvae, [id])
"refn": [], # (ckpt_name, ckpt_model, clip, bvae, [id])
"vae": [], # (vae_name, vae, [id])
"lora": [] # ([(lora_name, strength_model, strength_clip)], ckpt_name, lora_model, clip_lora, [id])
}
# Cache for Efficient Ksamplers
last_helds = {
"latent": [], # (latent, [parameters], id) # Base sampling latent results
"image": [], # (image, id) # Base sampling image results
"cnet_img": [] # (cnet_img, [parameters], id) # HiRes-Fix control net preprocessor image results
}
def store_ksampler_results(key: str, my_unique_id, value, parameters_list=None):
global last_helds
for i, data in enumerate(last_helds[key]):
id_ = data[-1] # ID will always be the last in the tuple
if id_ == my_unique_id:
# Check if parameters_list is provided or not
updated_data = (value, parameters_list, id_) if parameters_list is not None else (value, id_)
last_helds[key][i] = updated_data
return True
# If parameters_list is given
if parameters_list is not None:
last_helds[key].append((value, parameters_list, my_unique_id))
else:
last_helds[key].append((value, my_unique_id))
return True
# This function cleans global variables associated with nodes that are no longer detected on UI
def globals_cleanup(prompt):
global loaded_objects
global last_helds
# Step 1: Clean up last_helds
for key in list(last_helds.keys()):
original_length = len(last_helds[key])
last_helds[key] = [
(*values, id_)
for *values, id_ in last_helds[key]
if str(id_) in prompt.keys()
]
# Step 2: Clean up loaded_objects
for key in list(loaded_objects.keys()):
for i, tup in enumerate(list(loaded_objects[key])):
# Remove ids from id array in each tuple that don't exist in prompt
id_array = [id for id in tup[-1] if str(id) in prompt.keys()]
if len(id_array) != len(tup[-1]):
if id_array:
loaded_objects[key][i] = tup[:-1] + (id_array,)
#print(f'Updated tuple at index {i} in {key} in loaded_objects: {loaded_objects[key][i]}')
else:
# If id array becomes empty, delete the corresponding tuple
loaded_objects[key].remove(tup)
#print(f'Deleted tuple at index {i} in {key} in loaded_objects because its id array became empty.')
# Copied from ComfyUI Wanvideo Wrapper
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff