diff --git a/nodes.py b/nodes.py index 21f77cf..7930f96 100644 --- a/nodes.py +++ b/nodes.py @@ -3,7 +3,7 @@ import torchvision.transforms.functional as F import io import os import matplotlib -matplotlib.use('Agg') +matplotlib.use('Agg') import matplotlib.pyplot as plt import matplotlib.patches as patches from PIL import Image, ImageDraw, ImageColor, ImageFont @@ -11,7 +11,6 @@ import random import numpy as np import re - #workaround for unnecessary flash_attn requirement from unittest.mock import patch from transformers.dynamic_module_utils import get_imports @@ -127,8 +126,8 @@ class Florence2Run: } } - RETURN_TYPES = ("IMAGE", "MASK", "STRING",) - RETURN_NAMES =("image", "mask", "caption",) + RETURN_TYPES = ("IMAGE", "MASK", "STRING", "JSON") + RETURN_NAMES =("image", "mask", "caption", "data") FUNCTION = "encode" CATEGORY = "Florence2" @@ -174,6 +173,7 @@ class Florence2Run: out = [] out_masks = [] out_results = [] + out_datas = [] pbar = ProgressBar(len(image)) for img in image: image_pil = F.to_pil_image(img) @@ -206,6 +206,7 @@ class Florence2Run: out_results.append(clean_results) W, H = image_pil.size + data = [] 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': @@ -364,7 +365,9 @@ class Florence2Run: scale = 1 draw = ImageDraw.Draw(image_pil) bboxes, labels = predictions['quad_boxes'], predictions['labels'] + for box, label in zip(bboxes, labels): + data.append({"label": label, "box": box}) color = random.choice(colormap) new_box = (np.array(box) * scale).tolist() draw.polygon(new_box, width=3, outline=color) @@ -373,6 +376,7 @@ class Florence2Run: 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) @@ -402,6 +406,8 @@ class Florence2Run: out.append(F.to_tensor(image_pil).unsqueeze(0).permute(0, 2, 3, 1).cpu().float()) pbar.update(1) + # data for task to out_datas + out_datas.append(data) if len(out) > 0: out_tensor = torch.cat(out, dim=0) @@ -417,13 +423,5 @@ class Florence2Run: model.to(offload_device) mm.soft_empty_cache() - return (out_tensor, out_mask_tensor, out_results,) + return (out_tensor, out_mask_tensor, out_results, out_datas) -NODE_CLASS_MAPPINGS = { - "DownloadAndLoadFlorence2Model": DownloadAndLoadFlorence2Model, - "Florence2Run": Florence2Run, -} -NODE_DISPLAY_NAME_MAPPINGS = { - "DownloadAndLoadFlorence2Model": "DownloadAndLoadFlorence2Model", - "Florence2Run": "Florence2Run", -}