Update nodes.py

This commit is contained in:
Shashank Shekhar
2024-07-16 21:35:59 +05:30
committed by GitHub
parent 3fb3b63f9a
commit 2fa5207f02
+11 -13
View File
@@ -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",
}