From dc38ccbea800b62e7861a083f1de933135c50532 Mon Sep 17 00:00:00 2001 From: chflame163 Date: Tue, 17 Dec 2024 19:33:22 +0800 Subject: [PATCH] fix ObjectDetectorGemini return empty bbox --- py/gemini.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/py/gemini.py b/py/gemini.py index f50133c..8ca7a57 100644 --- a/py/gemini.py +++ b/py/gemini.py @@ -1,9 +1,11 @@ # layerstyle advance import json +import re from .imagefunc import * - +def is_only_digits_and_spaces(s:str) -> bool: + return bool(re.fullmatch(r'[0-9\s]*', s)) class LS_GeminiNode: @@ -141,6 +143,11 @@ class LS_OBJECT_DETECTOR_Gemini: _image = tensor2pil(img.unsqueeze(0)).convert('RGB') response = g_model.generate_content([_image, g_prompt]) ret_text = response.text + if not is_only_digits_and_spaces(ret_text): + ret_bboxes.append([(-1, -1, 0, 0)]) + ret_previews.append(pil2tensor(_image)) + log(f"{self.NODE_NAME} no object found", message_type='warning') + continue y1,x1,y2,x2 = [int(x) for x in ret_text.split()] @@ -153,10 +160,7 @@ class LS_OBJECT_DETECTOR_Gemini: preview = draw_bounding_boxes(_image.convert("RGB"), bboxes, color="random", line_width=-1) ret_previews.append(pil2tensor(preview)) - if len(bboxes) == 0: - log(f"{self.NODE_NAME} no object found", message_type='warning') - else: - log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info') + log(f"{self.NODE_NAME} found {len(bboxes)} object(s)", message_type='info') ret_bboxes.append(bboxes) return (ret_bboxes, torch.cat(ret_previews, dim=0),)