Initial support for the CIL (cond-image-leakage) models

https://github.com/thu-ml/cond-image-leakage
This commit is contained in:
kijai
2024-07-01 02:12:08 +03:00
parent f31af60d7d
commit 954f114564
6 changed files with 1360 additions and 32 deletions
File diff suppressed because it is too large Load Diff
Binary file not shown.
Binary file not shown.
+4
View File
@@ -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:
+5 -4
View File
@@ -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:
+111 -28
View File
@@ -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"
}