commit 2970a27d4cf066de17c8de217986d218d2d14039 Author: jsonL <306049768@qq.com> Date: Mon Aug 26 01:09:46 2024 +0800 first commit diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 0000000..dfe0770 --- /dev/null +++ b/.gitattributes @@ -0,0 +1,2 @@ +# Auto detect text files and perform LF normalization +* text=auto diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..bd13e80 --- /dev/null +++ b/.gitignore @@ -0,0 +1,9 @@ +.DS_Store +*pyc +.vscode +__pycache__ +*.egg-info +*.bak +checkpoints +results +backup \ No newline at end of file diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..061b495 --- /dev/null +++ b/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2024 Jukka Seppänen + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md new file mode 100644 index 0000000..4eb8914 --- /dev/null +++ b/README.md @@ -0,0 +1 @@ +#Comfyui Batch Tagger \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..2e96bd6 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..7d44e5f --- /dev/null +++ b/nodes.py @@ -0,0 +1,598 @@ +import torch +import torchvision.transforms.functional as F +import io +import os +import matplotlib +matplotlib.use('Agg') +import matplotlib.pyplot as plt +import matplotlib.patches as patches +from PIL import Image, ImageDraw, ImageColor, ImageFont +import random +import numpy as np +import re +from pathlib import Path + +#workaround for unnecessary flash_attn requirement +from unittest.mock import patch +from transformers.dynamic_module_utils import get_imports + +def fixed_get_imports(filename: str | os.PathLike) -> list[str]: + if not str(filename).endswith("modeling_florence2.py"): + return get_imports(filename) + imports = get_imports(filename) + imports.remove("flash_attn") + return imports + + +import comfy.model_management as mm +from comfy.utils import ProgressBar +import folder_paths + +script_directory = os.path.dirname(os.path.abspath(__file__)) + +from transformers import AutoModelForCausalLM, AutoProcessor, set_seed + +class DownloadAndLoadFlorence2Model: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "model": ( + [ + 'microsoft/Florence-2-base', + 'microsoft/Florence-2-base-ft', + 'microsoft/Florence-2-large', + 'microsoft/Florence-2-large-ft', + 'HuggingFaceM4/Florence-2-DocVQA', + 'thwri/CogFlorence-2.1-Large', + 'thwri/CogFlorence-2.2-Large', + 'gokaygokay/Florence-2-SD3-Captioner', + 'MiaoshouAI/Florence-2-base-PromptGen' + ], + { + "default": 'microsoft/Florence-2-base' + }), + "precision": ([ 'fp16','bf16','fp32'], + { + "default": 'fp16' + }), + "attention": ( + [ 'flash_attention_2', 'sdpa', 'eager'], + { + "default": 'sdpa' + }), + }, + "optional": { + "lora": ("PEFTLORA",), + } + } + + RETURN_TYPES = ("FL2MODEL",) + RETURN_NAMES = ("florence2_model",) + FUNCTION = "loadmodel" + CATEGORY = "Florence2" + + def loadmodel(self, model, precision, attention, lora=None): + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + + model_name = model.rsplit('/', 1)[-1] + model_path = os.path.join(folder_paths.models_dir, "LLM", model_name) + + if not os.path.exists(model_path): + print(f"Downloading Florence2 model to: {model_path}") + from huggingface_hub import snapshot_download + snapshot_download(repo_id=model, + local_dir=model_path, + local_dir_use_symlinks=False) + + print(f"using {attention} for attention") + with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): #workaround for unnecessary flash_attn requirement + model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, device_map=device, torch_dtype=dtype,trust_remote_code=True) + processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True) + + if lora is not None: + from peft import PeftModel + adapter_name = lora + model = PeftModel.from_pretrained(model, adapter_name, trust_remote_code=True) + + florence2_model = { + 'model': model, + 'processor': processor, + 'dtype': dtype + } + + return (florence2_model,) + +class DownloadAndLoadFlorence2Lora: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "model": ( + [ + 'NikshepShetty/Florence-2-pixelprose', + ], + ), + }, + + } + + RETURN_TYPES = ("PEFTLORA",) + RETURN_NAMES = ("lora",) + FUNCTION = "loadmodel" + CATEGORY = "Florence2" + + def loadmodel(self, model): + model_name = model.rsplit('/', 1)[-1] + model_path = os.path.join(folder_paths.models_dir, "LLM", model_name) + + if not os.path.exists(model_path): + print(f"Downloading Florence2 lora model to: {model_path}") + from huggingface_hub import snapshot_download + snapshot_download(repo_id=model, + local_dir=model_path, + local_dir_use_symlinks=False) + return (model_path,) + +class Florence2ModelLoader: + + @classmethod + def INPUT_TYPES(s): + return {"required": { + "model": ([item.name for item in Path(folder_paths.models_dir, "LLM").iterdir() if item.is_dir()],), + "precision": (['fp16','bf16','fp32'],), + "attention": ( + [ 'flash_attention_2', 'sdpa', 'eager'], + { + "default": 'sdpa' + }), + }, + "optional": { + "lora": ("PEFTLORA",), + } + } + + RETURN_TYPES = ("FL2MODEL",) + RETURN_NAMES = ("florence2_model",) + FUNCTION = "loadmodel" + CATEGORY = "Florence2" + + def loadmodel(self, model, precision, attention, lora=None): + device = mm.get_torch_device() + dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + model_path = Path(folder_paths.models_dir, "LLM", model) + print(f"Loading model from {model_path}") + print(f"using {attention} for attention") + with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): #workaround for unnecessary flash_attn requirement + model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, device_map=device, torch_dtype=dtype,trust_remote_code=True) + processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True) + + if lora is not None: + from peft import PeftModel + adapter_name = lora + model = PeftModel.from_pretrained(model, adapter_name, trust_remote_code=True) + + florence2_model = { + 'model': model, + 'processor': processor, + 'dtype': dtype + } + + return (florence2_model,) + +class Florence2Run: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE", ), + "florence2_model": ("FL2MODEL", ), + "text_input": ("STRING", {"default": "", "multiline": True}), + "task": ( + [ + 'region_caption', + 'dense_region_caption', + 'region_proposal', + 'caption', + 'detailed_caption', + 'more_detailed_caption', + 'caption_to_phrase_grounding', + 'referring_expression_segmentation', + 'ocr', + 'ocr_with_region', + 'docvqa', + 'prompt_gen' + ], + ), + "fill_mask": ("BOOLEAN", {"default": True}), + }, + "optional": { + "keep_model_loaded": ("BOOLEAN", {"default": False}), + "max_new_tokens": ("INT", {"default": 1024, "min": 1, "max": 4096}), + "num_beams": ("INT", {"default": 3, "min": 1, "max": 64}), + "do_sample": ("BOOLEAN", {"default": True}), + "output_mask_select": ("STRING", {"default": ""}), + "seed": ("INT", {"default": 1, "min": 1, "max": 0xffffffffffffffff}), + } + } + + RETURN_TYPES = ("IMAGE", "MASK", "STRING", "JSON") + RETURN_NAMES =("image", "mask", "caption", "data") + FUNCTION = "encode" + CATEGORY = "Florence2" + + def hash_seed(self, seed): + import hashlib + # Convert the seed to a string and then to bytes + seed_bytes = str(seed).encode('utf-8') + # Create a SHA-256 hash of the seed bytes + hash_object = hashlib.sha256(seed_bytes) + # Convert the hash to an integer + hashed_seed = int(hash_object.hexdigest(), 16) + # Ensure the hashed seed is within the acceptable range for set_seed + return hashed_seed % (2**32) + + def encode(self, image, text_input, florence2_model, task, fill_mask, keep_model_loaded=False, + num_beams=3, max_new_tokens=1024, do_sample=True, output_mask_select="", seed=None): + device = mm.get_torch_device() + _, height, width, _ = image.shape + offload_device = mm.unet_offload_device() + annotated_image_tensor = None + mask_tensor = None + processor = florence2_model['processor'] + model = florence2_model['model'] + dtype = florence2_model['dtype'] + model.to(device) + + if seed: + set_seed(self.hash_seed(seed)) + + colormap = ['blue','orange','green','purple','brown','pink','olive','cyan','red', + 'lime','indigo','violet','aqua','magenta','gold','tan','skyblue'] + + prompts = { + 'region_caption': '', + 'dense_region_caption': '', + 'region_proposal': '', + 'caption': '', + 'detailed_caption': '', + 'more_detailed_caption': '', + 'caption_to_phrase_grounding': '', + 'referring_expression_segmentation': '', + 'ocr': '', + 'ocr_with_region': '', + 'docvqa': '', + 'prompt_gen': '', + } + task_prompt = prompts.get(task, '') + + if (task not in ['referring_expression_segmentation', 'caption_to_phrase_grounding', 'docvqa']) and text_input: + raise ValueError("Text input (prompt) is only supported for 'referring_expression_segmentation', 'caption_to_phrase_grounding', and 'docvqa'") + + if text_input != "": + prompt = task_prompt + " " + text_input + else: + prompt = task_prompt + + image = image.permute(0, 3, 1, 2) + + out = [] + out_masks = [] + out_results = [] + out_data = [] + pbar = ProgressBar(len(image)) + for img in image: + image_pil = F.to_pil_image(img) + inputs = processor(text=prompt, images=image_pil, return_tensors="pt", do_rescale=False).to(dtype).to(device) + + generated_ids = model.generate( + input_ids=inputs["input_ids"], + pixel_values=inputs["pixel_values"], + max_new_tokens=max_new_tokens, + do_sample=do_sample, + num_beams=num_beams, + ) + + results = processor.batch_decode(generated_ids, skip_special_tokens=False)[0] + print(results) + # cleanup the special tokens from the final list + if task == 'ocr_with_region': + clean_results = str(results) + cleaned_string = re.sub(r'|<[^>]*>', '\n', clean_results) + clean_results = re.sub(r'\n+', '\n', cleaned_string) + else: + clean_results = str(results) + clean_results = clean_results.replace('', '') + clean_results = clean_results.replace('', '') + + #return single string if only one image for compatibility with nodes that can't handle string lists + if len(image) == 1: + out_results = clean_results + else: + out_results.append(clean_results) + + W, H = image_pil.size + + parsed_answer = processor.post_process_generation(results, task=task_prompt, image_size=(W, H)) + + if task == 'region_caption' or task == 'dense_region_caption' or task == 'caption_to_phrase_grounding' or task == 'region_proposal': + fig, ax = plt.subplots(figsize=(W / 100, H / 100), dpi=100) + fig.subplots_adjust(left=0, right=1, top=1, bottom=0) + ax.imshow(image_pil) + bboxes = parsed_answer[task_prompt]['bboxes'] + labels = parsed_answer[task_prompt]['labels'] + + mask_indexes = [] + # Determine mask indexes outside the loop + if output_mask_select != "": + mask_indexes = [n for n in output_mask_select.split(",")] + print(mask_indexes) + else: + mask_indexes = [str(i) for i in range(len(bboxes))] + + # Initialize mask_layer only if needed + if fill_mask: + mask_layer = Image.new('RGB', image_pil.size, (0, 0, 0)) + mask_draw = ImageDraw.Draw(mask_layer) + + for index, (bbox, label) in enumerate(zip(bboxes, labels)): + # Modify the label to include the index + indexed_label = f"{index}.{label}" + + if fill_mask: + if str(index) in mask_indexes: + print("match index:", str(index), "in mask_indexes:", mask_indexes) + mask_draw.rectangle([bbox[0], bbox[1], bbox[2], bbox[3]], fill=(255, 255, 255)) + if label in mask_indexes: + print("match label") + mask_draw.rectangle([bbox[0], bbox[1], bbox[2], bbox[3]], fill=(255, 255, 255)) + + # Create a Rectangle patch + rect = patches.Rectangle( + (bbox[0], bbox[1]), # (x,y) - lower left corner + bbox[2] - bbox[0], # Width + bbox[3] - bbox[1], # Height + linewidth=1, + edgecolor='r', + facecolor='none', + label=indexed_label + ) + # Calculate text width with a rough estimation + text_width = len(label) * 6 # Adjust multiplier based on your font size + text_height = 12 # Adjust based on your font size + + # Initial text position + text_x = bbox[0] + text_y = bbox[1] - text_height # Position text above the top-left of the bbox + + # Adjust text_x if text is going off the left or right edge + if text_x < 0: + text_x = 0 + elif text_x + text_width > W: + text_x = W - text_width + + # Adjust text_y if text is going off the top edge + if text_y < 0: + text_y = bbox[3] # Move text below the bottom-left of the bbox if it doesn't overlap with bbox + + # Add the rectangle to the plot + ax.add_patch(rect) + facecolor = random.choice(colormap) if len(image) == 1 else 'red' + # Add the label + plt.text( + text_x, + text_y, + indexed_label, + color='white', + fontsize=12, + bbox=dict(facecolor=facecolor, alpha=0.5) + ) + if fill_mask: + mask_tensor = F.to_tensor(mask_layer) + mask_tensor = mask_tensor.unsqueeze(0).permute(0, 2, 3, 1).cpu().float() + mask_tensor = mask_tensor.mean(dim=0, keepdim=True) + mask_tensor = mask_tensor.repeat(1, 1, 1, 3) + mask_tensor = mask_tensor[:, :, :, 0] + out_masks.append(mask_tensor) + + # Remove axis and padding around the image + ax.axis('off') + ax.margins(0,0) + ax.get_xaxis().set_major_locator(plt.NullLocator()) + ax.get_yaxis().set_major_locator(plt.NullLocator()) + fig.canvas.draw() + buf = io.BytesIO() + plt.savefig(buf, format='png', pad_inches=0) + buf.seek(0) + annotated_image_pil = Image.open(buf) + + annotated_image_tensor = F.to_tensor(annotated_image_pil) + out_tensor = annotated_image_tensor[:3, :, :].unsqueeze(0).permute(0, 2, 3, 1).cpu().float() + out.append(out_tensor) + + out_data.append(bboxes) + + + pbar.update(1) + + plt.close(fig) + + elif task == 'referring_expression_segmentation': + # Create a new black image + mask_image = Image.new('RGB', (W, H), 'black') + mask_draw = ImageDraw.Draw(mask_image) + + predictions = parsed_answer[task_prompt] + + # Iterate over polygons and labels + for polygons, label in zip(predictions['polygons'], predictions['labels']): + color = random.choice(colormap) + for _polygon in polygons: + _polygon = np.array(_polygon).reshape(-1, 2) + # Clamp polygon points to image boundaries + _polygon = np.clip(_polygon, [0, 0], [W - 1, H - 1]) + if len(_polygon) < 3: + print('Invalid polygon:', _polygon) + continue + + _polygon = _polygon.reshape(-1).tolist() + + # Draw the polygon + if fill_mask: + overlay = Image.new('RGBA', image_pil.size, (255, 255, 255, 0)) + image_pil = image_pil.convert('RGBA') + draw = ImageDraw.Draw(overlay) + color_with_opacity = ImageColor.getrgb(color) + (180,) + draw.polygon(_polygon, outline=color, fill=color_with_opacity, width=3) + image_pil = Image.alpha_composite(image_pil, overlay) + else: + draw = ImageDraw.Draw(image_pil) + draw.polygon(_polygon, outline=color, width=3) + + #draw mask + mask_draw.polygon(_polygon, outline="white", fill="white") + + image_tensor = F.to_tensor(image_pil) + image_tensor = image_tensor[:3, :, :].unsqueeze(0).permute(0, 2, 3, 1).cpu().float() + out.append(image_tensor) + + mask_tensor = F.to_tensor(mask_image) + mask_tensor = mask_tensor.unsqueeze(0).permute(0, 2, 3, 1).cpu().float() + mask_tensor = mask_tensor.mean(dim=0, keepdim=True) + mask_tensor = mask_tensor.repeat(1, 1, 1, 3) + mask_tensor = mask_tensor[:, :, :, 0] + out_masks.append(mask_tensor) + pbar.update(1) + + elif task == 'ocr_with_region': + try: + font = ImageFont.load_default().font_variant(size=24) + except: + font = ImageFont.load_default() + predictions = parsed_answer[task_prompt] + scale = 1 + draw = ImageDraw.Draw(image_pil) + bboxes, labels = predictions['quad_boxes'], predictions['labels'] + + for box, label in zip(bboxes, labels): + scaled_box = [ v / (width if idx % 2 == 0 else height) for idx, v in enumerate(box)] + out_data.append({"label": label, "box": scaled_box}) + + color = random.choice(colormap) + new_box = (np.array(box) * scale).tolist() + draw.polygon(new_box, width=3, outline=color) + draw.text((new_box[0]+8, new_box[1]+2), + "{}".format(label), + align="right", + font=font, + fill=color) + + image_tensor = F.to_tensor(image_pil) + image_tensor = image_tensor[:3, :, :].unsqueeze(0).permute(0, 2, 3, 1).cpu().float() + out.append(image_tensor) + + elif task == 'docvqa': + if text_input == "": + raise ValueError("Text input (prompt) is required for 'docvqa'") + prompt = " " + text_input + + inputs = processor(text=prompt, images=image_pil, return_tensors="pt", do_rescale=False).to(dtype).to(device) + generated_ids = model.generate( + input_ids=inputs["input_ids"], + pixel_values=inputs["pixel_values"], + max_new_tokens=max_new_tokens, + do_sample=do_sample, + num_beams=num_beams, + ) + + results = processor.batch_decode(generated_ids, skip_special_tokens=False)[0] + clean_results = results.replace('', '').replace('', '') + + if len(image) == 1: + out_results = clean_results + else: + out_results.append(clean_results) + + out.append(F.to_tensor(image_pil).unsqueeze(0).permute(0, 2, 3, 1).cpu().float()) + + pbar.update(1) + + if len(out) > 0: + out_tensor = torch.cat(out, dim=0) + else: + out_tensor = torch.zeros((1, 64,64, 3), dtype=torch.float32, device="cpu") + if len(out_masks) > 0: + out_mask_tensor = torch.cat(out_masks, dim=0) + else: + out_mask_tensor = torch.zeros((1,64,64), dtype=torch.float32, device="cpu") + + if not keep_model_loaded: + print("Offloading model...") + model.to(offload_device) + mm.soft_empty_cache() + + return (out_tensor, out_mask_tensor, out_results, out_data) + +class Florence2Run_save: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "text": ("STRING", {"forceInput": True}), + "path": ("STRING", {"default": './ComfyUI/output/[time(%Y-%m-%d)]', "multiline": False}), + + } + } + + OUTPUT_NODE = True + RETURN_TYPES = () + FUNCTION = "save_text_file" + CATEGORY = "file" + + def save_text_file(self, text, path): + + if not os.path.exists(path): + print(f"The path `{path}` doesn't exist! Creating it...").warning.print() + try: + os.makedirs(path, exist_ok=True) + except OSError as e: + print(f"The path `{path}` could not be created! Is there write access?\n{e}").error.print() + + # print(type(text)) + image_name_list = [f for f in os.listdir(path) if os.path.isfile(os.path.join(path, f))] + listText = text + for index in range(len(listText)): + text = listText[index] + file_extension = '.txt' + fileName = str(image_name_list[index]).replace(".png","").replace(".jpg","").replace(".jpeg","") + # 获取path下所有文件名 + file_path = os.path.join(path, fileName+file_extension) + + self.writeTextFile(file_path, text) + + return (text, { "ui": { "string": text } } ) + + def writeTextFile(self, file, content): + try: + with open(file, 'w', encoding='utf-8', newline='\n') as f: + f.write(content) + except OSError: + print(f"Unable to save file `{file}`").error.print() + +NODE_CLASS_MAPPINGS = { + "DownloadAndLoadFlorence2Model_jsonL": DownloadAndLoadFlorence2Model, + "DownloadAndLoadFlorence2Lora_jsonL": DownloadAndLoadFlorence2Lora, + "Florence2ModelLoader_jsonL": Florence2ModelLoader, + "Florence2Run_jsonL": Florence2Run, + "Florence2Run_save_file_jsonL":Florence2Run_save +} +NODE_DISPLAY_NAME_MAPPINGS = { + "DownloadAndLoadFlorence2Model_jsonL": "DownloadAndLoadFlorence2Model_jsonL", + "DownloadAndLoadFlorence2Lora_jsonL": "DownloadAndLoadFlorence2Lora_jsonL", + "Florence2ModelLoader_jsonL": "Florence2ModelLoader_jsonL", + "Florence2Run_jsonL": "Florence2Run_jsonL", + "Florence2Run_save_file_jsonL": "Florence2Run_save_file_jsonL", +} diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..e6cdfce --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,15 @@ +[project] +name = "comfyui-tagger" +description = "Nodes to use Florence2 VLM for image vision tasks: object detection, captioning, segmentation and ocr" +version = "1.0.0" +license = "MIT" +dependencies = ["transformers>=4.38.0"] + +[project.urls] +Repository = "https://github.com/StarMagicAI/comfyui_tagger" +# Used by Comfy Registry https://comfyregistry.org + +[tool.comfy] +PublisherId = "jsonL" +DisplayName = "ComfyUI-tagger" +Icon = "" \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..d969961 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +transformers>=4.39.0 +matplotlib +timm +pillow >= 10.2.0