Allow latent input

This commit is contained in:
kijai
2024-06-30 17:37:38 +03:00
parent 04d27efaa3
commit f31af60d7d
2 changed files with 31 additions and 10 deletions
+14 -7
View File
@@ -89,6 +89,7 @@ class DDIMSampler(object):
fs=None,
timestep_spacing='uniform', #uniform_trailing for starting from last timestep
guidance_rescale=0.0,
noise_multiplier=0,
**kwargs
):
@@ -134,6 +135,7 @@ class DDIMSampler(object):
precision=precision,
fs=fs,
guidance_rescale=guidance_rescale,
noise_multiplier=noise_multiplier,
**kwargs)
return samples, intermediates
@@ -143,13 +145,15 @@ class DDIMSampler(object):
callback=None, timesteps=None, quantize_denoised=False,
mask=None, x0=None, img_callback=None, log_every_t=100,
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):
device = self.model.betas.device
b = shape[0]
if x_T is None:
print("Using random noise")
img = torch.randn(shape, device=device)
else:
print("Using input noise")
img = x_T
if precision is not None:
if precision == 16:
@@ -171,6 +175,12 @@ class DDIMSampler(object):
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)
pbar = comfy.utils.ProgressBar(total_steps)
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 = img_orig * mask + (1. - mask) * img # keep original & modify use img
outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps,
outs = self.p_sample_ddim(img, cond, ts, sigmas, index=index, use_original_steps=ddim_use_original_steps,
quantize_denoised=quantize_denoised, temperature=temperature,
noise_dropout=noise_dropout, score_corrector=score_corrector,
corrector_kwargs=corrector_kwargs,
@@ -210,7 +217,7 @@ class DDIMSampler(object):
return img, intermediates
@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,
unconditional_guidance_scale=1., unconditional_conditioning=None,
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
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.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
if is_video:
+17 -3
View File
@@ -537,7 +537,9 @@ class ToonCrafterInterpolation:
},
"optional": {
"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"
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()
offload_device = mm.unet_offload_device()
mm.unload_all_models()
@@ -670,6 +672,16 @@ class ToonCrafterInterpolation:
self.model.image_proj_model.to(offload_device)
#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)
ddim_sampler = DDIMSampler(self.model)
@@ -683,7 +695,7 @@ class ToonCrafterInterpolation:
eta=eta,
temporal_length=noise_shape[2],
conditional_guidance_scale_temporal=None,
x_T=None,
x_T=samples_in,
fs=fs,
timestep_spacing=timestep_spacing,
guidance_rescale=guidance_rescale,
@@ -692,6 +704,7 @@ class ToonCrafterInterpolation:
x0=None,
frame_window_size = 16,
frame_window_stride = 4,
noise_multiplier = latent_noise_multiplier
)
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."
@@ -709,6 +722,7 @@ class ToonCrafterInterpolation:
"samples": samples,
"hidden_states": hidden_states,
}
return (latent,)
class ToonCrafterDecode: