feat: 新增AI生图(批量版)节点,模型改名,超时优化与友好报错

- 新增 BatchAsyncImageGenerator 节点(全并发+即时落盘,不怕中途失败丢图)
- 原版 AsyncImageGenerator 移除批量提示词功能,单节点只处理单提示词
- 模型改名:限时特价→次卡,gemini→nano-banana-官方
- 异步节点过滤 官方计费 渠道,仅保留次卡和官方模型
- 单任务超时提升至900s,批量超时改为动态计算(批次数×900s)
- No available channel for model 错误转化为中文友好提示
- 新增 base_async_provider / gemini_async_provider 异步客户端基类

Co-Authored-By: Claude Opus 4.6 <[email protected]>
This commit is contained in:
o1key
2026-04-29 18:24:55 +08:00
co-authored by Claude Opus 4.6
parent b6de4e49ab
commit 2e93a34434
12 changed files with 1405 additions and 741 deletions
+8 -5
View File
@@ -12,8 +12,8 @@ Comfyui_o1key - ComfyUI 自定义节点集合
import ssl import ssl
from .nodes import NanoBananaPro, NanoBananaProAsync, BatchNanoBananaPro, GoogleGemini, LoadFile, ImageStitchPro, SaveCleanImage, BatchCleanMetadata, VideoPreview, GoogleVeo, FluxImageEdit, UniversalLLMChat, KlingVideo, KlingFirstLastFrame, KlingMotionControlTest, AspectRatioPreset, MultiResPreview, BatchImagesO1key, Seedance, SeedanceMultiModal, StreamPreview, DoubaoImage, O1keyGPTImage, KVideo from .nodes import NanoBananaPro, BatchNanoBananaPro, GoogleGemini, LoadFile, ImageStitchPro, SaveCleanImage, BatchCleanMetadata, VideoPreview, GoogleVeo, FluxImageEdit, UniversalLLMChat, KlingVideo, KlingFirstLastFrame, KlingMotionControlTest, AspectRatioPreset, MultiResPreview, BatchImagesO1key, Seedance, SeedanceMultiModal, StreamPreview, DoubaoImage, O1keyGPTImage, KVideo
from .nodes import K3Video, K3VideoFirstLast, K3MotionControl, K3MotionVideoCheck from .nodes import K3Video, K3VideoFirstLast, K3MotionControl, K3MotionVideoCheck, AsyncImageGenerator, BatchAsyncImageGenerator
# 报错弹框友好文案(不修改原节点代码,仅在外层统一处理) # 报错弹框友好文案(不修改原节点代码,仅在外层统一处理)
_MSG_TIMEOUT = "API 请求超时,请稍后重试或检查网络。" _MSG_TIMEOUT = "API 请求超时,请稍后重试或检查网络。"
@@ -46,13 +46,13 @@ def _wrap_generate_for_error_display(cls, attr="generate"):
setattr(cls, attr, wrapped) setattr(cls, attr, wrapped)
_wrap_generate_for_error_display(NanoBananaPro) _wrap_generate_for_error_display(NanoBananaPro)
_wrap_generate_for_error_display(NanoBananaProAsync)
_wrap_generate_for_error_display(BatchNanoBananaPro) _wrap_generate_for_error_display(BatchNanoBananaPro)
_wrap_generate_for_error_display(AsyncImageGenerator)
_wrap_generate_for_error_display(BatchAsyncImageGenerator)
# ComfyUI 节点注册 # ComfyUI 节点注册
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"NanoBananaPro": NanoBananaPro, "NanoBananaPro": NanoBananaPro,
"NanoBananaProAsync": NanoBananaProAsync,
"BatchNanoBananaPro": BatchNanoBananaPro, "BatchNanoBananaPro": BatchNanoBananaPro,
"GoogleGemini": GoogleGemini, "GoogleGemini": GoogleGemini,
"LoadFile": LoadFile, "LoadFile": LoadFile,
@@ -79,11 +79,12 @@ NODE_CLASS_MAPPINGS = {
"K3VideoFirstLast": K3VideoFirstLast, "K3VideoFirstLast": K3VideoFirstLast,
"K3MotionControl": K3MotionControl, "K3MotionControl": K3MotionControl,
"K3MotionVideoCheck": K3MotionVideoCheck, "K3MotionVideoCheck": K3MotionVideoCheck,
"AsyncImageGenerator": AsyncImageGenerator,
"BatchAsyncImageGenerator": BatchAsyncImageGenerator,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"NanoBananaPro": "Nano Banana", "NanoBananaPro": "Nano Banana",
"NanoBananaProAsync": "Nano Banana(异步)",
"BatchNanoBananaPro": "批量 Nano Banana", "BatchNanoBananaPro": "批量 Nano Banana",
"GoogleGemini": "Google Gemini", "GoogleGemini": "Google Gemini",
"LoadFile": "加载文件", "LoadFile": "加载文件",
@@ -110,6 +111,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"K3VideoFirstLast": "首尾帧 K3 自研", "K3VideoFirstLast": "首尾帧 K3 自研",
"K3MotionControl": "动作控制 K3 自研", "K3MotionControl": "动作控制 K3 自研",
"K3MotionVideoCheck": "视频时长检测 K3", "K3MotionVideoCheck": "视频时长检测 K3",
"AsyncImageGenerator": "AI生图",
"BatchAsyncImageGenerator": "AI生图(批量版)",
} }
WEB_DIRECTORY = "./web" WEB_DIRECTORY = "./web"
+183
View File
@@ -0,0 +1,183 @@
"""
异步生图 Provider 抽象基类
定义异步提交+轮询模式的统一接口,支持多种生图模型后端
每个 Provider 封装一种 API 后端的通信协议:
- 如何提交任务(端点、请求体格式)
- 如何轮询状态(端点、状态字段语义)
- 如何解析结果(响应格式、图片提取方式)
新增第三方生图模型时,只需实现此接口即可接入异步节点。
"""
from abc import ABC, abstractmethod
from typing import Dict, List, Optional
from PIL import Image
class BaseAsyncImageProvider(ABC):
"""异步生图 Provider 抽象基类"""
def __init__(self, api_key: str, proxy_url: Optional[str] = None):
self.api_key = api_key
self.proxy_url = proxy_url
# ========================================================================
# 必须实现的抽象方法
# ========================================================================
@property
@abstractmethod
def api_base_url(self) -> str:
"""异步 API 的基础 URL,如 https://cf-api.o1key.com"""
...
@abstractmethod
def get_submit_endpoint(self, model: str, resolution: str) -> str:
"""获取提交任务的 API 端点路径(不含 base_url"""
...
@abstractmethod
def build_submit_body(
self,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
images: Optional[List[Image.Image]] = None,
**kwargs
) -> dict:
"""构建提交任务的请求体"""
...
@abstractmethod
def extract_task_id(self, response: dict) -> str:
"""从提交响应中提取 task_id"""
...
@abstractmethod
def extract_status(self, response: dict) -> str:
"""从轮询响应中提取任务状态(如 SUBMITTED / IN_PROGRESS / SUCCESS / FAILURE"""
...
@abstractmethod
async def parse_result(
self,
result_data: dict,
session
) -> List[Image.Image]:
"""从任务完成后的 result data 中解析生成的图像列表"""
...
@abstractmethod
def get_models(self) -> List[str]:
"""获取此 Provider 支持的模型 ID 列表"""
...
@abstractmethod
def get_model_aspect_ratios(self, model_id: str) -> List[str]:
"""获取指定模型支持的宽高比"""
...
@abstractmethod
def get_model_resolutions(self, model_id: str) -> List[str]:
"""获取指定模型支持的分辨率"""
...
# ========================================================================
# 可选的覆盖方法
# ========================================================================
def get_poll_endpoint(self, task_id: str) -> str:
"""获取轮询任务状态的 API 端点路径(默认实现适用于 o1key 异步 API)"""
return f"/async/v1/tasks/{task_id}"
def get_headers(self) -> dict:
"""获取 HTTP 请求头"""
return {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json"
}
def get_all_aspect_ratios(self) -> List[str]:
"""获取所有模型支持的宽高比(去重合并)"""
seen = set()
result = []
for model_id in self.get_models():
for ratio in self.get_model_aspect_ratios(model_id):
if ratio not in seen:
seen.add(ratio)
result.append(ratio)
return result
def get_all_resolutions(self) -> List[str]:
"""获取所有模型支持的分辨率(去重,按固定顺序排列)"""
_ORDER = ["512px", "1K", "2K", "4K"]
seen = set()
for model_id in self.get_models():
for res in self.get_model_resolutions(model_id):
seen.add(res)
return [r for r in _ORDER if r in seen]
def get_extra_inputs(self) -> dict:
"""
返回此 Provider 特有的额外 ComfyUI 输入参数。
子类重写以声明 Provider 专有的选项(如 Google Search Grounding)。
Returns:
dict,格式与 ComfyUI INPUT_TYPES 的 optional 字段一致
"""
return {}
def get_extra_kwargs(self, **kwargs) -> dict:
"""
从 ComfyUI kwargs 中提取此 Provider 特有的参数,
转换为 build_submit_body 可接收的 kwargs。
子类重写以处理 Provider 专有参数。
"""
return {}
def extract_progress(self, response: dict) -> Optional[float]:
"""
从轮询响应中提取生成进度。
Args:
response: 轮询接口返回的完整响应字典
Returns:
0.0-1.0 之间的进度值,或 None 表示该响应不含进度信息
"""
return None
def query_balance_sync(self) -> Optional[dict]:
"""
同步查询账户余额(可选)。
返回 None 表示不支持。
"""
return None
def format_balance_info(self, balance_data: dict) -> str:
"""格式化余额信息为展示文本"""
return ""
# ========================================================================
# 工具方法
# ========================================================================
@staticmethod
def build_proxy_url(port: str) -> Optional[str]:
"""
将端口号字符串转为 aiohttp 可用的 HTTP 代理 URL。
兼容 v2rayN (10808)、Clash Verge (7897) 等。
Args:
port: 用户填写的端口号,如 "7897",空字符串返回 None
Returns:
代理 URL 或 None
"""
port = (port or "").strip()
if not port or not port.isdigit():
return None
return f"http://127.0.0.1:{port}"
+149
View File
@@ -0,0 +1,149 @@
"""
Gemini 异步生图 Provider
通过 cf-api.o1key.com 的异步提交+轮询接口调用 Gemini 图像生成模型
协议说明:
- 提交:POST {base}/async{gemini_endpoint}?image_format=url
- 轮询:GET {base}/async/v1/tasks/{task_id}
- 结果:可能直接返回 image_url,也可能返回 Gemini 标准 candidates 格式
"""
from io import BytesIO
from typing import Dict, List, Optional
from PIL import Image
from .base_async_provider import BaseAsyncImageProvider
from .gemini_client import GeminiAPIClient
from ..utils.config import get_async_api_base_url, get_api_key_or_raise
from ..models_config import (
get_enabled_models,
get_model_supported_aspect_ratios,
get_model_supported_resolutions,
)
class GeminiAsyncImageProvider(BaseAsyncImageProvider):
"""
Gemini 异步生图 Provider
委托 GeminiAPIClient 处理:
- 端点构造(get_endpoint
- 请求体构建(build_request_body
- 响应解析(parse_response_async
"""
def __init__(self, api_key: str = None, proxy_url: str = None):
if api_key is None:
api_key = get_api_key_or_raise("O1KEY_API_KEY")
super().__init__(api_key=api_key, proxy_url=proxy_url)
self._client = GeminiAPIClient(api_key=api_key)
# ========================================================================
# 抽象方法实现
# ========================================================================
@property
def api_base_url(self) -> str:
return get_async_api_base_url()
def get_submit_endpoint(self, model: str, resolution: str) -> str:
gemini_endpoint = self._client.get_endpoint(
model=model, resolution=resolution, image_format="url"
)
base = gemini_endpoint.split("?")[0]
async_endpoint = f"/async{base}"
if "?" in gemini_endpoint:
async_endpoint += "?" + gemini_endpoint.split("?", 1)[1]
return async_endpoint
def build_submit_body(
self,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
images: Optional[List[Image.Image]] = None,
**kwargs
) -> dict:
return self._client.build_request_body(
prompt=prompt,
images=images,
aspect_ratio=aspect_ratio,
resolution=resolution,
enable_grounding=kwargs.get("enable_grounding", False),
enable_image_search=kwargs.get("enable_image_search", False),
)
def extract_task_id(self, response: dict) -> str:
task_id = response.get("task_id")
if not task_id:
raise RuntimeError(f"提交响应中未找到 task_id: {response}")
return task_id
def extract_status(self, response: dict) -> str:
return response.get("status", "UNKNOWN")
async def parse_result(self, result_data: dict, session) -> List[Image.Image]:
# 异步接口可能直接返回 image_url
image_url = result_data.get("image_url", "") if isinstance(result_data, dict) else ""
if image_url:
async with session.get(image_url) as img_resp:
if img_resp.status == 200:
img_bytes = await img_resp.read()
return [Image.open(BytesIO(img_bytes))]
raise RuntimeError(f"下载图片失败 ({img_resp.status}): {image_url}")
# 否则按 Gemini 标准格式解析
images_list, _ = await self._client.parse_response_async(result_data, session=session)
return images_list
def get_models(self) -> List[str]:
return get_enabled_models()
def get_model_aspect_ratios(self, model_id: str) -> List[str]:
return get_model_supported_aspect_ratios(model_id)
def get_model_resolutions(self, model_id: str) -> List[str]:
return get_model_supported_resolutions(model_id)
# ========================================================================
# 可选方法覆盖
# ========================================================================
def get_extra_inputs(self) -> dict:
"""Gemini 专有:Google Search Grounding"""
return {
"联网功能": (["关闭", "打开"], {"default": "关闭"}),
}
def get_extra_kwargs(self, **kwargs) -> dict:
return {
"enable_grounding": kwargs.pop("联网功能", "关闭") == "打开",
}
def extract_progress(self, response: dict) -> Optional[float]:
"""从轮询响应中提取进度(0.0-1.0"""
# 直接字段:progress / percentage
for field in ("progress", "percentage"):
val = response.get(field)
if val is not None and isinstance(val, (int, float)):
return val / 100.0 if val > 1 else float(val)
# 嵌套字段:progressInfo / progress_info
progress_info = response.get("progressInfo") or response.get("progress_info")
if isinstance(progress_info, dict):
for field in ("progress", "percentage"):
val = progress_info.get(field)
if val is not None and isinstance(val, (int, float)):
return val / 100.0 if val > 1 else float(val)
return None
def query_balance_sync(self) -> Optional[dict]:
try:
return self._client.query_balance_sync()
except Exception:
return None
def format_balance_info(self, balance_data: dict) -> str:
return self._client.format_balance_info(balance_data)
+2 -2
View File
@@ -70,7 +70,7 @@ class GeminiAPIClient(BaseAPIClient):
from ..models_config import get_model_endpoint from ..models_config import get_model_endpoint
# 特殊处理:动态端点模型(根据分辨率选择) # 特殊处理:动态端点模型(根据分辨率选择)
if model == "nano-banana-pro-限时特价": if model == "nano-banana-pro-次卡":
if resolution == "1K": if resolution == "1K":
endpoint = "/v1beta/models/nano-banana-pro:generateContent" endpoint = "/v1beta/models/nano-banana-pro:generateContent"
elif resolution == "2K": elif resolution == "2K":
@@ -80,7 +80,7 @@ class GeminiAPIClient(BaseAPIClient):
else: else:
endpoint = "/v1beta/models/nano-banana-pro-2k:generateContent" endpoint = "/v1beta/models/nano-banana-pro-2k:generateContent"
elif model == "nano-banana-2-限时特价": elif model == "nano-banana-2-次卡":
if resolution == "512px": if resolution == "512px":
endpoint = "/v1beta/models/nano-banana-2-0.5k:generateContent" endpoint = "/v1beta/models/nano-banana-2-0.5k:generateContent"
elif resolution == "1K": elif resolution == "1K":
+7 -19
View File
@@ -33,8 +33,7 @@ _ENDPOINT_EDITS = "/v1/images/edits/"
# ── 模型名映射(UI 显示名 → API 实际参数名)───────────────────────────────── # ── 模型名映射(UI 显示名 → API 实际参数名)─────────────────────────────────
_MODEL_NAME_MAP = { _MODEL_NAME_MAP = {
"gpt-image-1.5-特价": "gpt-image-1.5-special", "gpt-image-2-次卡": "gpt-image-2-special",
"gpt-image-2-特价": "gpt-image-2-special",
} }
# ── 超时 ────────────────────────────────────────────────────────────────────── # ── 超时 ──────────────────────────────────────────────────────────────────────
@@ -46,7 +45,7 @@ class GptImageClient:
GPT Image API 客户端 GPT Image API 客户端
接口说明: 接口说明:
generationsJSON body,支持 background / quality / size / n / model generationsJSON body,支持 quality / size / n / model
editsmultipart/form-data,必须包含 imagePNG),可选 maskPNG editsmultipart/form-data,必须包含 imagePNG),可选 maskPNG
两个接口的响应格式相同: 两个接口的响应格式相同:
@@ -227,7 +226,6 @@ class GptImageClient:
prompt: str, prompt: str,
model: str, model: str,
quality: str, quality: str,
background: str,
size: str, size: str,
n: int, n: int,
seed: int, seed: int,
@@ -244,7 +242,6 @@ class GptImageClient:
"model": api_model, "model": api_model,
"prompt": prompt, "prompt": prompt,
"quality": quality, "quality": quality,
"background": background,
"n": n, "n": n,
"moderation": "low", "moderation": "low",
} }
@@ -278,7 +275,7 @@ class GptImageClient:
url = f"{self.base_url}{_ENDPOINT_GENERATIONS}" url = f"{self.base_url}{_ENDPOINT_GENERATIONS}"
print(f"[o1key GPT Image] {mode} | 模型={model} | quality={quality} | " print(f"[o1key GPT Image] {mode} | 模型={model} | quality={quality} | "
f"background={background} | size={size} | n={n}") f"size={size} | n={n}")
connector = aiohttp.TCPConnector(ssl=False, force_close=True) connector = aiohttp.TCPConnector(ssl=False, force_close=True)
timeout = aiohttp.ClientTimeout(total=_REQUEST_TIMEOUT) timeout = aiohttp.ClientTimeout(total=_REQUEST_TIMEOUT)
@@ -311,15 +308,12 @@ class GptImageClient:
return await self._parse_response(resp_json, session) return await self._parse_response(resp_json, session)
# ── 图像编辑(edits 接口,multipart/form-data)────────────────────────── # ── 图像编辑(edits 接口,multipart/form-data)──────────────────────────
# 注意:o1key 中转服务的 edits 接口暂不支持 quality / background / moderation 参数,
# 这些字段暂时不传递,待服务方更新后可恢复。
async def _edit_async( async def _edit_async(
self, self,
prompt: str, prompt: str,
model: str, model: str,
quality: str, quality: str,
background: str,
size: str, size: str,
n: int, n: int,
seed: int, seed: int,
@@ -328,8 +322,6 @@ class GptImageClient:
) -> List[Image.Image]: ) -> List[Image.Image]:
""" """
调用 /v1/images/edits/ 接口(multipart/form-data)。 调用 /v1/images/edits/ 接口(multipart/form-data)。
当前仅传递 model / prompt / n / size / image / mask
quality / background / moderation 暂不支持(o1key 服务端限制)。
""" """
# 模型名映射:UI 显示名 → API 参数名 # 模型名映射:UI 显示名 → API 参数名
api_model = _MODEL_NAME_MAP.get(model, model) api_model = _MODEL_NAME_MAP.get(model, model)
@@ -342,8 +334,6 @@ class GptImageClient:
normalized_tensors.append(t) normalized_tensors.append(t)
num_images = len(normalized_tensors) num_images = len(normalized_tensors)
# o1key 中转服务的 edits 接口暂不支持 background / moderation / seed
# 待服务方更新后可重新加入。
form = aiohttp.FormData() form = aiohttp.FormData()
form.add_field("model", api_model) form.add_field("model", api_model)
form.add_field("prompt", prompt) form.add_field("prompt", prompt)
@@ -387,7 +377,8 @@ class GptImageClient:
mode = "图像编辑(无蒙版)" mode = "图像编辑(无蒙版)"
url = f"{self.base_url}{_ENDPOINT_EDITS}" url = f"{self.base_url}{_ENDPOINT_EDITS}"
print(f"[o1key GPT Image] {mode} | 模型={model} | 参考图={num_images}张 | size={size} | n={n}") print(f"[o1key GPT Image] {mode} | 模型={model} | 参考图={num_images}张 | "
f"quality={quality} | size={size} | n={n}")
connector = aiohttp.TCPConnector(ssl=False, force_close=True) connector = aiohttp.TCPConnector(ssl=False, force_close=True)
timeout = aiohttp.ClientTimeout(total=_REQUEST_TIMEOUT) timeout = aiohttp.ClientTimeout(total=_REQUEST_TIMEOUT)
@@ -430,7 +421,6 @@ class GptImageClient:
prompt: str, prompt: str,
model: str, model: str,
quality: str, quality: str,
background: str,
size: str, size: str,
n: int, n: int,
seed: int, seed: int,
@@ -443,21 +433,19 @@ class GptImageClient:
路由逻辑: 路由逻辑:
- 无 image_tensor → generations 接口(文生图,JSON body - 无 image_tensor → generations 接口(文生图,JSON body
- 有 image_tensor → edits 接口(图生图/编辑,multipart/form-data - 有 image_tensor → edits 接口(图生图/编辑,multipart/form-data
所有模型统一走 multipartquality/background 通过表单字段传递,
new-api 开启"透传请求体"后原样转发给上游。
""" """
use_edits = (image_tensor is not None) use_edits = (image_tensor is not None)
if use_edits: if use_edits:
coro = self._edit_async( coro = self._edit_async(
prompt=prompt, model=model, quality=quality, prompt=prompt, model=model, quality=quality,
background=background, size=size, n=n, seed=seed, size=size, n=n, seed=seed,
image_list=image_tensor, mask_tensor=mask_tensor, image_list=image_tensor, mask_tensor=mask_tensor,
) )
else: else:
coro = self._generate_async( coro = self._generate_async(
prompt=prompt, model=model, quality=quality, prompt=prompt, model=model, quality=quality,
background=background, size=size, n=n, seed=seed, size=size, n=n, seed=seed,
image_list=image_tensor, image_list=image_tensor,
) )
+2 -2
View File
@@ -56,14 +56,14 @@ class OpenAIAPIClient(BaseAPIClient):
只是把拼在 URL 路径里的模型段提取出来单独返回。 只是把拼在 URL 路径里的模型段提取出来单独返回。
Args: Args:
model: 节点下拉框中的模型 ID,如 "nano-banana-pro-限时特价" model: 节点下拉框中的模型 ID,如 "nano-banana-pro-次卡"
resolution: 分辨率字符串,如 "1K" / "2K" / "4K" / "512" resolution: 分辨率字符串,如 "1K" / "2K" / "4K" / "512"
Returns: Returns:
实际模型名,如 "nano-banana-pro-2k" 实际模型名,如 "nano-banana-pro-2k"
""" """
# ── 动态端点模型 ────────────────────────────────────────────────── # ── 动态端点模型 ──────────────────────────────────────────────────
if model == "nano-banana-pro-限时特价": if model == "nano-banana-pro-次卡":
if resolution == "1K": if resolution == "1K":
return "nano-banana-pro" return "nano-banana-pro"
elif resolution == "4K": elif resolution == "4K":
+45 -10
View File
@@ -44,9 +44,10 @@ from typing import List, Dict, Optional, Tuple
GEMINI_MODELS = [ GEMINI_MODELS = [
{ {
"id": "nano-banana-pro-限时特价", "id": "nano-banana-pro-次卡",
"description": "Nano Banana Pro 限时特价,根据分辨率自动选择端点 (1K/2K/4K),高性能图像生成模型", "description": "Nano Banana Pro 次卡,根据分辨率自动选择端点 (1K/2K/4K),高性能图像生成模型",
"enabled": True, "enabled": True,
"provider": "gemini_async",
"endpoint_type": "dynamic", "endpoint_type": "dynamic",
"endpoint": None, # 动态端点,由代码根据分辨率选择 "endpoint": None, # 动态端点,由代码根据分辨率选择
"supported_aspect_ratios": [ "supported_aspect_ratios": [
@@ -58,6 +59,7 @@ GEMINI_MODELS = [
"id": "nano-banana-pro-官方计费", "id": "nano-banana-pro-官方计费",
"description": "Nano Banana Pro 官方计费,按分辨率路由 (1K/2K/4K),使用官方计费通道", "description": "Nano Banana Pro 官方计费,按分辨率路由 (1K/2K/4K),使用官方计费通道",
"enabled": True, "enabled": True,
"provider": "gemini_async",
"endpoint_type": "dynamic", "endpoint_type": "dynamic",
"endpoint": None, # 动态端点,由代码根据分辨率选择 "endpoint": None, # 动态端点,由代码根据分辨率选择
"supported_aspect_ratios": [ "supported_aspect_ratios": [
@@ -66,9 +68,10 @@ GEMINI_MODELS = [
"supported_resolutions": ["1K", "2K", "4K"] "supported_resolutions": ["1K", "2K", "4K"]
}, },
{ {
"id": "nano-banana-2-限时特价", "id": "nano-banana-2-次卡",
"description": "Nano Banana 2 限时特价,根据分辨率自动选择端点 (512px/1K/2K/4K),图像生成模型", "description": "Nano Banana 2 次卡,根据分辨率自动选择端点 (512px/1K/2K/4K),图像生成模型",
"enabled": True, "enabled": True,
"provider": "gemini_async",
"endpoint_type": "dynamic", "endpoint_type": "dynamic",
"endpoint": None, # 动态端点,由代码根据分辨率选择 "endpoint": None, # 动态端点,由代码根据分辨率选择
"supported_aspect_ratios": [ "supported_aspect_ratios": [
@@ -81,6 +84,7 @@ GEMINI_MODELS = [
"id": "nano-banana-2-官方计费", "id": "nano-banana-2-官方计费",
"description": "Nano Banana 2 官方计费,按分辨率路由 (512/1K/2K/4K),使用官方计费通道", "description": "Nano Banana 2 官方计费,按分辨率路由 (512/1K/2K/4K),使用官方计费通道",
"enabled": True, "enabled": True,
"provider": "gemini_async",
"endpoint_type": "dynamic", "endpoint_type": "dynamic",
"endpoint": None, # 动态端点,由代码根据分辨率选择 "endpoint": None, # 动态端点,由代码根据分辨率选择
"supported_aspect_ratios": [ "supported_aspect_ratios": [
@@ -90,9 +94,10 @@ GEMINI_MODELS = [
"supported_resolutions": ["512px", "1K", "2K", "4K"] "supported_resolutions": ["512px", "1K", "2K", "4K"]
}, },
{ {
"id": "gemini-3-pro-image-preview", "id": "nano-banana-pro-官方",
"description": "标准模式,固定端点,适用于常规图像生成", "description": "Nano Banana Pro 官方,固定端点适用于常规图像生成",
"enabled": False, "enabled": True,
"provider": "gemini_async",
"endpoint_type": "standard", "endpoint_type": "standard",
"endpoint": "/v1beta/models/gemini-3-pro-image-preview:generateContent", "endpoint": "/v1beta/models/gemini-3-pro-image-preview:generateContent",
"supported_aspect_ratios": [ "supported_aspect_ratios": [
@@ -101,9 +106,10 @@ GEMINI_MODELS = [
"supported_resolutions": ["1K", "2K", "4K"] "supported_resolutions": ["1K", "2K", "4K"]
}, },
{ {
"id": "gemini-3.1-flash-image-preview", "id": "nano-banana-2-官方",
"description": "Gemini 3.1 Flash 图像生成,固定端点,快速图像生成模型", "description": "Nano Banana 2 官方,固定端点快速图像生成模型",
"enabled": False, "enabled": True,
"provider": "gemini_async",
"endpoint_type": "standard", "endpoint_type": "standard",
"endpoint": "/v1beta/models/gemini-3.1-flash-image-preview:generateContent", "endpoint": "/v1beta/models/gemini-3.1-flash-image-preview:generateContent",
"supported_aspect_ratios": [ "supported_aspect_ratios": [
@@ -341,6 +347,35 @@ def get_all_supported_resolutions() -> List[str]:
return [res for res in _ORDER if res in seen] return [res for res in _ORDER if res in seen]
def get_model_provider(model_id: str) -> Optional[str]:
"""
获取模型的异步 Provider 名称
Args:
model_id: 模型 ID
Returns:
Provider 名称(如 "gemini_async"),如果模型未配置 provider 则返回 None
"""
config = get_model_config(model_id)
if config is None:
return None
return config.get("provider")
def get_enabled_async_models() -> List[str]:
"""
获取所有启用的、支持异步模式的模型 ID 列表
Returns:
模型 ID 列表(仅包含配置了 provider 且 enabled 的模型)
"""
return [
model["id"] for model in GEMINI_MODELS
if model.get("enabled", False) and model.get("provider")
]
def get_endpoint_type(model_id: str) -> Optional[str]: def get_endpoint_type(model_id: str) -> Optional[str]:
""" """
获取模型的端点类型 获取模型的端点类型
+2 -2
View File
@@ -5,7 +5,6 @@
from .stream_preview import StreamPreview from .stream_preview import StreamPreview
from .nano_banana_pro import NanoBananaPro from .nano_banana_pro import NanoBananaPro
from .nano_banana_pro_async import NanoBananaProAsync
from .batch_nano_banana_pro import BatchNanoBananaPro from .batch_nano_banana_pro import BatchNanoBananaPro
from .google_gemini import GoogleGemini from .google_gemini import GoogleGemini
from .load_file import LoadFile from .load_file import LoadFile
@@ -19,6 +18,7 @@ from .universal_llm import UniversalLLMChat
from .multi_res_preview import MultiResPreview from .multi_res_preview import MultiResPreview
from .batch_images_o1key import BatchImagesO1key from .batch_images_o1key import BatchImagesO1key
from .seedance_video import Seedance, SeedanceMultiModal from .seedance_video import Seedance, SeedanceMultiModal
from .async_image_generation import AsyncImageGenerator, BatchAsyncImageGenerator
from .doubao_image import DoubaoImage from .doubao_image import DoubaoImage
from .gpt_image import O1keyGPTImage from .gpt_image import O1keyGPTImage
from .K_video import KVideo from .K_video import KVideo
@@ -26,4 +26,4 @@ from .K3_video import K3Video
from .K3_video_firstlast import K3VideoFirstLast from .K3_video_firstlast import K3VideoFirstLast
from .K3_motion_control import K3MotionControl, K3MotionVideoCheck from .K3_motion_control import K3MotionControl, K3MotionVideoCheck
__all__ = ['NanoBananaPro', 'NanoBananaProAsync', 'BatchNanoBananaPro', 'GoogleGemini', 'LoadFile', 'ImageStitchPro', 'SaveCleanImage', 'BatchCleanMetadata', 'VideoPreview', 'KlingVideo', 'KlingFirstLastFrame', 'KlingMotionControlTest', 'AspectRatioPreset', 'GoogleVeo', 'FluxImageEdit', 'UniversalLLMChat', 'MultiResPreview', 'BatchImagesO1key', 'Seedance', 'SeedanceMultiModal', 'StreamPreview', 'DoubaoImage', 'O1keyGPTImage', 'KVideo', 'K3Video', 'K3VideoFirstLast', 'K3MotionControl', 'K3MotionVideoCheck'] __all__ = ['NanoBananaPro', 'BatchNanoBananaPro', 'GoogleGemini', 'LoadFile', 'ImageStitchPro', 'SaveCleanImage', 'BatchCleanMetadata', 'VideoPreview', 'KlingVideo', 'KlingFirstLastFrame', 'KlingMotionControlTest', 'AspectRatioPreset', 'GoogleVeo', 'FluxImageEdit', 'UniversalLLMChat', 'MultiResPreview', 'BatchImagesO1key', 'Seedance', 'SeedanceMultiModal', 'StreamPreview', 'DoubaoImage', 'O1keyGPTImage', 'KVideo', 'K3Video', 'K3VideoFirstLast', 'K3MotionControl', 'K3MotionVideoCheck', 'AsyncImageGenerator', 'BatchAsyncImageGenerator']
+992
View File
@@ -0,0 +1,992 @@
"""
异步生图节点(通用)
ComfyUI 自定义节点,通过异步提交+轮询模式调用多种生图模型
架构:
- 节点层(本文件):批量调度、进度条、ComfyUI 集成,不关心具体 API 协议
- Provider 层:封装每种 API 后端的通信协议(端点、请求体格式、响应解析)
新增第三方生图模型时,只需实现 BaseAsyncImageProvider 并注册即可。
"""
import os
import time
import math
import random
import asyncio
import aiohttp
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, List, Optional, Tuple
try:
from comfy.model_management import processing_interrupted, InterruptProcessingException
INTERRUPT_AVAILABLE = True
except ImportError:
INTERRUPT_AVAILABLE = False
InterruptProcessingException = RuntimeError # fallback
import torch
import numpy as np
from PIL import Image
from ..utils.image_utils import tensor_to_pil, pil_to_tensor, parse_batch_prompts
from ..utils.config import get_api_key_or_raise
from ..models_config import (
get_enabled_async_models,
get_model_provider,
get_model_supported_aspect_ratios,
get_all_supported_aspect_ratios,
get_model_supported_resolutions,
get_all_supported_resolutions,
)
from ..clients.base_async_provider import BaseAsyncImageProvider
try:
import folder_paths # noqa: F401
FOLDER_PATHS_AVAILABLE = True
except ImportError:
FOLDER_PATHS_AVAILABLE = False
try:
from comfy.utils import ProgressBar
PROGRESS_BAR_AVAILABLE = True
except ImportError:
PROGRESS_BAR_AVAILABLE = False
try:
import psutil
MEMORY_MONITOR_AVAILABLE = True
except ImportError:
MEMORY_MONITOR_AVAILABLE = False
DEBUG_LOG_ENABLED = False
REQUEST_LOG_ENABLED = True
_POLL_INTERVAL = 2 # 轮询间隔(秒)
_MAX_WAIT_TIME = 900 # 单任务最大等待时间(秒)
_MAX_CONCURRENT = 50 # 最大并发提交数
def _images_to_tensor_safe(images: List[Image.Image], node_label: str) -> torch.Tensor:
"""将图像列表转为 tensor,过滤不同尺寸的图"""
if not images:
placeholder = Image.new('RGB', (512, 512), color=(128, 128, 128))
return pil_to_tensor([placeholder])
base_size = max(images, key=lambda img: img.size[0] * img.size[1]).size
matched = [img for img in images if img.size == base_size]
skipped = [img for img in images if img.size != base_size]
if skipped:
sizes_str = ", ".join(f"{img.size[0]}x{img.size[1]}" for img in skipped)
print(
f"{node_label}: 丢弃 {len(skipped)} 张较小尺寸的图 ({sizes_str})"
f"仅输出最大尺寸 {base_size[0]}x{base_size[1]}{len(matched)}"
)
return pil_to_tensor(matched)
class AsyncImageGenerator:
"""
异步生图节点(通用)
功能:
- 异步提交 + 轮询模式,避免 ComfyUI 主线程阻塞
- 支持多种生图模型后端(通过 Provider 扩展)
- 支持批量提示词、多参考图、代理端口
"""
NODE_LABEL = "AI生图"
# Provider 注册表:provider 名称 → 类路径
PROVIDER_CLASSES = {
"gemini_async": "..clients.gemini_async_provider.GeminiAsyncImageProvider",
}
# Provider 专有输入参数声明(用于 INPUT_TYPES 合并)
_PROVIDER_EXTRA_INPUTS: Dict[str, dict] = {
"gemini_async": {
"联网功能": (["关闭", "打开"], {"default": "关闭"}),
},
}
def __init__(self):
self._provider: Optional[BaseAsyncImageProvider] = None
self._provider_name: Optional[str] = None
@classmethod
def INPUT_TYPES(cls):
# 模型列表(排除官方计费渠道)
models = [m for m in get_enabled_async_models() if "官方计费" not in m]
if not models:
models = ["请在 models_config.py 中启用至少一个异步模型"]
# 宽高比 / 分辨率(取所有模型的并集,运行时验证)
all_aspect_ratios = get_all_supported_aspect_ratios()
if not all_aspect_ratios:
all_aspect_ratios = [
"1:1", "4:3", "3:4", "16:9", "9:16",
"2:3", "3:2", "4:5", "5:4", "21:9",
"1:4", "4:1", "1:8", "8:1"
]
all_resolutions = get_all_supported_resolutions()
if not all_resolutions:
all_resolutions = ["512px", "1K", "2K", "4K"]
# 可选输入(按展示顺序)
optional = {}
# Provider 专有参数(紧接 required 参数下方)
for provider_extra in cls._PROVIDER_EXTRA_INPUTS.values():
optional.update(provider_extra)
for i in range(1, 10):
optional[f"参考图{i}"] = ("IMAGE",)
optional["seed"] = ("INT", {
"default": 0,
"min": 0,
"max": 0xffffffffffffffff
})
optional["代理端口"] = ("STRING", {
"default": "",
"multiline": False,
"placeholder": "本地代理端口,如 7897Clash Verge)或 10808v2rayN),留空不使用"
})
return {
"required": {
"prompt": ("STRING", {
"default": "一个中国女子的OOTD",
"multiline": True
}),
"模型": (models, {"default": models[0]}),
"宽高比": (all_aspect_ratios, {"default": "1:1"}),
"分辨率": (all_resolutions, {"default": "2K"}),
"生图数量": ("INT", {
"default": 1,
"min": 1,
"max": 1000,
"step": 1
})
},
"optional": optional
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("输出图像",)
FUNCTION = "generate"
CATEGORY = "image/generation"
OUTPUT_NODE = True
# ========================================================================
# Provider 工厂
# ========================================================================
def _get_provider(self, model_id: str, proxy_url: Optional[str] = None) -> BaseAsyncImageProvider:
"""根据模型 ID 获取或创建对应的 Provider 实例"""
provider_name = get_model_provider(model_id)
if not provider_name:
raise ValueError(f"模型 \"{model_id}\" 不支持异步模式")
# 同类型 Provider 复用,只更新代理
if self._provider is not None and self._provider_name == provider_name:
self._provider.proxy_url = proxy_url
return self._provider
# 创建新 Provider
class_path = self.PROVIDER_CLASSES.get(provider_name)
if not class_path:
raise ValueError(f"未注册的 Provider: {provider_name}")
module_path, class_name = class_path.rsplit(".", 1)
if module_path.startswith(".."):
import importlib
module = importlib.import_module(module_path, package=__package__)
else:
import importlib
module = importlib.import_module(module_path)
provider_class = getattr(module, class_name)
api_key = get_api_key_or_raise("O1KEY_API_KEY")
self._provider = provider_class(api_key=api_key, proxy_url=proxy_url)
self._provider_name = provider_name
return self._provider
# ========================================================================
# 工具方法
# ========================================================================
@staticmethod
def _friendly_error(error_msg: str) -> str:
"""将上游错误转化为用户友好的提示"""
if "No available channel for model" in error_msg:
return (
"当前分组下模型不可用,请检查分组是否正确。"
"若是正常出图过程中遇到该报错,说明该错误只是暂时的,稍后重试即可。或切换其他模型。"
f"\n(原始错误: {error_msg}"
)
return error_msg
@staticmethod
def _check_interrupt():
"""检查 ComfyUI 是否点击了取消按钮,是则抛出 InterruptProcessingException"""
if INTERRUPT_AVAILABLE and processing_interrupted():
raise InterruptProcessingException()
# ========================================================================
# 核心异步逻辑
# ========================================================================
async def _submit_one(
self,
session: aiohttp.ClientSession,
provider: BaseAsyncImageProvider,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
input_images: List[Image.Image],
**extra_kwargs,
) -> str:
"""提交单个异步任务,返回 task_id"""
endpoint = provider.get_submit_endpoint(model, resolution)
request_body = provider.build_submit_body(
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
images=input_images if input_images else None,
**extra_kwargs,
)
url = f"{provider.api_base_url}{endpoint}"
headers = provider.get_headers()
if REQUEST_LOG_ENABLED:
import json
_log_body = {k: v for k, v in request_body.items()}
print(f"[异步提交] URL: {url}")
print(f"[异步提交] 请求体: {json.dumps(_log_body, ensure_ascii=False)[:500]}")
async with session.post(url, json=request_body, headers=headers, proxy=provider.proxy_url) as response:
if response.status != 200:
error_text = await response.text()
raise RuntimeError(f"提交任务失败 ({response.status}): {error_text}")
data = await response.json()
if DEBUG_LOG_ENABLED:
import json
print(f"[异步提交] 响应: {json.dumps(data, ensure_ascii=False)[:500]}")
return provider.extract_task_id(data)
async def _poll_one(
self,
session: aiohttp.ClientSession,
provider: BaseAsyncImageProvider,
task_id: str,
on_progress=None,
) -> dict:
"""轮询单个任务直到完成,返回 result dataon_progress(delta) 可选,用于驱动进度条"""
poll_endpoint = provider.get_poll_endpoint(task_id)
url = f"{provider.api_base_url}{poll_endpoint}"
headers = provider.get_headers()
start_time = time.time()
poll_count = 0
last_progress = 0.0
while True:
self._check_interrupt()
elapsed = time.time() - start_time
if elapsed > _MAX_WAIT_TIME:
raise RuntimeError(f"任务 {task_id} 超时({_MAX_WAIT_TIME}秒)")
poll_count += 1
async with session.get(url, headers=headers, proxy=provider.proxy_url) as response:
if response.status != 200:
error_text = await response.text()
raise RuntimeError(f"查询任务失败 ({response.status}): {error_text}")
result = await response.json()
status = provider.extract_status(result)
# 提取进度并回调(封顶 1.0 防止异常值导致进度条溢出)
if on_progress and status in ("SUBMITTED", "IN_PROGRESS"):
p = provider.extract_progress(result)
if p is not None:
p = min(p, 1.0)
if p > last_progress:
on_progress(p - last_progress)
last_progress = p
print(f"{self.NODE_LABEL}: 任务{task_id[:8]}... 进度 {p * 100:.0f}%")
if DEBUG_LOG_ENABLED:
import json
print(f"[轮询 #{poll_count}] {task_id}: status={status}")
if status == "SUCCESS":
# 补足剩余进度
if on_progress and last_progress < 1.0:
on_progress(1.0 - last_progress)
return result.get("data", {})
elif status == "FAILURE":
error_msg = result.get("error", "未知错误")
friendly_msg = self._friendly_error(error_msg)
raise RuntimeError(f"任务失败: {friendly_msg}")
elif status in ("SUBMITTED", "IN_PROGRESS"):
await asyncio.sleep(_POLL_INTERVAL)
else:
raise RuntimeError(f"未知任务状态: {status}")
async def _execute_one(
self,
session: aiohttp.ClientSession,
provider: BaseAsyncImageProvider,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
input_images: List[Image.Image],
global_task_index: int,
on_progress=None,
**extra_kwargs,
) -> dict:
"""执行单个异步生成任务(提交 + 轮询 + 解析);on_progress(delta) 可选"""
result = {
"global_task_index": global_task_index,
"prompt": prompt,
"success": False,
"generated_count": 0,
"output_images": [],
"error": None,
}
contributed = [0.0] # mutable container,追踪本任务已贡献的 pbar 进度
def _track_progress(delta):
contributed[0] += delta
if on_progress:
on_progress(delta)
try:
task_id = await self._submit_one(
session, provider, prompt, model,
resolution, aspect_ratio, input_images,
**extra_kwargs,
)
response_data = await self._poll_one(
session, provider, task_id, on_progress=_track_progress,
)
images_list = await provider.parse_result(response_data, session)
result["success"] = True
result["generated_count"] = len(images_list)
result["output_images"] = images_list
except InterruptProcessingException:
# 用户取消:补齐进度后向上传播,不吞掉
if contributed[0] < 1.0 and on_progress:
on_progress(1.0 - contributed[0])
raise
except Exception as e:
result["error"] = str(e)
# 失败也补齐 1.0 进度,保证进度条总数正确
if contributed[0] < 1.0 and on_progress:
on_progress(1.0 - contributed[0])
return result
async def _process_batch(
self,
provider: BaseAsyncImageProvider,
prompts: List[str],
model: str,
resolution: str,
aspect_ratio: str,
images_per_prompt: int,
input_images: List[Image.Image],
pbar=None,
**extra_kwargs,
) -> List[dict]:
"""异步批量处理"""
# 构建任务定义
tasks_def = []
for p_idx, prompt in enumerate(prompts):
for sub_idx in range(images_per_prompt):
tasks_def.append((p_idx, sub_idx, prompt))
total_tasks = len(tasks_def)
num_batches = math.ceil(total_tasks / _MAX_CONCURRENT)
all_results: List[dict] = []
completed = 0
# 所有任务共享同一个进度回调,驱动同一个进度条
_on_progress = (lambda delta: pbar.update(delta)) if pbar is not None else None
connector = aiohttp.TCPConnector(ssl=False, limit=0, limit_per_host=0)
async with aiohttp.ClientSession(connector=connector) as session:
for batch_idx in range(num_batches):
self._check_interrupt()
start_idx = batch_idx * _MAX_CONCURRENT
end_idx = min(start_idx + _MAX_CONCURRENT, total_tasks)
batch_tasks = []
for i in range(start_idx, end_idx):
_, _, prompt = tasks_def[i]
task = asyncio.create_task(
self._execute_one(
session=session,
provider=provider,
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
input_images=input_images,
global_task_index=i,
on_progress=_on_progress,
**extra_kwargs,
)
)
batch_tasks.append(task)
batch_results = []
for coro in asyncio.as_completed(batch_tasks):
result_data = None
try:
result_data = await coro
except InterruptProcessingException:
# 用户取消:终止所有未完成任务
for t in batch_tasks:
t.cancel()
raise
except Exception as e2:
# 意外错误(不应发生,_execute_one 内部已捕获常规异常)
result_data = {
"success": False,
"error": str(e2),
"generated_count": 0,
"output_images": [],
"prompt": "",
}
batch_results.append(result_data)
completed += 1
prompt_snippet = (result_data.get("prompt", "") or "")[:30]
if result_data and result_data.get("success"):
count = result_data.get("generated_count", 1)
print(f"{self.NODE_LABEL}: [{completed}/{total_tasks}] {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} -> OK({count}张)")
else:
error_msg = result_data.get("error", "未知错误") if result_data else "未知错误"
print(f"{self.NODE_LABEL}: [{completed}/{total_tasks}] {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} -> FAIL: {error_msg}")
all_results.extend(batch_results)
import gc
gc.collect()
await asyncio.sleep(0.1)
return all_results
# ========================================================================
# ComfyUI 入口
# ========================================================================
def generate(
self,
prompt: str,
模型: str,
宽高比: str,
分辨率: str,
生图数量: int,
**kwargs
) -> Tuple[torch.Tensor]:
"""生成图像(异步模式)"""
start_time = time.time()
# 提取通用可选参数
seed: int = kwargs.pop("seed", 0)
proxy_port: str = kwargs.pop("代理端口", "")
# 初始化 Provider
proxy_url = BaseAsyncImageProvider.build_proxy_url(proxy_port)
provider = self._get_provider(模型, proxy_url=proxy_url)
if proxy_url:
print(f"{self.NODE_LABEL}: 已启用代理加速 -> {proxy_url}")
# 提取 Provider 专有参数
extra_kwargs = provider.get_extra_kwargs(**kwargs)
# 进度条
pbar = None
if PROGRESS_BAR_AVAILABLE:
pbar = ProgressBar(生图数量)
try:
# 初始化随机种子
random.seed(seed)
np.random.seed(seed % (2**32))
# 内存监控
if MEMORY_MONITOR_AVAILABLE and 生图数量 > 50:
process = psutil.Process()
initial_memory = process.memory_info().rss / 1024 / 1024
print(f"{self.NODE_LABEL}: 初始内存使用: {initial_memory:.1f} MB")
# 运行时验证分辨率
supported_resolutions = provider.get_model_resolutions(模型)
if supported_resolutions and 分辨率 not in supported_resolutions:
raise ValueError(
f"分辨率 \"{分辨率}\" 与模型 \"{模型}\" 不兼容!\n"
f"该模型支持的分辨率:{', '.join(supported_resolutions)}"
)
# 运行时验证宽高比
supported_ratios = provider.get_model_aspect_ratios(模型)
if supported_ratios and 宽高比 not in supported_ratios:
raise ValueError(
f"宽高比 \"{宽高比}\" 与模型 \"{模型}\" 不兼容!\n"
f"该模型支持的宽高比:{', '.join(supported_ratios)}"
)
# 收集参考图
input_images = []
for i in range(1, 10):
key = f"参考图{i}"
if key in kwargs and kwargs[key] is not None:
pil_imgs = tensor_to_pil(kwargs[key])
input_images.extend(pil_imgs)
if input_images and len(input_images) > 14:
raise ValueError(
f"输入图像数量 {len(input_images)} 超过限制 14 张"
)
# 单提示词模式
mode_str = f"图生图模式 (输入{len(input_images)}张)" if input_images else "文生图模式"
print(f"{self.NODE_LABEL}: {mode_str} | {分辨率} {宽高比} | {生图数量}")
prompts_list = [prompt]
images_per_prompt = 生图数量
total_tasks = 生图数量
if total_tasks > 100:
print(f" {self.NODE_LABEL}: 警告!批量生成 {total_tasks} 张图片,内存占用可能较高")
if pbar is not None:
pbar = ProgressBar(total_tasks)
# 在独立线程中运行异步批量处理
def run_async():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
return loop.run_until_complete(
self._process_batch(
provider=provider,
prompts=prompts_list,
model=模型,
resolution=分辨率,
aspect_ratio=宽高比,
images_per_prompt=images_per_prompt,
input_images=input_images,
pbar=pbar,
**extra_kwargs,
)
)
finally:
loop.close()
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(run_async)
# 总超时 = 批次数 × 单任务超时,保证每批都有完整的时间窗口
num_batches = math.ceil(total_tasks / _MAX_CONCURRENT)
batch_timeout = num_batches * _MAX_WAIT_TIME
try:
results = future.result(timeout=batch_timeout)
except TimeoutError:
raise RuntimeError(f"任务执行超时({batch_timeout}秒),请减少数量或检查网络")
# 统计结果
success_count = sum(1 for r in results if r.get("success"))
fail_count = len(results) - success_count
elapsed = time.time() - start_time
time_str = f"{elapsed:.3f}s" if elapsed < 1 else f"{elapsed:.2f}s"
print(f"{self.NODE_LABEL}: 完成!总耗时 {time_str} | 成功: {success_count}/{total_tasks} | 失败: {fail_count}")
# 打印失败详情
failed = [r for r in results if not r.get("success")]
if failed:
for fr in failed:
idx = fr.get("global_task_index", -1) + 1
prompt_snippet = (fr.get("prompt", "") or "")[:30]
error_msg = fr.get("error", "未知错误")
print(f" FAIL #{idx}: {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} -> {error_msg}")
# 收集输出图像
output_images = []
for r in results:
output_images.extend(r.get("output_images", []))
if not output_images:
# 收集所有错误原因
error_details = "\n".join(
f" - {r.get('prompt', '未知提示词')[:40]}: {r.get('error', '未知错误')}"
for r in results if not r.get("success")
)
raise RuntimeError(
f"所有任务均失败 ({fail_count}/{total_tasks})\n{error_details}"
)
output_tensor = _images_to_tensor_safe(output_images, self.NODE_LABEL)
import gc
gc.collect()
return (output_tensor,)
except InterruptProcessingException:
print(f"{self.NODE_LABEL}: 用户取消")
raise
except ValueError as e:
if str(e) == "未授权!":
print(f"{self.NODE_LABEL}: 请联系作者授权后方可使用!")
raise ValueError("未授权!") from None
raise ValueError(str(e)) from None
except RuntimeError as e:
raise RuntimeError(str(e)) from None
except Exception as e:
raise type(e)(str(e)) from None
finally:
# 查询并打印余额
try:
balance_data = provider.query_balance_sync()
if balance_data:
print(f"{self.NODE_LABEL}: {provider.format_balance_info(balance_data)}")
except Exception:
pass
import gc
gc.collect()
class BatchAsyncImageGenerator(AsyncImageGenerator):
"""
AI生图(批量版)- 全并发提交 + 即时落盘
与原版区别:
- 所有任务一次性全并发提交,不分批次
- 每完成一个任务立即将图像保存到磁盘,不会因中途失败丢失已完成图片
- 最终从磁盘加载所有已保存的图像输出
"""
NODE_LABEL = "AI生图(批量版)"
RETURN_NAMES = ("输出图像",)
def __init__(self):
super().__init__()
self._output_file_paths: List[str] = []
self._output_dir: str = ""
def _get_output_dir(self) -> str:
"""获取本次运行的输出目录(带时间戳)"""
if FOLDER_PATHS_AVAILABLE:
base = folder_paths.get_output_directory()
else:
base = os.path.join(os.path.dirname(__file__), "..", "output")
run_id = time.strftime("%Y%m%d_%H%M%S")
run_dir = os.path.join(base, f"batch_{run_id}")
os.makedirs(run_dir, exist_ok=True)
return run_dir
async def _process_batch(
self,
provider: BaseAsyncImageProvider,
prompts: List[str],
model: str,
resolution: str,
aspect_ratio: str,
images_per_prompt: int,
input_images: List[Image.Image],
pbar=None,
**extra_kwargs,
) -> List[dict]:
"""全并发处理:所有任务一次性提交,谁先完成谁先落盘"""
# 构建任务定义
tasks_def = []
for p_idx, prompt in enumerate(prompts):
for sub_idx in range(images_per_prompt):
tasks_def.append((p_idx, sub_idx, prompt))
total_tasks = len(tasks_def)
all_results: List[dict] = []
completed = 0
# 初始化输出目录
self._output_file_paths = []
self._output_dir = self._get_output_dir()
print(f"{self.NODE_LABEL}: 输出目录: {self._output_dir}")
_on_progress = (lambda delta: pbar.update(delta)) if pbar is not None else None
connector = aiohttp.TCPConnector(ssl=False, limit=0, limit_per_host=0)
async with aiohttp.ClientSession(connector=connector) as session:
# 一次性提交所有任务(全并发)
batch_tasks = []
for i, (_, _, prompt) in enumerate(tasks_def):
task = asyncio.create_task(
self._execute_one(
session=session,
provider=provider,
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
input_images=input_images,
global_task_index=i,
on_progress=_on_progress,
**extra_kwargs,
)
)
batch_tasks.append(task)
# 谁先完成先处理谁
for coro in asyncio.as_completed(batch_tasks):
result_data = None
try:
result_data = await coro
except InterruptProcessingException:
for t in batch_tasks:
t.cancel()
raise
except Exception as e2:
result_data = {
"success": False,
"error": str(e2),
"generated_count": 0,
"output_images": [],
"prompt": "",
}
# 即时落盘
if result_data and result_data.get("success"):
for img_idx, img in enumerate(result_data.get("output_images", [])):
filepath = os.path.join(
self._output_dir,
f"task_{completed:04d}_{img_idx:02d}.png"
)
img.save(filepath)
self._output_file_paths.append(filepath)
# 释放内存中的图像对象
result_data["saved_count"] = len(result_data.get("output_images", []))
result_data["output_images"] = []
all_results.append(result_data)
completed += 1
prompt_snippet = (result_data.get("prompt", "") or "")[:30]
if result_data and result_data.get("success"):
count = result_data.get("saved_count", result_data.get("generated_count", 1))
print(f"{self.NODE_LABEL}: [{completed}/{total_tasks}] {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} -> OK({count}张) [已落盘]")
else:
error_msg = result_data.get("error", "未知错误") if result_data else "未知错误"
print(f"{self.NODE_LABEL}: [{completed}/{total_tasks}] {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} -> FAIL: {error_msg}")
import gc
gc.collect()
return all_results
def generate(
self,
prompt: str,
模型: str,
宽高比: str,
分辨率: str,
生图数量: int,
**kwargs
) -> Tuple[torch.Tensor]:
"""生成图像(异步模式 - 批量版:全并发 + 即时落盘)"""
start_time = time.time()
seed: int = kwargs.pop("seed", 0)
proxy_port: str = kwargs.pop("代理端口", "")
proxy_url = BaseAsyncImageProvider.build_proxy_url(proxy_port)
provider = self._get_provider(模型, proxy_url=proxy_url)
if proxy_url:
print(f"{self.NODE_LABEL}: 已启用代理加速 -> {proxy_url}")
extra_kwargs = provider.get_extra_kwargs(**kwargs)
pbar = None
if PROGRESS_BAR_AVAILABLE:
pbar = ProgressBar(生图数量)
try:
random.seed(seed)
np.random.seed(seed % (2**32))
if MEMORY_MONITOR_AVAILABLE and 生图数量 > 50:
process = psutil.Process()
initial_memory = process.memory_info().rss / 1024 / 1024
print(f"{self.NODE_LABEL}: 初始内存使用: {initial_memory:.1f} MB")
supported_resolutions = provider.get_model_resolutions(模型)
if supported_resolutions and 分辨率 not in supported_resolutions:
raise ValueError(
f"分辨率 \"{分辨率}\" 与模型 \"{模型}\" 不兼容!\n"
f"该模型支持的分辨率:{', '.join(supported_resolutions)}"
)
supported_ratios = provider.get_model_aspect_ratios(模型)
if supported_ratios and 宽高比 not in supported_ratios:
raise ValueError(
f"宽高比 \"{宽高比}\" 与模型 \"{模型}\" 不兼容!\n"
f"该模型支持的宽高比:{', '.join(supported_ratios)}"
)
input_images = []
for i in range(1, 10):
key = f"参考图{i}"
if key in kwargs and kwargs[key] is not None:
pil_imgs = tensor_to_pil(kwargs[key])
input_images.extend(pil_imgs)
if input_images and len(input_images) > 14:
raise ValueError(f"输入图像数量 {len(input_images)} 超过限制 14 张")
batch_prompts = parse_batch_prompts(prompt)
if batch_prompts:
num_prompts = len(batch_prompts)
total_images = num_prompts * 生图数量
mode_str = f"批量提示词模式 ({num_prompts}个提示词)"
if input_images:
mode_str += f" (输入{len(input_images)}张)"
print(f"{self.NODE_LABEL}: {mode_str} | {分辨率} {宽高比} | 共{total_images}")
prompts_list = batch_prompts
images_per_prompt = 生图数量
total_tasks = total_images
else:
mode_str = f"图生图模式 (输入{len(input_images)}张)" if input_images else "文生图模式"
print(f"{self.NODE_LABEL}: {mode_str} | {分辨率} {宽高比} | {生图数量}")
prompts_list = [prompt]
images_per_prompt = 生图数量
total_tasks = 生图数量
if total_tasks > 100:
print(f" {self.NODE_LABEL}: 全并发模式,{total_tasks} 张图片将同时提交")
if pbar is not None:
pbar = ProgressBar(total_tasks)
def run_async():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
return loop.run_until_complete(
self._process_batch(
provider=provider,
prompts=prompts_list,
model=模型,
resolution=分辨率,
aspect_ratio=宽高比,
images_per_prompt=images_per_prompt,
input_images=input_images,
pbar=pbar,
**extra_kwargs,
)
)
finally:
loop.close()
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(run_async)
# 全并发:所有任务并行,总超时 = 单任务超时
try:
results = future.result(timeout=_MAX_WAIT_TIME)
except TimeoutError:
raise RuntimeError(f"任务执行超时({_MAX_WAIT_TIME}秒),请减少数量或检查网络")
success_count = sum(1 for r in results if r.get("success"))
fail_count = len(results) - success_count
elapsed = time.time() - start_time
time_str = f"{elapsed:.3f}s" if elapsed < 1 else f"{elapsed:.2f}s"
print(f"{self.NODE_LABEL}: 完成!总耗时 {time_str} | 成功: {success_count}/{total_tasks} | 失败: {fail_count} | 已落盘: {len(self._output_file_paths)}")
failed = [r for r in results if not r.get("success")]
if failed:
for fr in failed:
idx = fr.get("global_task_index", -1) + 1
prompt_snippet = (fr.get("prompt", "") or "")[:30]
error_msg = fr.get("error", "未知错误")
print(f" FAIL #{idx}: {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} -> {error_msg}")
# 从磁盘加载已保存的图像
output_images = []
for fp in self._output_file_paths:
try:
img = Image.open(fp)
output_images.append(img)
except Exception as e:
print(f"{self.NODE_LABEL}: 加载图像失败 {fp}: {e}")
if not output_images:
error_details = "\n".join(
f" - {r.get('prompt', '未知提示词')[:40]}: {r.get('error', '未知错误')}"
for r in results if not r.get("success")
)
raise RuntimeError(
f"所有任务均失败 ({fail_count}/{total_tasks})\n{error_details}"
)
output_tensor = _images_to_tensor_safe(output_images, self.NODE_LABEL)
import gc
gc.collect()
return (output_tensor,)
except InterruptProcessingException:
# 即使被取消,已落盘的图片路径仍然保留
if self._output_file_paths:
print(f"{self.NODE_LABEL}: 用户取消,但 {len(self._output_file_paths)} 张已完成的图片已保存至: {self._output_dir}")
else:
print(f"{self.NODE_LABEL}: 用户取消")
raise
except ValueError as e:
if str(e) == "未授权!":
print(f"{self.NODE_LABEL}: 请联系作者授权后方可使用!")
raise ValueError("未授权!") from None
raise ValueError(str(e)) from None
except RuntimeError as e:
raise RuntimeError(str(e)) from None
except Exception as e:
raise type(e)(str(e)) from None
finally:
try:
balance_data = provider.query_balance_sync()
if balance_data:
print(f"{self.NODE_LABEL}: {provider.format_balance_info(balance_data)}")
except Exception:
pass
import gc
gc.collect()
+10 -13
View File
@@ -23,8 +23,8 @@ class O1keyGPTImage:
- 模型 : 模型选择 - 模型 : 模型选择
- 分辨率 : 图像尺寸(auto 让 API 自动决定) - 分辨率 : 图像尺寸(auto 让 API 自动决定)
- 生图数量 : 每条提示词生成数量 1-8 - 生图数量 : 每条提示词生成数量 1-8
- seed : 随机种子(0 表示不指定)
- 质量 : 生成质量 - 质量 : 生成质量
- seed : 随机种子(0 表示不指定)
- 图片 : 可选参考图(用于图生图或编辑) - 图片 : 可选参考图(用于图生图或编辑)
- 遮罩 : 可选蒙版(白色区域将被替换) - 遮罩 : 可选蒙版(白色区域将被替换)
""" """
@@ -39,12 +39,10 @@ class O1keyGPTImage:
}) })
optional_inputs["模型"] = ([ optional_inputs["模型"] = ([
"gpt-image-2", "gpt-image-2-按量",
"gpt-image-1.5", "gpt-image-2-次卡",
"gpt-image-2-特价",
"gpt-image-1.5-特价",
], { ], {
"default": "gpt-image-2", "default": "gpt-image-2-次卡",
}) })
optional_inputs["分辨率"] = ([ optional_inputs["分辨率"] = ([
"auto(默认)", "auto(默认)",
@@ -53,6 +51,7 @@ class O1keyGPTImage:
"1024x1536(肖像)", "1024x1536(肖像)",
"2048x20482K 平方)", "2048x20482K 平方)",
"2048x11522K 横屏)", "2048x11522K 横屏)",
"1152x20482K 竖屏)",
"3840x21604K 横屏)", "3840x21604K 横屏)",
"2160x38404K 竖屏)", "2160x38404K 竖屏)",
], { ], {
@@ -67,6 +66,10 @@ class O1keyGPTImage:
"display": "number", "display": "number",
"tooltip": "How many images to generate per prompt", "tooltip": "How many images to generate per prompt",
}) })
optional_inputs["质量"] = (["", "", "", "自动"], {
"default": "自动",
"tooltip": "Image quality: 高=high, 中=medium, 低=low, 自动=auto",
})
optional_inputs["seed"] = ("INT", { optional_inputs["seed"] = ("INT", {
"default": 0, "default": 0,
"min": 0, "min": 0,
@@ -76,10 +79,6 @@ class O1keyGPTImage:
"control_after_generate": True, "control_after_generate": True,
"tooltip": "Random seed (0 = not specified)", "tooltip": "Random seed (0 = not specified)",
}) })
optional_inputs["质量"] = (["", "", "", "自动"], {
"default": "自动",
"tooltip": "Image quality: 高=high, 中=medium, 低=low, 自动=auto",
})
optional_inputs["遮罩"] = ("MASK", { optional_inputs["遮罩"] = ("MASK", {
"tooltip": "Optional mask for inpainting (white areas will be replaced)", "tooltip": "Optional mask for inpainting (white areas will be replaced)",
}) })
@@ -104,7 +103,7 @@ class O1keyGPTImage:
def generate( def generate(
self, self,
prompt: str, prompt: str,
模型: str = "gpt-image-2", 模型: str = "gpt-image-2-次卡",
分辨率: str = "auto", 分辨率: str = "auto",
质量: str = "自动", 质量: str = "自动",
生图数量: int = 1, 生图数量: int = 1,
@@ -169,7 +168,6 @@ class O1keyGPTImage:
prompt=p, prompt=p,
model=模型, model=模型,
quality=quality, quality=quality,
background="auto",
size=size, size=size,
n=生图数量, n=生图数量,
seed=seed, seed=seed,
@@ -192,7 +190,6 @@ class O1keyGPTImage:
prompt=prompt, prompt=prompt,
model=模型, model=模型,
quality=quality, quality=quality,
background="auto",
size=size, size=size,
n=生图数量, n=生图数量,
seed=seed, seed=seed,
-683
View File
@@ -1,683 +0,0 @@
"""
Nano Banana Pro(异步)节点
ComfyUI 自定义节点,用于调用 Gemini 模型生成图像(异步提交+轮询模式)
"""
import time
import math
import random
import asyncio
import aiohttp
from concurrent.futures import ThreadPoolExecutor
from typing import Optional, Tuple, List
import torch
import numpy as np
from PIL import Image
from ..utils.image_utils import tensor_to_pil, pil_to_tensor, parse_batch_prompts
from ..utils.file_utils import ImageInfo, generate_timestamp_filename, save_image
from ..utils.config import get_async_api_base_url
from ..clients.gemini_client import GeminiAPIClient
from ..models_config import (
get_enabled_models, get_model_description,
get_model_supported_aspect_ratios, get_all_supported_aspect_ratios,
get_model_supported_resolutions, get_all_supported_resolutions
)
try:
import folder_paths
FOLDER_PATHS_AVAILABLE = True
except ImportError:
FOLDER_PATHS_AVAILABLE = False
try:
from comfy.utils import ProgressBar
PROGRESS_BAR_AVAILABLE = True
except ImportError:
PROGRESS_BAR_AVAILABLE = False
print("⚠️ NanoBananaProAsync: comfy.utils.ProgressBar 不可用,将只使用终端进度显示")
try:
import psutil
MEMORY_MONITOR_AVAILABLE = True
except ImportError:
MEMORY_MONITOR_AVAILABLE = False
print("⚠️ NanoBananaProAsync: psutil 不可用,内存监控功能禁用")
DEBUG_LOG_ENABLED = False
REQUEST_LOG_ENABLED = False
_NODE = "Nano Banana Pro(异步)"
_POLL_INTERVAL = 2 # 轮询间隔(秒)
_MAX_WAIT_TIME = 300 # 最大等待时间(秒)
def _images_to_tensor_safe(images: List[Image.Image], node_label: str) -> torch.Tensor:
if not images:
placeholder = Image.new('RGB', (512, 512), color=(128, 128, 128))
return pil_to_tensor([placeholder])
base_size = max(images, key=lambda img: img.size[0] * img.size[1]).size
matched = [img for img in images if img.size == base_size]
skipped = [img for img in images if img.size != base_size]
if skipped:
sizes_str = ", ".join(f"{img.size[0]}×{img.size[1]}" for img in skipped)
print(
f"{node_label}: 丢弃 {len(skipped)} 张较小尺寸的图 ({sizes_str})"
f"仅输出最大尺寸 {base_size[0]}×{base_size[1]}{len(matched)}"
)
return pil_to_tensor(matched)
class NanoBananaProAsync:
"""
Nano Banana Pro(异步)节点
功能:
- 异步提交任务到 cf-api.o1key.com
- 轮询任务状态直到完成
- 支持批量并发生成
"""
MODELS = None
ASPECT_RATIOS = [
"1:1", "4:3", "3:4", "16:9", "9:16",
"2:3", "3:2", "4:5", "5:4", "21:9",
"1:4", "4:1", "1:8", "8:1"
]
RESOLUTIONS = ["512px", "1K", "2K", "4K"]
def __init__(self):
self.client = None
@classmethod
def INPUT_TYPES(cls):
enabled_models = get_enabled_models()
if not enabled_models:
enabled_models = ["请在 models_config.py 中启用至少一个模型"]
all_aspect_ratios = get_all_supported_aspect_ratios()
if not all_aspect_ratios:
all_aspect_ratios = cls.ASPECT_RATIOS
all_resolutions = get_all_supported_resolutions()
if not all_resolutions:
all_resolutions = cls.RESOLUTIONS
optional_inputs = {}
for i in range(1, 10):
optional_inputs[f"参考图{i}"] = ("IMAGE",)
optional_inputs["代理端口"] = ("STRING", {
"default": "",
"multiline": False,
"placeholder": "本地代理端口,如 7897Clash Verge)或 10808v2rayN),留空不使用"
})
return {
"required": {
"prompt": ("STRING", {
"default": "一个中国女子的OOTD",
"multiline": True
}),
"模型": (enabled_models, {
"default": enabled_models[0]
}),
"宽高比": (all_aspect_ratios, {
"default": "1:1"
}),
"分辨率": (all_resolutions, {
"default": "2K"
}),
"生图数量": ("INT", {
"default": 1,
"min": 1,
"max": 1000,
"step": 1
}),
"联网功能": (["关闭", "打开"], {
"default": "关闭"
}),
"seed": ("INT", {
"default": 0,
"min": 0,
"max": 0xffffffffffffffff
})
},
"optional": optional_inputs
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("输出图像",)
try:
import folder_paths
FOLDER_PATHS_AVAILABLE = True
except ImportError:
FOLDER_PATHS_AVAILABLE = False
FUNCTION = "generate"
CATEGORY = "image/generation"
async def _submit_task_async(
self,
session: aiohttp.ClientSession,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
images: List[Image.Image],
enable_grounding: bool = False,
) -> str:
"""提交异步任务,返回 task_id"""
endpoint = self.client.get_endpoint(model=model, resolution=resolution, image_format="url")
async_endpoint = f"/async{endpoint.split('?')[0]}"
if "?" in endpoint:
async_endpoint += "?" + endpoint.split("?")[1]
request_body = self.client.build_request_body(
prompt=prompt,
images=images if images else None,
aspect_ratio=aspect_ratio,
resolution=resolution,
enable_grounding=enable_grounding,
enable_image_search=False,
)
url = f"{get_async_api_base_url()}{async_endpoint}"
headers = {
"Authorization": f"Bearer {self.client.api_key}",
"Content-Type": "application/json"
}
if REQUEST_LOG_ENABLED:
import json
import copy
debug_body = copy.deepcopy(request_body)
for content in debug_body.get("contents", []):
for part in content.get("parts", []):
if "inline_data" in part and "data" in part["inline_data"]:
data_str = part["inline_data"]["data"]
part["inline_data"]["data"] = f"{data_str[:50]}...[截断]" if len(data_str) > 50 else data_str
print(f"\n{'='*60}")
print(f"[异步提交] URL: {url}")
print(f"[异步提交] 请求体:\n{json.dumps(debug_body, indent=2, ensure_ascii=False)}")
print(f"{'='*60}\n")
async with session.post(url, json=request_body, headers=headers, proxy=self.client.proxy_url) as response:
if response.status != 200:
error_text = await response.text()
raise RuntimeError(f"提交任务失败 ({response.status}): {error_text}")
data = await response.json()
if DEBUG_LOG_ENABLED:
import json
print(f"\n{'='*60}")
print(f"[异步提交] 响应:\n{json.dumps(data, indent=2, ensure_ascii=False)}")
print(f"{'='*60}\n")
task_id = data.get("task_id")
if not task_id:
raise RuntimeError(f"提交响应中未找到 task_id: {data}")
return task_id
async def _poll_task_async(
self,
session: aiohttp.ClientSession,
task_id: str,
) -> dict:
"""轮询任务状态直到完成"""
url = f"{get_async_api_base_url()}/async/v1/tasks/{task_id}"
headers = {
"Authorization": f"Bearer {self.client.api_key}",
"Content-Type": "application/json"
}
start_time = time.time()
poll_count = 0
while True:
elapsed = time.time() - start_time
if elapsed > _MAX_WAIT_TIME:
raise RuntimeError(f"任务 {task_id} 超时({_MAX_WAIT_TIME}秒),请稍后手动查询")
poll_count += 1
async with session.get(url, headers=headers, proxy=self.client.proxy_url) as response:
if response.status != 200:
error_text = await response.text()
raise RuntimeError(f"查询任务失败 ({response.status}): {error_text}")
result = await response.json()
status = result.get("status")
if DEBUG_LOG_ENABLED:
import json
print(f"\n{'='*60}")
print(f"[轮询 #{poll_count}] task_id: {task_id}")
print(f"[轮询 #{poll_count}] 响应:\n{json.dumps(result, indent=2, ensure_ascii=False)}")
print(f"{'='*60}\n")
if status == "SUCCESS":
return result.get("data", {})
elif status == "FAILURE":
error_msg = result.get("error", "未知错误")
raise RuntimeError(f"任务失败: {error_msg}")
elif status in ["SUBMITTED", "IN_PROGRESS"]:
await asyncio.sleep(_POLL_INTERVAL)
else:
raise RuntimeError(f"未知任务状态: {status}")
async def _generate_single_task(
self,
session: aiohttp.ClientSession,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
images: List[Image.Image],
output_folder: str,
global_task_index: int,
enable_grounding: bool = False,
save_to_disk: bool = True,
) -> dict:
"""执行单个异步生成任务"""
result = {
"global_task_index": global_task_index,
"prompt": prompt,
"success": False,
"generated_count": 0,
"saved_files": [],
"output_images": [],
"error": None
}
try:
task_id = await self._submit_task_async(
session=session,
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
images=images,
enable_grounding=enable_grounding,
)
response_data = await self._poll_task_async(session=session, task_id=task_id)
# 兼容两种响应格式:
# 1. 异步接口直接返回 image_url{"image_url": "https://..."}
# 2. Gemini 标准格式:{"candidates": [...]}
image_url = response_data.get("image_url", "") if isinstance(response_data, dict) else ""
if image_url:
from io import BytesIO
async with session.get(image_url) as img_resp:
if img_resp.status == 200:
img_bytes = await img_resp.read()
img = Image.open(BytesIO(img_bytes))
images_list = [img]
else:
raise RuntimeError(f"下载图片失败 ({img_resp.status}): {image_url}")
else:
images_list, _ = await self.client.parse_response_async(response_data, session=session)
if save_to_disk:
for gen_img in images_list:
output_path = generate_timestamp_filename(
output_folder=output_folder,
extension=".png"
)
save_image(gen_img, output_path)
result["saved_files"].append(output_path)
gen_img = None
else:
result["output_images"] = images_list
result["success"] = True
result["generated_count"] = len(images_list)
except Exception as e:
result["error"] = str(e)
return result
async def _process_batch_async(
self,
prompts: List[str],
model: str,
resolution: str,
aspect_ratio: str,
images_per_prompt: int,
input_images: List[Image.Image],
output_folder: str,
pbar=None,
enable_grounding: bool = False,
save_to_disk: bool = True,
) -> List[dict]:
"""异步批量处理"""
tasks_def = []
for p_idx, prompt in enumerate(prompts):
for sub_idx in range(images_per_prompt):
tasks_def.append((p_idx, sub_idx, prompt))
total_tasks = len(tasks_def)
max_concurrent = 50
num_batches = math.ceil(total_tasks / max_concurrent)
all_results = []
completed = 0
success_count = 0
fail_count = 0
connector = aiohttp.TCPConnector(ssl=False, limit=0, limit_per_host=0)
async with aiohttp.ClientSession(connector=connector) as session:
for batch_idx in range(num_batches):
start_idx = batch_idx * max_concurrent
end_idx = min(start_idx + max_concurrent, total_tasks)
tasks = []
for i in range(start_idx, end_idx):
_, _, prompt = tasks_def[i]
task = asyncio.create_task(
self._generate_single_task(
session=session,
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
images=input_images,
output_folder=output_folder,
global_task_index=i,
enable_grounding=enable_grounding,
save_to_disk=save_to_disk,
)
)
tasks.append(task)
batch_results = []
for coro in asyncio.as_completed(tasks):
result_data = None
try:
result = await coro
if isinstance(result, Exception):
result_data = {"success": False, "error": str(result), "generated_count": 0, "saved_files": [], "prompt": ""}
else:
result_data = result
batch_results.append(result_data)
except Exception as e:
result_data = {"success": False, "error": str(e), "generated_count": 0, "saved_files": [], "prompt": ""}
batch_results.append(result_data)
completed += 1
prompt_snippet = (result_data.get("prompt", "") or "")[:30]
if result_data and result_data.get("success", False):
success_count += 1
count = result_data.get("generated_count", 1)
print(f"{_NODE}: [{completed}/{total_tasks}] {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} → ✓成功({count}张)")
else:
fail_count += 1
error_msg = result_data.get("error", "未知错误") if result_data else "未知错误"
print(f"{_NODE}: [{completed}/{total_tasks}] {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} → ✗失败: {error_msg}")
if pbar is not None:
pbar.update(1)
all_results.extend(batch_results)
import gc
gc.collect()
await asyncio.sleep(0.1)
return all_results
def generate(
self,
prompt: str,
模型: str,
宽高比: str,
分辨率: str,
生图数量: int,
seed: int,
**kwargs
) -> Tuple[torch.Tensor]:
"""生成图像(异步模式)"""
start_time = time.time()
enable_grounding: bool = (kwargs.pop("联网功能", "关闭") == "打开")
proxy_port: str = kwargs.pop("代理端口", "")
pbar = None
if PROGRESS_BAR_AVAILABLE:
pbar = ProgressBar(生图数量)
try:
random.seed(seed)
np.random.seed(seed % (2**32))
if MEMORY_MONITOR_AVAILABLE and 生图数量 > 50:
import psutil
process = psutil.Process()
initial_memory = process.memory_info().rss / 1024 / 1024
print(f"{_NODE}: 初始内存使用: {initial_memory:.1f} MB")
if self.client is None:
try:
self.client = GeminiAPIClient()
except ValueError as e:
raise ValueError(f"初始化失败: {str(e)}")
self.client.proxy_url = GeminiAPIClient.build_proxy_url(proxy_port)
if self.client.proxy_url:
print(f"{_NODE}: 已启用代理加速 → {self.client.proxy_url}")
supported_resolutions = get_model_supported_resolutions(模型)
if supported_resolutions and 分辨率 not in supported_resolutions:
raise ValueError(
f"分辨率 \"{分辨率}\" 与模型 \"{模型}\" 不兼容!\n"
f"该模型支持的分辨率:{', '.join(supported_resolutions)}"
)
supported_ratios = get_model_supported_aspect_ratios(模型)
if supported_ratios and 宽高比 not in supported_ratios:
raise ValueError(
f"宽高比 \"{宽高比}\" 与模型 \"{模型}\" 不兼容!\n"
f"该模型支持的宽高比:{', '.join(supported_ratios)}"
)
input_images = []
for i in range(1, 10):
key = f"参考图{i}"
if key in kwargs and kwargs[key] is not None:
pil_imgs = tensor_to_pil(kwargs[key])
input_images.extend(pil_imgs)
if input_images:
if len(input_images) > 14:
raise ValueError(
f"输入图像数量 {len(input_images)} 超过限制 14 张,请减少输入图像数量"
)
batch_prompts = parse_batch_prompts(prompt)
grounding_str = ""
if enable_grounding:
grounding_str = " | 谷歌搜索接地"
if batch_prompts:
num_prompts = len(batch_prompts)
total_images = num_prompts * 生图数量
mode_str = f"批量提示词模式 ({num_prompts}个提示词)"
if input_images:
mode_str += f" (输入{len(input_images)}张)"
print(f"{_NODE}: {mode_str} | {分辨率} {宽高比} | 共{total_images}{grounding_str}")
if total_images > 100:
print(f"⚠️ {_NODE}: 警告!批量生成 {total_images} 张图片,内存占用可能较高")
print(f"⚠️ 建议:分批执行或减少生图数量")
else:
mode_str = f"图生图模式 (输入{len(input_images)}张)" if input_images else "文生图模式"
print(f"{_NODE}: {mode_str} | {分辨率} {宽高比} | {生图数量}{grounding_str}")
if 生图数量 > 100:
print(f"⚠️ {_NODE}: 警告!批量生成 {生图数量} 张图片,内存占用可能较高")
print(f"⚠️ 建议:分批执行或减少生图数量")
success_count = 0
fail_count = 0
if batch_prompts:
num_prompts = len(batch_prompts)
total_images = num_prompts * 生图数量
if pbar is not None:
pbar = ProgressBar(total_images)
def run_async_in_thread():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
return loop.run_until_complete(
self._process_batch_async(
prompts=batch_prompts,
model=模型,
resolution=分辨率,
aspect_ratio=宽高比,
images_per_prompt=生图数量,
input_images=input_images,
output_folder="",
pbar=pbar,
enable_grounding=enable_grounding,
save_to_disk=False,
)
)
finally:
loop.close()
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(run_async_in_thread)
try:
results = future.result(timeout=900)
except TimeoutError:
raise RuntimeError("任务执行超时(900秒),请减少提示词数量或检查网络连接")
success_count = sum(1 for r in results if r.get("success", False))
fail_count = len(results) - success_count
total_generated = sum(r.get("generated_count", 0) for r in results)
elapsed = time.time() - start_time
time_str = f"{elapsed:.3f}s" if elapsed < 1 else f"{elapsed:.2f}s"
print(f"完成!总耗时 {time_str} | 成功: {success_count}/{total_images} | 失败: {fail_count}")
failed_results = [r for r in results if not r.get("success", False)]
if failed_results:
for fr in failed_results:
idx = fr.get("global_task_index", -1) + 1
prompt_snippet = (fr.get("prompt", "") or "")[:30]
error_msg = fr.get("error", "未知错误")
print(f" 失败 #{idx}: {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''}{error_msg}")
output_images = []
for r in results:
output_images.extend(r.get("output_images", []))
if not output_images:
placeholder = Image.new('RGB', (512, 512), color=(128, 128, 128))
output_images = [placeholder]
output_tensor = _images_to_tensor_safe(output_images, _NODE)
import gc
gc.collect()
return (output_tensor,)
else:
if pbar is not None:
pbar = ProgressBar(生图数量)
def run_async_in_thread():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
return loop.run_until_complete(
self._process_batch_async(
prompts=[prompt],
model=模型,
resolution=分辨率,
aspect_ratio=宽高比,
images_per_prompt=生图数量,
input_images=input_images,
output_folder="",
pbar=pbar,
enable_grounding=enable_grounding,
save_to_disk=False,
)
)
finally:
loop.close()
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(run_async_in_thread)
try:
results = future.result(timeout=900)
except TimeoutError:
raise RuntimeError("任务执行超时(900秒),请减少生图数量或检查网络连接")
success_count = sum(1 for r in results if r.get("success", False))
fail_count = len(results) - success_count
total_generated = sum(r.get("generated_count", 0) for r in results)
elapsed = time.time() - start_time
time_str = f"{elapsed:.3f}s" if elapsed < 1 else f"{elapsed:.2f}s"
print(f"完成!总耗时 {time_str} | 成功: {success_count}/{生图数量} | 失败: {fail_count}")
failed_results = [r for r in results if not r.get("success", False)]
if failed_results:
for fr in failed_results:
idx = fr.get("global_task_index", -1) + 1
error_msg = fr.get("error", "未知错误")
print(f" 失败 #{idx}: {prompt[:30]}{'...' if len(prompt) >= 30 else ''}{error_msg}")
output_images = []
for r in results:
output_images.extend(r.get("output_images", []))
if not output_images:
placeholder = Image.new('RGB', (512, 512), color=(128, 128, 128))
output_images = [placeholder]
output_tensor = _images_to_tensor_safe(output_images, _NODE)
import gc
gc.collect()
return (output_tensor,)
except ValueError as e:
if str(e) == "未授权!":
print("请联系作者授权后方可使用!")
raise ValueError("未授权!") from None
raise ValueError(str(e)) from None
except RuntimeError as e:
raise RuntimeError(str(e)) from None
except Exception as e:
raise type(e)(str(e)) from None
finally:
if self.client is not None:
try:
balance_data = self.client.query_balance_sync()
balance_info = self.client.format_balance_info(balance_data)
print(f"{_NODE}: {balance_info}")
except Exception:
pass
import gc
gc.collect()
+1 -1
View File
@@ -25,10 +25,10 @@ from ..utils.file_types import FileList
# ============================================================================ # ============================================================================
SUPPORTED_MODELS = [ SUPPORTED_MODELS = [
"gpt-5.5",
"gpt-5.4", "gpt-5.4",
"gemini-3.1-flash-lite-preview", "gemini-3.1-flash-lite-preview",
"gemini-3.1-pro-preview", "gemini-3.1-pro-preview",
"deepseek-v3.2",
"deepseek-v4-pro", "deepseek-v4-pro",
"doubao-seed-2-0-pro-260215", "doubao-seed-2-0-pro-260215",
] ]