From 4939ec791a89b115120fd12ffaf02b2a966c9766 Mon Sep 17 00:00:00 2001 From: chflame163 Date: Wed, 8 Apr 2026 12:27:27 +0800 Subject: [PATCH] update readme, remove debug info from florence2_ultra.py --- README.md | 1 + README_CN.MD | 1 + florence2_ultra.py | 723 +++++++++++++++++++++++++++++++++++++++++++++ pyproject.toml | 2 +- 4 files changed, 726 insertions(+), 1 deletion(-) create mode 100644 florence2_ultra.py diff --git a/README.md b/README.md index b49f3a1..c388494 100644 --- a/README.md +++ b/README.md @@ -144,6 +144,7 @@ Please try downgrading the ```protobuf``` dependency package to 3.20.3, or set e **If the dependency package error after updating, please double clicking ```repair_dependency.bat``` (for Official ComfyUI Protable) or ```repair_dependency_aki.bat``` (for ComfyUI-aki-v1.x) in the plugin folder to reinstall the dependency packages. +* Fix Florence2 config compatibility with transformers 5.x. * Fix the issue where Florence2 run with higher versions of Transformers, this solution comes from [kijai](https://github.com/kijai/ComfyUI-Florence2), Thanks to @flybirdxx for feedback. After updating plugin, find ```modeling_florence2.py``` and ```configuration_florence2.py``` from the ```florence2_models``` folder, copy and overwrite them to the model folder in ```ComfyUI/models/florence2```. * Commit [JimengImageToImageAPI](#JimengImageToImageAPI) node, edit images using the Instant Dreaming Image 3.0 API. Create an account on [Volcano Engine](#https://console.volcengine.com/iam/keymanage) and apply for API AccessKeyID and SecretAccessKey. Fill them into the ```api_key.ini``` directory in the plugin directory. diff --git a/README_CN.MD b/README_CN.MD index 0afdf61..5c6c630 100644 --- a/README_CN.MD +++ b/README_CN.MD @@ -121,6 +121,7 @@ If this call came from a _pb2.py file, your generated code is out of date and mu ## 更新说明 **如果本插件更新后出现依赖包错误,请双击运行插件目录下的```install_requirements.bat```(官方便携包),或 ```install_requirements_aki.bat```(秋叶整合包) 重新安装依赖包。 +* 修复 Florence2 模型加载,使其兼容transformers 5.x。 * 修复 Florence2 无法在高版本Transformers运行的问题,解决方法来自[kijai](https://github.com/kijai/ComfyUI-Florence2)。感谢 @flybirdxx 的反馈。 更新本插件后,将 florence2_models 文件夹下的 ```modeling_florence2.py``` 和 ```configuration_florence2.py``` 这两个文件复制到 ```ComfyUI/models/florence2``` 里面的模型文件夹下,覆盖同名文件。 * 添加 [JimengImageToImageAPI](#JimengImageToImageAPI) 节点,使用即梦图生图3.0API对图片进行编辑。在[火山引擎](#https://console.volcengine.com/iam/keymanage) 创建账号,并申请API AccessKeyID 和 SecretAccessKey,将其填入插件目录下的```api_key.ini```。 diff --git a/florence2_ultra.py b/florence2_ultra.py new file mode 100644 index 0000000..597069a --- /dev/null +++ b/florence2_ultra.py @@ -0,0 +1,723 @@ +# layerstyle advance + +import io +from unittest.mock import patch +import matplotlib.pyplot as plt +import matplotlib.patches as patches +import colorsys +from transformers.dynamic_module_utils import get_imports +import transformers +transformers.logging.set_verbosity_error() +from packaging import version +import comfy.model_management +from .imagefunc import * + +colormap = ['blue', 'orange', 'green', 'purple', 'brown', 'pink', 'gray', 'olive', 'cyan', 'red', + 'lime', 'indigo', 'violet', 'aqua', 'magenta', 'coral', 'gold', 'tan', 'skyblue'] + +device = comfy.model_management.get_torch_device() + +fl2_model_repos = { + "base": "microsoft/Florence-2-base", + "base-ft": "microsoft/Florence-2-base-ft", + "large": "microsoft/Florence-2-large", + "large-ft": "microsoft/Florence-2-large-ft", + "DocVQA": "HuggingFaceM4/Florence-2-DocVQA", + "SD3-Captioner": "gokaygokay/Florence-2-SD3-Captioner", + "base-PromptGen": "MiaoshouAI/Florence-2-base-PromptGen", + "CogFlorence-2-Large-Freeze": "thwri/CogFlorence-2-Large-Freeze", + "CogFlorence-2.1-Large": "thwri/CogFlorence-2.1-Large", + "base-PromptGen-v1.5":"MiaoshouAI/Florence-2-base-PromptGen-v1.5", + "large-PromptGen-v1.5":"MiaoshouAI/Florence-2-large-PromptGen-v1.5", + "base-PromptGen-v2.0":"MiaoshouAI/Florence-2-base-PromptGen-v2.0", + "large-PromptGen-v2.0":"MiaoshouAI/Florence-2-large-PromptGen-v2.0", + "Florence-2-Flux":"gokaygokay/Florence-2-Flux", + "Florence-2-Flux-Large":"gokaygokay/Florence-2-Flux-Large" +} + +def fixed_get_imports(filename) -> list[str]: + """Workaround for FlashAttention""" + if os.path.basename(filename) != "modeling_florence2.py": + return get_imports(filename) + imports = get_imports(filename) + try: + imports.remove("flash_attn") + except: + pass + return imports + +def _load_model_v5(model_path, attention, dtype): + """Load Florence2 model for transformers >= 5.0.0""" + # log(f"[DEBUG] _load_model_v5 called with model_path={model_path}, attention={attention}, dtype={dtype}") + from ..florence2_models.modeling_florence2 import Florence2ForConditionalGeneration, Florence2Config + from transformers import CLIPImageProcessor, BartTokenizerFast + from ..florence2_models.processing_florence2 import Florence2Processor + from accelerate import init_empty_weights + from accelerate.utils import set_module_tensor_to_device + from comfy.utils import load_torch_file + + offload_device = comfy.model_management.unet_offload_device() + # log(f"[DEBUG] offload_device={offload_device}") + + # log(f"[DEBUG] Loading Florence2Config from {model_path}") + config = Florence2Config.from_pretrained(model_path) + config._attn_implementation = attention + # log(f"[DEBUG] Config loaded, initializing empty model") + with init_empty_weights(): + model = Florence2ForConditionalGeneration(config) + + checkpoint_path = os.path.join(model_path, "model.safetensors") + if not os.path.exists(checkpoint_path): + checkpoint_path = os.path.join(model_path, "pytorch_model.bin") + if os.path.exists(checkpoint_path): + # log(f"[DEBUG] Loading weights from {checkpoint_path}") + state_dict = load_torch_file(checkpoint_path) + # log(f"[DEBUG] Loaded {len(state_dict)} keys from checkpoint") + else: + raise FileNotFoundError(f"No model weights found at {model_path}") + + key_mapping = {} + if "language_model.model.shared.weight" in state_dict: + key_mapping["language_model.model.encoder.embed_tokens.weight"] = "language_model.model.shared.weight" + key_mapping["language_model.model.decoder.embed_tokens.weight"] = "language_model.model.shared.weight" + + missing_keys = [] + for name, param in model.named_parameters(): + actual_key = key_mapping.get(name, name) + if actual_key in state_dict: + set_module_tensor_to_device(model, name, offload_device, value=state_dict[actual_key].to(dtype)) + else: + missing_keys.append(name) + # if missing_keys: + # log(f"[DEBUG] {len(missing_keys)} parameters not found in state_dict: {missing_keys[:5]}{'...' if len(missing_keys) > 5 else ''}", message_type='warning') + + # log(f"[DEBUG] Tying weights and finalizing model") + model.language_model.tie_weights() + model = model.eval().to(dtype).to(offload_device) + + image_processor = CLIPImageProcessor( + do_resize=True, + size={"height": 768, "width": 768}, + resample=3, + do_center_crop=False, + do_rescale=True, + rescale_factor=1/255.0, + do_normalize=True, + image_mean=[0.485, 0.456, 0.406], + image_std=[0.229, 0.224, 0.225], + ) + image_processor.image_seq_length = 577 + + # log(f"[DEBUG] Loading tokenizer from {model_path}") + try: + tokenizer = BartTokenizerFast.from_pretrained(model_path) + except (TypeError, Exception) as e: + # log(f"[DEBUG] BartTokenizerFast failed ({e}), loading from tokenizer.json directly", message_type='warning') + from tokenizers import Tokenizer as TokenizerFast + from transformers import PreTrainedTokenizerFast + import json + tokenizer_json = os.path.join(model_path, "tokenizer.json") + base_tokenizer = TokenizerFast.from_file(tokenizer_json) + # Read special token config + tokenizer_config_path = os.path.join(model_path, "tokenizer_config.json") + special_tokens = {} + if os.path.exists(tokenizer_config_path): + with open(tokenizer_config_path, 'r') as f: + tc = json.load(f) + for key in ('bos_token', 'eos_token', 'unk_token', 'pad_token', 'sep_token', 'cls_token', 'mask_token'): + val = tc.get(key) + if isinstance(val, dict): + val = val.get('content', None) + if val is not None: + special_tokens[key] = val + tokenizer = PreTrainedTokenizerFast(tokenizer_object=base_tokenizer, **special_tokens) + # log(f"[DEBUG] Creating Florence2Processor") + processor = Florence2Processor(image_processor=image_processor, tokenizer=tokenizer) + # log(f"[DEBUG] _load_model_v5 completed successfully") + return model, processor + +def load_model(ver): + florence_path = os.path.join(folder_paths.models_dir, "florence2") + os.makedirs(florence_path, exist_ok=True) + + model_path = os.path.join(florence_path, ver) + attention = 'sdpa' + + if not os.path.exists(model_path): + log(f"Downloading Florence2 {ver} model...") + repo_id = fl2_model_repos[ver] + from huggingface_hub import snapshot_download + snapshot_download(repo_id=repo_id, local_dir=model_path, ignore_patterns=["*.md", "*.txt"]) + + # log(f"[DEBUG] transformers version: {transformers.__version__}, v5+ path: {version.parse(transformers.__version__) >= version.parse('5.0.0')}") + # log(f"[DEBUG] model_path: {model_path}, exists: {os.path.exists(model_path)}") + + if version.parse(transformers.__version__) >= version.parse('5.0.0'): + # log(f"[DEBUG] Using transformers v5 loading path") + model, processor = _load_model_v5(model_path, attention, torch.float32) + # log(f"[DEBUG] Model loaded, model type: {type(model)}, processor type: {type(processor)}") + return (model.to(device), processor) + + try: + with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): + model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, device_map=device, + torch_dtype=torch.float32, trust_remote_code=True) + processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True) + except Exception as e: + try: + model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, device_map=device, + torch_dtype=torch.float32, trust_remote_code=True) + processor = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) + except Exception as e: + sys.path.append(model_path) + # Import the Florence modules + if ver == 'large-PromptGen-v1.5': + from florence2_large.modeling_florence2 import Florence2ForConditionalGeneration + from florence2_large.configuration_florence2 import Florence2Config + elif ver == 'base-PromptGen-v1.5': + from florence2_base_ft.modeling_florence2 import Florence2ForConditionalGeneration + from florence2_base_ft.configuration_florence2 import Florence2Config + else: + log(f"Error loading model or tokenizer: {str(e)}", message_type='error') + return (None, None) + + # Load the model configuration + model_config = Florence2Config.from_pretrained(model_path) + # Load the model + with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): + model = Florence2ForConditionalGeneration.from_pretrained( + model_path, + config=model_config, + attn_implementation=attention, + device_map=device + ).to(device) + + processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True) + + return (model.to(device), processor) + +def fig_to_pil(fig): + buf = io.BytesIO() + fig.savefig(buf, format='png', dpi=100, bbox_inches='tight', pad_inches=0) + buf.seek(0) + pil = Image.open(buf) + plt.close() + return pil + +def plot_bbox(image, data): + fig, ax = plt.subplots() + fig.set_size_inches(image.width / 100, image.height / 100) + ax.imshow(image) + for i, (bbox, label) in enumerate(zip(data['bboxes'], data['labels'])): + x1, y1, x2, y2 = bbox + rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1, linewidth=1, edgecolor='r', facecolor='none') + ax.add_patch(rect) + enum_label = f"{i}: {label}" + plt.text(x1 + 7, y1 + 17, enum_label, color='white', fontsize=8, bbox=dict(facecolor='red', alpha=0.5)) + ax.axis('off') + return fig + +def generate_color(index, total_colors=25): + # Generate color by varying the hue to maximize difference between colors + hue = (index / total_colors) % 1.0 # Normalize hue to be between 0 and 1 + saturation = 0.65 # Keep saturation constant + lightness = 0.5 # Keep lightness constant + + # Convert HSL to RGB, then to hexadecimal + r, g, b = colorsys.hls_to_rgb(hue, lightness, saturation) + return f'#{int(r * 255):02X}{int(g * 255):02X}{int(b * 255):02X}' + +def plot_mask_bbox(image, data): + fig, ax = plt.subplots() + fig.set_size_inches(image.width / 100, image.height / 100) + ax.imshow(image) + num_bboxes = len(data['bboxes']) + for i, (bbox, label) in enumerate(list(zip(data['bboxes'], data['labels']))[1:], start=1): + x1, y1, x2, y2 = bbox + if x2 < x1: + x1, y1, x2, y2 = x2, y2, x1, y1 + color = generate_color(i, total_colors=num_bboxes) + rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1, linewidth=1, edgecolor=color, facecolor='none') + ax.add_patch(rect) + enum_label = f"{i}: {label}" + plt.text(x1 + 7, y1 + 17, enum_label, color='white', fontsize=8, bbox=dict(facecolor=color, alpha=0.5)) + ax.axis('off') + return fig + +def plot_mask(image, data, indexes): + # Create a black background image (mode "1" for binary, "L" for grayscale) + mask = Image.new("L", (image.width, image.height), 0) # Black background + fig, ax = plt.subplots() + fig.set_size_inches(mask.width / 100, mask.height / 100) + ax.imshow(mask, cmap='gray') # Display the mask in grayscale + ax.set_facecolor('black') # Set the axes background to black + fig.patch.set_facecolor('black') # Set the figure background to black + for i, (bbox, label) in enumerate(list(zip(data['bboxes'], data['labels']))[1:], start=1): + x1, y1, x2, y2 = bbox + if x2 < x1: + x1, y1, x2, y2 = x2, y2, x1, y1 + rect = patches.Rectangle((x1, y1), x2 - x1, y2 - y1, linewidth=1, edgecolor='w', facecolor='w') + if i in indexes: + ax.add_patch(rect) + ax.axis('off') + return fig + +def draw_polygons(image, prediction, fill_mask=False): + output_image = copy.deepcopy(image) + draw = ImageDraw.Draw(output_image) + scale = 1 + for polygons, label in zip(prediction['polygons'], prediction['labels']): + color = random.choice(colormap) + fill_color = color if fill_mask else None + for _polygon in polygons: + _polygon = np.array(_polygon).reshape(-1, 2) + if len(_polygon) < 3: + print('Invalid polygon:', _polygon) + continue + _polygon = (_polygon * scale).reshape(-1).tolist() + if fill_mask: + draw.polygon(_polygon, outline=color, fill=fill_color) + else: + draw.polygon(_polygon, outline=color) + draw.text((_polygon[0] + 8, _polygon[1] + 2), label, fill=color) + return output_image + + +def convert_to_od_format(data): + od_results = { + 'bboxes': data.get('bboxes', []), + 'labels': data.get('bboxes_labels', []) + } + return od_results + + +def draw_ocr_bboxes(image, prediction): + scale = 1 + output_image = copy.deepcopy(image) + draw = ImageDraw.Draw(output_image) + bboxes, labels = prediction['quad_boxes'], prediction['labels'] + for box, label in zip(bboxes, labels): + 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", + fill=color) + return output_image + + +def run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input=None): + if text_input is None: + prompt = task_prompt + else: + prompt = task_prompt + text_input + inputs = processor(text=prompt, images=image, return_tensors="pt").to(device) + # log(f"[DEBUG] run_example: image size={image.size if hasattr(image, 'size') else 'N/A'}, pixel_values shape={inputs['pixel_values'].shape}, input_ids shape={inputs['input_ids'].shape}") + generated_ids = model.generate( + input_ids=inputs["input_ids"], + pixel_values=inputs["pixel_values"], + max_new_tokens=max_new_tokens, + early_stopping=False, + do_sample=do_sample, + num_beams=num_beams, + use_cache=False, + ) + generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0] + parsed_answer = processor.post_process_generation( + generated_text, + task=task_prompt, + image_size=(image.width, image.height) + ) + return parsed_answer + + +def process_image(model, processor, image, task_prompt, max_new_tokens, num_beams, do_sample, fill_mask, text_input=None): + if task_prompt == 'caption': + task_prompt = '' + result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + return result[task_prompt], None + elif task_prompt == 'detailed caption': + task_prompt = '' + result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + return result[task_prompt], None + elif task_prompt == 'more detailed caption': + task_prompt = '' + result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + return result[task_prompt], None + elif task_prompt == 'object detection': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + fig = plot_bbox(image, results['']) + return results[task_prompt], fig_to_pil(fig) + elif task_prompt == 'dense region caption': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + fig = plot_bbox(image, results['']) + return results[task_prompt], fig_to_pil(fig) + elif task_prompt == 'region proposal': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + fig = plot_bbox(image, results['']) + return results[task_prompt], fig_to_pil(fig) + elif task_prompt == 'region proposal (mask)': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + indexes = [] + if isinstance(text_input, str): + for i in text_input.split(','): + try: + indexes.append(int(i)) + except ValueError: + print(f"{i} is nit an instance of int") + if len(indexes) > 0: + fig = plot_mask(image, results[''], indexes) + pil = fig_to_pil(fig).resize((image.width, image.height), Image.Resampling.LANCZOS) + else: + fig = plot_mask_bbox(image, results['']) + pil = fig_to_pil(fig) + return results[task_prompt], pil + elif task_prompt == 'caption to phrase grounding': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input) + fig = plot_bbox(image, results['']) + return results[task_prompt], fig_to_pil(fig) + elif task_prompt == 'referring expression segmentation': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input) + output_image = draw_polygons(image, results[''], fill_mask) + return results[task_prompt], output_image + elif task_prompt == 'region to segmentation': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input) + output_image = draw_polygons(image, results[''], fill_mask) + return results[task_prompt], output_image + elif task_prompt == 'open vocabulary detection': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input) + bbox_results = convert_to_od_format(results['']) + fig = plot_bbox(image, bbox_results) + return bbox_results, fig_to_pil(fig) + elif task_prompt == 'region to category': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input) + return results[task_prompt], None + elif task_prompt == 'region to description': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample, text_input) + return results[task_prompt], None + elif task_prompt == 'OCR': + task_prompt = '' + result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + return result[task_prompt], None + elif task_prompt == 'OCR with region': + task_prompt = '' + results = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + output_image = draw_ocr_bboxes(image, results['']) + output_results = {'bboxes': results[task_prompt].get('quad_boxes', []), + 'labels': results[task_prompt].get('labels', [])} + return output_results, output_image + # gokaygokay/Florence-2-SD3-Captioner task + elif task_prompt == 'description': + task_prompt = '' + result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + return result[task_prompt], None + # MiaoshouAI/Florence-2-large-PromptGen-v1.5 task + elif task_prompt == 'generate tags(PromptGen 1.5)': + task_prompt = '' + result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + return result[task_prompt], None + elif task_prompt == 'mixed caption(PromptGen 1.5)': + task_prompt = '' + result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + return result[task_prompt], None + elif task_prompt == 'mixed caption plus(PromptGen 2.0)': + task_prompt = '' + result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + return result[task_prompt], None + elif task_prompt == 'analyze(PromptGen 2.0)': + task_prompt = '<>' + result = run_example(model, processor, task_prompt, image, max_new_tokens, num_beams, do_sample) + return result[task_prompt], None + + else: + return "", None # Return empty string and None for unknown task prompts + + +def remove_angle_bracket_content(text): + import re + # 正则表达式匹配 "<>" 包围的内容,包括尖括号本身 + pattern = r'<[^>]*>' + # 使用 re.sub 替换匹配的内容为空字符串 + cleaned_text = re.sub(pattern, '', text) + return cleaned_text + + +def decode_f_bboxes(F_BBOXES): + if isinstance(F_BBOXES, str): + return (torch.zeros(1, 512, 512, dtype=torch.float32), F_BBOXES) + + width = F_BBOXES["width"] + height = F_BBOXES["height"] + mask = np.zeros((height, width), dtype=np.uint8) + + x1_c = width + y1_c = height + x2_c = y2_c = 0 + label = "" + if "bboxes" in F_BBOXES: + for idx in range(len(F_BBOXES["bboxes"])): + bbox = F_BBOXES["bboxes"][idx] + + new_label = F_BBOXES["labels"][idx].removeprefix("") + if new_label not in label: + if idx > 0: + label = label + ", " + label = label + new_label + + if len(bbox) == 4: + x1, y1, x2, y2 = int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3]) + elif len(bbox) == 8: + x1 = int(min(bbox[0::2])) + x2 = int(max(bbox[0::2])) + y1 = int(min(bbox[1::2])) + y2 = int(max(bbox[1::2])) + else: + continue + + x1_c = min(x1_c, x1) + y1_c = min(y1_c, y1) + x2_c = max(x2_c, x2) + y2_c = max(y2_c, y2) + + mask[y1:y2, x1:x2] = 1 + + else: + image = Image.new('RGB', (width, height), color='black') + draw = ImageDraw.Draw(image) + + x1_c = width + y1_c = height + x2_c = y2_c = 0 + + if len(F_BBOXES["polygons"]) == 0: + raise ValueError("Invalid bounding box, LARGE model cannot work in Transformers 5.x, Switch to BASE model, or downgrade the Transformers package") + + for polygon in F_BBOXES["polygons"][0]: + _polygon = np.array(polygon).reshape(-1, 2) + if len(_polygon) < 3: + print('Invalid polygon:', _polygon) + continue + + draw.polygon(_polygon.flatten().tolist(), outline='white', fill='white') + + x1_c = min(x1_c, int(min(polygon[0::2]))) + x2_c = max(x2_c, int(max(polygon[0::2]))) + y1_c = min(y1_c, int(min(polygon[1::2]))) + y2_c = max(y2_c, int(max(polygon[1::2]))) + + mask = np.asarray(image)[..., 0].astype(np.float32) / 255 + + mask = torch.from_numpy(mask.astype(np.float32)).unsqueeze(0) + # label = remove_angle_bracket_content(label) + return (mask, label) + + +class LS_LoadFlorence2Model: + def __init__(self): + self.model = None + self.processor = None + self.version = None + + @classmethod + def INPUT_TYPES(s): + model_list = list(fl2_model_repos.keys()) + return { + "required": { + "version": (model_list,{"default": model_list[0]}), + }, + } + + RETURN_TYPES = ("FLORENCE2",) + RETURN_NAMES = ("florence2_model",) + FUNCTION = "load" + CATEGORY = '😺dzNodes/LayerMask' + + def load(self, version): + if self.version != version: + self.model, self.processor = load_model(version) + self.version = version + + return ({'model': self.model, 'processor': self.processor, 'version': self.version, 'device': device},) + + +class Florence2Ultra: + def __init__(self): + self.NODE_NAME = 'Florence2Ultra' + + @classmethod + def INPUT_TYPES(s): + segment_task_list = [ + "region to segmentation", + "referring expression segmentation", + "open vocabulary detection", + ] + method_list = ['VITMatte', 'VITMatte(local)', 'vitmatte-base-composition-1k', 'PyMatting', 'GuidedFilter', ] + device_list = ['cuda','cpu'] + return { + "required": { + "florence2_model": ("FLORENCE2",), + "image": ("IMAGE",), + "task": (segment_task_list,{"default": segment_task_list[0]}), + "text_input": ("STRING", {"default": "subject"}), + "detail_method": (method_list,), + "detail_erode": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}), + "detail_dilate": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}), + "black_point": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 0.98, "step": 0.01, "display": "slider"}), + "white_point": ("FLOAT", {"default": 0.99, "min": 0.02, "max": 0.99, "step": 0.01, "display": "slider"}), + "process_detail": ("BOOLEAN", {"default": True}), + "device": (device_list,), + "max_megapixels": ("FLOAT", {"default": 2.0, "min": 1, "max": 999, "step": 0.1}), + }, + } + + RETURN_TYPES = ("IMAGE", "MASK",) + RETURN_NAMES = ("image", "mask",) + FUNCTION = "florence2_ultra" + CATEGORY = '😺dzNodes/LayerMask' + + def florence2_ultra(self, florence2_model, image, task, text_input, + detail_method, detail_erode, detail_dilate, + black_point, white_point, process_detail, device, max_megapixels): + max_new_tokens = 512 + num_beams = 3 + do_sample = False + fill_mask = False + + ret_images = [] + ret_masks = [] + + if detail_method == 'VITMatte(local)': + local_files_only = True + else: + local_files_only = False + + model = florence2_model['model'] + processor = florence2_model['processor'] + + for i in image: + img = tensor2pil(i).convert("RGB") + + results, _ = process_image(model, processor, img, task, + max_new_tokens, num_beams, do_sample, + fill_mask, text_input) + + if isinstance(results, dict): + results["width"] = img.width + results["height"] = img.height + + _mask, _ = decode_f_bboxes(results) + + if process_detail: + detail_range = detail_erode + detail_dilate + if detail_method == 'GuidedFilter': + _mask = guided_filter_alpha(i, _mask, detail_range // 6 + 1) + _mask = tensor2pil(histogram_remap(_mask, black_point, white_point)) + elif detail_method == 'PyMatting': + _mask = tensor2pil(mask_edge_detail(i, _mask, detail_range // 8 + 1, black_point, white_point)) + else: + _trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate) + _mask = generate_VITMatte(img, _trimap, local_files_only=local_files_only, device=device, max_megapixels=max_megapixels, method=detail_method) + _mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point)) + else: + _mask = tensor2pil(_mask) + + ret_image = RGB2RGBA(img, _mask.convert('L')) + ret_images.append(pil2tensor(ret_image)) + ret_masks.append(image2mask(_mask)) + + return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),) + + +class Florence2Image2Prompt: + + def __init__(self): + self.NODE_NAME = 'Florence2Image2Prompt' + + @classmethod + def INPUT_TYPES(s): + caption_task_list = [ + "caption", + "detailed caption", + "more detailed caption", + 'description', + 'generate tags(PromptGen 1.5)', + 'mixed caption(PromptGen 1.5)', + 'mixed caption plus(PromptGen 2.0)', + 'analyze(PromptGen 2.0)', + "object detection", + "dense region caption", + "region proposal", + "region proposal (mask)", + "caption to phrase grounding", + "open vocabulary detection", + "region to category", + "region to description", + "OCR", + "OCR with region", + ] + return { + "required": { + "florence2_model": ("FLORENCE2",), + "image": ("IMAGE",), + "task": (caption_task_list,{"default": caption_task_list[2]}), + "text_input": ("STRING", {"default": ""}), + "max_new_tokens": ("INT", {"default": 1024, "step": 1}), + "num_beams": ("INT", {"default": 3, "min": 1, "step": 1}), + "do_sample": ('BOOLEAN', {"default": False}), + "fill_mask": ('BOOLEAN', {"default": False}), + }, + } + + RETURN_TYPES = ("STRING", "IMAGE",) + RETURN_NAMES = ("text", "preview_image",) + FUNCTION = "florence2_image2prompt" + CATEGORY = '😺dzNodes/LayerUtility/Prompt' + + def florence2_image2prompt(self, florence2_model, image, task, text_input, + max_new_tokens, num_beams, do_sample, fill_mask): + + model = florence2_model['model'] + processor = florence2_model['processor'] + + img = tensor2pil(image[0]) + caption = "" + results, output_image = process_image(model, processor, img, task, max_new_tokens, num_beams, + do_sample, fill_mask, + text_input) + + if isinstance(results, dict): + results["width"] = img.width + results["height"] = img.height + + if output_image == None: + output_image = image[0].detach().clone().unsqueeze(0) + else: + output_image = np.asarray(output_image).astype(np.float32) / 255 + output_image = torch.from_numpy(output_image).unsqueeze(0) + + _, caption = decode_f_bboxes(results) + + return (remove_angle_bracket_content(caption), output_image,) + +NODE_CLASS_MAPPINGS = { + "LayerMask: Florence2Ultra": Florence2Ultra, + "LayerMask: LoadFlorence2Model": LS_LoadFlorence2Model, + "LayerUtility: Florence2Image2Prompt": Florence2Image2Prompt +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LayerMask: Florence2Ultra": "LayerMask: Florence2 Ultra(Advance)", + "LayerMask: LoadFlorence2Model": "LayerMask: Load Florence2 Model(Advance)", + "LayerUtility: Florence2Image2Prompt": "LayerUtility: Florence2 Image2Prompt(Advance)" +} diff --git a/pyproject.toml b/pyproject.toml index d19dc4d..5acac39 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "ComfyUI_LayerStyle_Advance" description = "The nodes detached from ComfyUI Layer Style are mainly those with complex requirements for dependency packages." -version = "2.0.37" +version = "2.0.38" license = { text = "MIT License" } dependencies = ["numpy", "matplotlib", "scikit_image", "scikit_learn", "opencv-contrib-python", "pymatting", "timm", "blend_modes", "transformers", "diffusers", "loguru", "colour-science", "huggingface_hub", "segment_anything", "addict", "omegaconf", "yapf", "wget", "iopath", "mediapipe", "typer_config", "fastapi", "rich", "google-generativeai", "ultralytics", "transparent-background", "accelerate", "onnxruntime", "bitsandbytes", "peft", "protobuf", "hydra-core", "blind-watermark", "qrcode", "pyzbar", "psd-tools", "wandb", "zhipuai", "openai","google-genai", "fastapi","typer-config"]