Merge pull request #81 from Naomi-Ken-Korem/fix-DynamiCrafterBatchInterpolation

Use external clip vision model in DynamiCrafterBatchInterpolation
This commit is contained in:
Jukka Seppänen
2024-06-28 13:58:10 +03:00
committed by GitHub
+12 -18
View File
@@ -786,12 +786,14 @@ class DynamiCrafterBatchInterpolation:
def INPUT_TYPES(s):
return {"required": {
"model": ("DCMODEL",),
"clip_vision": ("CLIP_VISION",),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
"images": ("IMAGE",),
"steps": ("INT", {"default": 50, "min": 1, "max": 200, "step": 1}),
"cfg": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 20.0, "step": 0.01}),
"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 20.0, "step": 0.01}),
"frames": ("INT", {"default": 16, "min": 1, "max": 100, "step": 1}),
"prompt": ("STRING", {"multiline": True, "default": "",}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"fs": ("INT", {"default": 10, "min": 2, "max": 100, "step": 1}),
"keep_model_loaded": ("BOOLEAN", {"default": True}),
@@ -813,7 +815,8 @@ class DynamiCrafterBatchInterpolation:
FUNCTION = "process"
CATEGORY = "DynamiCrafterWrapper"
def process(self, model, images, prompt, cfg, steps, eta, seed, fs, keep_model_loaded, frames, vae_dtype, cut_near_keyframes):
def process(self, model, images, clip_vision, positive, negative, cfg, steps, eta, seed, fs, keep_model_loaded,
frames, vae_dtype, cut_near_keyframes):
assert images.shape[0] > 1, "DynamiCrafterBatchInterpolation needs at least 2 images"
device = mm.get_torch_device()
mm.unload_all_models()
@@ -847,7 +850,6 @@ class DynamiCrafterBatchInterpolation:
if orig_H % 64 != 0 or orig_W % 64 != 0:
images = F.interpolate(images, size=(H, W), mode="bicubic")
split_prompt = split_and_trim(prompt)
out = []
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():
@@ -869,18 +871,11 @@ class DynamiCrafterBatchInterpolation:
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)
try:
text_emb = self.model.get_learned_conditioning([split_prompt[i]])
print("Prompt: ", split_prompt[i])
except:
text_emb = self.model.get_learned_conditioning([split_prompt[0]])
print("Prompt: ", split_prompt[0])
text_emb = positive[0][0].to(device)
cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))['last_hidden_state'].to(device)
cond_images = self.model.embedder(image)
img_emb = self.model.image_proj_model(cond_images)
imtext_cond = torch.cat([text_emb, img_emb], dim=1)
@@ -895,13 +890,14 @@ class DynamiCrafterBatchInterpolation:
guidance_rescale = 0.7
## construct unconditional guidance
if cfg != 1.0:
uc_emb = self.model.get_learned_conditioning([""])
if cfg != 1.0:
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)
uc_img = torch.rand(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 = clip_vision.encode_image(uc_img.permute(0, 2, 3, 1))['last_hidden_state'].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):
@@ -911,8 +907,6 @@ class DynamiCrafterBatchInterpolation:
uc = uc_emb
else:
uc = None
self.model.cond_stage_model.to('cpu')
self.model.embedder.to('cpu')
self.model.image_proj_model.to('cpu')