From 9183f102d3e6cbb876a3d9d9a65d24fb8cc2e07c Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 20 Jun 2024 10:04:38 +0300 Subject: [PATCH] better caption output --- nodes.py | 30 ++++++++++++++++++++++-------- 1 file changed, 22 insertions(+), 8 deletions(-) diff --git a/nodes.py b/nodes.py index 3bd2a82..41610b7 100644 --- a/nodes.py +++ b/nodes.py @@ -6,7 +6,7 @@ import matplotlib matplotlib.use('Agg') import matplotlib.pyplot as plt import matplotlib.patches as patches -from PIL import Image, ImageDraw, ImageFont +from PIL import Image, ImageDraw import random import numpy as np @@ -23,7 +23,7 @@ def fixed_get_imports(filename: str | os.PathLike) -> list[str]: import comfy.model_management as mm -from comfy.utils import ProgressBar, load_torch_file +from comfy.utils import ProgressBar import folder_paths script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -100,7 +100,7 @@ class Florence2Run: "text_input": ("STRING", {"default": "", "multiline": True}), "task": ( [ - 'annotate', + 'region_caption', 'dense_region_caption', 'caption', 'detailed_caption', @@ -134,7 +134,7 @@ class Florence2Run: colormap = ['blue','orange','green','purple','brown','pink','gray','olive','cyan','red', 'lime','indigo','violet','aqua','magenta','coral','gold','tan','skyblue'] - if task == 'annotate': + if task == 'region_caption': prompt = "" elif task == 'dense_region_caption': prompt = '' @@ -148,7 +148,10 @@ class Florence2Run: prompt = '' elif task == 'referring_expression_segmentation': prompt = '' - + + if (task!= 'referring_expression_segmentation' and task!= 'caption_to_phrase_grounding') and text_input: + raise ValueError("text_input is only supported for 'referring_expression_segmentation' and 'caption_to_phrase_grounding'") + if text_input is not None: prompt = prompt + text_input @@ -171,8 +174,19 @@ class Florence2Run: ) results = processor.batch_decode(generated_ids, skip_special_tokens=False)[0] - out_results.append(results) - if task == 'annotate' or task == 'dense_region_caption' or task == 'caption_to_phrase_grounding': + + # cleanup the special tokens from the final list + 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) + + if task == 'region_caption' or task == 'dense_region_caption' or task == 'caption_to_phrase_grounding': parsed_answer = processor.post_process_generation(results, task="", image_size=(image_pil.width, image_pil.height)) fig, ax = plt.subplots(figsize=(image_pil.width / 100, image_pil.height / 100), dpi=100) @@ -284,7 +298,6 @@ class Florence2Run: out_tensor = torch.cat(out, dim=0) else: out_tensor = torch.zeros((1, 64,64, 3), dtype=torch.float32, device="cpu") - print(out_tensor.shape) if len(out_masks) > 0: out_mask_tensor = torch.cat(out_masks, dim=0) else: @@ -293,6 +306,7 @@ class Florence2Run: if not keep_model_loaded: print("Offloading model...") model.to(offload_device) + mm.soft_empty_cache() return (out_tensor, out_mask_tensor, out_results,)