Files
Aryan185-ComfyUI-ExternalAP…/imagen_edit.py
T

125 lines
5.7 KiB
Python

import io
import json
import base64
import torch
import numpy as np
from PIL import Image
from google import genai
from google.genai import types
from google.oauth2 import service_account
class GoogleImagenEditNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"mask": ("MASK",),
"prompt": ("STRING", {"multiline": True, "default": "Edit this image"}),
"project_id": ("STRING", {"multiline": False, "default": ""}),
"location": (["global", "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": True, "default": "", "tooltip": "Paste service account JSON content"}),
"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": 69, "min": 1, "max": 2147483646, "step": 1}),
"base_steps": ("INT", {"default": 50, "min": 10, "max": 100, "step": 1}),
"guidance_scale": ("FLOAT", {"default": 7.5, "min": 1.0, "max": 20.0, "step": 0.1}),
"mask_dilation": ("FLOAT", {"default": 0.03, "min": 0.0, "max": 1.0, "step": 0.01}),
},
"optional": {
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("edited_images",)
FUNCTION = "edit_image"
CATEGORY = "image/edit"
def setup_client(self, service_account_json, project_id, location):
if not service_account_json.strip():
raise ValueError("Service account JSON content is required.")
if not project_id.strip():
raise ValueError("Project ID is required.")
try:
sa_info = json.loads(service_account_json)
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON content: {str(e)}")
credentials = service_account.Credentials.from_service_account_info(
sa_info,
scopes=["https://www.googleapis.com/auth/cloud-platform"]
)
return genai.Client(
vertexai=True,
project=project_id.strip(),
location=location.strip(),
credentials=credentials,
http_options=types.HttpOptions(
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
)
)
def edit_image(self, image, mask, prompt, project_id, location, service_account,
edit_mode, number_of_images, seed, base_steps, guidance_scale, mask_dilation, negative_prompt=""):
client = self.setup_client(service_account, project_id, location)
def to_b64(img):
b = io.BytesIO()
img.save(b, format='PNG')
return base64.b64encode(b.getvalue()).decode('utf-8')
img_pil = Image.fromarray((image[0].cpu().numpy() * 255).astype(np.uint8))
mask_np = mask.cpu().numpy()
if mask_np.ndim == 3: mask_np = mask_np[0]
mask_pil = Image.fromarray((mask_np * 255).astype(np.uint8), mode='L')
config_dict = {
"edit_mode": edit_mode,
"number_of_images": number_of_images,
"base_steps": base_steps,
"seed": seed,
"guidance_scale": guidance_scale,
"output_mime_type": "image/jpeg",
"include_rai_reason": True,
}
if negative_prompt.strip():
config_dict["negative_prompt"] = negative_prompt.strip()
try:
response = client.models.edit_image(
model="imagen-3.0-capability-001",
prompt=prompt,
reference_images=[
types.RawReferenceImage(reference_id=0, reference_image={'image_bytes': to_b64(img_pil)}),
types.MaskReferenceImage(reference_id=1, reference_image={'image_bytes': to_b64(mask_pil)},
config=types.MaskReferenceConfig(mask_mode="MASK_MODE_USER_PROVIDED", mask_dilation=mask_dilation))
],
config=types.EditImageConfig(**config_dict)
)
if not response.generated_images:
raise ValueError("No images generated")
output_tensors = []
for item in response.generated_images:
res_img = Image.open(io.BytesIO(item.image.image_bytes)).convert("RGB")
output_tensors.append(torch.from_numpy(np.array(res_img).astype(np.float32) / 255.0))
return (torch.stack(output_tensors),)
except Exception as e:
raise RuntimeError(f"Google Imagen Edit Error: {e}")
@classmethod
def IS_CHANGED(cls, **kwargs):
return float("nan")
NODE_CLASS_MAPPINGS = {"GoogleImagenEditNode": GoogleImagenEditNode}
NODE_DISPLAY_NAME_MAPPINGS = {"GoogleImagenEditNode": "Google Imagen Edit (Vertex AI only)"}