Added node for imagen-edit
This commit is contained in:
+3
-2
@@ -3,9 +3,10 @@ from .flux_kontext_max_node import NODE_CLASS_MAPPINGS as MAX_MAPPINGS, NODE_DIS
|
||||
from .gemini_node import NODE_CLASS_MAPPINGS as GEMINI_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as GEMINI_DISPLAY
|
||||
from .gpt_image1 import NODE_CLASS_MAPPINGS as GPT_IMAGE_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as GPT_IMAGE_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
|
||||
|
||||
# Combine both mappings
|
||||
NODE_CLASS_MAPPINGS = {**PRO_MAPPINGS, **MAX_MAPPINGS, **GEMINI_MAPPINGS, **GPT_IMAGE_MAPPINGS, **IMAGEN_IMAGE_MAPPINGS}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {**PRO_DISPLAY, **MAX_DISPLAY, **GEMINI_DISPLAY, **GPT_IMAGE_DISPLAY, **IMAGEN_IMAGE_DISPLAY}
|
||||
NODE_CLASS_MAPPINGS = {**PRO_MAPPINGS, **MAX_MAPPINGS, **GEMINI_MAPPINGS, **GPT_IMAGE_MAPPINGS, **IMAGEN_IMAGE_MAPPINGS, **IMAGEN_EDIT_MAPPINGS}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {**PRO_DISPLAY, **MAX_DISPLAY, **GEMINI_DISPLAY, **GPT_IMAGE_DISPLAY, **IMAGEN_IMAGE_DISPLAY, **IMAGEN_EDIT_DISPLAY}
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
+154
@@ -0,0 +1,154 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import base64
|
||||
from io import BytesIO
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
from google.genai.types import RawReferenceImage, MaskReferenceImage
|
||||
|
||||
class GoogleImagenEditNode:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
"prompt": ("STRING", {"multiline": True, "default": "Edit this image"}),
|
||||
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
|
||||
"project_id": ("STRING", {"multiline": False, "default": ""}),
|
||||
"location": (["us-central1", "us-east1", "us-east4", "us-east5", "us-south1", "us-west1", "us-west2", "us-west3", "us-west4", "northamerica-northeast1", "northamerica-northeast2", "southamerica-east1", "southamerica-west1", "africa-south1", "europe-west1", "europe-north1", "europe-west2", "europe-west3", "europe-west4", "europe-west6", "europe-west8", "europe-west9", "europe-west12", "europe-southwest1", "europe-central2", "asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2", "asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1", "asia-southeast2", "australia-southeast1", "australia-southeast2", "me-central1", "me-central2", "me-west1"], {"default": "us-central1"}),
|
||||
"service_account": ("STRING", {"multiline": False, "default": ""}),
|
||||
"edit_mode": (["EDIT_MODE_INPAINT_INSERTION", "EDIT_MODE_INPAINT_REMOVAL", "EDIT_MODE_OUTPAINT", "EDIT_MODE_BGSWAP"], {"default": "EDIT_MODE_INPAINT_INSERTION"}),
|
||||
"number_of_images": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}),
|
||||
"seed": ("INT", {"default": 12345, "min": 1, "max": 4294967295, "step": 1}),
|
||||
"base_steps": ("INT", {"default": 50, "min": 10, "max": 100, "step": 1})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("edited_images",)
|
||||
FUNCTION = "edit_image"
|
||||
CATEGORY = "image/ai"
|
||||
|
||||
def tensor_to_pil(self, tensor):
|
||||
array = (tensor.cpu().numpy() * 255).astype(np.uint8)
|
||||
return Image.fromarray(array)
|
||||
|
||||
def mask_to_pil(self, mask):
|
||||
if len(mask.shape) == 3 and mask.shape[0] == 1:
|
||||
mask = mask.squeeze(0)
|
||||
array = (mask.cpu().numpy() * 255).astype(np.uint8)
|
||||
return Image.fromarray(array, mode='L')
|
||||
|
||||
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
|
||||
tensors.append(torch.from_numpy(array))
|
||||
return torch.stack(tensors)
|
||||
|
||||
def edit_image(self, image, mask, prompt, project_id, location, service_account, edit_mode, number_of_images, negative_prompt, seed, base_steps):
|
||||
try:
|
||||
if service_account.strip():
|
||||
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = service_account.strip()
|
||||
|
||||
if not os.environ.get('GOOGLE_APPLICATION_CREDENTIALS'):
|
||||
raise ValueError("No authentication provided.")
|
||||
|
||||
if not project_id.strip():
|
||||
raise ValueError("Project ID is required.")
|
||||
|
||||
client = genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
|
||||
|
||||
input_image = self.tensor_to_pil(image[0])
|
||||
input_mask = self.mask_to_pil(mask)
|
||||
|
||||
img_buffer = BytesIO()
|
||||
input_image.save(img_buffer, format='PNG')
|
||||
img_b64 = base64.b64encode(img_buffer.getvalue()).decode('utf-8')
|
||||
|
||||
mask_buffer = BytesIO()
|
||||
input_mask.save(mask_buffer, format='PNG')
|
||||
mask_b64 = base64.b64encode(mask_buffer.getvalue()).decode('utf-8')
|
||||
|
||||
raw_ref_image = RawReferenceImage(
|
||||
reference_image={'image_bytes': img_b64},
|
||||
reference_id=0
|
||||
)
|
||||
|
||||
mask_ref_image = MaskReferenceImage(
|
||||
reference_id=1,
|
||||
reference_image={'image_bytes': mask_b64},
|
||||
config=types.MaskReferenceConfig(
|
||||
mask_mode="MASK_MODE_USER_PROVIDED",
|
||||
mask_dilation=0.03,
|
||||
),
|
||||
)
|
||||
|
||||
config_params = {
|
||||
"edit_mode": edit_mode,
|
||||
"number_of_images": number_of_images,
|
||||
"include_rai_reason": True,
|
||||
"output_mime_type": "image/jpeg",
|
||||
"base_steps": base_steps,
|
||||
"seed": seed
|
||||
}
|
||||
|
||||
if negative_prompt.strip():
|
||||
config_params["negative_prompt"] = negative_prompt.strip()
|
||||
|
||||
response = client.models.edit_image(
|
||||
model="imagen-3.0-capability-001",
|
||||
prompt=prompt,
|
||||
reference_images=[raw_ref_image, mask_ref_image],
|
||||
config=types.EditImageConfig(**config_params),
|
||||
)
|
||||
|
||||
if not response.generated_images:
|
||||
raise ValueError("No images generated by the API")
|
||||
|
||||
pil_images = []
|
||||
for generated_image in response.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:
|
||||
buffer = BytesIO()
|
||||
image_data.save(buffer, format='PNG')
|
||||
buffer.seek(0)
|
||||
pil_images.append(Image.open(buffer))
|
||||
except:
|
||||
pil_images.append(Image.new('RGB', (512, 512), color='gray'))
|
||||
elif hasattr(image_data, 'read') or isinstance(image_data, bytes):
|
||||
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 Edit 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('negative_prompt', '')}-{kwargs.get('edit_mode', '')}-{kwargs.get('number_of_images', 1)}-{kwargs.get('seed', 12345)}-{kwargs.get('base_steps', 50)}"
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"GoogleImagenEditNode": GoogleImagenEditNode}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {"GoogleImagenEditNode": "Google Imagen Edit (Vertex AI only)"}
|
||||
Reference in New Issue
Block a user