support flux and Ideogram node

This commit is contained in:
Jimmy Wong
2024-10-31 17:22:43 +08:00
parent 2eea88b841
commit 094b28ac89
6 changed files with 99 additions and 100 deletions
+3 -3
View File
@@ -4,7 +4,7 @@ import logging
from .types import STRING
from .api_key_manager import save_api_key
# 设置日志
# Set up logging
logging.basicConfig(level=logging.DEBUG)
logger = logging.getLogger(__name__)
@@ -20,7 +20,7 @@ from .nodes_json import FlowyPreviewJSON, FlowyExtractJSON, ComflowyLoadJSON
from .nodes_http import FlowyHttpRequest
from .nodes_llm import FlowyLLM
from .nodes_upscale import FlowyUpscale
from .nodes_flux import ComflowyFlux
from .nodes_flux import FlowyFlux
from .nodes_ideogram import FlowyIdeogram
@@ -72,7 +72,7 @@ NODE_CLASS_MAPPINGS = {
"Comflowy_Set_API_Key": ComflowySetAPIKey,
"Comflowy_Upscale": FlowyUpscale,
"Comflowy_Ideogram": FlowyIdeogram,
"Comflowy_Flux": ComflowyFlux,
"Comflowy_Flux": FlowyFlux,
}
NODE_DISPLAY_NAME_MAPPINGS = {
+34 -37
View File
@@ -7,13 +7,13 @@ import torch
import numpy as np
import logging
import json
from .types import STRING, INT, API_HOST
from .types import STRING, INT, API_HOST, SAFETY_TOLERANCE, BOOLEAN
from .utils import logger, get_nested_value
from .api_key_manager import load_api_key
logger = logging.getLogger(__name__)
class ComflowyFlux:
class FlowyFlux:
@classmethod
def INPUT_TYPES(s):
return {
@@ -34,20 +34,15 @@ class ComflowyFlux:
],),
"height": ("INT", {"default": 256, "min": 256, "max": 1440}),
"width": ("INT", {"default": 256, "min": 256, "max": 1440}),
"prompt_upsampling": (["Off", "On"],),
"safety_tolerance": ([
"1",
"2",
"3",
"4",
"5"
]),
"seed": ("INT", {"default": 0, "min": 0, "max": 2147483647}),
"prompt_upsampling": BOOLEAN,
"safety_tolerance": (SAFETY_TOLERANCE,),
"output_quality": ("INT", {"default": 80, "min": 1, "max": 100}),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "generate_image_with_flux"
FUNCTION = "generate"
CATEGORY = "Comflowy"
DESCRIPTION = """
Nodes from https://comflowy.com:
@@ -58,11 +53,12 @@ Nodes from https://comflowy.com:
- Height and width are only used when aspect_ratio=custom. Must be a multiple of 32 (if it's not, it will be rounded to nearest multiple of 32).
- Prompt Upsampling: Automatically modify the prompt for more creative generation.
- Safety tolerance, 1 is most strict and 5 is most permissive.
- Quality when saving the output images, from 0 to 100. 100 is best quality, 0 is lowest quality. Not relevant for .png outputs.
- Make sure to set your API Key using the 'Comflowy Set API Key' node before using this node.
- Output: Returns the generated image.
"""
def generate_image_with_flux(self, prompt, version, aspect_ratio, height, width, seed, prompt_upsampling, safety_tolerance):
def generate(self, prompt, version, aspect_ratio, height, width, seed, prompt_upsampling, safety_tolerance, output_quality):
api_key = load_api_key()
if not api_key:
@@ -70,7 +66,7 @@ Nodes from https://comflowy.com:
logger.error(error_msg)
raise ValueError(error_msg)
logger.info(f"开始处理 Flux 图像生成请求。prompt: {prompt}, version: {version}, aspect_ratio: {aspect_ratio}, height: {height}, width: {width}, seed: {seed}, prompt_upsampling: {prompt_upsampling}, safety_tolerance: {safety_tolerance}")
logger.info(f"Starting Flux image generation request. prompt: {prompt}, version: {version}, aspect_ratio: {aspect_ratio}, height: {height}, width: {width}, seed: {seed}, prompt_upsampling: {prompt_upsampling}, safety_tolerance: {safety_tolerance}, output_quality: {output_quality}")
try:
response = requests.post(
@@ -82,68 +78,69 @@ Nodes from https://comflowy.com:
"aspect_ratio": aspect_ratio,
"height": height,
"width": width,
"seed": seed,
"prompt_upsampling": prompt_upsampling,
"safety_tolerance": safety_tolerance,
"seed": seed,
"output_quality": output_quality,
}
)
response.raise_for_status()
result = response.json()
logger.info(f"API 请求完成。状态码: {response.status_code}")
logger.debug(f"API 响应内容: {json.dumps(result, indent=2)}")
logger.info(f"API request completed. Status code: {response.status_code}")
logger.debug(f"API response content: {json.dumps(result, indent=2)}")
if not result.get('success'):
logger.error(f"API 请求失败。响应内容: {json.dumps(result, indent=2)}")
raise Exception(f"API 请求失败。响应内容: {json.dumps(result, indent=2)}")
logger.error(f"API request failed. Response content: {json.dumps(result, indent=2)}")
raise Exception(f"API request failed. Response content: {json.dumps(result, indent=2)}")
output_url = result.get('data', {}).get('output')
if not output_url or not isinstance(output_url, str):
logger.error(f"完整的 API 响应: {json.dumps(result, indent=2)}")
raise Exception(f"无法获取有效的输出图像 URL。API 响应中没有预期的数据结构。完整响应: {json.dumps(result, indent=2)}")
logger.error(f"Complete API response: {json.dumps(result, indent=2)}")
raise Exception(f"Unable to get valid output image URL. API response does not have expected data structure. Complete response: {json.dumps(result, indent=2)}")
logger.info(f"获取到的输出 URL: {output_url}")
logger.info(f"Obtained output URL: {output_url}")
# 验证 URL 是否可访问
# Verify URL is accessible
try:
url_check = requests.head(output_url)
url_check.raise_for_status()
except requests.RequestException as e:
logger.error(f"无法访问输出 URL: {str(e)}")
raise Exception(f"无法访问输出 URL: {str(e)}")
logger.error(f"Unable to access output URL: {str(e)}")
raise Exception(f"Unable to access output URL: {str(e)}")
# 添加延迟,等待 Replicate 处理完成
# Add delay, wait for Replicate to process
time.sleep(10)
img_response = requests.get(output_url, stream=True)
img_response.raise_for_status()
# 将图像数据转换为 PIL Image
# Convert image data to PIL Image
img = Image.open(img_response.raw)
# 转换为 numpy 数组
# Convert to numpy array
img_np = np.array(img)
# 确保图像是 3 通道 RGB
if len(img_np.shape) == 2: # 灰度图像
# Ensure image is 3 channel RGB
if len(img_np.shape) == 2: # Grayscale image
img_np = np.stack([img_np] * 3, axis=-1)
elif img_np.shape[-1] == 4: # RGBA 图像
elif img_np.shape[-1] == 4: # RGBA image
img_np = img_np[:, :, :3]
# 转换为 float32 并归一化到 0-1 范围
# Convert to float32 and normalize to 0-1 range
img_np = img_np.astype(np.float32) / 255.0
# 转换为 torch tensor,确保形状为 [B,H,W,C]
img_tensor = torch.from_numpy(img_np).unsqueeze(0) # 添加批次维度
# Convert to torch tensor, ensuring shape is [B,H,W,C]
img_tensor = torch.from_numpy(img_np).unsqueeze(0) # Add batch dimension
logger.info(f"图像处理完成。输出张量形状: {img_tensor.shape}")
logger.info(f"Image processing completed. Output tensor shape: {img_tensor.shape}")
return (img_tensor,)
except Exception as e:
error_msg = f"图像生成过程中出错: {str(e)}"
error_msg = f"Error during image generation: {str(e)}"
logger.error(error_msg)
logger.exception("详细错误信息:")
# 返回一个错误标记图像,确保形状为 [B,H,W,C]
logger.exception("Detailed error information:")
# Return an error marked image, ensuring shape is [B,H,W,C]
error_image = torch.zeros((1, 100, 400, 3), dtype=torch.float32)
return (error_image,)
+24 -24
View File
@@ -52,7 +52,7 @@ Nodes from https://comflowy.com:
logger.error(error_msg)
raise ValueError(error_msg)
logger.info(f"开始处理 Ideogram 图像生成请求。prompt: {prompt}, negative_prompt: {negative_prompt}, resolution: {resolution}, style_type: {style_type}, aspect_ratio: {aspect_ratio}, magic_prompt_option: {magic_prompt_option}, seed: {seed}")
logger.info(f"Starting Ideogram image generation request. prompt: {prompt}, negative_prompt: {negative_prompt}, resolution: {resolution}, style_type: {style_type}, aspect_ratio: {aspect_ratio}, magic_prompt_option: {magic_prompt_option}, seed: {seed}")
try:
response = requests.post(
@@ -72,60 +72,60 @@ Nodes from https://comflowy.com:
response.raise_for_status()
result = response.json()
logger.info(f"API 请求完成。状态码: {response.status_code}")
logger.debug(f"API 响应内容: {json.dumps(result, indent=2)}")
logger.info(f"API request completed. Status code: {response.status_code}")
logger.debug(f"API response content: {json.dumps(result, indent=2)}")
if not result.get('success'):
logger.error(f"API 请求失败。响应内容: {json.dumps(result, indent=2)}")
raise Exception(f"API 请求失败。响应内容: {json.dumps(result, indent=2)}")
logger.error(f"API request failed. Response content: {json.dumps(result, indent=2)}")
raise Exception(f"API request failed. Response content: {json.dumps(result, indent=2)}")
output_url = result.get('data', {}).get('output')
if not output_url or not isinstance(output_url, str):
logger.error(f"完整的 API 响应: {json.dumps(result, indent=2)}")
raise Exception(f"无法获取有效的输出图像 URL。API 响应中没有预期的数据结构。完整响应: {json.dumps(result, indent=2)}")
logger.error(f"Complete API response: {json.dumps(result, indent=2)}")
raise Exception(f"Unable to get valid output image URL. API response does not have expected data structure. Complete response: {json.dumps(result, indent=2)}")
logger.info(f"获取到的输出 URL: {output_url}")
logger.info(f"Obtained output URL: {output_url}")
# 验证 URL 是否可访问
# Verify URL is accessible
try:
url_check = requests.head(output_url)
url_check.raise_for_status()
except requests.RequestException as e:
logger.error(f"无法访问输出 URL: {str(e)}")
raise Exception(f"无法访问输出 URL: {str(e)}")
logger.error(f"Unable to access output URL: {str(e)}")
raise Exception(f"Unable to access output URL: {str(e)}")
# 添加延迟,等待 Replicate 处理完成
# Add delay, wait for Replicate to process
time.sleep(10)
img_response = requests.get(output_url, stream=True)
img_response.raise_for_status()
# 将图像数据转换为 PIL Image
# Convert image data to PIL Image
img = Image.open(img_response.raw)
# 转换为 numpy 数组
# Convert to numpy array
img_np = np.array(img)
# 确保图像是 3 通道 RGB
if len(img_np.shape) == 2: # 灰度图像
# Ensure image is 3 channel RGB
if len(img_np.shape) == 2: # Grayscale image
img_np = np.stack([img_np] * 3, axis=-1)
elif img_np.shape[-1] == 4: # RGBA 图像
elif img_np.shape[-1] == 4: # RGBA image
img_np = img_np[:, :, :3]
# 转换为 float32 并归一化到 0-1 范围
# Convert to float32 and normalize to 0-1 range
img_np = img_np.astype(np.float32) / 255.0
# 转换为 torch tensor,确保形状为 [B,H,W,C]
img_tensor = torch.from_numpy(img_np).unsqueeze(0) # 添加批次维度
# Convert to torch tensor, ensuring shape is [B,H,W,C]
img_tensor = torch.from_numpy(img_np).unsqueeze(0) # Add batch dimension
logger.info(f"图像处理完成。输出张量形状: {img_tensor.shape}")
logger.info(f"Image processing completed. Output tensor shape: {img_tensor.shape}")
return (img_tensor,)
except Exception as e:
error_msg = f"图像生成过程中出错: {str(e)}"
error_msg = f"Error during image generation: {str(e)}"
logger.error(error_msg)
logger.exception("详细错误信息:")
# 返回一个错误标记图像,确保形状为 [B,H,W,C]
logger.exception("Detailed error information:")
# Return an error marked image, ensuring shape is [B,H,W,C]
error_image = torch.zeros((1, 100, 400, 3), dtype=torch.float32)
return (error_image,)
+4 -4
View File
@@ -130,7 +130,7 @@ class OmostLLMNode:
try:
generated_text = llm_request(prompt=prompt, llm_model=llm_model, system_prompt=system_prompt, api_key=api_key, max_tokens=4000, timeout=10)
# 如果生成的字符中包含了多余的字符,比如 "```json" 或者 "```",则需要去掉改行
# If the generated text contains extra characters, such as "```json" or "```", remove the line
generated_text = generated_text.replace("```json", "").replace("```", "")
try:
@@ -301,7 +301,7 @@ class OmostToConditioning:
)
# 对于 LLM 动态生成的区域描述,该节点用于预览canvas的节点
# For LLM-generated region descriptions, this node is used to preview the canvas
class ComflowyOmostPreviewNode:
@classmethod
def INPUT_TYPES(s):
@@ -329,7 +329,7 @@ class ComflowyOmostPreviewNode:
)
# 对于高级用户,可以直接编辑python代码,然后加载到这个节点中
# For advanced users, you can directly edit the python code and load it into this node
class ComflowyOmostLoadCanvasPythonCodeNode:
"""Load python code generated by Omost demo app."""
@@ -350,7 +350,7 @@ class ComflowyOmostLoadCanvasPythonCodeNode:
canvas = OmostCanvas.from_python_code(python_str)
return (canvas.process(),)
# 定义这个节点可以在后续直接做一个前端编辑器,用于编辑基于区域的条件
# Define this node to allow for a frontend editor to edit the canvas conditions
class ComflowyOmostLoadCanvasConditioningNode:
@classmethod
def INPUT_TYPES(s):
+30 -30
View File
@@ -45,12 +45,12 @@ Nodes from https://comflowy.com:
logger.error(error_msg)
raise ValueError(error_msg)
logger.info(f"开始处理图像放大请求。scale_factor: {scale_factor}, model: {model}")
logger.info(f"Starting image upscale request. scale_factor: {scale_factor}, model: {model}")
# 处理输入图像
# Process input image
if isinstance(image, torch.Tensor):
if image.dim() == 4:
image = image.squeeze(0) # 移除批次维度
image = image.squeeze(0) # Remove batch dimension
if image.shape[-1] == 3:
image = (image.cpu().numpy() * 255).astype(np.uint8)
elif image.shape[0] == 3:
@@ -68,13 +68,13 @@ Nodes from https://comflowy.com:
else:
raise ValueError(f"Unsupported image type: {type(image)}")
# 将输入图像转换为 JPEG 格式并压缩
# Convert input image to JPEG format and compress
buffered = io.BytesIO()
Image.fromarray(image).save(buffered, format="JPEG", quality=85)
img_str = base64.b64encode(buffered.getvalue()).decode()
try:
# 使用 API_HOST 构建 API 请求的 URL
# Build the URL for the API request
response = requests.post(
f"{API_HOST}/api/open/v0/upscale",
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
@@ -87,62 +87,62 @@ Nodes from https://comflowy.com:
response.raise_for_status()
result = response.json()
logger.info(f"API 请求完成。状态码: {response.status_code}")
logger.debug(f"API 响应内容: {json.dumps(result, indent=2)}")
logger.info(f"API request completed. Status code: {response.status_code}")
logger.debug(f"API response content: {json.dumps(result, indent=2)}")
if not result.get('success'):
logger.error(f"API 请求失败。响应内容: {json.dumps(result, indent=2)}")
raise Exception(f"API 请求失败。响应内容: {json.dumps(result, indent=2)}")
logger.error(f"API request failed. Response content: {json.dumps(result, indent=2)}")
raise Exception(f"API request failed. Response content: {json.dumps(result, indent=2)}")
output_url = result.get('data', {}).get('output', [None])[0]
if not output_url:
logger.error(f"完整的 API 响应: {json.dumps(result, indent=2)}")
raise Exception(f"无法获取输出图像 URL。API 响应中没有预期的数据结构。完整响应: {json.dumps(result, indent=2)}")
logger.error(f"Complete API response: {json.dumps(result, indent=2)}")
raise Exception(f"Unable to get valid output image URL. API response does not have expected data structure. Complete response: {json.dumps(result, indent=2)}")
logger.info(f"获取到的输出 URL: {output_url}")
logger.info(f"Obtained output URL: {output_url}")
# 验证 URL 是否可访问
# Verify URL is accessible
try:
url_check = requests.head(output_url)
url_check.raise_for_status()
except requests.RequestException as e:
logger.error(f"无法访问输出 URL: {str(e)}")
raise Exception(f"无法访问输出 URL: {str(e)}")
logger.error(f"Unable to access output URL: {str(e)}")
raise Exception(f"Unable to access output URL: {str(e)}")
# 添加延迟,等待 Replicate 处理完成
# Add delay, wait for Replicate to process
time.sleep(10)
img_response = requests.get(output_url, stream=True)
img_response.raise_for_status()
# 将图像数据转换为 PIL Image
# Convert image data to PIL Image
img = Image.open(img_response.raw)
# 转换为 numpy 数组
# Convert image data to numpy array
img_np = np.array(img)
# 确保图像是 3 通道 RGB
if len(img_np.shape) == 2: # 灰度图像
# Ensure image is 3 channel RGB
if len(img_np.shape) == 2: # Grayscale image
img_np = np.stack([img_np] * 3, axis=-1)
elif img_np.shape[-1] == 4: # RGBA 图像
elif img_np.shape[-1] == 4: # RGBA image
img_np = img_np[:, :, :3]
# 转换为 float32 并归一化到 0-1 范围
# Convert to float32 and normalize to 0-1 range
img_np = img_np.astype(np.float32) / 255.0
# 转换为 torch tensor,确保形状为 [B,H,W,C]
img_tensor = torch.from_numpy(img_np).unsqueeze(0) # 添加批次维度
# Convert to torch tensor, ensuring shape is [B,H,W,C]
img_tensor = torch.from_numpy(img_np).unsqueeze(0) # Add batch dimension
logger.info(f"图像处理完成。输出张量形状: {img_tensor.shape}")
logger.info(f"API 请求完成。状态码: {response.status_code}")
logger.debug(f"API 响应内容: {response.text}")
logger.info(f"Image processing completed. Output tensor shape: {img_tensor.shape}")
logger.info(f"API request completed. Status code: {response.status_code}")
logger.debug(f"API response content: {response.text}")
return (img_tensor,)
except Exception as e:
error_msg = f"放大过程中出错: {str(e)}"
error_msg = f"Error during image upscale: {str(e)}"
logger.error(error_msg)
logger.exception("详细错误信息:")
# 返回一个错误标记图像,确保形状为 [B,H,W,C]
logger.exception("Detailed error information:")
# Return an error marked image, ensuring shape is [B,H,W,C]
error_image = torch.zeros((1, 100, 400, 3), dtype=torch.float32)
return (error_image,)
+4 -2
View File
@@ -1,7 +1,7 @@
import sys
# API_HOST = "https://app.comflowy.com"
API_HOST = "http://127.0.0.1:3000"
API_HOST = "https://app.comflowy.com"
# API_HOST = "http://127.0.0.1:3000"
FLOAT = (
"FLOAT",
{"default": 1, "min": -sys.float_info.max, "max": sys.float_info.max, "step": 0.01},
@@ -38,6 +38,8 @@ LLM_MODELS = [
"internlm/internlm2_5-7b-chat"
]
SAFETY_TOLERANCE = ["1", "2", "3", "4", "5"]
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""