111 lines
4.5 KiB
Python
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"} |