commit Gemini and ObjectDetectorGemini nodes

This commit is contained in:
chflame163
2024-12-17 18:43:07 +08:00
parent 662f10c151
commit 4dec7e9899
11 changed files with 839 additions and 51 deletions
+173
View File
@@ -0,0 +1,173 @@
# layerstyle advance
import json
from .imagefunc import *
class LS_GeminiNode:
CATEGORY = '😺dzNodes/LayerUtility'
FUNCTION = "run_gemini"
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("text",)
OUTPUT_IS_LIST = (True,)
def __init__(self):
self.NODE_NAME = 'Gemini'
@classmethod
def INPUT_TYPES(self):
gemini_model_list = [
"gemini-1.5-flash",
"gemini-1.5-pro",
"gemini-1.5-flash-8b",
"gemini-2.0-flash-exp",
"learnlm-1.5-pro-experimental"]
language_list = ['en', 'zh-CN']
return {
"required": {
"model": (gemini_model_list,),
"max_output_tokens": ("INT", {"default": 4096, "min": 1, "max": 8192, "step": 1}),
"temperature": ("FLOAT", {"default": 0.5, "min": 0, "max": 2, "step": 0.1}),
"words_limit": ("INT", {"default": 200, "min": 8, "max": 2048, "step": 1}),
"response_language": (language_list,),
"system_prompt": ("STRING",
{"default": "You are creating a prompt for Stable Diffusion to generate an image.",
"multiline": False}),
"user_prompt": ("STRING", {
"default": "Generate a prompt about a girl.",
"multiline": True}),
},
"optional": {
"image_1": ("IMAGE",),
"image_2": ("IMAGE",),
}
}
def run_gemini(self, model, system_prompt, user_prompt, max_output_tokens, temperature,
words_limit, response_language, image_1=None, image_2=None):
import google.generativeai as genai
ret_texts = []
g_model = genai.GenerativeModel(model,
generation_config=gemini_generate_config,
safety_settings=gemini_safety_settings)
g_cfg = genai.GenerationConfig(temperature=temperature,
max_output_tokens=max_output_tokens)
genai.configure(api_key=get_api_key('google_api_key'), transport='rest')
prompt = {
"USER_INPUT":user_prompt,
"action": f"{system_prompt}\n"
f"Follow the USER_INPUT to complete task, keep response length between {int(words_limit * 0.8)} to {int(words_limit * 1.2)} words.",
"output_format": {
"content": f"Only return the final result, not include any unnecessary content.",
"language": response_language,
}
}
prompt = json.dumps(prompt)
log(f"{self.NODE_NAME}: Request to {model}...")
if image_1 is not None and image_2 is not None:
for index,img in enumerate(image_1):
_image1 = tensor2pil(img.unsqueeze(0)).convert('RGB')
_image2 = tensor2pil(image_2[index].unsqueeze(0)).convert('RGB') if index < len(image_2) else tensor2pil(image_2[-1].unsqueeze(0)).convert('RGB')
response = g_model.generate_content([prompt, _image1, _image2], generation_config=g_cfg)
ret_text = response.text
log(f"{self.NODE_NAME}: Gemini response is:\n\033[1;36m{ret_text}\033[m")
ret_texts.append(ret_text)
elif (image_1 is not None and image_2 is None) or (image_2 is not None and image_1 is None):
_imgs = image_1 if image_1 is not None else image_2
for img in _imgs:
_image = tensor2pil(img.unsqueeze(0)).convert('RGB')
response = g_model.generate_content([prompt, _image], generation_config=g_cfg)
ret_text = response.text
log(f"{self.NODE_NAME}: Gemini response is:\n\033[1;36m{ret_text}\033[m")
ret_texts.append(ret_text)
else:
response = g_model.generate_content(prompt, generation_config=g_cfg)
ret_text = response.text
log(f"{self.NODE_NAME}: Gemini response is:\n\033[1;36m{ret_text}\033[m")
ret_texts.append(ret_text)
return (ret_texts,)
class LS_OBJECT_DETECTOR_Gemini:
CATEGORY = '😺dzNodes/LayerMask'
FUNCTION = "run_gemini_detect"
RETURN_TYPES = ("BBOXES", "IMAGE",)
RETURN_NAMES = ("bboxes", "preview",)
# OUTPUT_IS_LIST = (True,)
def __init__(self):
self.NODE_NAME = 'GeminiDetect'
@classmethod
def INPUT_TYPES(self):
gemini_model_list = [
"gemini-1.5-flash",
"gemini-1.5-pro",
"gemini-1.5-flash-8b",
"gemini-2.0-flash-exp"
]
return {
"required": {
"image": ("IMAGE",),
"model": (gemini_model_list,),
"prompt": ("STRING", {"default": "subject"}),
},
"optional": {
}
}
def run_gemini_detect(self, image, model, prompt):
import google.generativeai as genai
ret_bboxes = []
ret_previews = []
g_model = genai.GenerativeModel(model,
generation_config=gemini_generate_config,
safety_settings=gemini_safety_settings)
genai.configure(api_key=get_api_key('google_api_key'), transport='rest')
g_prompt = f"Return a bounding box of {prompt} in this image in [ymin, xmin, ymax, xmax] format."
log(f"{self.NODE_NAME}: Request to {model}...")
for img in image:
_image = tensor2pil(img.unsqueeze(0)).convert('RGB')
response = g_model.generate_content([_image, g_prompt])
ret_text = response.text
y1,x1,y2,x2 = [int(x) for x in ret_text.split()]
# Convert normalized coordinates to absolute coordinates
x1 = int(x1 / 1000 * _image.width)
y1 = int(y1 / 1000 * _image.height)
x2 = int(x2 / 1000 * _image.width)
y2 = int(y2 / 1000 * _image.height)
bboxes = standardize_bbox([(x1, y1, x2, y2)])
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')
ret_bboxes.append(bboxes)
return (ret_bboxes, torch.cat(ret_previews, dim=0),)
NODE_CLASS_MAPPINGS = {
"LayerUtility: Gemini": LS_GeminiNode,
"LayerMask: ObjectDetectorGemini": LS_OBJECT_DETECTOR_Gemini,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LayerUtility: Gemini": "LayerUtility: Gemini(Advance)",
"LayerMask: ObjectDetectorGemini": "LayerMask: Object Detector Gemini(Advance)",
}
+10 -37
View File
@@ -2321,43 +2321,16 @@ def get_resource_dir() -> list:
return (LUT_DICT, FONT_DICT)
# (LUT_DICT, FONT_DICT) = get_resource_dir()
# FONT_LIST = list(FONT_DICT.keys())
# LUT_LIST = list(LUT_DICT.keys())
# def get_models_dir() -> dict:
# models_dir_ini_file = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), "models_dir.ini")
# MODELS_DIR = {}
# model_dir_list = [
# "birefnet_dir",
# "evf-sam_dir",
# "florence2_dir",
# "lama_dir",
# "rmbg_dir",
# "segformerB2_dir",
# "segformerB3_clothes_dir",
# "segformerB3_fashion_dir",
# "sam2_dir",
# "transparent-background_dir",
# "yolo8_dir",
# "yolo_world_dir"
# ]
# try:
# with open(models_dir_ini_file, 'r') as f:
# ini = f.readlines()
# for line in ini:
# for model_dir in model_dir_list:
# if line.startswith(model_dir):
# path = line[line.find('=') + 1:].rstrip().lstrip()
# if os.path.exists(path):
# MODELS_DIR[model_dir] = path
# log(f'Find {len(MODELS_DIR)} path(s) in {models_dir_ini_file}.')
# except Exception as e:
# log(f'Warning: {models_dir_ini_file} not found' + f', default directory to be used.')
#
# return MODELS_DIR
#
# MODELS_DIR = get_models_dir()
# 规范bbox,保证x1 < x2, y1 < y2, 并返回int
def standardize_bbox(bboxes:list) -> list:
ret_bboxes = []
for bbox in bboxes:
x1 = int(min(bbox[0], bbox[2]))
y1 = int(min(bbox[1], bbox[3]))
x2 = int(max(bbox[0], bbox[2]))
y2 = int(max(bbox[1], bbox[3]))
ret_bboxes.append([x1, y1, x2, y2])
return ret_bboxes
def draw_bounding_boxes(image: Image, bboxes: list, color: str = "#FF0000", line_width: int = 5) -> Image:
"""
-11
View File
@@ -6,17 +6,6 @@ select_list = ["all", "first", "by_index"]
sort_method_list = ["left_to_right", "top_to_bottom", "big_to_small", "confidence"]
# 规范bbox,保证x1 < x2, y1 < y2, 并返回int
def standardize_bbox(bboxes:list) -> list:
ret_bboxes = []
for bbox in bboxes:
x1 = int(min(bbox[0], bbox[2]))
y1 = int(min(bbox[1], bbox[3]))
x2 = int(max(bbox[0], bbox[2]))
y2 = int(max(bbox[1], bbox[3]))
ret_bboxes.append([x1, y1, x2, y2])
return ret_bboxes
def sort_bboxes(bboxes:list, method:str) -> list:
sorted_bboxes = []
if method == "left_to_right":