Add TTMGuider class

Refactor KSamplerX0Inpaint and KSAMPLER classes to integrate TTMGuider for improved sampling functionality.
This commit is contained in:
Gius
2026-01-29 11:29:54 +01:00
committed by GitHub
parent ff27defc45
commit a4dccd9c6c
+206 -95
View File
@@ -1,109 +1,220 @@
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"]
start_sampler_step = model_options["start_sampler_step"]
skipped_sigmas = sigmas[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
i = torch.argmin(torch.abs(sigmas - timestep)).item()
# Cfg part 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)))
if i == 0: # Print cfg only at first step
print(f"Using per-step cfg list: {self.cfg}")
else:
model_k.noise = noise
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)
self.cfg = [self.cfg] * (steps + 1)
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
#---------------------------------------------------------
ttm_ref_latent = model_options["ttm_reference_latents"].to(x.device)
ttm_start_step = max(model_options["ttm_start_step"] - start_sampler_step, 0)
ttm_end_step = model_options["ttm_end_step"] - start_sampler_step
ttm_mask = model_options["motion_mask"].to(x.device)
# Time-to-move (TTM)
if i == 0: # First Step
if ttm_ref_latent is not None:
if ttm_start_step > steps:
raise ValueError("TTM start step is beyond the total number of steps")
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
if ttm_end_step > ttm_start_step:
print("Using Time-to-move (TTM)")
print(f"TTM reference latents shape: {ttm_ref_latent.shape}")
print(f"TTM motion mask shape: {ttm_mask.shape}")
print(f"Applying TTM from step {ttm_start_step} to {ttm_end_step}")
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
sigma_next = sigmas[ttm_start_step]
x = add_noise_at_step(ttm_ref_latent,
x,
sigma_next.to(x.device)
).to(x)
elif ttm_ref_latent is not None and (i + ttm_start_step) < ttm_end_step: # Following Steps if i > 0
if i + ttm_start_step < len(sigmas):
sigma_next = sigmas[i + ttm_start_step]
noisy_latents = add_noise_at_step(ttm_ref_latent,
x,
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)
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))
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)
return KSAMPLER(sampler_function, extra_options, inpaint_options)
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}
# Pass ttm options to KSAMPLER.sample
extra_args["model_options"]["ttm_reference_latents"] = self.ttm_reference_latents
extra_args["model_options"]["ttm_start_step"] = self.ttm_start_step
extra_args["model_options"]["ttm_end_step"] = self.ttm_end_step
extra_args["model_options"]["latent_image"] = self.latent_image
extra_args["model_options"]["motion_mask"] = self.motion_mask
extra_args["model_options"]["start_sampler_step"] = self.start_sampler_step
extra_args["model_options"]["sigmas"] = sigmas
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
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)
)
# 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 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
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())
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()
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
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