added Dalle Node
This commit is contained in:
@@ -55,6 +55,7 @@ from. nodes.FL_KSamplerXYZPlot import FL_KSamplerXYZPlot
|
||||
from .nodes.FL_SamplerStrings import FL_SamplerStrings
|
||||
from .nodes.FL_SchedulerStrings import FL_SchedulerStrings
|
||||
from .nodes.FL_ImageCaptionLayoutPDF import FL_ImageCaptionLayoutPDF
|
||||
from .nodes.FL_Dalle3 import FL_Dalle3
|
||||
|
||||
|
||||
|
||||
@@ -117,6 +118,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"FL_SamplerStrings": FL_SamplerStrings,
|
||||
"FL_SchedulerStrings": FL_SchedulerStrings,
|
||||
"FL_ImageCaptionLayoutPDF": FL_ImageCaptionLayoutPDF,
|
||||
"FL_Dalle3": FL_Dalle3,
|
||||
|
||||
}
|
||||
|
||||
@@ -178,6 +180,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FL_SamplerStrings": "FL Sampler String XYZ",
|
||||
"FL_SchedulerStrings": "FL Scheduler String XYZ",
|
||||
"FL_ImageCaptionLayoutPDF": "FL Image Caption Layout PDF",
|
||||
"FL_Dalle3": "FL Dalle 3",
|
||||
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
import openai
|
||||
import base64
|
||||
import io
|
||||
import os
|
||||
import json
|
||||
import asyncio
|
||||
import aiohttp
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torchvision.transforms import functional as TF
|
||||
|
||||
|
||||
class FL_Dalle3:
|
||||
def __init__(self):
|
||||
self.__client = openai.AsyncOpenAI()
|
||||
self.__previous_params = None
|
||||
self.__cache_images = None
|
||||
self.__cache_revised_prompts = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"resolution": (["1024x1024", "1024x1792", "1792x1024"],),
|
||||
"dummy_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"prompt": ("STRING", {
|
||||
"multiline": True,
|
||||
"default": "great picture"
|
||||
}),
|
||||
"quality": (["HD", "Standard"],),
|
||||
"style": (["vivid", "natural"],),
|
||||
"batch_size": ("INT", {"default": 1, "min": 1, "max": 10}),
|
||||
"retry": ("INT", {"default": 0, "min": 0, "max": 5}),
|
||||
},
|
||||
"optional": {
|
||||
"auto_save": ("BOOLEAN", {"default": False}),
|
||||
"auto_save_dir": ("STRING", {"default": "./output_dalle3"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "INT", "INT", "STRING")
|
||||
RETURN_NAMES = ("IMAGES", "WIDTH", "HEIGHT", "REVISED_PROMPTS")
|
||||
FUNCTION = "generate_images"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "🏵️Fill Nodes/GPT"
|
||||
|
||||
async def generate_single_image(self, prompt, resolution, quality, style, retry):
|
||||
for retry_count in range(retry + 1):
|
||||
try:
|
||||
response = await self.__client.images.generate(
|
||||
model="dall-e-3",
|
||||
prompt=prompt,
|
||||
size=resolution,
|
||||
quality="hd" if quality == "HD" else "standard",
|
||||
style="vivid" if style == "vivid" else "natural",
|
||||
n=1,
|
||||
response_format="b64_json"
|
||||
)
|
||||
return response
|
||||
except openai.BadRequestError as ex:
|
||||
if retry_count >= retry:
|
||||
raise ex
|
||||
print(
|
||||
f"FL_OpenAiDalle3: received BadRequestError, retrying... #{retry_count + 1} : {json.dumps(ex.response.json())}")
|
||||
return None
|
||||
|
||||
async def generate_batch(self, prompt, resolution, quality, style, batch_size, retry):
|
||||
tasks = [self.generate_single_image(prompt, resolution, quality, style, retry) for _ in range(batch_size)]
|
||||
return await asyncio.gather(*tasks)
|
||||
|
||||
def generate_images(self, resolution, dummy_seed, prompt, quality, style, batch_size, retry, auto_save=False,
|
||||
auto_save_dir="./output_dalle3"):
|
||||
current_params = (resolution, dummy_seed, prompt, quality, style, batch_size)
|
||||
|
||||
if self.__cache_images is None or self.__previous_params != current_params:
|
||||
responses = asyncio.run(self.generate_batch(prompt, resolution, quality, style, batch_size, retry))
|
||||
|
||||
images = []
|
||||
revised_prompts = []
|
||||
|
||||
for i, r0 in enumerate(responses):
|
||||
if r0 is None:
|
||||
continue
|
||||
|
||||
im0 = Image.open(io.BytesIO(base64.b64decode(r0.data[0].b64_json)))
|
||||
|
||||
if auto_save:
|
||||
os.makedirs(auto_save_dir, exist_ok=True)
|
||||
next_index = len([f for f in os.listdir(auto_save_dir) if f.endswith('.png')]) + 1
|
||||
image_file_name = os.path.join(auto_save_dir, f"dalle3_output_{next_index:06d}.png")
|
||||
state_file_name = os.path.join(auto_save_dir, f"dalle3_output_{next_index:06d}.json")
|
||||
im0.save(image_file_name)
|
||||
with open(state_file_name, "wt") as f:
|
||||
json.dump({
|
||||
"resolution": resolution,
|
||||
"prompt": prompt,
|
||||
"quality": quality,
|
||||
"style": style,
|
||||
"batch_index": i
|
||||
}, f, indent=2, ensure_ascii=False)
|
||||
|
||||
im1 = TF.to_tensor(im0.convert("RGBA"))
|
||||
im1[:3, im1[3, :, :] == 0] = 0
|
||||
images.append(im1)
|
||||
revised_prompts.append(r0.data[0].revised_prompt)
|
||||
|
||||
self.__previous_params = current_params
|
||||
self.__cache_images = images
|
||||
self.__cache_revised_prompts = revised_prompts
|
||||
else:
|
||||
images = self.__cache_images
|
||||
revised_prompts = self.__cache_revised_prompts
|
||||
|
||||
images_tensor = torch.stack(images)
|
||||
images_tensor = images_tensor.permute(0, 2, 3, 1)
|
||||
images_tensor = images_tensor[:, :, :, :3]
|
||||
width, height = map(int, resolution.split("x"))
|
||||
return images_tensor, width, height, ", ".join(revised_prompts)
|
||||
@@ -0,0 +1,106 @@
|
||||
import { app } from "../../../scripts/app.js";
|
||||
|
||||
// Animation parameters
|
||||
const ANIMATION_WIDTH = 120;
|
||||
const ANIMATION_HEIGHT = 50;
|
||||
const GHOST_SIZE = 15;
|
||||
const ANIMATION_X_OFFSET = -10;
|
||||
const ANIMATION_Y_OFFSET = 10;
|
||||
|
||||
app.registerExtension({
|
||||
name: "Ghost-API-Animation",
|
||||
async nodeCreated(node) {
|
||||
const animatedNodeClasses = [
|
||||
"FL_Dalle3",
|
||||
// Add other API-related node classes here
|
||||
];
|
||||
|
||||
if (animatedNodeClasses.includes(node.comfyClass)) {
|
||||
addGhostAPIAnimation(node);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
function addGhostAPIAnimation(node) {
|
||||
let ghosts = [];
|
||||
|
||||
function createGhost() {
|
||||
return {
|
||||
x: ANIMATION_WIDTH,
|
||||
y: Math.random() * ANIMATION_HEIGHT,
|
||||
speed: 0.5 + Math.random() * 1,
|
||||
size: GHOST_SIZE + Math.random() * 5,
|
||||
opacity: 1,
|
||||
waveFreq: 0.05 + Math.random() * 0.05,
|
||||
waveAmp: 1 + Math.random() * 2
|
||||
};
|
||||
}
|
||||
|
||||
function updateGhosts() {
|
||||
ghosts = ghosts.filter(ghost => ghost.x > -ghost.size && ghost.opacity > 0);
|
||||
ghosts.forEach(ghost => {
|
||||
ghost.x -= ghost.speed;
|
||||
ghost.opacity -= 0.01;
|
||||
});
|
||||
|
||||
if (Math.random() > 0.97) {
|
||||
ghosts.push(createGhost());
|
||||
}
|
||||
}
|
||||
|
||||
function drawGhost(ctx, ghost) {
|
||||
ctx.save();
|
||||
ctx.translate(ghost.x, ghost.y + Math.sin(ghost.x * ghost.waveFreq) * ghost.waveAmp);
|
||||
|
||||
// Ghost body
|
||||
ctx.beginPath();
|
||||
ctx.moveTo(0, 0);
|
||||
ctx.bezierCurveTo(-ghost.size/2, -ghost.size/2, -ghost.size/2, -ghost.size, 0, -ghost.size);
|
||||
ctx.bezierCurveTo(ghost.size/2, -ghost.size, ghost.size/2, -ghost.size/2, 0, 0);
|
||||
|
||||
// Ghost tail
|
||||
ctx.quadraticCurveTo(-ghost.size/4, ghost.size/2, -ghost.size/2, ghost.size);
|
||||
ctx.quadraticCurveTo(-ghost.size/8, ghost.size/2, 0, ghost.size);
|
||||
ctx.quadraticCurveTo(ghost.size/8, ghost.size/2, ghost.size/2, ghost.size);
|
||||
ctx.quadraticCurveTo(ghost.size/4, ghost.size/2, 0, 0);
|
||||
|
||||
ctx.fillStyle = `rgba(255, 255, 255, ${ghost.opacity})`;
|
||||
ctx.fill();
|
||||
|
||||
// Eyes
|
||||
ctx.fillStyle = `rgba(0, 0, 0, ${ghost.opacity})`;
|
||||
ctx.beginPath();
|
||||
ctx.arc(-ghost.size/4, -ghost.size/2, ghost.size/10, 0, Math.PI * 2);
|
||||
ctx.arc(ghost.size/4, -ghost.size/2, ghost.size/10, 0, Math.PI * 2);
|
||||
ctx.fill();
|
||||
|
||||
ctx.restore();
|
||||
}
|
||||
|
||||
node.onDrawBackground = function(ctx) {
|
||||
if (!this.flags.collapsed) {
|
||||
ctx.save();
|
||||
|
||||
const nodeWidth = this.size[0];
|
||||
const baseXOffset = (nodeWidth - ANIMATION_WIDTH) / 2;
|
||||
ctx.translate(baseXOffset + ANIMATION_X_OFFSET, ANIMATION_Y_OFFSET);
|
||||
|
||||
// Draw ghosts
|
||||
ghosts.forEach(ghost => drawGhost(ctx, ghost));
|
||||
|
||||
// Draw API text
|
||||
ctx.fillStyle = 'rgba(255, 255, 255, 0.7)';
|
||||
ctx.font = '12px Arial';
|
||||
ctx.fillText('', 5, ANIMATION_HEIGHT - 5);
|
||||
|
||||
ctx.restore();
|
||||
|
||||
updateGhosts();
|
||||
|
||||
this.setDirtyCanvas(true);
|
||||
requestAnimationFrame(() => this.setDirtyCanvas(true));
|
||||
}
|
||||
};
|
||||
|
||||
node.setDirtyCanvas(true);
|
||||
}
|
||||
Reference in New Issue
Block a user