Add node to caption a folder of images

This commit is contained in:
kijai
2024-08-22 21:39:13 +03:00
parent 358668282b
commit 58448eee87
+91
View File
@@ -14,6 +14,9 @@ import hashlib
import warnings
import comfy.model_management as mm
import folder_paths
from comfy_execution.graph_utils import GraphBuilder
from comfy.utils import ProgressBar
script_directory = os.path.dirname(os.path.abspath(__file__))
@@ -202,11 +205,99 @@ class LLaVA_OneVision_Run:
return (text_outputs[0],)
class OneVisionCaptionFolder:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"llava_model": ("LLAVAMODEL", ),
"folder_path": ("STRING", ),
"prompt": ("STRING", {"default": "", "multiline": True} ),
"max_tokens": ("INT", {"default": 4096, "min": 1, "max": 8192}),
"keep_model_loaded": ("BOOLEAN", {"default": True}),
"temperature": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"seed": ("INT", {"default": 1, "min": 1, "max": 0xffffffffffffffff}),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "caption"
OUTPUT_IS_LIST = (True,)
CATEGORY = "LLaVA-OneVision"
def caption(self, folder_path, llava_model, prompt, max_tokens, keep_model_loaded, temperature, seed):
from PIL import Image
image_files = []
for filename in os.listdir(folder_path):
if filename.endswith(('.png', '.jpg', '.jpeg')):
image_files.append(filename)
pbar = ProgressBar(len(image_files))
transform = transforms.ToTensor()
vision_node = LLaVA_OneVision_Run()
results_list = []
for filename in image_files:
img_path = os.path.join(folder_path, filename)
try:
img = Image.open(img_path).convert('RGB')
img_tensor = transform(img)
img_tensor = img_tensor.unsqueeze(0).permute(0, 2, 3, 1)
except IOError:
print(f"Cannot open image: {img_path}")
result, = vision_node.run(
llava_model=llava_model,
image=img_tensor, prompt=prompt,
max_tokens=max_tokens,
keep_model_loaded=keep_model_loaded,
temperature=temperature,
seed=seed)
results_list.append(result)
base_filename = os.path.splitext(img_path)[0]
with open(f'{base_filename}.txt', 'w') as file:
file.write(result)
pbar.update(1)
return (results_list,)
class SaveCaptionToTextFile:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", ),
"filename": ("STRING", ),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "caption"
CATEGORY = "LLaVA-OneVision"
def caption(self, txt, filename):
print("SaveCaptionToTextFile: ", txt)
return txt,
NODE_CLASS_MAPPINGS = {
"DownloadAndLoadLLaVAOneVisionModel": DownloadAndLoadLLaVAOneVisionModel,
"LLaVA_OneVision_Run": LLaVA_OneVision_Run,
"OneVisionCaptionFolder": OneVisionCaptionFolder,
"SaveCaptionToTextFile": SaveCaptionToTextFile,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadLLaVAOneVisionModel": "(Down)Load LLaVA-OneVision Model",
"LLaVA_OneVision_Run": "LLaVA-OneVision Run",
"OneVisionCaptionFolder": "OneVision Caption Folder",
"SaveCaptionToTextFile": "SaveCaptionToTextFile",
}