Allow latent input
This commit is contained in:
@@ -89,6 +89,7 @@ class DDIMSampler(object):
|
|||||||
fs=None,
|
fs=None,
|
||||||
timestep_spacing='uniform', #uniform_trailing for starting from last timestep
|
timestep_spacing='uniform', #uniform_trailing for starting from last timestep
|
||||||
guidance_rescale=0.0,
|
guidance_rescale=0.0,
|
||||||
|
noise_multiplier=0,
|
||||||
**kwargs
|
**kwargs
|
||||||
):
|
):
|
||||||
|
|
||||||
@@ -134,6 +135,7 @@ class DDIMSampler(object):
|
|||||||
precision=precision,
|
precision=precision,
|
||||||
fs=fs,
|
fs=fs,
|
||||||
guidance_rescale=guidance_rescale,
|
guidance_rescale=guidance_rescale,
|
||||||
|
noise_multiplier=noise_multiplier,
|
||||||
**kwargs)
|
**kwargs)
|
||||||
return samples, intermediates
|
return samples, intermediates
|
||||||
|
|
||||||
@@ -143,13 +145,15 @@ class DDIMSampler(object):
|
|||||||
callback=None, timesteps=None, quantize_denoised=False,
|
callback=None, timesteps=None, quantize_denoised=False,
|
||||||
mask=None, x0=None, img_callback=None, log_every_t=100,
|
mask=None, x0=None, img_callback=None, log_every_t=100,
|
||||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
||||||
unconditional_guidance_scale=1., unconditional_conditioning=None, verbose=True,precision=None,fs=None,guidance_rescale=0.0,
|
unconditional_guidance_scale=1., unconditional_conditioning=None, verbose=True,precision=None,fs=None,guidance_rescale=0.0, noise_multiplier=0,
|
||||||
**kwargs):
|
**kwargs):
|
||||||
device = self.model.betas.device
|
device = self.model.betas.device
|
||||||
b = shape[0]
|
b = shape[0]
|
||||||
if x_T is None:
|
if x_T is None:
|
||||||
|
print("Using random noise")
|
||||||
img = torch.randn(shape, device=device)
|
img = torch.randn(shape, device=device)
|
||||||
else:
|
else:
|
||||||
|
print("Using input noise")
|
||||||
img = x_T
|
img = x_T
|
||||||
if precision is not None:
|
if precision is not None:
|
||||||
if precision == 16:
|
if precision == 16:
|
||||||
@@ -171,6 +175,12 @@ class DDIMSampler(object):
|
|||||||
|
|
||||||
clean_cond = kwargs.pop("clean_cond", False)
|
clean_cond = kwargs.pop("clean_cond", False)
|
||||||
|
|
||||||
|
sigmas = self.ddim_sigmas_for_original_num_steps if ddim_use_original_steps else self.ddim_sigmas
|
||||||
|
|
||||||
|
if noise_multiplier != 1.0:
|
||||||
|
adjustment_factor = 1 + (noise_multiplier - 1) * 0.001
|
||||||
|
sigmas = sigmas * adjustment_factor
|
||||||
|
print("Sigmas:", sigmas)
|
||||||
# cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning)
|
# cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning)
|
||||||
pbar = comfy.utils.ProgressBar(total_steps)
|
pbar = comfy.utils.ProgressBar(total_steps)
|
||||||
for i, step in enumerate(iterator):
|
for i, step in enumerate(iterator):
|
||||||
@@ -186,10 +196,7 @@ class DDIMSampler(object):
|
|||||||
img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? <ddim inversion>
|
img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? <ddim inversion>
|
||||||
img = img_orig * mask + (1. - mask) * img # keep original & modify use img
|
img = img_orig * mask + (1. - mask) * img # keep original & modify use img
|
||||||
|
|
||||||
|
outs = self.p_sample_ddim(img, cond, ts, sigmas, index=index, use_original_steps=ddim_use_original_steps,
|
||||||
|
|
||||||
|
|
||||||
outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps,
|
|
||||||
quantize_denoised=quantize_denoised, temperature=temperature,
|
quantize_denoised=quantize_denoised, temperature=temperature,
|
||||||
noise_dropout=noise_dropout, score_corrector=score_corrector,
|
noise_dropout=noise_dropout, score_corrector=score_corrector,
|
||||||
corrector_kwargs=corrector_kwargs,
|
corrector_kwargs=corrector_kwargs,
|
||||||
@@ -210,7 +217,7 @@ class DDIMSampler(object):
|
|||||||
return img, intermediates
|
return img, intermediates
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def p_sample_ddim(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
|
def p_sample_ddim(self, x, c, t, sigmas, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
|
||||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
||||||
unconditional_guidance_scale=1., unconditional_conditioning=None,
|
unconditional_guidance_scale=1., unconditional_conditioning=None,
|
||||||
uc_type=None, conditional_guidance_scale_temporal=None,mask=None,x0=None,guidance_rescale=0.0,**kwargs):
|
uc_type=None, conditional_guidance_scale_temporal=None,mask=None,x0=None,guidance_rescale=0.0,**kwargs):
|
||||||
@@ -248,7 +255,7 @@ class DDIMSampler(object):
|
|||||||
alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev
|
alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev
|
||||||
sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas
|
sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas
|
||||||
# sigmas = self.model.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
# sigmas = self.model.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
||||||
sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
#sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
||||||
# select parameters corresponding to the currently considered timestep
|
# select parameters corresponding to the currently considered timestep
|
||||||
|
|
||||||
if is_video:
|
if is_video:
|
||||||
|
|||||||
@@ -537,7 +537,9 @@ class ToonCrafterInterpolation:
|
|||||||
},
|
},
|
||||||
"optional": {
|
"optional": {
|
||||||
"image_embed_ratio": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
"image_embed_ratio": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||||
"augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001})
|
"augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001}),
|
||||||
|
"optional_latents": ("LATENT",),
|
||||||
|
"latent_noise_multiplier": ("FLOAT", {"default": 1.0, "min": 0, "max": 100, "step": 0.001}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -546,7 +548,7 @@ class ToonCrafterInterpolation:
|
|||||||
FUNCTION = "process"
|
FUNCTION = "process"
|
||||||
CATEGORY = "DynamiCrafterWrapper"
|
CATEGORY = "DynamiCrafterWrapper"
|
||||||
|
|
||||||
def process(self, model, clip_vision, images, positive, negative, cfg, steps, eta, seed, fs, frames, vae_dtype, image_embed_ratio=1.0, augmentation_level=0):
|
def process(self, model, clip_vision, images, positive, negative, cfg, steps, eta, seed, fs, frames, vae_dtype, image_embed_ratio=1.0, augmentation_level=0, optional_latents=None, latent_noise_multiplier=1.0):
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
offload_device = mm.unet_offload_device()
|
offload_device = mm.unet_offload_device()
|
||||||
mm.unload_all_models()
|
mm.unload_all_models()
|
||||||
@@ -670,6 +672,16 @@ class ToonCrafterInterpolation:
|
|||||||
self.model.image_proj_model.to(offload_device)
|
self.model.image_proj_model.to(offload_device)
|
||||||
|
|
||||||
#inference
|
#inference
|
||||||
|
if optional_latents is not None:
|
||||||
|
samples_in = optional_latents['samples'].clone().to(device)
|
||||||
|
samples_in = samples_in * 0.18215
|
||||||
|
samples_in = samples_in.unsqueeze(0).permute(0, 2, 1, 3, 4)
|
||||||
|
noise = torch.randn(noise_shape, device=device)
|
||||||
|
samples_in[:, :, 0, :, :] = noise[:, :, 0, :, :]
|
||||||
|
samples_in[:, :, -1, :, :] = noise[:, :, -1, :, :]
|
||||||
|
samples_in = samples_in.to(dtype).to(device)
|
||||||
|
else:
|
||||||
|
samples_in = None
|
||||||
|
|
||||||
self.model.model.diffusion_model.to(device)
|
self.model.model.diffusion_model.to(device)
|
||||||
ddim_sampler = DDIMSampler(self.model)
|
ddim_sampler = DDIMSampler(self.model)
|
||||||
@@ -683,7 +695,7 @@ class ToonCrafterInterpolation:
|
|||||||
eta=eta,
|
eta=eta,
|
||||||
temporal_length=noise_shape[2],
|
temporal_length=noise_shape[2],
|
||||||
conditional_guidance_scale_temporal=None,
|
conditional_guidance_scale_temporal=None,
|
||||||
x_T=None,
|
x_T=samples_in,
|
||||||
fs=fs,
|
fs=fs,
|
||||||
timestep_spacing=timestep_spacing,
|
timestep_spacing=timestep_spacing,
|
||||||
guidance_rescale=guidance_rescale,
|
guidance_rescale=guidance_rescale,
|
||||||
@@ -692,6 +704,7 @@ class ToonCrafterInterpolation:
|
|||||||
x0=None,
|
x0=None,
|
||||||
frame_window_size = 16,
|
frame_window_size = 16,
|
||||||
frame_window_stride = 4,
|
frame_window_stride = 4,
|
||||||
|
noise_multiplier = latent_noise_multiplier
|
||||||
)
|
)
|
||||||
print(f"Sampled {i+1} out of {(len(images) - 1)}")
|
print(f"Sampled {i+1} out of {(len(images) - 1)}")
|
||||||
assert not torch.isnan(samples).any().item(), "Resulting tensor containts NaNs. I'm unsure why this happens, changing step count and/or image dimensions might help."
|
assert not torch.isnan(samples).any().item(), "Resulting tensor containts NaNs. I'm unsure why this happens, changing step count and/or image dimensions might help."
|
||||||
@@ -709,6 +722,7 @@ class ToonCrafterInterpolation:
|
|||||||
"samples": samples,
|
"samples": samples,
|
||||||
"hidden_states": hidden_states,
|
"hidden_states": hidden_states,
|
||||||
}
|
}
|
||||||
|
|
||||||
return (latent,)
|
return (latent,)
|
||||||
|
|
||||||
class ToonCrafterDecode:
|
class ToonCrafterDecode:
|
||||||
|
|||||||
Reference in New Issue
Block a user