Add TTMGuider class
Refactor KSamplerX0Inpaint and KSAMPLER classes to integrate TTMGuider for improved sampling functionality.
This commit is contained in:
+206
-95
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user