Use comfy text encoder and clip_vision
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user