Update nodes.py
This commit is contained in:
@@ -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",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user