Files
Aryan185-ComfyUI-ExternalAP…/imagen.py
T
2025-08-24 13:13:08 +05:30

111 lines
4.5 KiB
Python

import os
import torch
import numpy as np
from PIL import Image
from google import genai
from google.genai import types
class GoogleImagenNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompt": ("STRING", {"multiline": True, "default": "A majestic lion in the savanna"}),
"api_key": ("STRING", {"multiline": False, "default": ""}),
"model": (["models/imagen-4.0-ultra-generate-001", "models/imagen-4.0-generate-001", "models/imagen-4.0-fast-generate-001", "models/imagen-3.0-generate-002"],),
"number_of_images": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}),
"aspect_ratio": (["1:1", "9:16", "16:9", "4:3", "3:4"],),
"image_size": (["1K", "2K"],),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "generate_images"
CATEGORY = "image/ai"
def pil_to_tensor(self, images):
if not isinstance(images, list):
images = [images]
tensors = []
for image in images:
if image.mode != 'RGB':
image = image.convert('RGB')
array = np.array(image).astype(np.float32) / 255.0
tensor = torch.from_numpy(array)
tensors.append(tensor)
return torch.stack(tensors)
def generate_images(self, prompt, api_key, model, number_of_images, aspect_ratio, image_size):
try:
key = api_key.strip() or os.environ.get("GEMINI_API_KEY")
if not key:
raise ValueError("No API key provided.")
client = genai.Client(api_key=key)
config = types.GenerateImagesConfig(
number_of_images=number_of_images,
include_rai_reason=True,
output_mime_type="image/jpeg",
aspect_ratio=aspect_ratio,
)
if model in ["models/imagen-4.0-ultra-generate-001", "models/imagen-4.0-generate-001"]:
config.image_size = image_size
else:
print("2K resolution not supported for this model, using default image size")
result = client.models.generate_images(
model=model,
prompt=prompt,
config=config
)
if not result.generated_images:
raise ValueError("No images generated by the API")
pil_images = []
for generated_image in result.generated_images:
image_data = generated_image.image
if hasattr(image_data, 'mode') and hasattr(image_data, 'size'):
pil_images.append(image_data)
elif hasattr(image_data, '_pil_image'):
pil_images.append(image_data._pil_image)
elif hasattr(image_data, 'show'):
try:
from io import BytesIO
buffer = BytesIO()
image_data.save(buffer, format='PNG')
buffer.seek(0)
pil_image = Image.open(buffer)
pil_images.append(pil_image)
except:
pil_images.append(Image.new('RGB', (512, 512), color='gray'))
elif hasattr(image_data, 'read') or isinstance(image_data, bytes):
from io import BytesIO
image_bytes = image_data.read() if hasattr(image_data, 'read') else image_data
pil_images.append(Image.open(BytesIO(image_bytes)))
else:
try:
pil_images.append(Image.open(image_data))
except:
pil_images.append(Image.new('RGB', (512, 512), color='gray'))
return (self.pil_to_tensor(pil_images),)
except Exception as e:
print(f"Google Imagen Error: {str(e)}")
error_image = Image.new('RGB', (512, 512), color='black')
return (self.pil_to_tensor([error_image]),)
@classmethod
def IS_CHANGED(cls, **kwargs):
return f"{kwargs.get('prompt', '')}-{kwargs.get('model', '')}-{kwargs.get('number_of_images', 1)}-{kwargs.get('aspect_ratio', '1:1')}-{kwargs.get('image_size', '1K')}"
NODE_CLASS_MAPPINGS = {"GoogleImagenNode": GoogleImagenNode}
NODE_DISPLAY_NAME_MAPPINGS = {"GoogleImagenNode": "Google Imagen Generator"}