Files
hujuying-ComfyUI-ModelScope…/modelscope_image_caption_node.py
hujuying 15707d0c1d Add files via upload
对图片解析节点增加随机种子选项
2025-12-12 22:37:47 +08:00

234 lines
8.4 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import requests
import json
import time
import torch
import numpy as np
from PIL import Image
from io import BytesIO
import os
import base64
import re
from .modelscope_image_node import load_config, save_config, tensor_to_base64_url
# 检查openai库是否可用
try:
from openai import OpenAI
OPENAI_AVAILABLE = True
except ImportError:
OPENAI_AVAILABLE = False
# 仅与modelscope_config.json交互的API Token管理函数
def load_api_tokens():
try:
cfg = load_config()
tokens_from_cfg = cfg.get("api_tokens", [])
if tokens_from_cfg and isinstance(tokens_from_cfg, list):
return [token.strip() for token in tokens_from_cfg if token.strip()]
except Exception as e:
print(f"读取config中的tokens失败: {e}")
return []
def save_api_tokens(tokens):
try:
cfg = load_config()
cfg["api_tokens"] = tokens
return save_config(cfg)
except Exception as e:
print(f"保存tokens到config失败: {e}")
return False
class ModelScopeImageCaptionNode:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
if not OPENAI_AVAILABLE:
return {
"required": {
"error_message": ("STRING", {
"default": "请先安装openai库: pip install openai",
"multiline": True
}),
}
}
saved_tokens = load_api_tokens()
# 定义支持的模型列表
supported_models = [
"Qwen/Qwen3-VL-8B-Instruct",
"Qwen/Qwen3-VL-235B-A22B-Instruct"
]
return {
"required": {
"api_tokens": ("STRING", {
"default": f"***已保存{len(saved_tokens)}个Token***" if saved_tokens else "",
"placeholder": "请输入API Token(支持多个,用逗号/换行分隔)",
"multiline": True
}),
},
"optional": {
# 关键修改:将image设置为可选输入
"image": ("IMAGE", {"optional": True}),
"prompt1": ("STRING", {
"multiline": True,
"default": "详细描述这张图片的内容,包括主体、背景、颜色、风格等信息"
}),
"prompt2": ("STRING", {
"multiline": True,
"default": ""
}),
"model": (supported_models, {
"default": "Qwen/Qwen3-VL-8B-Instruct"
}),
"max_tokens": ("INT", {
"default": 1000,
"min": 100,
"max": 4000
}),
"temperature": ("FLOAT", {
"default": 0.7,
"min": 0.1,
"max": 2.0,
"step": 0.1
}),
# 新增seed选项(与生图节点保持一致:默认-1表示随机)
"seed": ("INT", {
"default": -1,
"min": -1,
"max": 2147483647
})
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("description",)
FUNCTION = "generate_caption"
CATEGORY = "ModelScopeAPI"
def parse_api_tokens(self, token_input):
"""解析输入的多个API Token(支持逗号、分号、换行分隔)"""
if not token_input or token_input.strip() in ["", f"***已保存{len(load_api_tokens())}个Token***"]:
return load_api_tokens()
# 支持多种分隔符拆分Token
tokens = re.split(r'[,;\n]+', token_input)
return [token.strip() for token in tokens if token.strip()]
def create_blank_image(self, width=64, height=64):
"""创建空白图像张量(符合ComfyUI的图像格式要求)"""
# 创建白色背景的RGB图像
blank_np = np.ones((height, width, 3), dtype=np.uint8) * 255
# 转换为ComfyUI格式的张量 (batch, height, width, channels)
blank_tensor = torch.from_numpy(blank_np).unsqueeze(0).float() / 255.0
return blank_tensor
def generate_caption(self, image=None, api_tokens="", prompt1="详细描述这张图片的内容", prompt2="", model="Qwen/Qwen3-VL-8B-Instruct", max_tokens=1000, temperature=0.7, seed=-1):
if not OPENAI_AVAILABLE:
return ("请先安装openai库: pip install openai",)
# 应用seed(-1则使用随机种子)
if seed == -1:
seed = np.random.randint(0, 2147483647)
np.random.seed(seed % (2**32 - 1))
# 关键修改:处理输入图像为空的情况
if image is None:
print("⚠️ 未输入图像,自动生成空白图像作为输入")
image = self.create_blank_image()
# 处理提示词合并
prompt_parts = []
if prompt1.strip():
prompt_parts.append(prompt1.strip())
if prompt2.strip():
prompt_parts.append(prompt2.strip())
if not prompt_parts:
prompt = "详细描述这张图片的内容,包括主体、背景、颜色、风格等信息"
else:
prompt = ", ".join(prompt_parts)
# 解析Token列表
tokens = self.parse_api_tokens(api_tokens)
if not tokens:
raise Exception("请提供至少一个有效的API Token")
# 保存新Token(如果有变化)
saved_tokens = load_api_tokens()
if api_tokens.strip() not in ["", f"***已保存{len(saved_tokens)}个Token***"]:
if save_api_tokens(tokens):
print(f"✅ 已保存 {len(tokens)} 个API Token")
else:
print("⚠️ API Token保存失败,但不影响当前使用")
try:
print(f"🔍 开始生成图像描述...")
print(f"📝 提示词: {prompt}")
print(f"🤖 模型: {model}")
print(f"🔑 可用Token数量: {len(tokens)}")
print(f"🌱 Seed: {seed}") # 打印seed信息
# 转换图像为base64格式
image_url = tensor_to_base64_url(image)
print(f"🖼️ 图像已转换为base64格式")
# 构建消息体
messages = [{
'role': 'user',
'content': [{
'type': 'text',
'text': prompt,
}, {
'type': 'image_url',
'image_url': {
'url': image_url,
},
}],
}]
# 轮询尝试每个Token
last_exception = None
for i, token in enumerate(tokens):
try:
print(f"🔄 尝试使用第 {i+1}/{len(tokens)} 个Token...")
client = OpenAI(
base_url='https://api-inference.modelscope.cn/v1',
api_key=token
)
response = client.chat.completions.create(
model=model,
messages=messages,
max_tokens=max_tokens,
temperature=temperature,
stream=False
)
description = response.choices[0].message.content
print(f"✅ 第 {i+1} 个Token调用成功!")
print(f"📄 结果预览: {description[:100]}...")
return (description,)
except Exception as e:
last_exception = e
print(f"❌ 第 {i+1} 个Token调用失败: {str(e)}")
if i < len(tokens) - 1:
print(f"⏳ 准备尝试下一个Token...")
# 所有Token都失败
raise Exception(f"所有Token均调用失败: {str(last_exception)}")
except Exception as e:
error_msg = f"图像描述生成失败: {str(e)}"
print(f"❌ {error_msg}")
return (error_msg,)
# 节点映射
NODE_CLASS_MAPPINGS = {
"ModelScopeImageCaptionNode": ModelScopeImageCaptionNode
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ModelScopeImageCaptionNode": "ModelScope 图像描述生成"
}