Use comfy text encoder and clip_vision

This commit is contained in:
kijai
2024-06-02 16:44:17 +03:00
parent 055e1c39f6
commit 86e789d35d
5 changed files with 55 additions and 41 deletions
+13 -12
View File
@@ -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
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
+6 -6
View File
@@ -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:
+1
View File
@@ -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
+26 -21
View File
@@ -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
+9 -2
View File
@@ -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