Files
2024-01-30 17:07:05 +08:00

283 lines
9.5 KiB
Python

import os
import io
import json
import requests
import torch
import dashscope
from dashscope import MultiModalConversation
from io import BytesIO
from PIL import Image, ImageChops
from datetime import datetime
import tempfile
import random
import platform
import hashlib
p = os.path.dirname(os.path.realpath(__file__))
def get_qwenvl_api_key():
try:
config_path = os.path.join(p, 'config.json')
with open(config_path, 'r') as f:
config = json.load(f)
api_key = config["QWENVL_API_KEY"]
except:
print("出错啦 Error: API key is required")
return ""
return api_key
class QWenVL_API_S_Zho:
def __init__(self):
self.api_key = get_qwenvl_api_key()
if self.api_key is not None:
dashscope.api_key=self.api_key
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"prompt": ("STRING", {"default": "Describe this image", "multiline": True}),
"model_name": (["qwen-vl-plus", "qwen-vl-max"],),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("text",)
FUNCTION = "qwen_vl_generation"
CATEGORY = "Zho模块组/💫QWenVL"
def tensor_to_image(self, tensor):
# 确保张量是在CPU上
tensor = tensor.cpu()
# 将张量数据转换为0-255范围并转换为整数
# 这里假设张量已经是H x W x C格式
image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy()
# 创建PIL图像
image = Image.fromarray(image_np, mode='RGB')
return image
def qwen_vl_generation(self, image, prompt, model_name, seed):
if not self.api_key:
raise ValueError("API key is required")
if image == None:
raise ValueError("qwen_vl needs a image")
else:
# 转换图像
pil_image = self.tensor_to_image(image)
# 生成临时文件路径
temp_directory = tempfile.gettempdir()
unique_suffix = "_temp_" + ''.join(random.choice("abcdefghijklmnopqrstuvwxyz") for _ in range(5))
filename = f"image{unique_suffix}.png"
temp_image_path = os.path.join(temp_directory, filename)
#temp_image_url = f"file://{temp_image_path}"
# 根据操作系统选择正确的文件URL格式
if platform.system() == 'Windows':
temp_image_url = f"file://{temp_image_path}"
else:
temp_image_url = f"file:///{temp_image_path}"
temp_image_url = temp_image_url.replace('\\', '/')
# 保存图像到临时路径
pil_image.save(temp_image_path)
messages = [
{
"role": "user",
"content": [
{"image": temp_image_url},
{"text": prompt}
]
}
]
#print("temp_image_url:", temp_image_url)
#print("prompt:", prompt)
torch.manual_seed(seed)
response = dashscope.MultiModalConversation.call(model=model_name, messages=messages, seed=seed)
#print(response)
response_json = response
if 'output' in response_json and 'choices' in response_json['output']:
choices = response_json['output']['choices']
if choices and 'message' in choices[0]:
message_content = choices[0]['message']['content']
if message_content and 'text' in message_content[0]:
text_output = message_content[0]['text']
#print(text_output)
else:
print("No text content found.")
else:
print("No message found in the first choice.")
else:
print("No choices found in the output.")
os.remove(temp_image_path)
#print("remove : done!" )
return (text_output, )
class QWenVL_API_S_Multi_Zho:
def __init__(self):
self.api_key = get_qwenvl_api_key()
self.messages = [] # 初始化对话历史为空
self.last_image_hash = None
if self.api_key is not None:
dashscope.api_key=self.api_key
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"prompt": ("STRING", {"default": "Describe this image", "multiline": True}),
"model_name": (["qwen-vl-plus", "qwen-vl-max"],),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("text",)
FUNCTION = "qwen_vl_generation"
CATEGORY = "Zho模块组/💫QWenVL"
def tensor_to_image(self, tensor):
# 确保张量是在CPU上
tensor = tensor.cpu()
# 将张量数据转换为0-255范围并转换为整数
# 这里假设张量已经是H x W x C格式
image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy()
# 创建PIL图像
image = Image.fromarray(image_np, mode='RGB')
return image
def format_qwchat_history(self):
formatted_history = []
for message in self.messages:
role = message['role']
contents = message['content']
for content in contents:
if 'text' in content:
text = content['text']
formatted_message = f"{role}: {text}"
formatted_history.append(formatted_message)
formatted_history.append("-" * 40) # 添加分隔线
return "\n".join(formatted_history)
def get_image_hash(self, pil_image):
# 将图像转换为字节
image_bytes = pil_image.tobytes()
# 使用哈希函数计算哈希值
return hashlib.md5(image_bytes).hexdigest()
def qwen_vl_generation(self, image, prompt, model_name, seed):
if not self.api_key:
raise ValueError("API key is required")
if image == None:
raise ValueError("qwen_vl needs a image")
else:
# 转换图像
pil_image = self.tensor_to_image(image)
# 在当前文件目录下创建 "qw" 文件夹(如果不存在)
qw_folder = os.path.join(p, "qw")
os.makedirs(qw_folder, exist_ok=True)
# 获取当前图像的哈希值
current_image_hash = self.get_image_hash(pil_image)
# 构建文件名
local_image_filename = f"image_{current_image_hash}.png"
local_image_path = os.path.join(qw_folder, local_image_filename)
# 根据操作系统选择正确的文件URL格式
if platform.system() == 'Windows':
local_image_url = f"file://{local_image_path}"
else:
local_image_url = f"file:///{local_image_path}"
# 保证路径中的反斜杠被替换为正斜杠
local_image_url = local_image_url.replace('\\', '/')
# 如果当前图像与上次的不同
if current_image_hash != self.last_image_hash:
pil_image.save(local_image_path)
# 更新last_image_hash
self.last_image_hash = current_image_hash
print(f"Image saved as {local_image_filename}")
else:
print("Image not saved as it is identical to the last one.")
self.messages.append({
"role": "user",
"content": [
{"image": local_image_url},
{"text": prompt}
]
})
#print("local_image_url:", local_image_url)
#print("prompt:", prompt)
torch.manual_seed(seed)
response = dashscope.MultiModalConversation.call(model=model_name, messages=self.messages, seed=seed)
#print(response)
# 更新对话历史
if response and response.output and response.output.choices:
choice = response.output.choices[0]
if choice and choice.message:
self.messages.append({'role': choice.message.role, 'content': choice.message.content})
response_json = response
if 'output' in response_json and 'choices' in response_json['output']:
choices = response_json['output']['choices']
if choices and 'message' in choices[0]:
message_content = choices[0]['message']['content']
if message_content and 'text' in message_content[0]:
text_output = message_content[0]['text']
#print(text_output)
else:
print("No text content found.")
else:
print("No message found in the first choice.")
else:
print("No choices found in the output.")
# 获取格式化的对话历史
chat_history = self.format_qwchat_history()
return (chat_history, )
NODE_CLASS_MAPPINGS = {
"QWenVL_API_S_Zho": QWenVL_API_S_Zho,
"QWenVL_API_S_Multi_Zho": QWenVL_API_S_Multi_Zho,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"QWenVL_API_S_Zho": "㊙️QWenVL_Zho",
"QWenVL_API_S_Multi_Zho": "㊙️QWenVL_Chat_Zho",
}