diff --git a/configs/tooncrafter_512_interp.yaml b/configs/tooncrafter_512_interp.yaml index 951175d..7a420b2 100644 --- a/configs/tooncrafter_512_interp.yaml +++ b/configs/tooncrafter_512_interp.yaml @@ -77,17 +77,6 @@ model: dropout: 0.0 lossconfig: target: torch.nn.Identity - - cond_stage_config: - target: .lvdm.modules.encoders.condition.FrozenOpenCLIPEmbedder - params: - freeze: true - layer: "penultimate" - - img_cond_stage_config: - target: .lvdm.modules.encoders.condition.FrozenOpenCLIPImageEmbedderV2 - params: - freeze: true image_proj_stage_config: target: .lvdm.modules.encoders.resampler.Resampler @@ -100,4 +89,16 @@ model: embedding_dim: 1280 output_dim: 1024 ff_mult: 4 - video_length: 16 \ No newline at end of file + video_length: 16 + # cond_stage_config: + # target: .lvdm.modules.encoders.condition.FrozenOpenCLIPEmbedder + # params: + # freeze: true + # layer: "penultimate" + + img_cond_stage_config: + target: .lvdm.modules.encoders.condition.FrozenOpenCLIPImageEmbedderV2 + params: + freeze: true + + \ No newline at end of file diff --git a/lvdm/models/ddpm3d.py b/lvdm/models/ddpm3d.py index 842680c..77e2d25 100644 --- a/lvdm/models/ddpm3d.py +++ b/lvdm/models/ddpm3d.py @@ -359,7 +359,7 @@ class LatentDiffusion(DDPM): """main class""" def __init__(self, first_stage_config, - cond_stage_config, + #cond_stage_config, num_timesteps_cond=None, cond_stage_key="caption", cond_stage_trainable=False, @@ -421,12 +421,12 @@ class LatentDiffusion(DDPM): self.register_buffer('scale_arr', to_torch(scale_arr)) self.instantiate_first_stage(first_stage_config) - self.instantiate_cond_stage(cond_stage_config) + #self.instantiate_cond_stage(cond_stage_config) self.first_stage_config = first_stage_config - self.cond_stage_config = cond_stage_config + #self.cond_stage_config = cond_stage_config self.clip_denoised = False - self.cond_stage_forward = cond_stage_forward + #self.cond_stage_forward = cond_stage_forward self.encoder_type = encoder_type assert(encoder_type in ["2d", "3d"]) self.uncond_prob = uncond_prob @@ -699,9 +699,9 @@ class LatentDiffusion(DDPM): class LatentVisualDiffusion(LatentDiffusion): def __init__(self, img_cond_stage_config, image_proj_stage_config, freeze_embedder=True, *args, **kwargs): super().__init__(*args, **kwargs) - self._init_embedder(img_cond_stage_config, freeze_embedder) + #self._init_embedder(img_cond_stage_config, freeze_embedder) self.image_proj_model = instantiate_from_config(image_proj_stage_config) - + self.embedder = None def _init_embedder(self, config, freeze=True): embedder = instantiate_from_config(config) if freeze: diff --git a/lvdm/modules/encoders/condition.py b/lvdm/modules/encoders/condition.py index 2bf0a2b..d2a55ce 100644 --- a/lvdm/modules/encoders/condition.py +++ b/lvdm/modules/encoders/condition.py @@ -300,6 +300,7 @@ class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder): def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", freeze=True, layer="pooled", antialias=True): super().__init__() + return model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=version, ) del model.transformer diff --git a/nodes.py b/nodes.py index d8e6b2e..ab9f848 100644 --- a/nodes.py +++ b/nodes.py @@ -70,9 +70,6 @@ class DownloadAndLoadDynamiCrafterModel: }), "fp8_unet": ("BOOLEAN", {"default": False}), }, - "optional": { - "opt_openclippath": ("OPENCLIPVISIONPATH",) - } } RETURN_TYPES = ("DCMODEL",) @@ -80,7 +77,7 @@ class DownloadAndLoadDynamiCrafterModel: FUNCTION = "loadmodel" CATEGORY = "DynamiCrafterWrapper" - def loadmodel(self, dtype, model, fp8_unet=False, opt_openclippath=None): + def loadmodel(self, dtype, model, fp8_unet=False): mm.soft_empty_cache() custom_config = { 'dtype': dtype, @@ -119,11 +116,6 @@ class DownloadAndLoadDynamiCrafterModel: print(f"No matching config for model: {model}") config = OmegaConf.load(config_file) - if opt_openclippath is not None: - print("Using open clip from: ", opt_openclippath) - config.model.params.cond_stage_config.params.version = opt_openclippath - config.model.params.img_cond_stage_config.params.version = opt_openclippath - model_config = config.pop("model", OmegaConf.create()) model_config['params']['unet_config']['params']['use_checkpoint']=False self.model = instantiate_from_config(model_config) @@ -429,6 +421,9 @@ class ToonCrafterInterpolation: def INPUT_TYPES(s): return {"required": { "model": ("DCMODEL",), + "clip_vision": ("CLIP_VISION", ), + "positive": ("CONDITIONING", ), + "negative": ("CONDITIONING", ), "images": ("IMAGE",), "steps": ("INT", {"default": 20, "min": 1, "max": 200, "step": 1}), "cfg": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 200.0, "step": 0.01}), @@ -457,7 +452,7 @@ class ToonCrafterInterpolation: FUNCTION = "process" CATEGORY = "DynamiCrafterWrapper" - def process(self, model, images, prompt, cfg, steps, eta, seed, fs, frames, vae_dtype, image_embed_ratio=1.0): + def process(self, model, clip_vision, images, positive, negative, prompt, cfg, steps, eta, seed, fs, frames, vae_dtype, image_embed_ratio=1.0): device = mm.get_torch_device() mm.unload_all_models() mm.soft_empty_cache() @@ -521,20 +516,27 @@ class ToonCrafterInterpolation: self.model.first_stage_model.to('cpu') - self.model.cond_stage_model.to(device) - self.model.embedder.to(device) - self.model.image_proj_model.to(device) + #self.model.cond_stage_model.to(device) + #self.model.embedder.to(device) + - text_emb = self.model.get_learned_conditioning([prompt]) - cond_images = self.model.embedder(image) - cond_images2 = self.model.embedder(image2) + #text_emb = self.model.get_learned_conditioning([prompt]) + + text_emb = positive[0][0].to(device) + #cond_images = self.model.embedder(image) + #cond_images2 = self.model.embedder(image2) + cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))['last_hidden_state'].to(device) + cond_images2 = clip_vision.encode_image(image2.permute(0, 2, 3, 1))['last_hidden_state'].to(device) + + self.model.image_proj_model.to(device) img_emb = self.model.image_proj_model(cond_images) img_emb2 = self.model.image_proj_model(cond_images2) + img_embeds = img_emb * image_embed_ratio + img_emb2 * (1.0 - image_embed_ratio) imtext_cond = torch.cat([text_emb, img_embeds], dim=1) - del cond_images, img_emb, text_emb + del cond_images, img_emb, img_emb2, text_emb fs = torch.tensor([fs], dtype=torch.long, device=self.model.device) cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]} @@ -548,12 +550,15 @@ class ToonCrafterInterpolation: ## construct unconditional guidance if cfg != 1.0: - uc_emb = self.model.get_learned_conditioning([""]) + #uc_emb = self.model.get_learned_conditioning([""]) + uc_emb = negative[0][0].to(device) ## process image embedding token if hasattr(self.model, 'embedder'): uc_img = torch.zeros(noise_shape[0],3,224,224).to(self.model.device) ## img: b c h w >> b l c - uc_img = self.model.embedder(uc_img) + #uc_img = self.model.embedder(uc_img) + uc_img = clip_vision.encode_image(uc_img.permute(0, 2, 3, 1))['last_hidden_state'] + uc_img = uc_img.to(self.model.device) uc_img = self.model.image_proj_model(uc_img) uc_emb = torch.cat([uc_emb, uc_img], dim=1) if isinstance(cond, dict): @@ -564,8 +569,8 @@ class ToonCrafterInterpolation: else: uc = None - self.model.cond_stage_model.to('cpu') - self.model.embedder.to('cpu') + #self.model.cond_stage_model.to('cpu') + #self.model.embedder.to('cpu') self.model.image_proj_model.to('cpu') #inference diff --git a/scripts/evaluation/funcs.py b/scripts/evaluation/funcs.py index 1744a73..bf61bc5 100644 --- a/scripts/evaluation/funcs.py +++ b/scripts/evaluation/funcs.py @@ -16,8 +16,15 @@ def load_model_checkpoint(model, ckpt): state_dict = torch.load(ckpt, map_location="cpu") if "state_dict" in list(state_dict.keys()): state_dict = state_dict["state_dict"] + + filtered_state_dict = { + k: v + for k, v in state_dict.items() + if not (k.startswith("cond_stage_model") or k.startswith("embedder")) + #if not (k.startswith("cond_stage_model")) + } # Filter out keys starting with "cond_stage_model" and "embedder" try: - model.load_state_dict(state_dict, strict=full_strict) + model.load_state_dict(filtered_state_dict, strict=full_strict) except: ## rename the keys for 256x256 model new_pl_sd = OrderedDict() @@ -38,7 +45,7 @@ def load_model_checkpoint(model, ckpt): # model.load_state_dict(new_pl_sd, strict=full_strict) return model - load_checkpoint(model, ckpt, full_strict=True) + load_checkpoint(model, ckpt, full_strict=False) print('>>> model checkpoint loaded.') return model