added Dalle Node

This commit is contained in:
Fillip
2024-08-22 00:15:01 -07:00
parent 1077ec298c
commit c3345f93a4
3 changed files with 228 additions and 0 deletions
+3
View File
@@ -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",
}
+119
View File
@@ -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)
+106
View File
@@ -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);
}