Initial support for the CIL (cond-image-leakage) models
https://github.com/thu-ml/cond-image-leakage
This commit is contained in:
File diff suppressed because it is too large
Load Diff
Binary file not shown.
Binary file not shown.
@@ -222,6 +222,9 @@ class DDPM(pl.LightningModule):
|
||||
variance = extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape)
|
||||
log_variance = extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape)
|
||||
return mean, variance, log_variance
|
||||
|
||||
def get_sqrt_alpha_t_bar(self,x_start,t):
|
||||
return extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape)
|
||||
|
||||
def predict_start_from_noise(self, x_t, t, noise):
|
||||
return (
|
||||
@@ -703,6 +706,7 @@ class LatentVisualDiffusion(LatentDiffusion):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._init_embedder(img_cond_stage_config, freeze_embedder)
|
||||
self.image_proj_model = instantiate_from_config(image_proj_stage_config)
|
||||
|
||||
def _init_embedder(self, config, freeze=True):
|
||||
embedder = instantiate_from_config(config)
|
||||
if freeze:
|
||||
|
||||
@@ -27,9 +27,9 @@ class DDIMSampler(object):
|
||||
attr = attr.to(torch.device(device))
|
||||
setattr(self, name, attr)
|
||||
|
||||
def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True):
|
||||
def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., ddpm_from=1000, verbose=True):
|
||||
self.ddim_timesteps = make_ddim_timesteps(ddim_discr_method=ddim_discretize, num_ddim_timesteps=ddim_num_steps,
|
||||
num_ddpm_timesteps=self.ddpm_num_timesteps,verbose=verbose)
|
||||
num_ddpm_timesteps=ddpm_from,verbose=verbose)
|
||||
alphas_cumprod = self.model.alphas_cumprod
|
||||
assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, 'alphas have to be defined for each timestep'
|
||||
to_torch = lambda x: x.clone().detach().to(torch.float32).to(self.model.device)
|
||||
@@ -89,7 +89,8 @@ class DDIMSampler(object):
|
||||
fs=None,
|
||||
timestep_spacing='uniform', #uniform_trailing for starting from last timestep
|
||||
guidance_rescale=0.0,
|
||||
noise_multiplier=0,
|
||||
noise_multiplier=1.0,
|
||||
ddpm_from=1000,
|
||||
**kwargs
|
||||
):
|
||||
|
||||
@@ -107,7 +108,7 @@ class DDIMSampler(object):
|
||||
if conditioning.shape[0] != batch_size:
|
||||
print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}")
|
||||
|
||||
self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, verbose=schedule_verbose)
|
||||
self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, ddpm_from=ddpm_from, verbose=schedule_verbose)
|
||||
|
||||
# make shape
|
||||
if len(shape) == 3:
|
||||
|
||||
@@ -46,7 +46,8 @@ class DownloadAndLoadDynamiCrafterModel:
|
||||
"model": (
|
||||
[ 'tooncrafter_512_interp-fp16.safetensors',
|
||||
'dynamicrafter_512_interp_v1_bf16.safetensors',
|
||||
'dynamicrafter_1024_v1_bf16.safetensors'
|
||||
'dynamicrafter_1024_v1_bf16.safetensors',
|
||||
'DynamiCrafter-CIL-512-no-watermark-fp16.safetensors',
|
||||
],
|
||||
{
|
||||
"default": 'tooncrafter_512_interp-fp16.safetensors'
|
||||
@@ -133,7 +134,12 @@ class DownloadAndLoadDynamiCrafterModel:
|
||||
if fp8_unet:
|
||||
self.model.model.diffusion_model = self.model.model.diffusion_model.to(torch.float8_e4m3fn)
|
||||
print(f"Model using dtype: {self.model.dtype}")
|
||||
return (self.model,)
|
||||
|
||||
dcmodel = {
|
||||
'model': self.model,
|
||||
'model_name': model,
|
||||
}
|
||||
return (dcmodel,)
|
||||
|
||||
class DownloadAndLoadCLIPModel:
|
||||
@classmethod
|
||||
@@ -298,7 +304,11 @@ class DynamiCrafterModelLoader:
|
||||
if fp8_unet:
|
||||
self.model.model.diffusion_model = self.model.model.diffusion_model.to(torch.float8_e4m3fn)
|
||||
print(f"Model using dtype: {self.model.dtype}")
|
||||
return (self.model,)
|
||||
dcmodel = {
|
||||
'model': self.model,
|
||||
'model_name': ckpt_name,
|
||||
}
|
||||
return (dcmodel,)
|
||||
|
||||
class DynamiCrafterI2V:
|
||||
@classmethod
|
||||
@@ -332,7 +342,8 @@ class DynamiCrafterI2V:
|
||||
"mask": ("MASK",),
|
||||
"frame_window_size": ("INT", {"default": 16, "min": 1, "max": 200, "step": 1}),
|
||||
"frame_window_stride": ("INT", {"default": 4, "min": 1, "max": 200, "step": 1}),
|
||||
"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}),
|
||||
"init_noise": ("DCNOISE",),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -342,27 +353,29 @@ class DynamiCrafterI2V:
|
||||
CATEGORY = "DynamiCrafterWrapper"
|
||||
|
||||
def process(self, model, image, clip_vision, positive, negative, cfg, steps, eta, seed, fs, keep_model_loaded,
|
||||
frames, vae_dtype, frame_window_size=16, frame_window_stride=4, mask=None, image2=None, augmentation_level=0):
|
||||
frames, vae_dtype, frame_window_size=16, frame_window_stride=4, mask=None, image2=None, augmentation_level=0, init_noise=None):
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
mm.unload_all_models()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
self.model = model['model']
|
||||
|
||||
torch.manual_seed(seed)
|
||||
dtype = model.dtype
|
||||
dtype = self.model.dtype
|
||||
if vae_dtype == "auto":
|
||||
try:
|
||||
if mm.should_use_bf16():
|
||||
model.first_stage_model.to(convert_dtype('bf16'))
|
||||
self.model.first_stage_model.to(convert_dtype('bf16'))
|
||||
else:
|
||||
model.first_stage_model.to(convert_dtype('fp32'))
|
||||
self.model.first_stage_model.to(convert_dtype('fp32'))
|
||||
except:
|
||||
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtype manually.")
|
||||
else:
|
||||
model.first_stage_model.to(convert_dtype(vae_dtype))
|
||||
print(f"VAE using dtype: {model.first_stage_model.dtype}")
|
||||
self.model.first_stage_model.to(convert_dtype(vae_dtype))
|
||||
print(f"VAE using dtype: {self.model.first_stage_model.dtype}")
|
||||
|
||||
self.model = model
|
||||
|
||||
self.model.to(device)
|
||||
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
|
||||
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||
@@ -461,6 +474,20 @@ class DynamiCrafterI2V:
|
||||
mask = mask.permute(0, 2, 1, 3, 4)
|
||||
mask = torch.where(mask < 1.0, torch.tensor(0.0, device=device, dtype=dtype), torch.tensor(1.0, device=device, dtype=dtype))
|
||||
|
||||
if init_noise is not None:
|
||||
init = init_noise['noise'].to(dtype).to(device)
|
||||
timestep_spacing = "uniform_trailing"
|
||||
guidance_rescale = 0.0
|
||||
ddpm_from = init_noise['M']
|
||||
|
||||
if noise_shape[2] % init.shape[2] == 0:
|
||||
init = init.repeat(1, 1, noise_shape[2] // init.shape[2], 1, 1)
|
||||
else:
|
||||
raise ValueError("The target dimension size is not an integral multiple of the original dimension size.")
|
||||
else:
|
||||
init = None
|
||||
ddpm_from = 1000
|
||||
|
||||
#inference
|
||||
ddim_sampler = DDIMSampler(self.model)
|
||||
samples, _ = ddim_sampler.sample(S=steps,
|
||||
@@ -473,7 +500,7 @@ class DynamiCrafterI2V:
|
||||
eta=eta,
|
||||
temporal_length=noise_shape[2],
|
||||
conditional_guidance_scale_temporal=None,
|
||||
x_T=None,
|
||||
x_T=init,
|
||||
fs=fs,
|
||||
timestep_spacing=timestep_spacing,
|
||||
guidance_rescale=guidance_rescale,
|
||||
@@ -481,7 +508,9 @@ class DynamiCrafterI2V:
|
||||
mask=mask,
|
||||
x0=img_tensor_repeat.clone() if mask is not None else None,
|
||||
frame_window_size = frame_window_size,
|
||||
frame_window_stride = frame_window_stride
|
||||
frame_window_stride = frame_window_stride,
|
||||
noise_multiplier=1.0,
|
||||
ddpm_from=ddpm_from
|
||||
)
|
||||
|
||||
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."
|
||||
@@ -509,6 +538,56 @@ class DynamiCrafterI2V:
|
||||
video = F.interpolate(video.permute(0, 3, 1, 2), size=(final_H, final_W), mode="bicubic").permute(0, 2, 3, 1)
|
||||
last_image = video[-1].unsqueeze(0)
|
||||
return (video, last_image)
|
||||
|
||||
class DynamiCrafterLoadInitNoise:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model": ("DCMODEL",),
|
||||
"M": ("INT", {"default": 1000, "min": 1, "max": 1000, "step": 1}),
|
||||
"analytic_init": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("DCNOISE", "INT", "INT",)
|
||||
RETURN_NAMES = ("init_noise", "width", "height",)
|
||||
FUNCTION = "load"
|
||||
CATEGORY = "DynamiCrafterWrapper"
|
||||
|
||||
def load(self, model, M, analytic_init):
|
||||
device = mm.get_torch_device()
|
||||
|
||||
model_name = model['model_name']
|
||||
if '512' in model_name:
|
||||
analytic_noise = "initial_noise_512.safetensors"
|
||||
elif '1024' in model_name:
|
||||
analytic_noise = "initial_noise_1024.safetensors"
|
||||
else:
|
||||
print("Can't find matching init_noise for model: ", model_name)
|
||||
model_path = os.path.join(script_directory, 'init_noises', analytic_noise)
|
||||
|
||||
# Analytic-Init:load initial noise
|
||||
#dic=torch.load(model_path)
|
||||
dic = comfy.utils.load_torch_file(model_path)
|
||||
expectation_X_0=dic["Expectation_X0"].to(device)
|
||||
tr_Cov_d=dic["Tr_Cov_d"].to(device)
|
||||
sqrt_alpha_t=model['model'].get_sqrt_alpha_t_bar(expectation_X_0,torch.tensor([M-1]).to(device))
|
||||
mu_p=sqrt_alpha_t*expectation_X_0
|
||||
alpha_t=sqrt_alpha_t**2
|
||||
sigma_p=torch.sqrt(1-alpha_t + alpha_t*tr_Cov_d)
|
||||
eps=torch.randn_like(mu_p)
|
||||
|
||||
if analytic_init:
|
||||
init=mu_p+sigma_p*eps
|
||||
else :
|
||||
init=torch.randn_like(mu_p)
|
||||
print("init noise shape: ",init.shape)
|
||||
|
||||
init_noise = {"noise": init, "M": M}
|
||||
width = init.shape[4] * 8
|
||||
height = init.shape[3] * 8
|
||||
|
||||
return (init_noise, width, height)
|
||||
|
||||
class ToonCrafterInterpolation:
|
||||
@classmethod
|
||||
@@ -555,18 +634,21 @@ class ToonCrafterInterpolation:
|
||||
mm.soft_empty_cache()
|
||||
|
||||
torch.manual_seed(seed)
|
||||
dtype = model.dtype
|
||||
|
||||
self.model = model['model']
|
||||
|
||||
dtype = self.model.dtype
|
||||
if vae_dtype == "auto":
|
||||
try:
|
||||
if mm.should_use_bf16():
|
||||
model.first_stage_model.to(convert_dtype('bf16'))
|
||||
self.model.first_stage_model.to(convert_dtype('bf16'))
|
||||
else:
|
||||
model.first_stage_model.to(convert_dtype('fp32'))
|
||||
self.model.first_stage_model.to(convert_dtype('fp32'))
|
||||
except:
|
||||
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtype manually.")
|
||||
else:
|
||||
model.first_stage_model.to(convert_dtype(vae_dtype))
|
||||
print(f"VAE using dtype: {model.first_stage_model.dtype}")
|
||||
self.model.first_stage_model.to(convert_dtype(vae_dtype))
|
||||
print(f"VAE using dtype: {self.model.first_stage_model.dtype}")
|
||||
|
||||
images = images.permute(0, 3, 1, 2).to(dtype).to(device)
|
||||
|
||||
@@ -579,7 +661,6 @@ class ToonCrafterInterpolation:
|
||||
if orig_H % 64 != 0 or orig_W % 64 != 0:
|
||||
images = F.interpolate(images, size=(H, W), mode="bicubic")
|
||||
|
||||
self.model = model
|
||||
self.model.to(device)
|
||||
|
||||
out = []
|
||||
@@ -755,7 +836,7 @@ class ToonCrafterDecode:
|
||||
device = mm.get_torch_device()
|
||||
mm.unload_all_models()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
model = model['model']
|
||||
samples = latent["samples"]
|
||||
num_samples = samples.shape[0]
|
||||
samples = samples * 0.18215
|
||||
@@ -857,20 +938,20 @@ class DynamiCrafterBatchInterpolation:
|
||||
|
||||
torch.manual_seed(seed)
|
||||
dtype = model.dtype
|
||||
self.model = model['model']
|
||||
|
||||
if vae_dtype == "auto":
|
||||
try:
|
||||
if mm.should_use_bf16():
|
||||
model.first_stage_model.to(convert_dtype('bf16'))
|
||||
self.model.first_stage_model.to(convert_dtype('bf16'))
|
||||
else:
|
||||
model.first_stage_model.to(convert_dtype('fp32'))
|
||||
self.model.first_stage_model.to(convert_dtype('fp32'))
|
||||
except:
|
||||
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtype manually.")
|
||||
else:
|
||||
model.first_stage_model.to(convert_dtype(vae_dtype))
|
||||
print(f"VAE using dtype: {model.first_stage_model.dtype}")
|
||||
|
||||
self.model = model
|
||||
self.model.first_stage_model.to(convert_dtype(vae_dtype))
|
||||
print(f"VAE using dtype: {self.model.first_stage_model.dtype}")
|
||||
|
||||
self.model.to(device)
|
||||
images = images * 2 - 1
|
||||
images = images.permute(0, 3, 1, 2).to(dtype).to(device)
|
||||
@@ -1020,7 +1101,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ToonCrafterDecode": ToonCrafterDecode,
|
||||
"DownloadAndLoadDynamiCrafterModel": DownloadAndLoadDynamiCrafterModel,
|
||||
"DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel,
|
||||
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel
|
||||
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel,
|
||||
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise
|
||||
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -1031,5 +1113,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ToonCrafterDecode": "ToonCrafterDecode",
|
||||
"DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel",
|
||||
"DownloadAndLoadCLIPModel": "DownloadAndLoadCLIPModel",
|
||||
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel"
|
||||
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel",
|
||||
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user