From fb237a1e02a9eb04e5b71e00230492a37e7dc6eb Mon Sep 17 00:00:00 2001 From: Aryan185 Date: Fri, 29 Aug 2025 06:35:03 +0000 Subject: [PATCH] Added node for mask segmentation with Gemini --- __init__.py | 5 +- gemini_node.py | 5 +- gemini_segment.py | 185 ++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 192 insertions(+), 3 deletions(-) create mode 100644 gemini_segment.py diff --git a/__init__.py b/__init__.py index e58108c..ac92ac3 100644 --- a/__init__.py +++ b/__init__.py @@ -5,9 +5,10 @@ from .gpt_image1 import NODE_CLASS_MAPPINGS as GPT_IMAGE_MAPPINGS, NODE_DISPLAY_ from .imagen import NODE_CLASS_MAPPINGS as IMAGEN_IMAGE_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as IMAGEN_IMAGE_DISPLAY from .imagen_edit import NODE_CLASS_MAPPINGS as IMAGEN_EDIT_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as IMAGEN_EDIT_DISPLAY from .veo import NODE_CLASS_MAPPINGS as VEO_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as VEO_DISPLAY +from .gemini_segment import NODE_CLASS_MAPPINGS as GEMINI_SEGMENT_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as GEMINI_SEGMENT_DISPLAY # Combine both mappings -NODE_CLASS_MAPPINGS = {**PRO_MAPPINGS, **MAX_MAPPINGS, **GEMINI_MAPPINGS, **GPT_IMAGE_MAPPINGS, **IMAGEN_IMAGE_MAPPINGS, **IMAGEN_EDIT_MAPPINGS, **VEO_MAPPINGS} -NODE_DISPLAY_NAME_MAPPINGS = {**PRO_DISPLAY, **MAX_DISPLAY, **GEMINI_DISPLAY, **GPT_IMAGE_DISPLAY, **IMAGEN_IMAGE_DISPLAY, **IMAGEN_EDIT_DISPLAY, **VEO_DISPLAY} +NODE_CLASS_MAPPINGS = {**PRO_MAPPINGS, **MAX_MAPPINGS, **GEMINI_MAPPINGS, **GPT_IMAGE_MAPPINGS, **IMAGEN_IMAGE_MAPPINGS, **IMAGEN_EDIT_MAPPINGS, **VEO_MAPPINGS, **GEMINI_SEGMENT_MAPPINGS} +NODE_DISPLAY_NAME_MAPPINGS = {**PRO_DISPLAY, **MAX_DISPLAY, **GEMINI_DISPLAY, **GPT_IMAGE_DISPLAY, **IMAGEN_IMAGE_DISPLAY, **IMAGEN_EDIT_DISPLAY, **VEO_DISPLAY, **GEMINI_SEGMENT_DISPLAY} __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/gemini_node.py b/gemini_node.py index 326d73c..0f08caf 100644 --- a/gemini_node.py +++ b/gemini_node.py @@ -44,7 +44,7 @@ class GeminiChatNode: # Initialize client and build parts - client = genai.Client(api_key=key) + client = genai.Client(api_key=key, http_options=types.HttpOptions(retry_options=types.HttpRetryOptions(attempts=3, jitter=10))) parts = [types.Part.from_text(text=prompt)] # Handle image input @@ -62,14 +62,17 @@ class GeminiChatNode: model_lower = model.lower() if "gemini-2.0" in model_lower: + print("Gemini-2.0 models do not support thinking - disabling thinking config") final_thinking_budget = None elif not thinking: final_thinking_budget = 0 if "gemini-2.5-pro" in model_lower: + print("Gemini-2.5-Pro cannot have thinking turned off - defaulting thinking budget to -1") final_thinking_budget = -1 else: final_thinking_budget = thinking_budget if "gemini-2.5-pro" in model_lower and final_thinking_budget == 0: + print("Gemini-2.5-Pro cannot have thinking turned off - defaulting thinking budget to -1") final_thinking_budget = -1 config = types.GenerateContentConfig( diff --git a/gemini_segment.py b/gemini_segment.py new file mode 100644 index 0000000..e70764c --- /dev/null +++ b/gemini_segment.py @@ -0,0 +1,185 @@ +import base64 +import os +import io +import json +import numpy as np +import torch +from PIL import Image +from google import genai +from google.genai import types + +class GeminiSegmentationNode: + """ComfyUI Node for Gemini API Image Segmentation""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "segment_prompt": ("STRING", {"default": "all objects", "multiline": True}), + "model": ("STRING", {"default": "gemini-2.5-flash", "multiline": False}), + "temperature": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 2.0, "step": 0.1}), + "thinking": ("BOOLEAN", {"default": True}), + "seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}), + "api_key": ("STRING", {"default": "", "multiline": False}) + }, + "optional": { + "thinking_budget": ("INT", {"default": 0, "min": -1, "max": 24576, "step": 1}), + } + } + + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("mask",) + FUNCTION = "generate_segmentation" + CATEGORY = "AI/Gemini" + + def generate_segmentation(self, image: torch.Tensor, segment_prompt: str, model: str, + temperature: float, thinking: bool, seed: int, api_key: str, + thinking_budget: int = 0) -> tuple: + + key = api_key.strip() or os.environ.get("GEMINI_API_KEY") or os.environ.get("GOOGLE_API_KEY") + if not key: + raise ValueError("Error: No API key provided. Set GEMINI_API_KEY or GOOGLE_API_KEY environment variable, or provide it in the node.") + + img_array = image.cpu().numpy() if isinstance(image, torch.Tensor) else image + if len(img_array.shape) == 4: + img_array = img_array[0] # Remove batch dimension + if img_array.dtype in [np.float32, np.float64]: + img_array = (img_array * 255).astype(np.uint8) + + original_image = Image.fromarray(img_array).convert('RGB') + original_width, original_height = original_image.size + + max_size = 1024 + scale = min(max_size / original_width, max_size / original_height) + + if scale < 1: + new_width = int(original_width * scale) + new_height = int(original_height * scale) + processed_image = original_image.resize((new_width, new_height), Image.Resampling.LANCZOS) + else: + processed_image = original_image + + buffer = io.BytesIO() + processed_image.save(buffer, format='PNG') + image_data = buffer.getvalue() + + base_prompt = f"Give the segmentation masks for {segment_prompt}. Output a JSON list of segmentation masks where each entry contains the 2D bounding box in the key \"box_2d\", the segmentation mask in key \"mask\", and the text label in the key \"label\". Use descriptive labels." + + client = genai.Client(api_key=key, http_options=types.HttpOptions(retry_options=types.HttpRetryOptions(attempts=3, jitter=10))) + + parts = [ + types.Part.from_bytes(mime_type="image/png", data=image_data), + types.Part.from_text(text=base_prompt) + ] + + model_lower = model.lower() + + if "gemini-2.0" in model_lower: + print("Gemini-2.0 models do not support thinking - disabling thinking config") + final_thinking_budget = None + elif not thinking: + final_thinking_budget = 0 + if "gemini-2.5-pro" in model_lower: + print("Gemini-2.5-Pro cannot have thinking turned off - defaulting thinking budget to -1") + final_thinking_budget = -1 + else: + final_thinking_budget = thinking_budget + if "gemini-2.5-pro" in model_lower and final_thinking_budget == 0: + print("Gemini-2.5-Pro cannot have thinking turned off - defaulting thinking budget to -1") + final_thinking_budget = -1 + + config = types.GenerateContentConfig( + temperature=temperature, + seed=seed, + response_mime_type="text/plain" + ) + + if "gemini-2.0" not in model_lower: + config.thinking_config = types.ThinkingConfig(thinking_budget=final_thinking_budget) + + # Generate content + try: + response = client.models.generate_content( + model=model, + contents=[types.Content(role="user", parts=parts)], + config=config + ) + + response_text = response.text + if '```json' in response_text: + response_text = response_text.split('```json')[1].split('```')[0] + + segments = json.loads(response_text) + + except Exception as e: + raise RuntimeError(f"Error calling Gemini API: {str(e)}") + + # Create mask from segments + proc_width, proc_height = processed_image.size + mask_image = Image.new('L', (proc_width, proc_height), 0) + + # Sort segments by size (largest first) + segments_with_size = [] + for segment in segments: + box_2d = segment['box_2d'] + ymin, xmin, ymax, xmax = box_2d + w = (xmax - xmin) / 1000 + h = (ymax - ymin) / 1000 + segments_with_size.append((segment, w * h)) + + segments_with_size.sort(key=lambda x: x[1], reverse=True) + + # Process each segment + for i, (segment, _) in enumerate(segments_with_size): + try: + box_2d = segment['box_2d'] + ymin, xmin, ymax, xmax = box_2d + + x = int(xmin / 1000 * proc_width) + y = int(ymin / 1000 * proc_height) + w = int((xmax - xmin) / 1000 * proc_width) + h = int((ymax - ymin) / 1000 * proc_height) + + mask_data = segment['mask'] + + if isinstance(mask_data, str): + if mask_data.startswith('data:image'): + mask_data = mask_data.split(',')[1] + mask_bytes = base64.b64decode(mask_data) + mask_img = Image.open(io.BytesIO(mask_bytes)).convert('L') + else: + continue + + if mask_img.size != (w, h): + mask_img = mask_img.resize((w, h), Image.Resampling.LANCZOS) + + mask_array = list(mask_img.getdata()) + final_pixels = [255 if alpha > 128 else 0 for alpha in mask_array] + segment_mask = Image.new('L', (w, h)) + segment_mask.putdata(final_pixels) + + if x + w <= proc_width and y + h <= proc_height and x >= 0 and y >= 0: + region = mask_image.crop((x, y, x + w, y + h)) + region_pixels = list(region.getdata()) + segment_pixels = list(segment_mask.getdata()) + combined_pixels = [max(r, s) for r, s in zip(region_pixels, segment_pixels)] + combined_region = Image.new('L', (w, h)) + combined_region.putdata(combined_pixels) + mask_image.paste(combined_region, (x, y)) + + except Exception: + continue + + if processed_image.size != original_image.size: + mask_image = mask_image.resize(original_image.size, Image.Resampling.LANCZOS) + + # Convert PIL mask to ComfyUI mask format + mask_array = np.array(mask_image, dtype=np.float32) / 255.0 + mask_tensor = torch.from_numpy(mask_array).unsqueeze(0) # Add batch dimension + + return (mask_tensor,) + +# Node mappings +NODE_CLASS_MAPPINGS = {"GeminiSegmentationNode": GeminiSegmentationNode} +NODE_DISPLAY_NAME_MAPPINGS = {"GeminiSegmentationNode": "Gemini Segmentation"} \ No newline at end of file