Add model unloading option

This commit is contained in:
kijai
2024-02-29 10:52:40 +02:00
parent 4fcf81e3c3
commit 1427a29689
+29 -18
View File
@@ -19,6 +19,7 @@ class MoondreamQuery:
return {"required": { return {"required": {
"images": ("IMAGE", ), "images": ("IMAGE", ),
"question": ("STRING", {"multiline": True, "default": "What is this?",}), "question": ("STRING", {"multiline": True, "default": "What is this?",}),
"keep_model_loaded": ("BOOLEAN", {"default": True}),
}, },
} }
@@ -28,43 +29,53 @@ class MoondreamQuery:
CATEGORY = "Moondream" CATEGORY = "Moondream"
def process(self, images, question): def process(self, images, question, keep_model_loaded):
batch_size = images.shape[0] batch_size = images.shape[0]
device = comfy.model_management.get_torch_device() device = comfy.model_management.get_torch_device()
dtype = torch.float16 if comfy.model_management.should_use_fp16() and not comfy.model_management.is_device_mps(device) else torch.float32 dtype = torch.float16 if comfy.model_management.should_use_fp16() and not comfy.model_management.is_device_mps(device) else torch.float32
checkpoint_path = os.path.join(script_directory, f"checkpoints/moondream1") checkpoint_path = os.path.join(script_directory, f"checkpoints/moondream1")
if os.path.exists(checkpoint_path): if not hasattr(self, "moondream") or self.moondream is None:
checkpoint_path = checkpoint_path if os.path.exists(checkpoint_path):
else: checkpoint_path = checkpoint_path
try: else:
from huggingface_hub import snapshot_download try:
snapshot_download(repo_id=f"vikhyatk/moondream1", ignore_patterns=["*.jpg","*.pt","*.bin", "*0000*"],local_dir=checkpoint_path, local_dir_use_symlinks=False) from huggingface_hub import snapshot_download
except: snapshot_download(repo_id=f"vikhyatk/moondream1", ignore_patterns=["*.jpg","*.pt","*.bin", "*0000*"],local_dir=checkpoint_path, local_dir_use_symlinks=False)
raise FileNotFoundError("No model found.") except:
raise FileNotFoundError("No model found.")
tokenizer = Tokenizer.from_pretrained(checkpoint_path)
moondream = Moondream.from_pretrained(checkpoint_path).to(device=device, dtype=dtype) self.tokenizer = Tokenizer.from_pretrained(checkpoint_path)
moondream.eval() self.moondream = Moondream.from_pretrained(checkpoint_path).to(device=device, dtype=dtype)
self.moondream.eval()
answer_dict = {} answer_dict = {}
if batch_size > 1: if batch_size > 1:
for i in range(batch_size): for i in range(batch_size):
image = Image.fromarray(np.clip(255. * images[i].cpu().numpy(),0,255).astype(np.uint8)) image = Image.fromarray(np.clip(255. * images[i].cpu().numpy(),0,255).astype(np.uint8))
image_embeds = moondream.encode_image(image) image_embeds = self.moondream.encode_image(image)
answer = moondream.answer_question(image_embeds, question, tokenizer) answer = self.moondream.answer_question(image_embeds, question, self.tokenizer)
answer_dict[str(i)] = answer answer_dict[str(i)] = answer
formatted_answers = ",\n".join([f'"{frame}" : "{answer}"' for frame, answer in answer_dict.items()]) formatted_answers = ",\n".join([f'"{frame}" : "{answer}"' for frame, answer in answer_dict.items()])
formatted_output = "{\n" + formatted_answers + "\n}" formatted_output = "{\n" + formatted_answers + "\n}"
print(formatted_output) print(formatted_output)
answer = formatted_output
return formatted_output, return formatted_output,
else: else:
image = Image.fromarray(np.clip(255. * images[0].cpu().numpy(),0,255).astype(np.uint8)) image = Image.fromarray(np.clip(255. * images[0].cpu().numpy(),0,255).astype(np.uint8))
image_embeds = moondream.encode_image(image) image_embeds = self.moondream.encode_image(image)
answer = moondream.answer_question(image_embeds, question, tokenizer) answer = self.moondream.answer_question(image_embeds, question, self.tokenizer)
return answer,
if not keep_model_loaded:
self.moondream = None
self.tokenizer = None
comfy.model_management.soft_empty_cache()
return answer,