245 lines
9.0 KiB
Python
245 lines
9.0 KiB
Python
import os
|
|
import json
|
|
import base64
|
|
import urllib.request
|
|
import urllib.error
|
|
from PIL import Image
|
|
import io
|
|
import numpy as np
|
|
import torch
|
|
|
|
### ===== Gemini Image 2 Prompt Node ===== ###
|
|
|
|
class GeminiImageToPrompt:
|
|
CATEGORY = "flux/prompt"
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"api_key": ("STRING", {"multiline": False, "default": ""}),
|
|
"main_image": ("IMAGE", {"label": "Main Subject Image"}),
|
|
"background_image": ("IMAGE", {"label": "Background Scene Image"})
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("STRING",)
|
|
RETURN_NAMES = ("prompt",)
|
|
FUNCTION = "generate_prompt"
|
|
OUTPUT_NODE = False
|
|
DESCRIPTION = "Extract main subject and background scene from two images using Gemini API and generate a descriptive prompt with style and language."
|
|
|
|
def encode_image(self, image_data):
|
|
if isinstance(image_data, torch.Tensor):
|
|
image_data = image_data.cpu().numpy()
|
|
image_data = (image_data * 255).clip(0, 255).astype(np.uint8)
|
|
img = Image.fromarray(image_data.squeeze())
|
|
buffered = io.BytesIO()
|
|
img.save(buffered, format="PNG")
|
|
return base64.b64encode(buffered.getvalue()).decode("utf-8")
|
|
|
|
def call_gemini_api(self, api_key, img_base64, task):
|
|
headers = {
|
|
"Content-Type": "application/json",
|
|
"x-goog-api-key": api_key
|
|
}
|
|
|
|
body = {
|
|
"contents": [
|
|
{
|
|
"parts": [
|
|
{
|
|
"text": (
|
|
f"Analyze the image and describe the {task} for use in AI image generation. "
|
|
f"Include:\n- Description:\n- Style:\n- Language:\n"
|
|
)
|
|
},
|
|
{
|
|
"inline_data": {
|
|
"mime_type": "image/png",
|
|
"data": img_base64
|
|
}
|
|
}
|
|
]
|
|
}
|
|
]
|
|
}
|
|
|
|
try:
|
|
req = urllib.request.Request(
|
|
url="https://generativelanguage.googleapis.com/v1beta/models/gemini-pro-vision:generateContent",
|
|
data=json.dumps(body).encode("utf-8"),
|
|
headers=headers,
|
|
method="POST"
|
|
)
|
|
with urllib.request.urlopen(req) as response:
|
|
res = json.loads(response.read().decode("utf-8"))
|
|
return res['candidates'][0]['content']['parts'][0]['text']
|
|
except Exception as e:
|
|
return f"[Gemini Error] {e}"
|
|
|
|
def extract_info(self, raw):
|
|
info = {"description": "", "style": "", "language": ""}
|
|
for line in raw.strip().split("\n"):
|
|
if line.lower().startswith("style:"):
|
|
info["style"] = line.split(":", 1)[-1].strip()
|
|
elif line.lower().startswith("language:"):
|
|
info["language"] = line.split(":", 1)[-1].strip()
|
|
elif line.lower().startswith("description:"):
|
|
info["description"] = line.split(":", 1)[-1].strip()
|
|
else:
|
|
info["description"] += " " + line.strip()
|
|
return info
|
|
|
|
def generate_prompt(self, api_key, main_image, background_image):
|
|
main_b64 = self.encode_image(main_image)
|
|
bg_b64 = self.encode_image(background_image)
|
|
|
|
main_raw = self.call_gemini_api(api_key, main_b64, "main subject")
|
|
bg_raw = self.call_gemini_api(api_key, bg_b64, "background scene")
|
|
|
|
main_info = self.extract_info(main_raw)
|
|
bg_info = self.extract_info(bg_raw)
|
|
|
|
prompt = (
|
|
f"{main_info['description'].strip()} placed within {bg_info['description'].strip()}. "
|
|
f"This image is rendered in {main_info['style'] or bg_info['style'] or 'cinematic'} style, "
|
|
f"using {main_info['language'] or bg_info['language'] or 'natural'} language. "
|
|
f"Focus on mood, lighting, and artistic detail."
|
|
)
|
|
|
|
return (prompt,)
|
|
|
|
|
|
### ===== Smart Auto Prompt Node ===== ###
|
|
|
|
class SmartAutoPromptNode:
|
|
CATEGORY = "flux/prompt"
|
|
|
|
def __init__(self):
|
|
self.state = {}
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"api_key": ("STRING", {"multiline": False, "default": "sk-xxx"}),
|
|
"model": (["gpt-4o", "gpt-4", "gpt-3.5-turbo"],),
|
|
"num_prompts": ("INT", {"default": 25, "min": 1, "max": 100}),
|
|
"prompt_input": ("STRING", {"multiline": True, "default": ""}),
|
|
"subject": ("STRING", {"default": ""}),
|
|
"obj": ("STRING", {"default": ""}),
|
|
"lora_trigger": ("STRING", {"default": ""}),
|
|
"setting": ("STRING", {"default": ""}),
|
|
"interaction": ("STRING", {"default": ""}),
|
|
"style": ("STRING", {"default": ""}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("STRING",)
|
|
RETURN_NAMES = ("prompt",)
|
|
FUNCTION = "generate_prompt"
|
|
OUTPUT_NODE = False
|
|
DESCRIPTION = "Generate a single-line prompt for AI image generation. Accepts raw or structured input."
|
|
|
|
def parse_prompt(self, api_key, model, prompt_input):
|
|
system = (
|
|
"You are an expert prompt parser. Given a detailed AI image generation prompt, "
|
|
"extract and return a JSON object with the following fields: subject, obj, lora_trigger, setting, interaction, style."
|
|
)
|
|
user = f"Prompt: {prompt_input}\nExtract to JSON:"
|
|
payload = {
|
|
"model": model,
|
|
"temperature": 0.3,
|
|
"messages": [
|
|
{"role": "system", "content": system},
|
|
{"role": "user", "content": user}
|
|
]
|
|
}
|
|
headers = {
|
|
"Authorization": f"Bearer {api_key}",
|
|
"Content-Type": "application/json"
|
|
}
|
|
try:
|
|
req = urllib.request.Request(
|
|
url="https://openai-api.codejoyai.com:8003/openai/v1/chat/completions",
|
|
data=json.dumps(payload).encode("utf-8"),
|
|
headers=headers,
|
|
method="POST"
|
|
)
|
|
with urllib.request.urlopen(req) as res:
|
|
content = json.loads(res.read().decode("utf-8"))
|
|
return json.loads(content['choices'][0]['message']['content'])
|
|
except Exception as e:
|
|
return {"error": str(e)}
|
|
|
|
def generate_prompt(self, api_key, model, num_prompts, prompt_input,
|
|
subject, obj, lora_trigger, setting, interaction, style):
|
|
|
|
if prompt_input.strip():
|
|
parsed = self.parse_prompt(api_key, model, prompt_input)
|
|
if "error" in parsed:
|
|
return (f"[Parse Error] {parsed['error']}",)
|
|
subject = parsed.get("subject", subject)
|
|
obj = parsed.get("obj", obj)
|
|
lora_trigger = parsed.get("lora_trigger", lora_trigger)
|
|
setting = parsed.get("setting", setting)
|
|
interaction = parsed.get("interaction", interaction)
|
|
style = parsed.get("style", style)
|
|
|
|
user_prompt = (
|
|
f"Generate {num_prompts} image generation prompts using:\n"
|
|
f"1. Subject: {subject}\n"
|
|
f"2. Object: {obj}\n"
|
|
f"3. LoRA Trigger: '{lora_trigger}'\n"
|
|
f"4. Setting: {setting}\n"
|
|
f"5. Interaction: {interaction}\n"
|
|
f"6. Style: {style}\n"
|
|
f"Start each prompt with the LoRA trigger. Line-separated output only."
|
|
)
|
|
|
|
system_prompt = (
|
|
"You are a professional prompt writer. Generate distinct, high-quality prompts with these elements. "
|
|
"Output only the prompts in plain text, no numbering or commentary."
|
|
)
|
|
|
|
payload = {
|
|
"model": model,
|
|
"temperature": 0.88,
|
|
"messages": [
|
|
{"role": "system", "content": system_prompt},
|
|
{"role": "user", "content": user_prompt}
|
|
]
|
|
}
|
|
headers = {
|
|
"Authorization": f"Bearer {api_key}",
|
|
"Content-Type": "application/json"
|
|
}
|
|
|
|
try:
|
|
request = urllib.request.Request(
|
|
"https://openai-api.codejoyai.com:8003/openai/v1/chat/completions",
|
|
data=json.dumps(payload).encode("utf-8"),
|
|
headers=headers,
|
|
method='POST'
|
|
)
|
|
with urllib.request.urlopen(request) as response:
|
|
result = json.loads(response.read().decode("utf-8"))
|
|
lines = [line.strip() for line in result['choices'][0]['message']['content'].split("\n") if line.strip()]
|
|
return (lines[0] if lines else "",)
|
|
except Exception as e:
|
|
return (f"[Generation Error] {str(e)}",)
|
|
|
|
|
|
### ===== Node Registration ===== ###
|
|
|
|
NODE_CLASS_MAPPINGS = {
|
|
"SmartAutoPromptNode": SmartAutoPromptNode,
|
|
"GeminiImageToPrompt": GeminiImageToPrompt
|
|
}
|
|
|
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
|
"SmartAutoPromptNode": "Prompt Generator (Smart Input)",
|
|
"GeminiImageToPrompt": "Gemini Image 2 Prompt"
|
|
}
|