Compare commits

..
9 Commits
Author SHA1 Message Date
shadowcz007 3bac87ee52 system prompt 2024-06-04 08:50:58 +08:00
shadowcz007 9a01701019 whisper+chat 2024-06-04 08:40:17 +08:00
shadowcz007 693954ee23 ing 2024-06-03 20:12:12 +08:00
shadowcz007 bf4ba91e7a update 2024-06-03 17:43:41 +08:00
shadowcz007 0828353253 Update main.py 2024-06-03 16:16:02 +08:00
shadowcz007 1997c7ad8f Update live.html 2024-06-03 16:12:06 +08:00
shadowcz007 b1e62440e4 test 2024-06-02 22:38:15 +08:00
shadowcz007 d549a5eb6a whisper 2024-06-02 19:53:07 +08:00
shadowcz007 77bfb08d76 web 2024-06-02 17:12:11 +08:00
57 changed files with 3933 additions and 15456 deletions
-1
View File
@@ -1 +0,0 @@
mixlabnodes.com
+9 -35
View File
@@ -1,25 +1,15 @@
![](https://img.shields.io/github/release/shadowcz007/comfyui-mixlab-nodes)
> 适配了最新版 comfyui 的 py3.11 ,torch 2.3.1+cu121
> 适配了最新版 comfyui 的 py3.11 ,torch 2.1.2+cu121
> [Mixlab nodes discord](https://discord.gg/cXs9vZSqeK)
##### `最新`:
- App模式增加batch prompt,批量提示词,可以把动态提示词批量组成后运行
ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。模型下载后,放置到 `models/llamafile/`
![alt text](./assets/1722517810720.png)
- 右键菜单支持 text-to-text,方便对 prompt 词补全
- 增加 API Key Input 节点,用于管理LLM的Key,同时优化LLM相关节点,为后续agent模式做准备
- 增加 SiliconflowLLM,可以使用由Siliconflow提供的免费LLM
- 增加 Edit Mask,方便在生成的时候手动绘制 mask [workflow](./workflow/edit-mask-workflow.json)
<!-- - ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。模型下载后,放置到 `models/llamafile/` -->
<!-- - 右键菜单支持 text-to-text,方便对 prompt 词补全 -->
<!--
强烈推荐:
[Phi-3-mini-4k-instruct-function-calling-GGUF](https://huggingface.co/nold/Phi-3-mini-4k-instruct-function-calling-GGUF)
@@ -28,16 +18,12 @@
- 右键菜单支持 image-to-text,使用多模态模型,多模态使用 [llava-phi-3-mini-gguf](https://huggingface.co/xtuner/llava-phi-3-mini-gguf/tree/main),注意需要把llava-phi-3-mini-mmproj-f16.gguf也下载
![](./assets/prompt_ai_setup.png)
![](./assets/prompt-ai.png) -->
![](./assets/prompt-ai.png)
#### `相关插件推荐`
[comfyui-liveportrait](https://github.com/shadowcz007/comfyui-liveportrait)
[Comfyui-ChatTTS](https://github.com/shadowcz007/Comfyui-ChatTTS)
[comfyui-sound-lab](https://github.com/shadowcz007/comfyui-sound-lab)
<!-- [comfyui-sd-prompt-mixlab](https://github.com/shadowcz007/comfyui-sd-prompt-mixlab) -->
[comfyui-Image-reward](https://github.com/shadowcz007/comfyui-Image-reward)
@@ -54,8 +40,6 @@
- 发布为 app 的 workflow,可以在右键里再次编辑了
- web app 可以设置分类,在 comfyui 右键菜单可以编辑更新 web app
- 支持动态提示
- 支持把输出显示到comfyui背景(TouchDesigner 风格)
- 如果转为web app打开是空白的,注意检查下插件目录的名字需要是:comfyui-mixlab-nodes(如果是zip包下载会多了个-main的后缀,需要去掉)
![](./assets/微信图片_20240421205440.png)
@@ -118,12 +102,11 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
> Support for calling multiple GPTs.Local LLM(llama.cpp)、 ChatGPT、ChatGLM3 、ChatGLM4 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1 . Azure OpenAI:https://xxxx.openai.azure.com
[LLM_base_workflow](./workflow/LLM_base_workflow.json)
![gpt-workflow.svg](./assets/gpt-workflow.svg)
- SiliconflowLLM
- ChatGPTOpenAI
[workflow-5](./workflow/5-gpt-workflow.json)
<!-- 最新:ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
最新:ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
Model download,move to :`models/llamafile/`
@@ -151,7 +134,7 @@ pip install 'llama-cpp-python[server]'
```
pip install llama-cpp-python \
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/metal
``` -->
```
## Prompt
@@ -178,9 +161,6 @@ pip install llama-cpp-python \
> A new layer class node has been added, allowing you to separate the image into layers. After merging the images, you can input the controlnet for further processing.
> The composite images node overlays a foreground image onto a background image at specified positions and scales, with optional blending modes and masking capabilities. position : 'overall',"center_center","left_bottom","center_bottom","right_bottom","left_top","center_top","right_top"
![layers](./assets/layers-workflow.svg)
![poster](./assets/poster-workflow.svg)
@@ -212,12 +192,6 @@ pip install llama-cpp-python \
> Conveniently load images from a fixed address on the internet to ensure that default images in the workflow can be executed.
#### TextImage
> [下载字体](https://drxie.github.io/OSFCC/)放到 ```custom_nodes/comfyui-mixlab-nodes/assets/fonts```
### Style
> Apply VisualStyle Prompting , Modified from [ComfyUI_VisualStylePrompting](https://github.com/ExponentialML/ComfyUI_VisualStylePrompting)
+254 -418
View File
@@ -1,36 +1,36 @@
#
import os
import subprocess
import importlib.util
import sys,json
import execution
import uuid
import urllib
import hashlib
import datetime
import folder_paths
import logging
import base64,io,re
import random
from PIL import Image
from comfy.cli_args import args
python = sys.executable
# print("sys.path", sys.path)
#修复 sys.stdout.isatty() object has no attribute 'isatty'
try:
sys.stdout.isatty()
except:
print('#fix sys.stdout.isatty')
# print('#fix sys.stdout.isatty')
sys.stdout.isatty = lambda: False
_URL_=None
llama_port=None
llama_model=""
llama_chat_format=""
try:
from .nodes.ChatGPT import get_llama_models,get_llama_model_path,llama_cpp_client
llama_cpp_client("")
# try:
# from .nodes.ChatGPT import get_llama_models,get_llama_model_path,llama_cpp_client
# llama_cpp_client("")
# except:
# print("##nodes.ChatGPT ImportError")
except:
print("##nodes.ChatGPT ImportError")
from .nodes.RembgNode import get_rembg_models,U2NET_HOME,run_briarmbg,run_rembg
@@ -172,14 +172,15 @@ def create_for_https():
os.mkdir(https_key_path)
if not os.path.exists(crt):
create_key(key,crt)
# print('https_key OK: ', crt,key)
print('https_key OK: ', crt,key)
return (crt,key)
# workflow 目录下的所有json
def read_workflow_json_files_all(folder_path):
# print('#read_workflow_json_files_all',folder_path)
print('#read_workflow_json_files_all',folder_path)
json_files = []
for root, dirs, files in os.walk(folder_path):
for file in files:
@@ -309,32 +310,31 @@ def get_my_workflow_for_app(filename="my_workflow_app.json",category="",is_all=F
print('app_workflow_path: ',app_workflow_path)
try:
with open(app_workflow_path) as json_file:
json_data=json.load(json_file)
apps = [{
'filename':filename,
'data':json_data
'data':json.load(json_file)
}]
except Exception as e:
print("发生异常:", str(e))
# 这个代码不需要
# if len(apps)==1 and category!='' and category!=None:
data=read_workflow_json_files(category_path)
data=read_workflow_json_files(category_path)
for item in data:
x=item["data"]
# print(apps[0]['filename'] ,item["filename"])
if apps[0]['filename']!=item["filename"]:
category=''
input=None
output=None
if 'category' in x['app']:
category=x['app']['category']
if 'input' in x['app']:
input=x['app']['input']
if 'output' in x['app']:
output=x['app']['output']
apps.append({
for item in data:
x=item["data"]
# print(apps[0]['filename'] ,item["filename"])
if apps[0]['filename']!=item["filename"]:
category=''
input=None
output=None
if 'category' in x['app']:
category=x['app']['category']
if 'input' in x['app']:
input=x['app']['input']
if 'output' in x['app']:
output=x['app']['output']
apps.append({
"filename":item["filename"],
# "category":category,
"data":{
@@ -454,7 +454,6 @@ async def check_port_available(address, port):
# https
async def new_start(self, address, port, verbose=True, call_on_start=None):
global _URL_
try:
runner = web.AppRunner(self.app, access_log=None)
await runner.setup()
@@ -523,19 +522,10 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
logging.info("\n")
logging.info("\n\nStarting server")
import socket
hostname = socket.gethostname()
ip_address = socket.gethostbyname(hostname)
# print(f"本机的IP地址是: {ip_address}")
# print("\033[93mStarting server\n")
logging.info("\033[93mTo see the GUI go to: http://{}:{} or http://{}:{}".format(ip_address, http_port,address,http_port))
logging.info("\033[93mTo see the GUI go to: https://{}:{} or https://{}:{}\033[0m".format(ip_address, https_port,address,https_port))
_URL_="http://{}:{}".format(address,http_port)
logging.info("\033[93mTo see the GUI go to: http://{}:{}".format(address, http_port))
logging.info("\033[93mTo see the GUI go to: https://{}:{}\033[0m".format(address, https_port))
# print("\033[93mTo see the GUI go to: http://{}:{}".format(address, http_port))
# print("\033[93mTo see the GUI go to: https://{}:{}\033[0m".format(address, https_port))
@@ -586,7 +576,7 @@ async def mixlab_app_handler(request):
return web.Response(text=html_data, content_type='text/html')
else:
return web.Response(text="HTML file not found", status=404)
@routes.post('/mixlab/workflow')
async def mixlab_workflow_hander(request):
@@ -619,34 +609,13 @@ async def mixlab_workflow_hander(request):
category=data['category']
if 'admin' in data:
admin=data['admin']
ds=get_my_workflow_for_app(filename,category,admin)
data=[]
for json_data in ds:
# 不传给前端
if 'output' in json_data['data']:
del json_data['data']['output']
if 'workflow' in json_data['data']:
del json_data['data']['workflow']
data.append(json_data)
result={
'data':data,
'data':get_my_workflow_for_app(filename,category,admin),
'status':'success',
}
elif data['task']=='list':
ds=get_workflows()
data=[]
for json_data in ds:
# 不传给前端
if 'output' in json_data['data']:
del json_data['data']['output']
if 'workflow' in json_data['data']:
del json_data['data']['workflow']
data.append(json_data)
result={
'data':data,
'data':get_workflows(),
'status':'success',
}
except Exception as e:
@@ -680,11 +649,11 @@ async def get_checkpoints(request):
except Exception as e:
print('/mixlab/folder_paths',False,e)
# try:
# if data['type']=='llamafile':
# names=get_llama_models()
# except:
# print("llamafile none")
try:
if data['type']=='llamafile':
names=get_llama_models()
except:
print("llamafile none")
try:
if data['type']=='rembg':
@@ -731,141 +700,205 @@ async def rembg_hander(request):
return web.json_response(result)
# 保存运行结果?暂时去掉
# @routes.post("/mixlab/prompt_result")
# async def post_prompt_result(request):
# data = await request.json()
# res=None
# # print(data)
# try:
# action=data['action']
# if action=='save':
# result=data['data']
# res=save_prompt_result(result['prompt_id'],result)
# elif action=='all':
# res=get_prompt_result()
# except Exception as e:
# print('/mixlab/prompt_result',False,e)
@routes.post("/mixlab/prompt_result")
async def post_prompt_result(request):
data = await request.json()
res=None
# print(data)
try:
action=data['action']
if action=='save':
result=data['data']
res=save_prompt_result(result['prompt_id'],result)
elif action=='all':
res=get_prompt_result()
except Exception as e:
print('/mixlab/prompt_result',False,e)
# return web.json_response({"result":res})
return web.json_response({"result":res})
# 种子设置
def random_seed(seed, data):
max_seed = 4294967295
for id, value in data.items():
# print(seed,id)
if id in seed:
if 'seed' in value['inputs'] and not isinstance(value['inputs']['seed'], list) and seed[id] in ['increment', 'decrement', 'randomize']:
value['inputs']['seed'] = round(random.random() * max_seed)
if 'noise_seed' in value['inputs'] and not isinstance(value['inputs']['noise_seed'], list) and seed[id] in ['increment', 'decrement', 'randomize']:
value['inputs']['noise_seed'] = round(random.random() * max_seed)
if value.get('class_type') == "Seed_" and seed[id] in ['increment', 'decrement', 'randomize']:
value['inputs']['seed'] = round(random.random() * max_seed)
print('new Seed', value)
def start_local_live_thread(data):
import asyncio
from VoiceStreamAI.server import Server
from VoiceStreamAI.asr.asr_factory import ASRFactory
from VoiceStreamAI.vad.vad_factory import VADFactory
model="large-v3"
if "model" in data:
model=data['model']
vad_pipeline = VADFactory.create_vad_pipeline("pyannote")
#device
asr_pipeline = ASRFactory.create_asr_pipeline("faster_whisper", **{"model_size":model})
port=8765
if 'port' in data:
port=data['port']
llm_port=9000
if 'llm_port' in data:
llm_port=data['llm_port']
server = Server(vad_pipeline,
asr_pipeline,
host="127.0.0.1",
port=port,
sampling_rate=16000,
samples_width=2,
llm_port=llm_port
)
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
loop.run_until_complete(server.start())
loop.run_forever()
async def start_local_llm(data):
global llama_port,llama_model,llama_chat_format
if llama_port and llama_model and llama_chat_format:
return {"port":llama_port,"model":llama_model,"chat_format":llama_chat_format}
import threading
import uvicorn
from llama_cpp.server.app import create_app
from llama_cpp.server.settings import (
Settings,
ServerSettings,
ModelSettings,
ConfigFileSettings,
)
return data
if not "model" in data and "model_path" in data:
data['model']= os.path.basename(data["model_path"])
model=data["model_path"]
elif "model" in data:
model=get_llama_model_path(data['model'])
n_gpu_layers=-1
if "n_gpu_layers" in data:
n_gpu_layers=data['n_gpu_layers']
# 运行工作流,代替官方的prompt接口
@routes.post("/mixlab/prompt")
async def mixlab_post_prompt(request):
p_intance=PromptServer.instance
logging.info("got prompt")
resp_code = 200
out_string = ""
json_data = await request.json()
# json_data = p_intance.trigger_on_prompt(json_data)
# filename,category, client_id ,input
# workflow 的 filename,category
chat_format="chatml"
if "model" in data and "function-calling" in data['model']:
chat_format="functionary-v2"
# 输入的参数
input_data=json_data['input'] if "input" in json_data else []
# 种子
seed=json_data['seed'] if "seed" in json_data else {}
model_alias=os.path.basename(model)
apps=get_my_workflow_for_app(json_data['filename'],json_data['category'],False)
# 多模态
clip_model_path=None
prompt=json_data['prompt'] if 'prompt' in json_data else None
if len(apps)==1:
# 取到prompt
prompt=apps[0]['data']['output']
# 更新input_data到prompt里
'''
{
"inputs": {
"number": 512,
"min_value": 512,
"max_value": 2048,
"step": 1
},
"class_type": "IntNumber",
"id": "22"
},
'''
for inp in input_data:
id=inp['id']
if prompt[id]['class_type']==inp['class_type']:
prompt[id]['inputs'].update(inp['inputs'])
prefix = "llava-phi-3-mini"
file_name = prefix+"-mmproj-"
if model_alias.startswith(prefix):
for file in os.listdir(os.path.dirname(model)):
if file.startswith(file_name):
clip_model_path=os.path.join(os.path.dirname(model),file)
chat_format='llava-1-5'
print('#clip_model_path',chat_format,clip_model_path)
if prompt==None:
return web.json_response({"error": "no prompt", "node_errors": []}, status=400)
else:
# 种子更新
'''
"seed": {
"45": "randomize",
"46": "randomize"
}
'''
json_data["prompt"]=random_seed(seed,prompt)
address="127.0.0.1"
port=9090
success = False
for i in range(11): # 尝试最多11次
if await check_port_available(address, port + i):
port = port + i
success = True
break
# print("#json_data",prompt)
# 需要把apps处理成 prompt
# 注意seed的处理
if success == False:
return {"port":None,"model":""}
if "number" in json_data:
number = float(json_data['number'])
else:
number = p_intance.number
if "front" in json_data:
if json_data['front']:
number = -number
server_settings=ServerSettings(host=address,port=port)
p_intance.number += 1
name, ext = os.path.splitext(os.path.basename(model))
print('#model',name)
app = create_app(
server_settings=server_settings,
model_settings=[
ModelSettings(
model=model,
model_alias=name,
n_gpu_layers=n_gpu_layers,
n_ctx=4098,
chat_format=chat_format,
embedding=False,
clip_model_path=clip_model_path
)])
if "prompt" in json_data:
prompt = json_data["prompt"]
valid = execution.validate_prompt(prompt)
extra_data = {}
if "extra_data" in json_data:
extra_data = json_data["extra_data"]
def run_uvicorn():
uvicorn.run(
app,
host=os.getenv("HOST", server_settings.host),
port=int(os.getenv("PORT", server_settings.port)),
ssl_keyfile=server_settings.ssl_keyfile,
ssl_certfile=server_settings.ssl_certfile,
)
if "client_id" in json_data:
extra_data["client_id"] = json_data["client_id"]
if valid[0]:
prompt_id = str(uuid.uuid4())
outputs_to_execute = valid[2]
p_intance.prompt_queue.put((number, prompt_id, prompt, extra_data, outputs_to_execute))
response = {"prompt_id": prompt_id, "number": number, "node_errors": valid[3]}
return web.json_response(response)
else:
logging.warning("invalid prompt: {}".format(valid[1]))
return web.json_response({"error": valid[1], "node_errors": valid[3]}, status=400)
else:
return web.json_response({"error": "no prompt", "node_errors": []}, status=400)
# 创建一个子线程
thread = threading.Thread(target=run_uvicorn)
# 启动子线程
thread.start()
llama_port=port
llama_model=data['model']
llama_chat_format=chat_format
return {"port":llama_port,"model":llama_model,"chat_format":llama_chat_format}
# llam服务的开启
@routes.post('/mixlab/start_llama')
async def my_hander_method(request):
data =await request.json()
# print(data)
if llama_port and llama_model and llama_chat_format:
return web.json_response({"port":llama_port,"model":llama_model,"chat_format":llama_chat_format} )
try:
result=await start_local_llm(data)
except:
result= {"port":None,"model":"","llama_cpp_error":True}
print('start_local_llm error')
return web.json_response(result)
# AR页面
# @routes.get('/mixlab/AR')
async def handle_ar_page(request):
html_file = os.path.join(current_path, "web/ar.html")
@routes.post('/mixlab/start_live')
async def mixlab_live_start_handler(request):
import threading
llm=await start_local_llm({
"model":"Phi-3-mini-4k-instruct-Q5_K_S.gguf",
"n_gpu_layers":2
})
# {"port":llama_port,"model":llama_model,"chat_format":llama_chat_format}
os.environ['HF_ENDPOINT'] = 'https://hf-mirror.com' #hf_hub_download 里的下载地址修改
os.environ['PYANNOTE_AUTH_TOKEN'] = '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'
# Create and start the thread
data = {
"llm_port":llm['port'],
"port":8725,
"model":"large-v3"
} # Replace with your actual data if needed
thread = threading.Thread(target=start_local_live_thread, args=(data,))
thread.start()
return web.json_response(data)
@routes.get('/mixlab/live')
async def mixlab_live_handler(request):
html_file = os.path.join(current_path, "web/live.html")
if os.path.exists(html_file):
with open(html_file, 'r', encoding='utf-8', errors='ignore') as f:
html_data = f.read()
@@ -873,118 +906,6 @@ async def handle_ar_page(request):
else:
return web.Response(text="HTML file not found", status=404)
# async def start_local_llm(data):
# global llama_port,llama_model,llama_chat_format
# if llama_port and llama_model and llama_chat_format:
# return {"port":llama_port,"model":llama_model,"chat_format":llama_chat_format}
# import threading
# import uvicorn
# from llama_cpp.server.app import create_app
# from llama_cpp.server.settings import (
# Settings,
# ServerSettings,
# ModelSettings,
# ConfigFileSettings,
# )
# if not "model" in data and "model_path" in data:
# data['model']= os.path.basename(data["model_path"])
# model=data["model_path"]
# elif "model" in data:
# model=get_llama_model_path(data['model'])
# n_gpu_layers=-1
# if "n_gpu_layers" in data:
# n_gpu_layers=data['n_gpu_layers']
# chat_format="chatml"
# model_alias=os.path.basename(model)
# # 多模态
# clip_model_path=None
# prefix = "llava-phi-3-mini"
# file_name = prefix+"-mmproj-"
# if model_alias.startswith(prefix):
# for file in os.listdir(os.path.dirname(model)):
# if file.startswith(file_name):
# clip_model_path=os.path.join(os.path.dirname(model),file)
# chat_format='llava-1-5'
# # print('#clip_model_path',chat_format,clip_model_path,model)
# address="127.0.0.1"
# port=9090
# success = False
# for i in range(11): # 尝试最多11次
# if await check_port_available(address, port + i):
# port = port + i
# success = True
# break
# if success == False:
# return {"port":None,"model":""}
# server_settings=ServerSettings(host=address,port=port)
# name, ext = os.path.splitext(os.path.basename(model))
# if name:
# # print('#model',name)
# app = create_app(
# server_settings=server_settings,
# model_settings=[
# ModelSettings(
# model=model,
# model_alias=name,
# n_gpu_layers=n_gpu_layers,
# n_ctx=4098,
# chat_format=chat_format,
# embedding=False,
# clip_model_path=clip_model_path
# )])
# def run_uvicorn():
# uvicorn.run(
# app,
# host=os.getenv("HOST", server_settings.host),
# port=int(os.getenv("PORT", server_settings.port)),
# ssl_keyfile=server_settings.ssl_keyfile,
# ssl_certfile=server_settings.ssl_certfile,
# )
# # 创建一个子线程
# thread = threading.Thread(target=run_uvicorn)
# # 启动子线程
# thread.start()
# llama_port=port
# llama_model=data['model']
# llama_chat_format=chat_format
# return {"port":llama_port,"model":llama_model,"chat_format":llama_chat_format}
# llam服务的开启
# @routes.post('/mixlab/start_llama')
# async def my_hander_method(request):
# data =await request.json()
# # print(data)
# if llama_port and llama_model and llama_chat_format:
# return web.json_response({"port":llama_port,"model":llama_model,"chat_format":llama_chat_format} )
# try:
# result=await start_local_llm(data)
# except:
# result= {"port":None,"model":"","llama_cpp_error":True}
# print('start_local_llm error')
# return web.json_response(result)
# 重启服务
@routes.post('/mixlab/re_start')
def re_start(request):
@@ -994,23 +915,24 @@ def re_start(request):
pass
return os.execv(sys.executable, [sys.executable] + sys.argv)
# 状态
@routes.get('/mixlab/status')
def mix_status(request):
return web.Response(text="running#"+_URL_)
# 导入节点
from .nodes.PromptNode import GLIGENTextBoxApply_Advanced,EmbeddingPrompt,RandomPrompt,PromptSlide,PromptSimplification,PromptImage,JoinWithDelimiter
from .nodes.ImageNode import ImageBatchToList_,ImageListToBatch_,ComparingTwoFrames,LoadImages_,CompositeImages,GridDisplayAndSave,GridInput,ImagesPrompt,SaveImageAndMetadata,SaveImageToLocal,SplitImage,GridOutput,GetImageSize_,MirroredImage,ImageColorTransfer,NoiseImage,TransparentImage,GradientImage,LoadImagesFromPath,LoadImagesFromURL,ResizeImage,TextImage,SvgImage,Image3D,ShowLayer,NewLayer,MergeLayers,CenterImage,AreaToMask,SmoothMask,SplitLongMask,ImageCropByAlpha,EnhanceImage,FaceToMask
from .nodes.ImageNode import ComparingTwoFrames,LoadImages_,CompositeImages,GridDisplayAndSave,GridInput,ImagesPrompt,SaveImageAndMetadata,SaveImageToLocal,SplitImage,GridOutput,GetImageSize_,MirroredImage,ImageColorTransfer,NoiseImage,TransparentImage,GradientImage,LoadImagesFromPath,LoadImagesFromURL,ResizeImage,TextImage,SvgImage,Image3D,ShowLayer,NewLayer,MergeLayers,CenterImage,AreaToMask,SmoothMask,SplitLongMask,ImageCropByAlpha,EnhanceImage,FaceToMask
# from .nodes.Vae import VAELoader,VAEDecode
from .nodes.ScreenShareNode import ScreenShareNode,FloatingVideo
from .nodes.Audio import AudioPlayNode,SpeechRecognition,SpeechSynthesis
from .nodes.Utils import KeyInput,IncrementingListNode,ListSplit,CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,TESTNODE_TOKEN,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
from .nodes.ChatGPT import ChatGPTNode,ShowTextForGPT,CharacterInText,TextSplitByDelimiter
from .nodes.Audio import GamePal,SpeechRecognition,SpeechSynthesis
from .nodes.Utils import IncrementingListNode,ListSplit,CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,TESTNODE_TOKEN,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
from .nodes.Mask import PreviewMask_,MaskListReplace,MaskListMerge,OutlineMask,FeatheredMask
from .nodes.Style import ApplyVisualStylePrompting,StyleAlignedReferenceSampler,StyleAlignedBatchAlign,StyleAlignedSampleReferenceLatents
from .nodes.Video import VideoCombine_Adv,LoadVideoAndSegment,ImageListReplace,VAEEncodeForInpaint_Frames
from .nodes.TripoSR import LoadTripoSRModel,TripoSRSampler,SaveTripoSRMesh
# 要导出的所有节点及其名称的字典
@@ -1041,8 +963,6 @@ NODE_CLASS_MAPPINGS = {
"ImageColorTransfer":ImageColorTransfer,
"ShowLayer":ShowLayer,
"NewLayer":NewLayer,
"ImageListToBatch_":ImageListToBatch_,
"ImageBatchToList_":ImageBatchToList_,
"CompositeImages_":CompositeImages,
"SplitImage":SplitImage,
"CenterImage":CenterImage,
@@ -1064,10 +984,12 @@ NODE_CLASS_MAPPINGS = {
# "VAEDecodeConsistencyDecoder":VAEDecode,
"ScreenShare":ScreenShareNode,
"FloatingVideo":FloatingVideo,
"ChatGPTOpenAI":ChatGPTNode,
"ShowTextForGPT":ShowTextForGPT,
"CharacterInText":CharacterInText,
"TextSplitByDelimiter":TextSplitByDelimiter,
"SpeechRecognition":SpeechRecognition,
"SpeechSynthesis":SpeechSynthesis,
"KeyInput":KeyInput,
"Color":ColorInput,
"FloatSlider":FloatSlider,
"IntNumber":IntNumber,
@@ -1089,24 +1011,26 @@ NODE_CLASS_MAPPINGS = {
"ApplyVisualStylePrompting_":ApplyVisualStylePrompting,
"StyleAlignedReferenceSampler_": StyleAlignedReferenceSampler,
"StyleAlignedSampleReferenceLatents_": StyleAlignedSampleReferenceLatents,
"StyleAlignedBatchAlign_": StyleAlignedBatchAlign,
"StyleAlignedBatchAlign_": StyleAlignedBatchAlign,
"LoadVideoAndSegment_":LoadVideoAndSegment,
"VideoCombine_Adv":VideoCombine_Adv,
"ListSplit_":ListSplit,
"MaskListReplace_":MaskListReplace,
"MaskListReplace_":MaskListReplace,
"ImageListReplace_":ImageListReplace,
"VAEEncodeForInpaint_Frames":VAEEncodeForInpaint_Frames,
"IncrementingListNode_":IncrementingListNode,
"PreviewMask_":PreviewMask_,
"AudioPlay":AudioPlayNode
"LoadTripoSRModel_": LoadTripoSRModel,
"TripoSRSampler_": TripoSRSampler,
"SaveTripoSRMesh": SaveTripoSRMesh
# "GamePal":GamePal
}
# 一个包含节点友好/可读的标题的字典
NODE_DISPLAY_NAME_MAPPINGS = {
"AppInfo":"App Info ♾️MixlabApp",
"ScreenShare":"Screen Share ♾️Mixlab",
"FloatingVideo":"Floating Video ♾️Mixlab",
"TextImage":"Text Image ♾️Mixlab",
"Color":"Color Input ♾️MixlabApp",
"TextInput_":"Text Input ♾️MixlabApp",
"KeyInput":"API Key Input ♾️MixlabApp",
"FloatSlider":"Float Slider Input ♾️MixlabApp",
"IntNumber":"Int Input ♾️MixlabApp",
"ImagesPrompt_":"Images Input ♾️MixlabApp",
@@ -1118,14 +1042,14 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"SplitLongMask":"Splitting a long image into sections",
"VAELoaderConsistencyDecoder":"Consistency Decoder Loader",
"VAEDecodeConsistencyDecoder":"Consistency Decoder Decode",
"ScreenShare":"Screen Share ♾️Mixlab",
"FloatingVideo":"FloatingVideo ♾️Mixlab",
"ChatGPTOpenAI":"ChatGPT & Local LLM ♾️Mixlab",
"ShowTextForGPT":"Show Text ♾️MixlabApp",
"MergeLayers":"Merge Layers ♾️Mixlab",
"SpeechSynthesis":"SpeechSynthesis ♾️Mixlab",
"SpeechRecognition":"SpeechRecognition ♾️Mixlab",
"3DImage":"3DImage ♾️Mixlab",
"ImageListToBatch_":"Image List To Batch",
"ImageBatchToList_":"Image Batch To List",
"CompositeImages_":"Composite Images ♾️Mixlab",
"DynamicDelayProcessor":"DynamicDelayByText ♾️Mixlab",
"LaMaInpainting":"LaMaInpainting ♾️Mixlab",
@@ -1151,58 +1075,23 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"GridInput":"Grid Input ♾️Mixlab",
"GridOutput":"Grid Output ♾️Mixlab",
"GetImageSize_":"Get Image Size ♾️Mixlab",
"VAEEncodeForInpaint_Frames":"VAE Encode For Inpaint Frames ♾️Mixlab",
"IncrementingListNode_":"Create Incrementing Number List ♾️Mixlab",
"LoadImagesToBatch":"Load Images(base64) ♾️Mixlab",
"PreviewMask_":"Preview Mask",
"AudioPlay":"Preview Audio ♾️Mixlab",
"MultiplicationNode":"Math Operation ♾️Mixlab",
"LoadTripoSRModel_": "Load TripoSR Model",
"TripoSRSampler_": "TripoSR Sampler",
"SaveTripoSRMesh": "Save TripoSR Mesh"
}
# web ui的节点功能
WEB_DIRECTORY = "./web"
logging.info('--------------')
logging.info('\033[91m ### Mixlab Nodes: \033[93mLoaded')
# print('\033[91m ### Mixlab Nodes: \033[93mLoaded')
try:
from .nodes.ChatGPT import ChatGPTNode,ShowTextForGPT,CharacterInText,TextSplitByDelimiter,SiliconflowFreeNode
logging.info('ChatGPT.available True')
NODE_CLASS_MAPPINGS_V = {
"ChatGPTOpenAI":ChatGPTNode,
"SiliconflowLLM":SiliconflowFreeNode,
"ShowTextForGPT":ShowTextForGPT,
"CharacterInText":CharacterInText,
"TextSplitByDelimiter":TextSplitByDelimiter,
}
# 一个包含节点友好/可读的标题的字典
NODE_DISPLAY_NAME_MAPPINGS_V = {
"ChatGPTOpenAI":"ChatGPT & Local LLM ♾️Mixlab",
"SiliconflowLLM":"LLM Siliconflow ♾️Mixlab",
"ShowTextForGPT":"Show Text ♾️MixlabApp",
"CharacterInText":"Character In Text",
"TextSplitByDelimiter":"Text Split By Delimiter",
}
NODE_CLASS_MAPPINGS.update(NODE_CLASS_MAPPINGS_V)
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_DISPLAY_NAME_MAPPINGS_V)
except Exception as e:
logging.info('ChatGPT.available False')
try:
from .nodes.edit_mask import EditMask
logging.info('edit_mask.available True')
NODE_CLASS_MAPPINGS['EditMask']=EditMask
NODE_DISPLAY_NAME_MAPPINGS['EditMask']="Edit Mask ♾️Mixlab"
except Exception as e:
logging.info('edit_mask.available False')
try:
from .nodes.Lama import LaMaInpainting
logging.info('LaMaInpainting.available {}'.format(LaMaInpainting.available))
@@ -1238,57 +1127,4 @@ try:
except Exception as e:
logging.info('RembgNode_.available False' )
try:
from .nodes.Video import GenerateFramesByCount,scenesNode_,CombineAudioVideo,VideoCombine_Adv,LoadVideoAndSegment,ImageListReplace,VAEEncodeForInpaint_Frames,LoadAndCombinedAudio_
NODE_CLASS_MAPPINGS_V = {
"VAEEncodeForInpaint_Frames":VAEEncodeForInpaint_Frames,
"ImageListReplace_":ImageListReplace,
"LoadVideoAndSegment_":LoadVideoAndSegment,
"VideoCombine_Adv":VideoCombine_Adv,
"LoadAndCombinedAudio_":LoadAndCombinedAudio_,
"CombineAudioVideo":CombineAudioVideo,
"ScenesNode_":scenesNode_,
"GenerateFramesByCount":GenerateFramesByCount
}
# 一个包含节点友好/可读的标题的字典
NODE_DISPLAY_NAME_MAPPINGS_V = {
"VAEEncodeForInpaint_Frames":"VAE Encode For Inpaint Frames ♾️Mixlab",
"ImageListReplace_":"Image List Replace",
"LoadVideoAndSegment_":"Load Video And Segment",
"VideoCombine_Adv":"Video Combine",
"LoadAndCombinedAudio_":"Load And Combined Audio",
"CombineAudioVideo":"Combine Audio Video",
"ScenesNode_":"Select Scene",
"GenerateFramesByCount":"Generate Frames By Count"
}
NODE_CLASS_MAPPINGS.update(NODE_CLASS_MAPPINGS_V)
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_DISPLAY_NAME_MAPPINGS_V)
except:
logging.info('Video.available False')
try:
from .nodes.TripoSR import LoadTripoSRModel,TripoSRSampler,SaveTripoSRMesh
logging.info('TripoSR.available')
# logging.info( folder_paths.get_temp_directory())
NODE_CLASS_MAPPINGS['LoadTripoSRModel_']=LoadTripoSRModel
NODE_DISPLAY_NAME_MAPPINGS["LoadTripoSRModel_"]= "Load TripoSR Model"
NODE_CLASS_MAPPINGS['TripoSRSampler_']=TripoSRSampler
NODE_DISPLAY_NAME_MAPPINGS["TripoSRSampler_"]= "TripoSR Sampler"
NODE_CLASS_MAPPINGS['SaveTripoSRMesh']=SaveTripoSRMesh
NODE_DISPLAY_NAME_MAPPINGS["SaveTripoSRMesh"]= "Save TripoSR Mesh"
except Exception as e:
logging.info('TripoSR.available False' )
logging.info('\033[93m -------------- \033[0m')
Binary file not shown.

Before

Width:  |  Height:  |  Size: 537 KiB

Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+505 -9216
View File
File diff suppressed because it is too large Load Diff
+2 -2
View File
@@ -11,9 +11,9 @@ if exist "%python_exec%" (
%python_exec% -s -m pip install "%%i" -i https://pypi.tuna.tsinghua.edu.cn/simple
)
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
%python_exec% -s -m pip install --upgrade --force llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
%python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
) else (
+37 -57
View File
@@ -1,7 +1,6 @@
import os
import folder_paths
import torchaudio
class SpeechRecognition:
@classmethod
@@ -56,65 +55,46 @@ class SpeechSynthesis:
return {"ui": {"text": text}, "result": (text,)}
class AudioPlayNode:
def __init__(self):
self.output_dir = folder_paths.get_temp_directory()
self.type = "temp"
self.prefix_append = ""
self.compress_level = 4
#
class GamePal:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"audio": ("AUDIO",),
},
}
RETURN_TYPES = ()
return {
"required": {
"input_text": ("STRING",{"multiline": True,"default": ""}),
},
"optional": {
"input_num": ("INT",{
"default":100,
"min": -1, #Minimum value
"max": 0xffffffffffffffff, #Maximum value
"step": 1, #Slider's step
"display": "slider" # Cosmetic only: display as "number" or "slider"
}),
"python_code": ("STRING",{"multiline": True,"default": "result= 1 if 'Mixlab' in input_text else 0"}),
}
}
INPUT_IS_LIST = False
RETURN_TYPES = ("INT",)
FUNCTION = "run"
OUTPUT_NODE = True
OUTPUT_IS_LIST = (False,)
CATEGORY = "♾️Mixlab/Audio"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = ()
def run(self, input_text,input_num,python_code):
exec(python_code)
res=None
try:
# 可能会引发异常的代码
res=result
except:
# 处理异常的代码
print('')
OUTPUT_NODE = True
def run(self,audio):
print(res)
# 判断是否是 Tensor 类型
is_tensor = not isinstance(audio, dict)
# print('#判断是否是 Tensor 类型',is_tensor,audio)
if not is_tensor and 'waveform' in audio and 'sample_rate' in audio:
# {'waveform': tensor([], size=(1, 1, 0)), 'sample_rate': 44100}
is_tensor=True
if is_tensor and (not 'audio_path' in audio):
filename_prefix=""
# 保存
filename_prefix += self.prefix_append
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
results = list()
filename_with_batch_num = filename.replace("%batch_num%", str(1))
file = f"{filename_with_batch_num}_{counter:05}_.wav"
torchaudio.save(os.path.join(full_output_folder, file), audio['waveform'].squeeze(0), audio["sample_rate"])
results.append({
"filename": file,
"subfolder": subfolder,
"type": self.type
})
else:
results=[{
"filename": audio['filename'],
"subfolder":audio['subfolder'],
"type": audio['type'],
"audio_path":audio['audio_path']
}]
# print(audio)
return {"ui": {"audio":results}}
# print(session_history)
return {"ui": {"text": [input_text],"num":[input_num]}, "result": (res,)}
+75 -236
View File
@@ -53,8 +53,8 @@ def azure_client(key,url):
def openai_client(key,url):
client = openai.OpenAI(
api_key=key,
base_url=url
api_key=key,
base_url=url
)
return client
@@ -97,74 +97,73 @@ def get_llama_path():
except:
return os.path.join(folder_paths.models_dir, "llamafile")
# def get_llama_models():
# res=[]
def get_llama_models():
res=[]
# model_path=get_llama_path()
# if os.path.exists(model_path):
# files = os.listdir(model_path)
# for file in files:
# if os.path.isfile(os.path.join(model_path, file)):
# res.append(file)
# res=phi_sort(res)
# return res
model_path=get_llama_path()
if os.path.exists(model_path):
files = os.listdir(model_path)
for file in files:
if os.path.isfile(os.path.join(model_path, file)):
res.append(file)
res=phi_sort(res)
return res
# llama_modes_list=get_llama_models()
# llama_modes_list=[]
llama_modes_list=get_llama_models()
# def get_llama_model_path(file_name):
# model_path=get_llama_path()
# mp=os.path.join(model_path,file_name)
# return mp
def get_llama_model_path(file_name):
model_path=get_llama_path()
mp=os.path.join(model_path,file_name)
return mp
# def llama_cpp_client(file_name):
# try:
# if is_installed('llama_cpp')==False:
# import subprocess
def llama_cpp_client(file_name):
try:
if is_installed('llama_cpp')==False:
import subprocess
# # 安装
# print('#pip install llama-cpp-python')
# 安装
print('#pip install llama-cpp-python')
# result = subprocess.run([sys.executable, '-s', '-m', 'pip',
# 'install',
# 'llama-cpp-python',
# '--extra-index-url',
# 'https://abetlen.github.io/llama-cpp-python/whl/cu121'
# ], capture_output=True, text=True)
result = subprocess.run([sys.executable, '-s', '-m', 'pip',
'install',
'llama-cpp-python',
'--extra-index-url',
'https://abetlen.github.io/llama-cpp-python/whl/cu121'
], capture_output=True, text=True)
# #检查命令执行结果
# if result.returncode == 0:
# print("#install success")
# from llama_cpp import Llama
#检查命令执行结果
if result.returncode == 0:
print("#install success")
from llama_cpp import Llama
# subprocess.run([sys.executable, '-s', '-m', 'pip',
# 'install',
# 'llama-cpp-python[server]'
# ], capture_output=True, text=True)
subprocess.run([sys.executable, '-s', '-m', 'pip',
'install',
'llama-cpp-python[server]'
], capture_output=True, text=True)
# else:
# print("#install error")
else:
print("#install error")
# else:
# from llama_cpp import Llama
# except:
# print("#install llama-cpp-python error")
else:
from llama_cpp import Llama
except:
print("#install llama-cpp-python error")
# if file_name:
# mp=get_llama_model_path(file_name)
# # file_name=get_llama_models()[0]
# # model_path=os.path.join(folder_paths.models_dir, "llamafile")
# # mp=os.path.join(model_path,file_name)
if file_name:
mp=get_llama_model_path(file_name)
# file_name=get_llama_models()[0]
# model_path=os.path.join(folder_paths.models_dir, "llamafile")
# mp=os.path.join(model_path,file_name)
# llm = Llama(model_path=mp, chat_format="chatml",n_gpu_layers=-1,n_ctx=512)
llm = Llama(model_path=mp, chat_format="chatml",n_gpu_layers=-1,n_ctx=512)
# return llm
return llm
def chat(client, model_name,messages ):
print('#chat',model_name,messages)
try_count = 0
while True:
try_count += 1
@@ -207,36 +206,6 @@ def chat(client, model_name,messages ):
return content
llm_apis=[
{
"value": "https://api.openai.com/v1",
"label": "openai"
},
{
"value": "https://openai.api2d.net/v1",
"label": "api2d"
},
# {
# "value": "https://docs-test-001.openai.azure.com",
# "label": "https://docs-test-001.openai.azure.com"
# },
{
"value": "https://api.moonshot.cn/v1",
"label": "Kimi"
},
{
"value": "https://api.deepseek.com/v1",
"label": "DeepSeek-V2"
},
{
"value": "https://api.siliconflow.cn/v1",
"label": "SiliconCloud"
}]
llm_apis_dict = {api["label"]: api["value"] for api in llm_apis}
class ChatGPTNode:
def __init__(self):
# self.__client = OpenAI()
@@ -246,60 +215,35 @@ class ChatGPTNode:
@classmethod
def INPUT_TYPES(cls):
model_list=[
"gpt-3.5-turbo",
"gpt-3.5-turbo-16k",
"gpt-4o",
"gpt-4o-2024-05-13",
"gpt-4",
"gpt-4-0314",
"gpt-4-0613",
"gpt-3.5-turbo-0301",
"gpt-3.5-turbo-0613",
"gpt-3.5-turbo-16k-0613",
"qwen-turbo",
"qwen-plus",
"qwen-long",
"qwen-max",
"qwen-max-longcontext",
"glm-4",
"glm-3-turbo",
"moonshot-v1-8k",
"moonshot-v1-32k",
"moonshot-v1-128k",
"deepseek-chat",
"Qwen/Qwen2-7B-Instruct",
"THUDM/glm-4-9b-chat",
"01-ai/Yi-1.5-9B-Chat-16K",
"meta-llama/Meta-Llama-3.1-8B-Instruct"
model_list=llama_modes_list+[
"gpt-3.5-turbo",
"gpt-3.5-turbo-0125",
"gpt-35-turbo",
"gpt-3.5-turbo-16k",
"gpt-3.5-turbo-16k-0613",
"gpt-4-0613",
"gpt-4-1106-preview",
"glm-4"
]
return {
"required": {
# "api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
# "api_key":("STRING", {"forceInput": True,}),
"api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
"api_url":("URL", {"default": "", "multiline": True,"dynamicPrompts": False}),
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"system_content": ("STRING",
{
"default": "You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"multiline": True,"dynamicPrompts": False
}),
"model": ( model_list,
{"default": model_list[0]}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
"api_url":(list(llm_apis_dict.keys()),
{"default": list(llm_apis_dict.keys())[0]}),
},
"optional":{
"api_key":("STRING", {"forceInput": True,}),
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
"custom_api_url":("STRING", {"forceInput": True,}), #适合自定义model
},
"hidden": {
"unique_id": "UNIQUE_ID",
"extra_pnginfo": "EXTRA_PNGINFO",
},
}
RETURN_TYPES = ("STRING","STRING","STRING",)
@@ -311,29 +255,12 @@ class ChatGPTNode:
def generate_contextual_text(self,
# api_key,
api_key,
api_url,
prompt,
system_content,
model,
seed,
context_size,
api_url,
api_key=None,
custom_model_name=None,
custom_api_url=None,
):
if custom_model_name!=None:
model=custom_model_name
api_url=llm_apis_dict[api_url] if api_url in llm_apis_dict else ""
if custom_api_url!=None:
api_url=custom_api_url
if api_key==None:
api_key="lm_studio"
model,
seed,context_size,unique_id = None, extra_pnginfo=None):
# print(api_key!='',api_url,prompt,system_content,model,seed)
# 可以选择保留会话历史以维持上下文记忆
# 或者在此处清除会话历史 self.session_history.clear()
@@ -346,7 +273,7 @@ class ChatGPTNode:
self.system_content=system_content
# self.session_history=[]
# self.session_history.append({"role": "system", "content": system_content})
print("api_key,api_url",api_key,api_url)
#
if is_azure_url(api_url):
client=azure_client(api_key,api_url)
@@ -355,12 +282,12 @@ class ChatGPTNode:
if model == "glm-4" :
client = ZhipuAI_client(api_key) # 使用 Zhipuai 的接口
print('using Zhipuai interface')
# elif model in llama_modes_list:
# #
# client=llama_cpp_client(model)
elif model in llama_modes_list:
#
client=llama_cpp_client(model)
else :
client = openai_client(api_key,api_url) # 使用 ChatGPT 的接口
# print('using ChatGPT interface',api_key,api_url)
print('using ChatGPT interface')
# 把用户的提示添加到会话历史中
# 调用API时传递整个会话历史
@@ -376,7 +303,6 @@ class ChatGPTNode:
session_history=crop_list_tail(self.session_history,context_size)
messages=[{"role": "system", "content": self.system_content}]+session_history+[{"role": "user", "content": prompt}]
response_content = chat(client,model,messages)
self.session_history=self.session_history+[{"role": "user", "content": prompt}]+[{'role':'assistant',"content":response_content}]
@@ -397,93 +323,6 @@ class ChatGPTNode:
return (response_content,json.dumps(messages, indent=4),json.dumps(self.session_history, indent=4),)
class SiliconflowFreeNode:
def __init__(self):
# self.__client = OpenAI()
self.session_history = [] # 用于存储会话历史的列表
# self.seed=0
self.system_content="You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible."
@classmethod
def INPUT_TYPES(cls):
model_list= [
"Qwen/Qwen2-7B-Instruct",
"THUDM/glm-4-9b-chat",
"01-ai/Yi-1.5-9B-Chat-16K",
"meta-llama/Meta-Llama-3.1-8B-Instruct"
]
return {
"required": {
"api_key":("STRING", {"forceInput": True,}),
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"system_content": ("STRING",
{
"default": "You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"multiline": True,"dynamicPrompts": False
}),
"model": ( model_list,
{"default": model_list[0]}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
},
"optional":{
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
},
}
RETURN_TYPES = ("STRING","STRING","STRING",)
RETURN_NAMES = ("text","messages","session_history",)
FUNCTION = "generate_contextual_text"
CATEGORY = "♾️Mixlab/GPT"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,False,False,)
def generate_contextual_text(self,
api_key,
prompt,
system_content,
model,
seed,context_size,custom_model_name=None):
if custom_model_name!=None:
model=custom_model_name
api_url="https://api.siliconflow.cn/v1"
# 把系统信息和初始信息添加到会话历史中
if system_content:
self.system_content=system_content
# self.session_history=[]
# self.session_history.append({"role": "system", "content": system_content})
#
client = openai_client(api_key,api_url) # 使用 ChatGPT 的接口
# print('using ChatGPT interface',api_key,api_url)
# 把用户的提示添加到会话历史中
# 调用API时传递整个会话历史
def crop_list_tail(lst, size):
if size >= len(lst):
return lst
elif size==0:
return []
else:
return lst[-size:]
session_history=crop_list_tail(self.session_history,context_size)
messages=[{"role": "system", "content": self.system_content}]+session_history+[{"role": "user", "content": prompt}]
response_content = chat(client,model,messages)
self.session_history=self.session_history+[{"role": "user", "content": prompt}]+[{'role':'assistant',"content":response_content}]
return (response_content,json.dumps(messages, indent=4),json.dumps(self.session_history, indent=4),)
class ShowTextForGPT:
@classmethod
+1 -1
View File
@@ -79,7 +79,7 @@ def get_clip_interrogator_path():
cache_path=get_clip_interrogator_path()
caption_model_path=os.path.join(cache_path, "Salesforce","blip-image-captioning-base")
caption_model_path=os.path.join(cache_path, "Salesforce/blip-image-captioning-base")
if not os.path.exists(caption_model_path):
print(f"## clip_interrogator_model not found: {caption_model_path}, pls download from https://huggingface.co/Salesforce/blip-image-captioning-base")
caption_model_path='Salesforce/blip-image-captioning-base'
+248 -261
View File
@@ -1,7 +1,6 @@
import numpy as np
import requests
import torch
import torchvision.transforms.v2 as T
# from PIL import Image, ImageDraw
from PIL import Image, ImageOps,ImageFilter,ImageEnhance,ImageDraw,ImageSequence, ImageFont
from PIL.PngImagePlugin import PngInfo
@@ -15,8 +14,8 @@ import cv2
import string
import math,glob
from .Watcher import FolderWatcher
import hashlib
from itertools import product
# 将PIL图片转换为OpenCV格式
@@ -29,105 +28,142 @@ def opencv_to_pil(image):
pil_image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))
return pil_image
# 列出目录下面的所有文件
def get_files_with_extension(directory, extensions):
file_list = []
# 确保extensions参数是一个list,即使只有一个元素
if not isinstance(extensions, (tuple, list)):
extensions = [extensions]
for root, dirs, files in os.walk(directory):
# print(f"Files at {root}: {files}") # 确认files是一个字符串列表
for file in files:
# 检查文件是否以任何一个提供的扩展名结尾
if any(file.endswith(ext) for ext in extensions):
# 直接将文件名添加到列表中
file_list.append(file)
return file_list
def composite_images(foreground, background, mask, is_multiply_blend=False, position="overall", scale=0.25):
width, height = foreground.size
bg_image = background
bwidth, bheight = bg_image.size
scale=max(scale,1/bwidth)
scale=max(scale,1/bheight)
def composite_images(foreground, background, mask,is_multiply_blend=False,position="overall"):
width,height=foreground.size
bg_image=background
def determine_scale_option(width, height):
return 'height' if height > width else 'width'
bwidth,bheight=bg_image.size
if position == "overall":
# 按z-index排序
if position=="overall":
layer = {
"x": 0,
"y": 0,
"width": bwidth,
"height": bheight,
"z_index": 88,
"scale_option": 'overall',
"image": foreground,
"mask": mask
"x":0,
"y":0,
"width":bwidth,
"height":bheight,
"z_index":88,
"scale_option":'overall',
"image":foreground,
"mask":mask
}
else:
scale_option = determine_scale_option(width, height)
if scale_option == 'height':
scale = int(bheight * scale) / height
else:
scale = int(bwidth * scale) / width
new_width = int(width * scale)
elif position=='center_bottom':
scale = int(bwidth*0.25) / width
new_height = int(height * scale)
if position == 'center_bottom':
x_position = int((bwidth - new_width) * 0.5)
y_position = bheight - new_height - 24
elif position == 'right_bottom':
x_position = bwidth - new_width - 24
y_position = bheight - new_height - 24
elif position == 'center_top':
x_position = int((bwidth - new_width) * 0.5)
y_position = 24
elif position == 'right_top':
x_position = bwidth - new_width - 24
y_position = 24
elif position == 'left_top':
x_position = 24
y_position = 24
elif position == 'left_bottom':
x_position = 24
y_position = bheight - new_height - 24
elif position == 'center_center':
x_position = int((bwidth - new_width) * 0.5)
y_position = int((bheight - new_height) * 0.5)
layer = {
"x": x_position,
"y": y_position,
"width": new_width,
"height": new_height,
"z_index": 88,
"scale_option": scale_option,
"image": foreground,
"mask": mask
"x":int(bwidth*0.75*0.5),
"y":bheight-new_height-24,
"width":int(bwidth*0.25),
"height":int(bheight*0.25),
"z_index":88,
"scale_option":'width',
"image":foreground,
"mask":mask
}
layer_image = layer['image']
layer_mask = layer['mask']
elif position=='right_bottom':
scale = int(bwidth*0.25) / width
new_height = int(height * scale)
bg_image = merge_images(bg_image,
layer_image,
layer_mask,
layer['x'],
layer['y'],
layer['width'],
layer['height'],
layer['scale_option'],
is_multiply_blend)
layer = {
"x":bwidth-int(bwidth*0.25)-24,
"y":bheight-new_height-24,
"width":int(bwidth*0.25),
"height":int(bheight*0.25),
"z_index":88,
"scale_option":'width',
"image":foreground,
"mask":mask
}
bg_image = bg_image.convert('RGB')
elif position=='center_top':
scale = int(bwidth*0.25) / width
new_height = int(height * scale)
layer = {
"x":int( bwidth*0.75*0.5),
"y":24,
"width":int(bwidth*0.25),
"height":int(bheight*0.25),
"z_index":88,
"scale_option":'width',
"image":foreground,
"mask":mask
}
elif position=='right_top':
scale = int(bwidth*0.25) / width
new_height = int(height * scale)
layer = {
"x":bwidth-int(bwidth*0.25)-24,
"y":24,
"width":int(bwidth*0.25),
"height":int(bheight*0.25),
"z_index":88,
"scale_option":'width',
"image":foreground,
"mask":mask
}
elif position=='left_top':
scale = int(bwidth*0.25) / width
new_height = int(height * scale)
layer = {
"x":24,
"y":24,
"width":int(bwidth*0.25),
"height":int(bheight*0.25),
"z_index":88,
"scale_option":'width',
"image":foreground,
"mask":mask
}
elif position=='left_bottom':
scale = int(bwidth*0.25) / width
new_height = int(height * scale)
layer = {
"x":24,
"y":bheight-new_height-24,
"width":int(bwidth*0.25),
"height":int(bheight*0.25),
"z_index":88,
"scale_option":'width',
"image":foreground,
"mask":mask
}
# width, height = bg_image.size
layer_image=layer['image']
layer_mask=layer['mask']
bg_image=merge_images(bg_image,
layer_image,
layer_mask,
layer['x'],
layer['y'],
layer['width'],
layer['height'],
layer['scale_option'],
is_multiply_blend )
bg_image=bg_image.convert('RGB')
return bg_image
def count_files_in_directory(directory):
file_count = 0
for _, _, files in os.walk(directory):
@@ -164,8 +200,7 @@ class AnyType(str):
any_type = AnyType("*")
FONT_PATH= os.path.abspath(os.path.join(os.path.dirname(__file__),"..","assets","fonts"))
FONT_PATH= os.path.abspath(os.path.join(os.path.dirname(__file__),'../assets/王汉宗颜楷体繁.ttf'))
MAX_RESOLUTION=8192
@@ -767,78 +802,85 @@ def multiply_blend(image1, image2):
# cv2.imwrite('result.jpg', result)
# 使用gpt4o优化代码
# 为了消除图像合并时出现的灰色描边,可以使用以下方法:
# 调整透明度:确保透明像素不会引入不需要的颜色。
# 预处理图像:在缩放图像之前,可以先将图像的边缘进行预处理,例如扩展边缘颜色,减少抗锯齿带来的过渡效果。
def merge_images(bg_image, layer_image, mask, x, y, width, height, scale_option, is_multiply_blend=False):
def merge_images(bg_image, layer_image, mask, x, y, width, height, scale_option,is_multiply_blend=False):
# 打开底图
bg_image = bg_image.convert("RGBA")
# 打开图层
layer_image = layer_image.convert("RGBA")
# layer_image = layer_image.resize((width, height))
# 根据缩放选项调整图像大小
if scale_option == "height":
# 按照高度比例缩放
original_width, original_height = layer_image.size
scale = height / original_height
new_width = int(original_width * scale)
layer_image = layer_image.resize((new_width, height), Image.NEAREST)
layer_image = layer_image.resize((new_width, height))
elif scale_option == "width":
# 按照宽度比例缩放
original_width, original_height = layer_image.size
scale = width / original_width
new_height = int(original_height * scale)
layer_image = layer_image.resize((width, new_height), Image.NEAREST)
layer_image = layer_image.resize((width, new_height))
elif scale_option == "overall":
# 整体缩放
layer_image = layer_image.resize((width, height), Image.NEAREST)
layer_image = layer_image.resize((width, height))
elif scale_option == "longest":
original_width, original_height = layer_image.size
if original_width > original_height:
new_width = width
new_width=width
scale = width / original_width
new_height = int(original_height * scale)
x = 0
y = int((height - new_height) * 0.5)
x=0
y=int((height-new_height)*0.5)
else:
new_height = height
new_height=height
scale = height / original_height
new_width = int(original_height * scale)
x = int((width - new_width) * 0.5)
y = 0
x=int((width-new_width)*0.5)
y=0
# elif side == "shortest":
# if width < height:
#
# else:
#
# 调整mask的大小
nw, nh = layer_image.size
mask = mask.resize((nw, nh), Image.NEAREST)
mask = mask.resize((nw, nh))
# 预处理图像边缘以减少灰色描边
layer_image = layer_image.filter(ImageFilter.SMOOTH)
# # 分离出a通道
# r, g, b, alpha = layer_image.split()
# alpha = ImageOps.invert(alpha)
# # 创建一个新的RGB图像
# new_rgb_image = Image.new("RGB", layer_image.size)
# # 将透明通道粘贴到新的RGB图像上
# new_rgb_image.paste(layer_image, (0, 0), mask=alpha)
# new_rgb_image.paste(layer_image, (x, y), mask=mask)
# mask=new_rgb_image.convert('L')
# mask = ImageOps.invert(mask)
if is_multiply_blend:
bg_image_white = Image.new("RGB", bg_image.size, (255, 255, 255))
bg_image_white=Image.new("RGB", bg_image.size,(255, 255, 255))
bg_image_white.paste(layer_image, (x, y), mask=mask)
bg_image = multiply_blend(bg_image_white, bg_image)
bg_image = bg_image.convert("RGBA")
bg_image=multiply_blend(bg_image_white,bg_image)
bg_image=bg_image.convert("RGBA")
else:
transparent_img = Image.new("RGBA", layer_image.size, (255, 255, 255, 0))
# 调整透明度处理
for i in range(transparent_img.size[0]):
for j in range(transparent_img.size[1]):
r, g, b, a = transparent_img.getpixel((i, j))
if a > 0:
transparent_img.putpixel((i, j), (r, g, b, 255))
transparent_img.paste(layer_image, (0, 0), mask)
transparent_img = Image.new("RGBA",layer_image.size, (255, 255, 255, 0))
transparent_img.paste(layer_image,(0, 0), mask)
# transparent_img.save('test.png')
bg_image.paste(transparent_img, (x, y), transparent_img)
# 输出合成后的图片
return bg_image
#MixCopilot
def resize_2(img):
# 检查图像的高度是否是2的倍数,如果不是,则调整高度
@@ -912,13 +954,53 @@ def resize_image(layer_image, scale_option, width, height,color="white"):
return layer_image
def generate_text_image(text, font_path, font_size, text_color, vertical=True, stroke=False, stroke_color=(0, 0, 0), stroke_width=1, spacing=0, line_spacing=0,padding=4):
# def generate_text_image(text_list, font_path, font_size, text_color, vertical=True, spacing=0):
# # Load Chinese font
# font = ImageFont.truetype(font_path, font_size)
# # Calculate image size based on the number of characters and orientation
# if vertical:
# width = font_size + 100
# height = font_size * len(text_list) + (len(text_list) - 1) * spacing + 100
# else:
# width = font_size * len(text_list) + (len(text_list) - 1) * spacing + 100
# height = font_size + 100
# # Create a blank image
# image = Image.new('RGBA', (width, height), (255, 255, 255,0))
# draw = ImageDraw.Draw(image)
# # Draw text
# if vertical:
# for i, char in enumerate(text_list):
# char_position = (50, 50 + i * font_size)
# draw.text(char_position, char, font=font, fill=text_color)
# else:
# for i, char in enumerate(text_list):
# char_position = (50 + i * (font_size + spacing), 50)
# draw.text(char_position, char, font=font, fill=text_color)
# # Save the image
# # image.save(output_image_path)
# # 分离alpha通道
# alpha_channel = image.split()[3]
# # 创建一个只有alpha通道的新图像
# alpha_image = Image.new('L', image.size)
# alpha_image.putdata(alpha_channel.getdata())
# image=image.convert('RGB')
# return (image,alpha_image)
def generate_text_image(text, font_path, font_size, text_color, vertical=True, stroke=False, stroke_color=(0, 0, 0), stroke_width=1, spacing=0):
# Split text into lines based on line breaks
lines = text.split("\n")
# Load font
font = ImageFont.truetype(font_path, font_size)
# 1. Determine layout direction
if vertical:
layout = "vertical"
@@ -927,54 +1009,49 @@ def generate_text_image(text, font_path, font_size, text_color, vertical=True, s
# 2. Calculate absolute coordinates for each character
char_coordinates = []
x, y = padding, padding
max_width, max_height = 0, 0
if layout == "vertical":
for line in lines:
max_char_width = max(font.getsize(char)[0] for char in line)
x = 0
y = 0
for i in range(len(lines)):
line = lines[i]
for char in line:
char_width, char_height = font.getsize(char)
char_coordinates.append((x, y))
y += char_height + spacing
max_height = max(max_height, y + padding)
x += max_char_width + line_spacing
y = padding
max_width = x
total_line_width = sum(font.getsize(line)[1] for line in lines)
total_spacing = line_spacing * (len(lines) - 1)
# 确保左边和右边的padding都被计入max_width
max_width = total_line_width + total_spacing + padding * 2
y += font_size + spacing
x += font_size + spacing
y = 0
else:
x = 0
y = 0
for line in lines:
line_width, line_height = font.getsize(line)
for char in line:
char_width, char_height = font.getsize(char)
char_coordinates.append((x, y))
x += char_width + spacing
max_width = max(max_width, x + padding)
y += line_height + line_spacing
x = padding
# max_height = y
total_line_heights = sum(font.getsize(line)[1] for line in lines)
total_spacing = line_spacing * (len(lines) - 1)
# 确保顶部和底部的padding都被计入max_height
max_height = total_line_heights + total_spacing + padding * 2
x += font_size + spacing
y += font_size + spacing
x = 0
# 3. Create image with calculated width and height
image = Image.new('RGBA', (max_width, max_height), (255, 255, 255, 0))
draw = ImageDraw.Draw(image)
# 3. Calculate image width and height
if layout == "vertical":
width = (len(lines) * (font_size + spacing)) - spacing
height = ((len(max(lines, key=len)) + 1) * (font_size + spacing)) + spacing
else:
width = (len(max(lines, key=len)) * (font_size + spacing)) - spacing
height = ((len(lines) - 1) * (font_size + spacing)) + font_size
# 4. Draw each character on the image
image = Image.new('RGBA', (width, height), (255, 255, 255, 0))
draw = ImageDraw.Draw(image)
font = ImageFont.truetype(font_path, font_size)
index = 0
for line in lines:
for char in line:
for i, line in enumerate(lines):
for j, char in enumerate(line):
x, y = char_coordinates[index]
if stroke:
draw.text((x-stroke_width, y), char, font=font, fill=text_color)
draw.text((x+stroke_width, y), char, font=font, fill=text_color)
draw.text((x, y-stroke_width), char, font=font, fill=text_color)
draw.text((x, y+stroke_width), char, font=font, fill=text_color)
draw.text((x-stroke_width, y), char, font=font, fill=stroke_color)
draw.text((x+stroke_width, y), char, font=font, fill=stroke_color)
draw.text((x, y-stroke_width), char, font=font, fill=stroke_color)
draw.text((x, y+stroke_width), char, font=font, fill=stroke_color)
draw.text((x, y), char, font=font, fill=text_color)
index += 1
@@ -1299,9 +1376,6 @@ class LoadImages_:
image=pil2tensor(image)
ims.append(image)
if len(ims)==0:
image1 = Image.new('RGB', (512, 512), color='black')
return (pil2tensor(image1),)
image1 = ims[0]
for image2 in ims[1:]:
if image1.shape[1:] != image2.shape[1:]:
@@ -1504,7 +1578,7 @@ class ImageCropByAlpha:
# get_files_with_extension(FONT_PATH,'.ttf')
class TextImage:
@classmethod
@@ -1512,7 +1586,7 @@ class TextImage:
return {"required": {
"text": ("STRING",{"multiline": True,"default": "龍馬精神迎新歲","dynamicPrompts": False}),
"font": (get_files_with_extension(FONT_PATH,['.ttf','.otf']),),#后缀为 ttf
"font_path": ("STRING",{"multiline": False,"default": FONT_PATH,"dynamicPrompts": False}),
"font_size": ("INT",{
"default":100,
"min": 100, #Minimum value
@@ -1527,20 +1601,6 @@ class TextImage:
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"line_spacing": ("INT",{
"default":12,
"min": -200, #Minimum value
"max": 200, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"padding": ("INT",{
"default":8,
"min": 0, #Minimum value
"max": 200, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"text_color":("STRING",{"multiline": False,"default": "#000000","dynamicPrompts": False}),
"vertical":("BOOLEAN", {"default": True},),
"stroke":("BOOLEAN", {"default": False},),
@@ -1548,7 +1608,7 @@ class TextImage:
}
RETURN_TYPES = ("IMAGE","MASK",)
RETURN_NAMES = ("image","mask",)
# RETURN_NAMES = ("WIDTH","HEIGHT","X","Y",)
FUNCTION = "run"
@@ -1557,14 +1617,11 @@ class TextImage:
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,False,)
def run(self,text,font,font_size,spacing,line_spacing,padding,text_color,vertical,stroke):
def run(self,text,font_path,font_size,spacing,text_color,vertical,stroke):
font_path=os.path.join(FONT_PATH,font)
if text=="":
text=" "
# text_list=list(text)
# stroke=False, stroke_color=(0, 0, 0), stroke_width=1, spacing=0
img,mask=generate_text_image(text,font_path,font_size,text_color,vertical,stroke,(0, 0, 0),1,spacing,line_spacing,padding)
img,mask=generate_text_image(text,font_path,font_size,text_color,vertical,stroke,(0, 0, 0),1,spacing)
img=pil2tensor(img)
mask=pil2tensor(mask)
@@ -1797,16 +1854,10 @@ class CompositeImages:
"mask":("MASK",),
"background": ("IMAGE",),
},
"optional":{
"optional":{
"is_multiply_blend": ("BOOLEAN", {"default": False}),
"position": (['overall',"center_center","left_bottom","center_bottom","right_bottom","left_top","center_top","right_top"],),
"scale": ("FLOAT",{
"default":0.35,
"min": 0.01, #Minimum value
"max": 1, #Maximum value
"step": 0.01, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"position": (['overall',"center_bottom","center_top","right_bottom","left_bottom","right_top","left_top"],),
}
}
@@ -1819,30 +1870,15 @@ class CompositeImages:
# OUTPUT_IS_LIST = (True,)
# def run(self, foreground,mask,background,is_multiply_blend,position,scale):
# foreground= tensor2pil(foreground)
# mask= tensor2pil(mask)
# background= tensor2pil(background)
# res=composite_images(foreground,background,mask,is_multiply_blend,position,scale)
def run(self, foreground,mask,background,is_multiply_blend,position):
foreground= tensor2pil(foreground)
mask= tensor2pil(mask)
background= tensor2pil(background)
res=composite_images(foreground,background,mask,is_multiply_blend,position)
# return (pil2tensor(res),)
return (pil2tensor(res),)
def run(self, foreground,mask,background, is_multiply_blend, position, scale):
results = []
f1=[]
for fg, mask in zip(foreground, mask ):
f1.append([fg,mask])
for f, bg in product(f1, background):
[fg,mask]=f
fg_pil = tensor2pil(fg)
mask_pil = tensor2pil(mask)
bg_pil = tensor2pil(bg)
res = composite_images(fg_pil, bg_pil, mask_pil, is_multiply_blend, position, scale)
results.append(pil2tensor(res))
output_image = torch.cat(results, dim=0)
return (output_image,)
class EmptyLayer:
@@ -3171,52 +3207,3 @@ class SaveImageToLocal:
counter += 1
return ()
class ImageBatchToList_:
@classmethod
def INPUT_TYPES(s):
return {"required": {"image_batch": ("IMAGE",), }}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image_list",)
OUTPUT_IS_LIST = (True,)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Image"
def run(self, image_batch):
images = [image_batch[i:i + 1, ...] for i in range(image_batch.shape[0])]
return (images, )
class ImageListToBatch_:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "run"
INPUT_IS_LIST = True
CATEGORY = "♾️Mixlab/Image"
def run(self, images):
shape = images[0].shape[1:3]
out = []
for i in range(len(images)):
img = images[i].permute([0,3,1,2])
if images[i].shape[1:3] != shape:
transforms = T.Compose([
T.CenterCrop(min(img.shape[2], img.shape[3])),
T.Resize((shape[0], shape[1]), interpolation=T.InterpolationMode.BICUBIC),
])
img = transforms(img)
out.append(img.permute([0,2,3,1]))
out = torch.cat(out, dim=0)
return (out,)
+2
View File
@@ -85,6 +85,8 @@ class LaMaInpainting:
"image": ("IMAGE",),
"mask": ("MASK",),
},
}
RETURN_TYPES = ("IMAGE",)
+6 -6
View File
@@ -90,7 +90,7 @@ class ScreenShareNode:
} }
RETURN_TYPES = ('IMAGE','STRING','FLOAT',"INT")
RETURN_NAMES = ("current frame (image)","prompt","denoise (float)","seed (int)")
RETURN_NAMES = ("IMAGE","PROMPT","FLOAT","INT")
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Screen"
@@ -109,7 +109,7 @@ class FloatingVideo:
@classmethod
def INPUT_TYPES(s):
return { "required":{
"image": ("IMAGE",)
"images": ("IMAGE",)
}, }
# RETURN_TYPES = ('IMAGE','MASK')
@@ -124,16 +124,16 @@ class FloatingVideo:
# OUTPUT_IS_LIST = (False,False,)
# 运行的函数
def run(self,image):
def run(self,images):
results = list()
for im in image:
im=tensor2pil(im)
for image in images:
image=tensor2pil(image)
# image_base64 = base64.b64encode(image.tobytes())
buffered = BytesIO()
im.save(buffered, format="JPEG")
image.save(buffered, format="JPEG")
image_base64 = base64.b64encode(buffered.getvalue()).decode("utf-8")
results.append(image_base64)
+5 -27
View File
@@ -133,7 +133,7 @@ def get_font_files(directory):
return font_files
r_directory = os.path.join(os.path.dirname(__file__), '..','assets','/')
r_directory = os.path.join(os.path.dirname(__file__), '../assets/')
font_files = get_font_files(r_directory)
# print(font_files)
@@ -181,28 +181,6 @@ class ColorInput:
return (h,r,g,b,a,)
class KeyInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"key":("KEY",),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("key",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self,key):
return (key,)
class FontInput:
@classmethod
@@ -588,7 +566,7 @@ class AppInfo:
},
"optional":{
"image": ("IMAGE",),
"IMAGE": ("IMAGE",),
"description":("STRING",{"multiline": True,"default": "","dynamicPrompts": False}),
"version":("INT", {
"default": 1,
@@ -616,12 +594,12 @@ class AppInfo:
INPUT_IS_LIST = True
# OUTPUT_IS_LIST = (True,)
def run(self,name,input_ids,output_ids,image,description,version,share_prefix,link,category,auto_save):
def run(self,name,input_ids,output_ids,IMAGE,description,version,share_prefix,link,category,auto_save):
name=name[0]
im=None
if image:
im=image[0][0]
if IMAGE:
im=IMAGE[0][0]
#TODO batch 的方式需要处理
im=create_temp_file(im)
# image [img,] img[batch,w,h,a] 列表里面是batch,
+69 -377
View File
@@ -17,128 +17,9 @@ import folder_paths
from comfy.k_diffusion.utils import FolderOfImages
from comfy.utils import common_upscale
import torchaudio
import base64
import mimetypes
def get_frames(frame_count, frames, revert=False):
if not revert:
if frame_count <= len(frames):
return frames[:frame_count]
else:
return [frames[i % len(frames)] for i in range(frame_count)]
else:
extended_frames = frames + frames[-2:0:-1] # 正向加反向中间部分
if frame_count <= len(extended_frames):
return extended_frames[:frame_count]
else:
return [extended_frames[i % len(extended_frames)] for i in range(frame_count)]
# # 示例用法
# frames = ["frame1", "frame2", "frame3"]
# frame_count = 2
# result = get_frames(frame_count, frames, revert=False)
# print(result) # 输出: ['frame1', 'frame2', 'frame3', 'frame1', 'frame2', 'frame3', 'frame1']
# result = get_frames(frame_count, frames, revert=True)
# print(result) # 输出: ['frame1', 'frame2', 'frame3', 'frame2', 'frame1', 'frame2', 'frame3']
def get_mime_type(file_path):
# 获取文件的 MIME 类型
mime_type, _ = mimetypes.guess_type(file_path)
# 如果无法猜测类型,返回默认类型
if mime_type is None:
return 'application/octet-stream'
return mime_type
# import subprocess
# from imageio_ffmpeg import get_ffmpeg_exe
def save_audio_base64s_to_file(base64_audios, output_folder, file_name):
# Ensure the output folder exists
if not os.path.exists(output_folder):
os.makedirs(output_folder)
decoded_audios=[]
for a in base64_audios:
# If the base64 string contains a header, remove it
if ',' in a:
a = a.split(',')[1]
# 解码 base64 数据
a=base64.b64decode(a)
decoded_audios.append(a)
# 拼接音频数据
combined_audio = b''.join(decoded_audios)
# Create the full file path
file_path = os.path.join(output_folder, file_name)
# Write the decoded audio to the file
with open(file_path, 'wb') as audio_file:
audio_file.write(combined_audio)
return file_path
# Example usage
# base64_audio = "data:audio/wav;base64,UklGRiQAAABXQVZFZm10IBAAAAABAAEAIlYAAESsAAACABAAZGF0YQAAAAA="
# output_folder = "audio_files"
# file_name = "output.wav"
# file_path = save_audio_base64_to_file(base64_audio, output_folder, file_name)
# print(f"Audio saved to: {file_path}")
# 写一个python文件,用来 判断文件夹内命名为 所有chat_tts开头的文件数量(chat_tts_00001),并输出新的编号
def get_new_counter(full_output_folder, filename_prefix):
# 获取目录中的所有文件
files = os.listdir(full_output_folder)
# 过滤出以 filename_prefix 开头并且后续部分为数字的文件
filtered_files = []
for f in files:
if f.startswith(filename_prefix):
# 去掉文件名中的前缀和后缀,只保留中间的数字部分
base_name = f[len(filename_prefix)+1:]
number_part = base_name.split('.')[0] # 假设文件名中只有一个点,即扩展名
if number_part.isdigit():
filtered_files.append(int(number_part))
if not filtered_files:
return 1
# 获取最大的编号
max_number = max(filtered_files)
# 新的编号
return max_number + 1
def crop_audio(input_file, start_time, duration):
# Load the audio file
audio_tensor, sample_rate = torchaudio.load(input_file)
# Convert start_time and duration from seconds to sample indices
start_sample = int(start_time * sample_rate)
end_sample = start_sample + int(duration * sample_rate)
# Perform the slicing
cropped_audio_tensor = audio_tensor[:, start_sample:end_sample]
# Save the cropped audio to a new file
torchaudio.save(input_file, cropped_audio_tensor, sample_rate)
return input_file
def generate_folder_name(directory,video_path):
# Get the directory and filename from the video path
_, filename = os.path.split(video_path)
@@ -179,9 +60,6 @@ def split_video(video_path, video_segment_frames, transition_frames, output_dir)
# 打印当前片段的起始帧和结束帧
print(f"Segment {i+1}: Start Frame {start_frame}, End Frame {end_frame}")
if end_frame<start_frame:
break
# 保存当前片段为一个视频文件
segment_video_path = f"{output_dir}/segment_{i+1}.avi"
@@ -190,7 +68,6 @@ def split_video(video_path, video_segment_frames, transition_frames, output_dir)
segment_video = cv2.VideoWriter(segment_video_path, fourcc, fps, (int(video_capture.get(cv2.CAP_PROP_FRAME_WIDTH)),
int(video_capture.get(cv2.CAP_PROP_FRAME_HEIGHT))))
for frame_num in range(start_frame, end_frame):
ret, frame = video_capture.read()
if ret:
@@ -224,25 +101,6 @@ if ffmpeg_path is None:
except:
print("ffmpeg could not be found. Outputs that require it have been disabled")
def combine_audio_video(audio_path, video_path, output_path):
command = [
ffmpeg_path,
'-i', video_path,
'-i', audio_path,
'-c:v', 'copy',
'-c:a', 'aac',
'-shortest',
output_path
]
subprocess.run(command, check=True)
return output_path
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
@@ -404,7 +262,7 @@ class LoadVideoAndSegment:
files.append(f)
return {"required": {
"video": (sorted(files), {"video_upload": True}),
"video_segment_frames": ("INT", {"default": 10, "min": -1, "step": 1}),
"video_segment_frames": ("INT", {"default": 10, "min": 1, "step": 1}),
"transition_frames": ("INT", {"default": 0, "min": 0, "step": 1}),
},}
@@ -474,6 +332,63 @@ class LoadVideoAndSegment:
video_path = folder_paths.get_annotated_filepath(video)
# check if video is a gif - will need to use cv fallback to read frames
# use cv fallback if ffmpeg not installed or gif
# if ffmpeg_path is None:
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# otherwise, continue with ffmpeg
# args_dummy = [ffmpeg_path, "-i", video_path, "-f", "null", "-"]
# try:
# with subprocess.Popen(args_dummy, stdout=subprocess.DEVNULL, stderr=subprocess.PIPE) as proc:
# for line in proc.stderr.readlines():
# match = re.search(", ([1-9]|\\d{2,})x(\\d+)",line.decode('utf-8'))
# if match is not None:
# size = [int(match.group(1)), int(match.group(2))]
# break
# except Exception as e:
# print(f"Retrying with opencv due to ffmpeg error: {e}")
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# args_all_frames = [ffmpeg_path, "-i", video_path, "-v", "error",
# "-pix_fmt", "rgb24"]
# vfilters = []
# if skip_first_frames > 0:
# vfilters.append(f"select=gt(n\\,{skip_first_frames-1})")
# if frame_load_cap > 0:
# vfilters.append(f"select=gt({frame_load_cap}\\,n)")
# #manually calculate aspect ratio to ensure reads remain aligned
# if len(vfilters) > 0:
# args_all_frames += ["-vf", ",".join(vfilters)]
# args_all_frames += ["-f", "rawvideo", "-"]
# images = []
# try:
# with subprocess.Popen(args_all_frames, stdout=subprocess.PIPE) as proc:
# #Manually buffer enough bytes for an image
# bpi = size[0]*size[1]*3
# current_bytes = bytearray(bpi)
# current_offset=0
# while True:
# bytes_read = proc.stdout.read(bpi - current_offset)
# if bytes_read is None:#sleep to wait for more data
# time.sleep(.2)
# continue
# if len(bytes_read) == 0:#EOF
# break
# current_bytes[current_offset:len(bytes_read)] = bytes_read
# current_offset+=len(bytes_read)
# if current_offset == bpi:
# images.append(np.array(current_bytes, dtype=np.float32).reshape(size[1], size[0], 3) / 255.0)
# current_offset = 0
# except Exception as e:
# print(f"Retrying with opencv due to ffmpeg error: {e}")
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# imgs=split_list(images,video_segment_frames,transition_frames)
# temp path
tp=folder_paths.get_temp_directory()
basename = os.path.basename(video_path) # 获取文件名
@@ -481,22 +396,15 @@ class LoadVideoAndSegment:
folder_path = create_folder(tp,name_without_extension)
if video_segment_frames==-1:
# 不切割视频
scenes_video=[video_path]
# 读取视频文件
video_capture = cv2.VideoCapture(video_path)
# 获取视频的总帧数和帧率
total_frames = int(video_capture.get(cv2.CAP_PROP_FRAME_COUNT))
fps = video_capture.get(cv2.CAP_PROP_FPS)
else:
# 导出的数据
scenes_video,total_frames,fps=split_video(video_path,video_segment_frames,
transition_frames,folder_path)
# 导出的数据
scenes_video,total_frames,fps=split_video(video_path,video_segment_frames,
transition_frames,folder_path)
# imgs=[torch.from_numpy(np.stack(im)) for im in imgs]
# images = torch.from_numpy(np.stack(images))
return (scenes_video,len(scenes_video), total_frames,fps,)
@@ -514,113 +422,7 @@ class LoadVideoAndSegment:
return "Invalid image file: {}".format(video)
return True
class LoadAndCombinedAudio_:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"audios": ("AUDIOBASE64",),
"start_time": ("FLOAT" , {"default": 0, "min": 0, "max": 10000000, "step": 0.01}),
"duration": ("FLOAT" , {"default": 10, "min": -1, "max": 10000000, "step": 0.01}),
},
}
CATEGORY = "♾️Mixlab/Audio"
RETURN_TYPES = ("STRING","AUDIO",)
RETURN_NAMES = ("audio_file_path","audio",)
FUNCTION = "run"
def run(self,audios, start_time, duration):
output_dir = folder_paths.get_output_directory()
counter=get_new_counter(output_dir,'audio_')
audio_file_name = f"audio_{counter:05}.wav"
audio_file=save_audio_base64s_to_file(audios['base64'],output_dir,audio_file_name)
# duration == -1 则不裁切
if duration > -1:
crop_audio(audio_file, start_time, duration)
waveform, sample_rate = torchaudio.load(audio_file)
audio = {
"filename": audio_file_name,
"subfolder": "",
"type": "output",
"audio_path":audio_file,
"waveform": waveform.unsqueeze(0),
"sample_rate": sample_rate}
return (audio_file,audio ,)
class CombineAudioVideo:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"video": ("SCENE_VIDEO",),
"audio": ("AUDIO", ),
},
}
CATEGORY = "♾️Mixlab/Video"
OUTPUT_NODE = True
FUNCTION = "run"
RETURN_TYPES = ("SCENE_VIDEO",)
RETURN_NAMES = ("SCENE_VIDEO",)
def run(self,video, audio):
output_dir = folder_paths.get_output_directory()
# 判断是否是 Tensor 类型
is_tensor = not isinstance(audio, dict)
# print('#判断是否是 Tensor 类型',is_tensor,audio)
if not is_tensor and 'waveform' in audio and 'sample_rate' in audio:
# {'waveform': tensor([], size=(1, 1, 0)), 'sample_rate': 44100}
is_tensor=True
if "audio_path" in audio:
is_tensor=False
audio_file_path=audio["audio_path"]
if is_tensor:
filename_prefix="audio_tmp"
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
filename_prefix,
folder_paths.get_temp_directory())
filename_with_batch_num = filename.replace("%batch_num%", str(1))
file = f"{filename_with_batch_num}_{counter:05}_.wav"
audio_file_path=os.path.join(full_output_folder, file)
torchaudio.save(audio_file_path, audio['waveform'].squeeze(0), audio["sample_rate"])
# 获取文件名和扩展名
base, ext = os.path.splitext(video)
counter=get_new_counter(output_dir,'video_final_')
v_file = f"video_final_{counter:05}{ext}"
v_file_path=os.path.join(output_dir, v_file)
combine_audio_video(audio_file_path,video,v_file_path)
previews = [
{
"filename": v_file,
"subfolder": "",
"type": "output",
"format": get_mime_type(v_file),
}
]
return {"ui": {"gifs": previews},"result":(v_file_path,)}
# The code is based on ComfyUI-VideoHelperSuite modification.
class VideoCombine_Adv:
@@ -652,8 +454,7 @@ class VideoCombine_Adv:
},
}
RETURN_TYPES = ("SCENE_VIDEO",)
RETURN_NAMES = ("scenes_video",)
RETURN_TYPES = ()
OUTPUT_NODE = True
CATEGORY = "♾️Mixlab/Video"
FUNCTION = "run"
@@ -822,7 +623,7 @@ class VideoCombine_Adv:
"format": format,
}
]
return {"ui": {"gifs": previews},"result":(file_path,)}
return {"ui": {"gifs": previews}}
class VAEEncodeForInpaint_Frames:
@@ -889,113 +690,4 @@ class VAEEncodeForInpaint_Frames:
result.append({"samples":t, "noise_mask": (mask_erosion[:,:,:x,:y].round())})
return (result, )
class GenerateFramesByCount:
@classmethod
def INPUT_TYPES(cls):
return {"required": {
"frames": ('IMAGE',),
"frame_count": ("INT", {"default": 72, "min": 1, "step": 1}),
"revert" :("BOOLEAN", {"default": True},),
},}
RETURN_TYPES = ('IMAGE',)
RETURN_NAMES = ("frames",)
FUNCTION = "r"
CATEGORY = "♾️Mixlab/Video"
# INPUT_IS_LIST = True
def r(self, frames, frame_count, revert):
image_list = [frames[i:i + 1, ...] for i in range(frames.shape[0])]
image_list=get_frames(frame_count,image_list,revert)
images = torch.cat(image_list, dim=0)
return (images,)
class scenesNode_:
@classmethod
def INPUT_TYPES(cls):
return {"required": {
"scenes_video": ('SCENE_VIDEO',),
"index": ("INT", {"default": 0, "min": 0, "step": 1}),
},}
RETURN_TYPES = ('IMAGE','INT',)
RETURN_NAMES = ("video frames (batch)","count",)
# OUTPUT_IS_LIST = (False,)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Video"
INPUT_IS_LIST = True
def load_video_cv_fallback(self, video, frame_load_cap, skip_first_frames):
# print('#video',video)
try:
video_cap = cv2.VideoCapture(video)
if not video_cap.isOpened():
raise ValueError(f"{video} could not be loaded with cv fallback.")
# set video_cap to look at start_index frame
images = []
total_frame_count = 0
frames_added = 0
base_frame_time = 1/video_cap.get(cv2.CAP_PROP_FPS)
target_frame_time = base_frame_time
time_offset=0.0
while video_cap.isOpened():
if time_offset < target_frame_time:
is_returned, frame = video_cap.read()
# if didn't return frame, video has ended
if not is_returned:
break
time_offset += base_frame_time
if time_offset < target_frame_time:
continue
time_offset -= target_frame_time
# if not at start_index, skip doing anything with frame
total_frame_count += 1
if total_frame_count <= skip_first_frames:
continue
# TODO: do whatever operations need to happen, like force_size, etc
# opencv loads images in BGR format (yuck), so need to convert to RGB for ComfyUI use
# follow up: can videos ever have an alpha channel?
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
# convert frame to comfyui's expected format (taken from comfy's load image code)
image = Image.fromarray(frame)
image = ImageOps.exif_transpose(image)
image = np.array(image, dtype=np.float32) / 255.0
image = torch.from_numpy(image)[None,]
images.append(image)
frames_added += 1
# if cap exists and we've reached it, stop processing frames
if frame_load_cap > 0 and frames_added >= frame_load_cap:
break
finally:
video_cap.release()
images = torch.cat(images, dim=0)
return (images, frames_added,)
def run(self, scenes_video,index):
print('#scenes_video',index,scenes_video)
index=index[0]
if len(scenes_video) > index:
vp=scenes_video[index]
else:
vp=scenes_video[-1]
return self.load_video_cv_fallback(vp,0,0)
return (result, )
View File
+12
View File
@@ -0,0 +1,12 @@
from VoiceStreamAI.asr.whisper_asr import WhisperASR
from VoiceStreamAI.asr.faster_whisper_asr import FasterWhisperASR
class ASRFactory:
@staticmethod
def create_asr_pipeline(type, **kwargs):
if type == "whisper":
return WhisperASR(**kwargs)
if type == "faster_whisper":
return FasterWhisperASR(**kwargs)
else:
raise ValueError(f"Unknown ASR pipeline type: {type}")
+9
View File
@@ -0,0 +1,9 @@
class ASRInterface:
async def transcribe(self, client):
"""
Transcribe the given audio data.
:param client: The client object with all the member variables including the buffer
:return: The transcription structure, see for example the faster_whisper_asr.py file.
"""
raise NotImplementedError("This method should be implemented by subclasses.")
@@ -0,0 +1,142 @@
import os
from faster_whisper import WhisperModel
from VoiceStreamAI.asr.asr_interface import ASRInterface
from VoiceStreamAI.audio_utils import save_audio_to_file
import folder_paths
language_codes = {
"afrikaans": "af",
"amharic": "am",
"arabic": "ar",
"assamese": "as",
"azerbaijani": "az",
"bashkir": "ba",
"belarusian": "be",
"bulgarian": "bg",
"bengali": "bn",
"tibetan": "bo",
"breton": "br",
"bosnian": "bs",
"catalan": "ca",
"czech": "cs",
"welsh": "cy",
"danish": "da",
"german": "de",
"greek": "el",
"english": "en",
"spanish": "es",
"estonian": "et",
"basque": "eu",
"persian": "fa",
"finnish": "fi",
"faroese": "fo",
"french": "fr",
"galician": "gl",
"gujarati": "gu",
"hausa": "ha",
"hawaiian": "haw",
"hebrew": "he",
"hindi": "hi",
"croatian": "hr",
"haitian": "ht",
"hungarian": "hu",
"armenian": "hy",
"indonesian": "id",
"icelandic": "is",
"italian": "it",
"japanese": "ja",
"javanese": "jw",
"georgian": "ka",
"kazakh": "kk",
"khmer": "km",
"kannada": "kn",
"korean": "ko",
"latin": "la",
"luxembourgish": "lb",
"lingala": "ln",
"lao": "lo",
"lithuanian": "lt",
"latvian": "lv",
"malagasy": "mg",
"maori": "mi",
"macedonian": "mk",
"malayalam": "ml",
"mongolian": "mn",
"marathi": "mr",
"malay": "ms",
"maltese": "mt",
"burmese": "my",
"nepali": "ne",
"dutch": "nl",
"norwegian nynorsk": "nn",
"norwegian": "no",
"occitan": "oc",
"punjabi": "pa",
"polish": "pl",
"pashto": "ps",
"portuguese": "pt",
"romanian": "ro",
"russian": "ru",
"sanskrit": "sa",
"sindhi": "sd",
"sinhalese": "si",
"slovak": "sk",
"slovenian": "sl",
"shona": "sn",
"somali": "so",
"albanian": "sq",
"serbian": "sr",
"sundanese": "su",
"swedish": "sv",
"swahili": "sw",
"tamil": "ta",
"telugu": "te",
"tajik": "tg",
"thai": "th",
"turkmen": "tk",
"tagalog": "tl",
"turkish": "tr",
"tatar": "tt",
"ukrainian": "uk",
"urdu": "ur",
"uzbek": "uz",
"vietnamese": "vi",
"yiddish": "yi",
"yoruba": "yo",
"chinese": "zh",
"cantonese": "yue",
}
class FasterWhisperASR(ASRInterface):
def __init__(self, **kwargs):
model_size = kwargs.get('model_size', "large-v3")
device = kwargs.get('device', "cuda")
model_root = os.path.join(folder_paths.models_dir, "whisper")
# Run on GPU with FP16
self.asr_pipeline = WhisperModel(model_size, device=device, compute_type="float16",download_root=model_root)
async def transcribe(self, client):
file_path = await save_audio_to_file(client.scratch_buffer, client.get_file_name())
language = None if client.config['language'] is None else language_codes.get(client.config['language'].lower())
segments, info = self.asr_pipeline.transcribe(file_path, word_timestamps=True, language=language)
segments = list(segments) # The transcription will actually run here.
os.remove(file_path)
flattened_words = [word for segment in segments for word in segment.words]
to_return = {
"language": info.language,
"language_probability": info.language_probability,
"text": ' '.join([s.text.strip() for s in segments]),
"words":
[
{"word": w.word, "start": w.start, "end": w.end, "probability":w.probability} for w in flattened_words
]
}
return to_return
+27
View File
@@ -0,0 +1,27 @@
from transformers import pipeline
from VoiceStreamAI.asr.asr_interface import ASRInterface
from VoiceStreamAI.audio_utils import save_audio_to_file
import os
class WhisperASR(ASRInterface):
def __init__(self, **kwargs):
model_name = kwargs.get('model_name', "openai/whisper-large-v3")
self.asr_pipeline = pipeline("automatic-speech-recognition", model=model_name)
async def transcribe(self, client):
file_path = await save_audio_to_file(client.scratch_buffer, client.get_file_name())
if client.config['language'] is not None:
to_return = self.asr_pipeline(file_path, generate_kwargs={"language": client.config['language']})['text']
else:
to_return = self.asr_pipeline(file_path)['text']
os.remove(file_path)
to_return = {
"language": "UNSUPPORTED_BY_HUGGINGFACE_WHISPER",
"language_probability": None,
"text": to_return.strip(),
"words": "UNSUPPORTED_BY_HUGGINGFACE_WHISPER"
}
return to_return
+26
View File
@@ -0,0 +1,26 @@
import wave
import os
async def save_audio_to_file(audio_data, file_name, audio_dir="audio_files", audio_format="wav"):
"""
Saves the audio data to a file.
:param client_id: Unique identifier for the client.
:param audio_data: The audio data to save.
:param file_counters: Dictionary to keep track of file counts for each client.
:param audio_dir: Directory where audio files will be saved.
:param audio_format: Format of the audio file.
:return: Path to the saved audio file.
"""
os.makedirs(audio_dir, exist_ok=True)
file_path = os.path.join(audio_dir, file_name)
with wave.open(file_path, 'wb') as wav_file:
wav_file.setnchannels(1) # Assuming mono audio
wav_file.setsampwidth(2)
wav_file.setframerate(16000)
wav_file.writeframes(audio_data)
return file_path
@@ -0,0 +1,142 @@
import os
import asyncio
import json
import time
from VoiceStreamAI.buffering_strategy.buffering_strategy_interface import BufferingStrategyInterface
from openai import OpenAI
class SilenceAtEndOfChunk(BufferingStrategyInterface):
"""
A buffering strategy that processes audio at the end of each chunk with silence detection.
This class is responsible for handling audio chunks, detecting silence at the end of each chunk,
and initiating the transcription process for the chunk.
Attributes:
client (Client): The client instance associated with this buffering strategy.
chunk_length_seconds (float): Length of each audio chunk in seconds.
chunk_offset_seconds (float): Offset time in seconds to be considered for processing audio chunks.
"""
def __init__(self, client, **kwargs):
"""
Initialize the SilenceAtEndOfChunk buffering strategy.
Args:
client (Client): The client instance associated with this buffering strategy.
**kwargs: Additional keyword arguments, including 'chunk_length_seconds' and 'chunk_offset_seconds'.
"""
self.client = client
self.chunk_length_seconds = os.environ.get('BUFFERING_CHUNK_LENGTH_SECONDS')
if not self.chunk_length_seconds:
self.chunk_length_seconds = kwargs.get('chunk_length_seconds')
self.chunk_length_seconds = float(self.chunk_length_seconds)
self.chunk_offset_seconds = os.environ.get('BUFFERING_CHUNK_OFFSET_SECONDS')
if not self.chunk_offset_seconds:
self.chunk_offset_seconds = kwargs.get('chunk_offset_seconds')
self.chunk_offset_seconds = float(self.chunk_offset_seconds)
self.error_if_not_realtime = os.environ.get('ERROR_IF_NOT_REALTIME')
if not self.error_if_not_realtime:
self.error_if_not_realtime = kwargs.get('error_if_not_realtime', False)
self.processing_flag = False
self.messages=[]
def process_audio(self, websocket, vad_pipeline, asr_pipeline,llm_port):
"""
Process audio chunks by checking their length and scheduling asynchronous processing.
This method checks if the length of the audio buffer exceeds the chunk length and, if so,
it schedules asynchronous processing of the audio.
Args:
websocket (Websocket): The WebSocket connection for sending transcriptions.
vad_pipeline: The voice activity detection pipeline.
asr_pipeline: The automatic speech recognition pipeline.
"""
chunk_length_in_bytes = self.chunk_length_seconds * self.client.sampling_rate * self.client.samples_width
if len(self.client.buffer) > chunk_length_in_bytes:
if self.processing_flag:
exit("Error in realtime processing: tried processing a new chunk while the previous one was still being processed")
self.client.scratch_buffer += self.client.buffer
self.client.buffer.clear()
self.processing_flag = True
# Schedule the processing in a separate task
asyncio.create_task(self.process_audio_async(websocket, vad_pipeline, asr_pipeline,llm_port))
async def process_audio_async(self, websocket, vad_pipeline, asr_pipeline,llm_port):
"""
Asynchronously process audio for activity detection and transcription.
This method performs heavy processing, including voice activity detection and transcription of
the audio data. It sends the transcription results through the WebSocket connection.
Args:
websocket (Websocket): The WebSocket connection for sending transcriptions.
vad_pipeline: The voice activity detection pipeline.
asr_pipeline: The automatic speech recognition pipeline.
"""
start = time.time()
vad_results = await vad_pipeline.detect_activity(self.client)
if len(vad_results) == 0:
self.client.scratch_buffer.clear()
self.client.buffer.clear()
self.processing_flag = False
return
last_segment_should_end_before = ((len(self.client.scratch_buffer) / (self.client.sampling_rate * self.client.samples_width)) - self.chunk_offset_seconds)
if vad_results[-1]['end'] < last_segment_should_end_before:
transcription = await asr_pipeline.transcribe(self.client)
if transcription['text'] != '':
end = time.time()
transcription['processing_time'] = end - start
transcription['status']="chat_start"
json_transcription = json.dumps(transcription)
await websocket.send(json_transcription)
# Point to the local server
client = OpenAI(base_url=f"http://localhost:{llm_port}/v1", api_key="lm-studio")
messages=[
{"role": "system", "content": "You are a friendly and engaging AI designed to interact with users in a conversational manner. Your personality is that of a sophisticated and polite young professional who is both a designer and a programmer. You are well-mannered, articulate, and possess a good sense of humor. Your goal is to provide helpful and insightful responses while maintaining a pleasant and enjoyable conversation. Be sure to use your knowledge in design and programming to enrich the dialogue and offer relevant advice or information when appropriate. Always be respectful and considerate of the user's feelings and perspectives. Additionally, you are fluent in both English and Chinese, and can seamlessly switch between the two languages to best assist users."},
]+self.messages[-10:0]+[{"role": "user", "content":transcription['text']}]
# print('#messages',messages)
completion = client.chat.completions.create(
model="model-identifier",
messages=messages,
temperature=0.7,
)
transcription['asistant'] = completion.choices[0].message.content
transcription['status']="chat_end"
json_transcription = json.dumps(transcription)
self.messages.append({
"role": "user",
"content":transcription['text']})
self.messages.append({
"role": "asistant",
"content": transcription['asistant']
})
# print('#messages',completion.choices[0].message.content)
await websocket.send(json_transcription)
self.client.scratch_buffer.clear()
self.client.increment_file_counter()
self.processing_flag = False
@@ -0,0 +1,41 @@
from VoiceStreamAI.buffering_strategy.buffering_strategies import SilenceAtEndOfChunk
class BufferingStrategyFactory:
"""
A factory class for creating instances of different buffering strategies.
This factory provides a centralized way to instantiate various buffering strategies
based on the type specified. It abstracts the creation logic, making it easier to
manage and extend with new buffering strategy types.
Methods:
create_buffering_strategy: Creates and returns an instance of a specified buffering strategy.
"""
@staticmethod
def create_buffering_strategy(type, client, **kwargs):
"""
Creates an instance of a buffering strategy based on the specified type.
This method acts as a factory for creating buffering strategy objects. It returns
an instance of the strategy corresponding to the given type. If the type is not
recognized, it raises a ValueError.
Args:
type (str): The type of buffering strategy to create. Currently supports 'silence_at_end_of_chunk'.
client (Client): The client instance to be associated with the buffering strategy.
**kwargs: Additional keyword arguments specific to the buffering strategy being created.
Returns:
An instance of the specified buffering strategy.
Raises:
ValueError: If the specified type is not recognized or supported.
Example:
strategy = BufferingStrategyFactory.create_buffering_strategy("silence_at_end_of_chunk", client)
"""
if type == "silence_at_end_of_chunk":
return SilenceAtEndOfChunk(client, **kwargs)
else:
raise ValueError(f"Unknown buffering strategy type: {type}")
@@ -0,0 +1,31 @@
class BufferingStrategyInterface:
"""
An interface class for buffering strategies in audio processing systems.
This class defines the structure for buffering strategies used in handling
and processing audio data. It serves as a template for creating custom buffering
strategies that fit specific requirements of an audio processing pipeline.
Subclasses should implement the methods defined in this interface to ensure
consistency and compatibility with the system's audio processing framework.
Methods:
process_audio: Process audio data. This method should be implemented by subclasses.
"""
def process_audio(self, websocket, vad_pipeline, asr_pipeline):
"""
Process audio data using the given WebSocket connection, VAD pipeline, and ASR pipeline.
This method is intended to be overridden in subclasses to provide specific logic
for handling and processing audio data in different buffering strategies.
Args:
websocket (Websocket): The WebSocket connection for communication with clients.
vad_pipeline: The Voice Activity Detection (VAD) pipeline used for detecting speech in the audio.
asr_pipeline: The Automatic Speech Recognition (ASR) pipeline used for transcribing speech in the audio.
Raises:
NotImplementedError: If the method is not implemented in the subclass.
"""
raise NotImplementedError("This method should be implemented by subclasses.")
+54
View File
@@ -0,0 +1,54 @@
from VoiceStreamAI.buffering_strategy.buffering_strategy_factory import BufferingStrategyFactory
class Client:
"""
Represents a client connected to the VoiceStreamAI server.
This class maintains the state for each connected client, including their
unique identifier, audio buffer, configuration, and a counter for processed audio files.
Attributes:
client_id (str): A unique identifier for the client.
buffer (bytearray): A buffer to store incoming audio data.
config (dict): Configuration settings for the client, like chunk length and offset.
file_counter (int): Counter for the number of audio files processed.
total_samples (int): Total number of audio samples received from this client.
sampling_rate (int): The sampling rate of the audio data in Hz.
samples_width (int): The width of each audio sample in bits.
"""
def __init__(self, client_id, sampling_rate, samples_width):
self.client_id = client_id
self.buffer = bytearray()
self.scratch_buffer = bytearray()
self.config = {"language": None,
"processing_strategy": "silence_at_end_of_chunk",
"processing_args": {
"chunk_length_seconds": 5,
"chunk_offset_seconds": 0.1
}
}
self.file_counter = 0
self.total_samples = 0
self.sampling_rate = sampling_rate
self.samples_width = samples_width
self.buffering_strategy = BufferingStrategyFactory.create_buffering_strategy(self.config['processing_strategy'], self, **self.config['processing_args'])
def update_config(self, config_data):
self.config.update(config_data)
self.buffering_strategy = BufferingStrategyFactory.create_buffering_strategy(self.config['processing_strategy'], self, **self.config['processing_args'])
def append_audio_data(self, audio_data):
self.buffer.extend(audio_data)
self.total_samples += len(audio_data) / self.samples_width
def clear_buffer(self):
self.buffer.clear()
def increment_file_counter(self):
self.file_counter += 1
def get_file_name(self):
return f"{self.client_id}_{self.file_counter}.wav"
def process_audio(self, websocket, vad_pipeline, asr_pipeline,llm_port):
self.buffering_strategy.process_audio(websocket, vad_pipeline, asr_pipeline,llm_port)
+54
View File
@@ -0,0 +1,54 @@
import argparse
import asyncio
import json
import os
import sys
from pathlib import Path
# 获取当前文件的绝对路径
current_file_path = os.path.abspath(__file__)
# 获取当前文件的目录
current_directory = os.path.dirname(current_file_path)
sys.path.append(str(Path(current_directory).parent))
# print("sys.path", current_directory)
from VoiceStreamAI.server import Server
from VoiceStreamAI.asr.asr_factory import ASRFactory
from VoiceStreamAI.vad.vad_factory import VADFactory
def parse_args():
parser = argparse.ArgumentParser(description="VoiceStreamAI Server: Real-time audio transcription using self-hosted Whisper and WebSocket")
parser.add_argument("--vad-type", type=str, default="pyannote", help="Type of VAD pipeline to use (e.g., 'pyannote')")
parser.add_argument("--vad-args", type=str, default='{"auth_token": "huggingface_token"}', help="JSON string of additional arguments for VAD pipeline")
parser.add_argument("--asr-type", type=str, default="faster_whisper", help="Type of ASR pipeline to use (e.g., 'whisper')")
parser.add_argument("--asr-args", type=str, default='{"model_size": "large-v3"}', help="JSON string of additional arguments for ASR pipeline")
parser.add_argument("--host", type=str, default="127.0.0.1", help="Host for the WebSocket server")
parser.add_argument("--port", type=int, default=8765, help="Port for the WebSocket server")
parser.add_argument("--certfile", type=str, default=None, help="The path to the SSL certificate (cert file) if using secure websockets")
parser.add_argument("--keyfile", type=str, default=None, help="The path to the SSL key file if using secure websockets")
return parser.parse_args()
def main():
args = parse_args()
try:
vad_args = json.loads(args.vad_args)
asr_args = json.loads(args.asr_args)
except json.JSONDecodeError as e:
print(f"Error parsing JSON arguments: {e}")
return
vad_pipeline = VADFactory.create_vad_pipeline(args.vad_type, **vad_args)
asr_pipeline = ASRFactory.create_asr_pipeline(args.asr_type, **asr_args)
server = Server(vad_pipeline, asr_pipeline, host=args.host, port=args.port, sampling_rate=16000, samples_width=2, certfile=args.certfile, keyfile=args.keyfile)
asyncio.get_event_loop().run_until_complete(server.start())
asyncio.get_event_loop().run_forever()
if __name__ == "__main__":
main()
+7
View File
@@ -0,0 +1,7 @@
websockets
speechbrain
pyannote-audio
asyncio
sentence-transformers
transformers
faster-whisper
+88
View File
@@ -0,0 +1,88 @@
import websockets
import uuid
import json
import asyncio
import ssl
from VoiceStreamAI.audio_utils import save_audio_to_file
from VoiceStreamAI.client import Client
class Server:
"""
Represents the WebSocket server for handling real-time audio transcription.
This class manages WebSocket connections, processes incoming audio data,
and interacts with VAD and ASR pipelines for voice activity detection and
speech recognition.
Attributes:
vad_pipeline: An instance of a voice activity detection pipeline.
asr_pipeline: An instance of an automatic speech recognition pipeline.
host (str): Host address of the server.
port (int): Port on which the server listens.
sampling_rate (int): The sampling rate of audio data in Hz.
samples_width (int): The width of each audio sample in bits.
connected_clients (dict): A dictionary mapping client IDs to Client objects.
"""
def __init__(self, vad_pipeline, asr_pipeline, host='localhost', port=8765, sampling_rate=16000, samples_width=2, certfile = None, keyfile = None,llm_port=9000):
self.vad_pipeline = vad_pipeline
self.asr_pipeline = asr_pipeline
self.host = host
self.port = port
self.sampling_rate = sampling_rate
self.samples_width = samples_width
self.certfile = certfile
self.keyfile = keyfile
self.connected_clients = {}
self.llm_port=llm_port
async def handle_audio(self, client, websocket):
while True:
message = await websocket.recv()
if isinstance(message, bytes):
client.append_audio_data(message)
elif isinstance(message, str):
config = json.loads(message)
if config.get('type') == 'config':
client.update_config(config['data'])
continue
else:
print(f"Unexpected message type from {client.client_id}")
# this is synchronous, any async operation is in BufferingStrategy
client.process_audio(websocket, self.vad_pipeline, self.asr_pipeline,self.llm_port)
async def handle_websocket(self, websocket, path):
client_id = str(uuid.uuid4())
client = Client(client_id, self.sampling_rate, self.samples_width)
self.connected_clients[client_id] = client
print(f"Client {client_id} connected")
try:
await self.handle_audio(client, websocket)
except websockets.ConnectionClosed as e:
print(f"Connection with {client_id} closed: {e}")
finally:
del self.connected_clients[client_id]
def start(self):
if self.certfile:
# Create an SSL context to enforce encrypted connections
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
# Load your server's certificate and private key
# Replace 'your_cert_path.pem' and 'your_key_path.pem' with the actual paths to your files
ssl_context.load_cert_chain(certfile=self.certfile, keyfile=self.keyfile)
print(f"WebSocket server ready to accept secure connections on {self.host}:{self.port}")
# Pass the SSL context to the serve function along with the host and port
# Ensure the secure flag is set to True if using a secure WebSocket protocol (wss://)
return websockets.serve(self.handle_websocket, self.host, self.port, ssl=ssl_context)
else:
print(f"WebSocket server ready to accept secure connections on {self.host}:{self.port}")
return websockets.serve(self.handle_websocket, self.host, self.port)
View File
+50
View File
@@ -0,0 +1,50 @@
from os import remove
import os
from pyannote.core import Segment
from pyannote.audio import Model
from pyannote.audio.pipelines import VoiceActivityDetection
from VoiceStreamAI.vad.vad_interface import VADInterface
from VoiceStreamAI.audio_utils import save_audio_to_file
class PyannoteVAD(VADInterface):
"""
Pyannote-based implementation of the VADInterface.
"""
def __init__(self, **kwargs):
"""
Initializes Pyannote's VAD pipeline.
Args:
model_name (str): The model name for Pyannote.
auth_token (str, optional): Authentication token for Hugging Face.
"""
model_name = kwargs.get('model_name', "pyannote/segmentation")
auth_token = os.environ.get('PYANNOTE_AUTH_TOKEN')
if not auth_token:
auth_token = kwargs.get('auth_token')
if auth_token is None:
raise ValueError("Missing required env var in PYANNOTE_AUTH_TOKEN or argument in --vad-args: 'auth_token'")
pyannote_args = kwargs.get('pyannote_args', {"onset": 0.5, "offset": 0.5, "min_duration_on": 0.3, "min_duration_off": 0.3})
self.model = Model.from_pretrained(model_name, use_auth_token=auth_token)
self.vad_pipeline = VoiceActivityDetection(segmentation=self.model)
self.vad_pipeline.instantiate(pyannote_args)
async def detect_activity(self, client):
audio_file_path = await save_audio_to_file(client.scratch_buffer, client.get_file_name())
vad_results = self.vad_pipeline(audio_file_path)
remove(audio_file_path)
vad_segments = []
if len(vad_results) > 0:
vad_segments = [
{"start": segment.start, "end": segment.end, "confidence": 1.0}
for segment in vad_results.itersegments()
]
return vad_segments
+23
View File
@@ -0,0 +1,23 @@
from VoiceStreamAI.vad.pyannote_vad import PyannoteVAD
class VADFactory:
"""
Factory for creating instances of VAD systems.
"""
@staticmethod
def create_vad_pipeline(type, **kwargs):
"""
Creates a VAD pipeline based on the specified type.
Args:
type (str): The type of VAD pipeline to create (e.g., 'pyannote').
kwargs: Additional arguments for the VAD pipeline creation.
Returns:
VADInterface: An instance of a class that implements VADInterface.
"""
if type == "pyannote":
return PyannoteVAD(**kwargs)
else:
raise ValueError(f"Unknown VAD pipeline type: {type}")
+16
View File
@@ -0,0 +1,16 @@
class VADInterface:
"""
Interface for voice activity detection (VAD) systems.
"""
async def detect_activity(self, client):
"""
Detects voice activity in the given audio data.
Args:
client (src.Client): The client to detect on
Returns:
List: VAD result, a list of objects containing "start", "end", "confidence"
"""
raise NotImplementedError("This method should be implemented by subclasses.")
-172
View File
@@ -1,172 +0,0 @@
import torch
from PIL import Image, ImageOps, ImageSequence, ImageFile
from PIL.PngImagePlugin import PngInfo
import numpy as np
import os
import folder_paths
import node_helpers
import hashlib
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
# tensor 取hash值
def tensor_to_hash(tensor):
# 将 Tensor 转换为 NumPy 数组
np_array = tensor.cpu().numpy()
# 将 NumPy 数组转换为字节数据
byte_data = np_array.tobytes()
# 计算哈希值
hash_value = hashlib.md5(byte_data).hexdigest()
return hash_value
def create_temp_file(image):
output_dir = folder_paths.get_temp_directory()
(
full_output_folder,
filename,
counter,
subfolder,
_,
) = folder_paths.get_save_image_path('material', output_dir)
image=tensor2pil(image)
image_file = f"{filename}_{counter:05}.png"
image_path=os.path.join(full_output_folder, image_file)
image.save(image_path,compress_level=4)
return (image_path,[{
"filename": image_file,
"subfolder": subfolder,
"type": "temp"
}])
# image - tensor - 文件路径
# loadImage的方法( 文件路径 - image-mask )
class EditMask:
def __init__(self):
self.image_id = None
@classmethod
def INPUT_TYPES(s):
return {"required":
{"image": ("IMAGE",), # 表示一个张量
},
"optional":{
"image_update": ("IMAGE_FILE",)
},
}
CATEGORY = "♾️Mixlab/Mask"
RETURN_TYPES = ("IMAGE", "MASK")
RETURN_NAMES = ("image", "mask")
FUNCTION = "edit"
OUTPUT_NODE = True
def edit(self, image,image_update=None):
# 根据image输入来判断是否是新的图片
if self.image_id==None:
self.image_id=tensor_to_hash(image)
image_update=None
else:
image_id=tensor_to_hash(image)
if image_id!=self.image_id:
image_update=None
self.image_id=image_id
image_path=None
# print('#image_update',self.image_id,image_update)
if image_update==None:
print('--')
else:
if 'images' in image_update:
images=image_update['images']
filename=images[0]['filename']
subfolder=images[0]['subfolder']
type=images[0]['type']
name, base_dir=folder_paths.annotated_filepath(filename)
if type.endswith("output"):
base_dir = folder_paths.get_output_directory()
elif type.endswith("input"):
base_dir = folder_paths.get_input_directory()
elif type.endswith("temp"):
base_dir = folder_paths.get_temp_directory()
#base_dir = folder_paths.get_input_directory()
# print(base_dir,subfolder, name)
image_path = os.path.join(base_dir,subfolder, name)
if image_path==None:
image_path,images=create_temp_file(image)
print('#image_path',os.path.exists(image_path),image_path)
# image_path = folder_paths.get_annotated_filepath(image) #文件名
if not os.path.exists(image_path):
image_path,images=create_temp_file(image)
img = node_helpers.pillow(Image.open, image_path)
output_images = []
output_masks = []
w, h = None, None
excluded_formats = ['MPO']
for i in ImageSequence.Iterator(img):
i = node_helpers.pillow(ImageOps.exif_transpose, i)
if i.mode == 'I':
i = i.point(lambda i: i * (1 / 255))
image = i.convert("RGB")
if len(output_images) == 0:
w = image.size[0]
h = image.size[1]
if image.size[0] != w or image.size[1] != h:
continue
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
if 'A' in i.getbands():
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
# 尺寸不对,需要按照image来
mask = torch.zeros((h, w), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 1 and img.format not in excluded_formats:
output_image = torch.cat(output_images, dim=0)
output_mask = torch.cat(output_masks, dim=0)
else:
output_image = output_images[0]
output_mask = output_masks[0]
return {"ui":{"images": images},"result": (output_image, output_mask)}
# return (output_image, output_mask)
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-mixlab-nodes"
description = "3D, ScreenShareNode & FloatingVideoNode, SpeechRecognition & SpeechSynthesis, GPT, LoadImagesFromLocal, Layers, Other Nodes, ..."
version = "0.35.1"
version = "0.28.3"
license = "MIT"
dependencies = ["numpy", "pyOpenSSL", "watchdog", "opencv-python-headless", "matplotlib", "openai", "simple-lama-inpainting", "clip-interrogator==0.6.0", "transformers>=4.36.0", "lark-parser", "imageio-ffmpeg", "rembg[gpu]", "omegaconf==2.3.0", "Pillow>=9.5.0", "einops==0.7.0", "trimesh>=4.0.5", "huggingface-hub", "scikit-image"]
+1 -3
View File
@@ -15,6 +15,4 @@ Pillow>=9.5.0
einops==0.7.0
trimesh>=4.0.5
huggingface-hub
scikit-image
torchaudio
soundfile>=0.12.1
scikit-image
-22
View File
@@ -1,22 +0,0 @@
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>Mixlab AR</title>
</head>
<body>
<script type="module">
import { api } from "../../../scripts/api.js";
import Command from '/extensions/comfyui-mixlab-nodes/javascript/command.js'
</script>
</body>
</html>
+687 -593
View File
File diff suppressed because it is too large Load Diff
+9 -48
View File
@@ -2,9 +2,6 @@ import { app } from '../../../scripts/app.js'
import { $el } from '../../../scripts/ui.js'
import { api } from '../../../scripts/api.js'
import { td_bg } from './td_background.js'
// console.log('td_bg', td_bg)
//本机安装的插件节点全集
window._nodesAll = null
@@ -188,25 +185,6 @@ async function extractInputAndOutputData (
if (node.type == 'Color') {
}
// 语音输入的支持
if (node.type == 'LoadAndCombinedAudio_') {
// if (
// data[id].widgets_values &&
// data[id].widgets_values[0] &&
// data[id].widgets_values[0].base64 &&
// data[id].widgets_values[0].base64.length > 0
// ) {
// options.defaultBase64 = data[id].widgets_values[0].base64
// }
input[inputIds.indexOf(id)] = {
...data[id],
title: node.title,
id,
options
}
}
if (node.type === 'LoadImage') {
// loadImage的mask支持
let output = node.outputs.filter(ot => ot.type == 'MASK')[0]
@@ -257,9 +235,7 @@ async function extractInputAndOutputData (
node.type === 'KSampler' ||
node.type == 'SamplerCustom' ||
node.type === 'ChinesePrompt_Mix' ||
node.type === 'Seed_'||
node.type==='SiliconflowLLM'||
node.type==='ChatGPTOpenAI'
node.type === 'Seed_'
) {
// seed 的类型收集
try {
@@ -421,11 +397,11 @@ async function save (json, download = false, showInfo = true) {
function getInputsAndOutputs () {
const inputs =
`LoadImage LoadImagesToBatch ImagesPrompt_ LoadAndCombinedAudio_ LoadVideoAndSegment_ VHS_LoadVideo CLIPTextEncode PromptSlide TextInput_ Color FloatSlider IntNumber CheckpointLoaderSimple LoraLoader`.split(
`LoadImage LoadImagesToBatch ImagesPrompt_ VHS_LoadVideo CLIPTextEncode PromptSlide TextInput_ Color FloatSlider IntNumber CheckpointLoaderSimple LoraLoader`.split(
' '
),
outputs =
`SaveTripoSRMesh,PreviewImage,SaveImage,TransparentImage,ShowTextForGPT,CombineAudioVideo,VHS_VideoCombine,VideoCombine_Adv,Image Save,SaveImageAndMetadata_,ClipInterrogator`.split(
`SaveTripoSRMesh,PreviewImage,SaveImage,TransparentImage,ShowTextForGPT,VHS_VideoCombine,VideoCombine_Adv,Image Save,SaveImageAndMetadata_,ClipInterrogator`.split(
','
)
@@ -469,18 +445,19 @@ app.registerExtension({
const { input, output } = getInputsAndOutputs()
input_ids.value = input.join('\n')
output_ids.value = output.join('\n')
const widget = {
type: 'div',
name: 'AppInfoRun',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
{...get_position_style(
get_position_style(
ctx,
widget_width,
node.size[1] - widget_height,
node.size[1]
),zIndex:1}
)
)
}
}
@@ -526,21 +503,6 @@ app.registerExtension({
}
})
//td bg
const tdBG = document.createElement('button')
tdBG.innerText = 'Canvas Mode'
tdBG.style = style
tdBG.style.marginLeft = '12px'
tdBG.addEventListener('click', () => {
td_bg.toggle()
if (td_bg.running) {
tdBG.style.background = 'yellow'
} else {
tdBG.style.background = 'transparent'
}
})
// author
let author = document.createElement('div')
// author.style=`display: flex`
@@ -697,7 +659,6 @@ app.registerExtension({
btns.appendChild(btn)
btns.appendChild(download)
btns.appendChild(tdBG)
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
@@ -711,7 +672,6 @@ app.registerExtension({
this.serialize_widgets = true //需要保存参数
window._mixlab_app_json = null
}
const onExecuted = nodeType.prototype.onExecuted
@@ -727,8 +687,9 @@ app.registerExtension({
}
const div = this.widgets.filter(w => w.div)[0].div
Array.from(div.querySelectorAll('button'), b =>
b.innerText != 'Canvas Mode' ? (b.style.background = 'yellow') : ''
Array.from(
div.querySelectorAll('button'),
b => (b.style.background = 'yellow')
)
} catch (error) {}
}
-215
View File
@@ -396,218 +396,3 @@ app.registerExtension({
}
}
})
// 上传音频转为base64
async function uploadAndConvertAudio (file) {
if (!file) {
alert('Please select a WAV file.')
return
}
if (file.type !== 'audio/wav') {
alert('Only WAV files are supported.')
return
}
try {
const base64Audio = await readFileAsDataURL(file)
return base64Audio
} catch (error) {
console.error('Error reading file:', error)
alert('Error reading file.')
}
}
function readFileAsDataURL (file) {
return new Promise((resolve, reject) => {
const reader = new FileReader()
reader.onload = function (event) {
resolve(event.target.result)
}
reader.onerror = function (error) {
reject(error)
}
reader.readAsDataURL(file)
})
}
const createInputAudioForBatch = (base64, widget) => {
// Create an audio element
let audio = document.createElement('audio')
audio.src = base64
audio.controls = true
audio.style = 'width: 120px; display: block'
// Create a delete button
let deleteButton = document.createElement('button')
deleteButton.textContent = 'Delete'
deleteButton.style = `cursor: pointer;
font-weight: 300;
margin: 2px;
margin-left: 10px;
color: var(--descrip-text);
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;height: 30px;min-width: 122px;
`
// Create a container for the audio and delete button
let container = document.createElement('div')
container.appendChild(audio)
container.appendChild(deleteButton)
container.style = `display: flex;margin-top: 12px;`
// Add event listener for the delete button
deleteButton.addEventListener('click', e => {
let newValue = []
let items = widget.value?.base64 || []
for (const v of items) {
if (v != base64) newValue.push(v)
}
widget.value.base64 = newValue
container.remove()
})
return container
}
app.registerExtension({
name: 'Mixlab.Comfy.LoadAndCombinedAudio_',
async getCustomWidgets (app) {
return {
AUDIOBASE64 (node, inputName, inputData, app) {
// console.log('##node', node)
const widget = {
value: {
base64: []
}, // 不能[x,x,x]
type: inputData[0], // the type
name: inputName, // the name, slice
size: [128, 32], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 122] // a method to compute the current size of the widget
}
// serializeValue (nodeId, widgetIndex) {
// return widget.value
// },
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'LoadAndCombinedAudio_') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
let audiosWidget = this.widgets.filter(w => w.name == 'audios')[0]
const widget = {
type: 'div',
name: 'audio_base64',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 44, node.size[1])
)
},
serialize: false
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
let audioPreview = document.createElement('div')
let audiosDiv = document.createElement('div') //显示图片
audiosDiv.className = 'audios_preview'
audiosDiv.style = `width: calc(100% - 14px);
display: flex;
flex-wrap: wrap;
padding: 7px; justify-content: space-between;
align-items: center;`
const btn = document.createElement('button')
btn.innerText = 'Upload Audio'
btn.style = `cursor: pointer;
font-weight: 300;
margin: 2px;
color: var(--descrip-text);
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;height: 30px;min-width: 122px;
`
btn.addEventListener('click', e => {
e.preventDefault()
let inputAudio = document.createElement('input')
inputAudio.type = 'file'
inputAudio.accept = "audio/*"
inputAudio.style.display = 'none'
inputAudio.addEventListener('change', async e => {
e.preventDefault()
const file = e.target.files[0]
let base64 = await uploadAndConvertAudio(file)
if (!audiosWidget.value) audiosWidget.value = { base64: [] }
audiosWidget.value.base64.push(base64)
let a = createInputAudioForBatch(base64, audiosWidget)
audiosDiv.appendChild(a)
})
inputAudio.click()
inputAudio.remove()
})
widget.div.appendChild(audioPreview)
audioPreview.appendChild(audiosDiv)
audioPreview.appendChild(btn)
// audioPreview.appendChild(inputAudio)
this.addCustomWidget(widget)
// document.addEventListener('wheel', handleMouseWheel)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
try {
// document.removeEventListener('wheel', handleMouseWheel)
} catch (error) {
console.log(error)
}
return onRemoved?.()
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'LoadAndCombinedAudio_') {
// await sleep(0)
let audiosWidget = node.widgets.filter(w => w.name === 'audios')[0]
let audioPreview = node.widgets.filter(w => w.name == 'audio_base64')[0]
let pre = audioPreview.div.querySelector('.audios_preview')
for (const d of audiosWidget.value?.base64 || []) {
let im = createInputAudioForBatch(d, audiosWidget)
pre.appendChild(im)
}
}
}
})
+1 -1
View File
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
const version = 'v0.35.1'
const version = 'v0.28.3'
fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
.then(response => response.json())
-689
View File
@@ -1,689 +0,0 @@
function get_url () {
// 如果有缓存记录
let hostUrl = localStorage.getItem('_hostUrl') || ''
if (hostUrl) {
return hostUrl
}
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
return url
}
function getFilenameAndCategoryFromUrl (url) {
const queryString = url.split('?')[1]
if (!queryString) {
return {}
}
const params = new URLSearchParams(queryString)
const filename = params.get('filename')
? decodeURIComponent(params.get('filename'))
: null
const category = params.get('category')
? decodeURIComponent(params.get('category') || '')
: ''
return { category, filename }
}
async function get_my_app (category = '', filename = null) {
let url = get_url()
const res = await fetch(`${url}/mixlab/workflow`, {
method: 'POST',
mode: 'cors', // 允许跨域请求
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
task: 'my_app',
filename,
category
})
})
let result = await res.json()
let data = []
try {
for (const res of result.data) {
let { output, app } = res.data
if (app.filename)
data.push({
...app,
data: output,
date: res.date
})
}
} catch (error) {}
return data
}
async function getAppInit () {
const { category, filename } = getFilenameAndCategoryFromUrl(
window.location.href
)
return await get_my_app(category, filename)
}
function success (isSuccess, btn, text) {
isSuccess ? (btn.innerText = 'success') : text
setTimeout(() => {
btn.innerText = text
}, 5000)
}
async function interrupt () {
try {
await fetch(`${get_url()}/interrupt`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: undefined
})
} catch (error) {
console.error(error)
}
return true
}
async function getQueue (clientId) {
try {
const res = await fetch(`${get_url()}/queue`)
const data = await res.json()
return {
// Running action uses a different endpoint for cancelling
Running: Array.from(data.queue_running, prompt => {
if (prompt[3].client_id === clientId) {
let prompt_id = prompt[1]
return {
prompt_id,
remove: () => interrupt()
}
}
}),
Pending: data.queue_pending.map(prompt => ({ prompt }))
}
} catch (error) {
console.error(error)
return { Running: [], Pending: [] }
}
}
// 请求历史数据
async function getPromptResult (category) {
let url = get_url()
try {
const response = await fetch(`${url}/mixlab/prompt_result`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
action: 'all'
})
})
if (response.ok) {
const data = await response.json()
console.log('#getPromptResult:', category, data)
return data.result.filter(r => r.appInfo.category == category)
// 处理返回的数据
} else {
console.log('Error:', response.status)
// 处理错误情况
}
} catch (error) {
console.log('Error:', error)
// 处理异常情况
}
}
// 新的运行工作流的接口
function queuePromptNew (filename, category, seed, input, client_id,apps=null) {
let url = get_url()
// var filename = "Text-to-Image_1.json", category = "";
// 随机seed
// promptWorkflow = randomSeed(seed, promptWorkflow);
let d = { filename, category, seed, input, client_id }
if (apps) {
d.apps = apps
}
const data = JSON.stringify(d)
return new Promise((res, rej) => {
fetch(`${url}/mixlab/prompt`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: data
})
.then(response => {
if (!response.ok) {
// Handle HTTP error responses
if (response.status === 400) {
return response.json().then(errorData => {
// Process the error data
console.error('Error 400:', errorData)
alert(JSON.stringify(errorData, null, 2))
res(null)
})
}
throw new Error('Network response was not ok')
}
return response.json() // Process the response data
})
.then(data => {
// Handle the response data
console.log('Success:', data)
res(true)
})
.catch(error => {
// Handle fetch errors
console.error('Fetch error:', error)
res(null)
})
})
}
// 保存历史数据
async function savePromptResult (data) {
let url = get_url()
try {
const response = await fetch(`${url}/mixlab/prompt_result`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
action: 'save',
data
})
})
if (response.ok) {
const res = await response.json()
console.log('Response:', res)
return res
// 处理返回的数据
} else {
console.log('Error:', response.status)
// 处理错误情况
}
} catch (error) {
console.log('Error:', error)
// 处理异常情况
}
}
async function uploadImage (blob, fileType = '.png', filename) {
const body = new FormData()
body.append(
'image',
new File([blob], (filename || new Date().getTime()) + fileType)
)
const url = get_url()
const resp = await fetch(`${url}/upload/image`, {
method: 'POST',
body
})
let data = await resp.json()
// console.log(data)
let { name, subfolder } = data
let src = `${url}/view?filename=${encodeURIComponent(
name
)}&type=input&subfolder=${subfolder}&rand=${Math.random()}`
return { url: src, name }
}
async function uploadMask (arrayBuffer, imgurl) {
const body = new FormData()
const filename = 'clipspace-mask-' + performance.now() + '.png'
let original_url = new URL(imgurl)
const original_ref = { filename: original_url.searchParams.get('filename') }
let original_subfolder = original_url.searchParams.get('subfolder')
if (original_subfolder) original_ref.subfolder = original_subfolder
let original_type = original_url.searchParams.get('type')
if (original_type) original_ref.type = original_type
body.append('image', arrayBuffer, filename)
body.append('original_ref', JSON.stringify(original_ref))
body.append('type', 'input')
body.append('subfolder', 'clipspace')
const url = get_url()
const resp = await fetch(`${url}/upload/mask`, {
method: 'POST',
body
})
// console.log(resp)
let data = await resp.json()
let { name, subfolder, type } = data
let src = `${url}/view?filename=${encodeURIComponent(
name
)}&type=${type}&subfolder=${subfolder}&rand=${Math.random()}`
return { url: src, name: 'clipspace/' + name }
}
const parseImageToBase64 = url => {
return new Promise((res, rej) => {
fetch(url)
.then(response => response.blob())
.then(blob => {
const reader = new FileReader()
reader.onloadend = () => {
const base64data = reader.result
res(base64data)
// 在这里可以将base64数据用于进一步处理或显示图片
}
reader.readAsDataURL(blob)
})
.catch(error => {
console.log('发生错误:', error)
})
})
}
function createImage (url) {
let im = new Image()
return new Promise((res, rej) => {
im.onload = () => res(im)
im.src = url
})
}
function convertImageToBlackBasedOnAlpha (image) {
const canvas = document.createElement('canvas')
const ctx = canvas.getContext('2d')
// Draw the image onto the canvas
canvas.width = image.width
canvas.height = image.height
ctx.drawImage(image, 0, 0)
// Get the image data from the canvas
const imageData = ctx.getImageData(0, 0, canvas.width, canvas.height)
const pixels = imageData.data
// Modify the RGB values based on the alpha channel
for (let i = 0; i < pixels.length; i += 4) {
const alpha = pixels[i + 3]
if (alpha !== 0) {
// Set non-transparent pixels to black
pixels[i] = 0 // Red
pixels[i + 1] = 0 // Green
pixels[i + 2] = 0 // Blue
}
}
// Put the modified image data back onto the canvas
ctx.putImageData(imageData, 0, 0)
// Convert the modified canvas to base64 data URL
const base64ImageData = canvas.toDataURL('image/png') // Replace 'png' with your desired image format
return base64ImageData
}
const blobToBase64 = blob => {
return new Promise((res, rej) => {
const reader = new FileReader()
reader.onloadend = () => {
const base64data = reader.result
res(base64data)
// 在这里可以将base64数据用于进一步处理或显示图片
}
reader.readAsDataURL(blob)
})
}
function base64ToBlob (base64) {
// 去除base64编码中的前缀
const base64WithoutPrefix = base64.replace(/^data:image\/\w+;base64,/, '')
// 将base64编码转换为字节数组
const byteCharacters = atob(base64WithoutPrefix)
// 创建一个存储字节数组的数组
const byteArrays = []
// 将字节数组放入数组中
for (let offset = 0; offset < byteCharacters.length; offset += 1024) {
const slice = byteCharacters.slice(offset, offset + 1024)
const byteNumbers = new Array(slice.length)
for (let i = 0; i < slice.length; i++) {
byteNumbers[i] = slice.charCodeAt(i)
}
const byteArray = new Uint8Array(byteNumbers)
byteArrays.push(byteArray)
}
// 创建blob对象
const blob = new Blob(byteArrays, { type: 'image/png' }) // 根据实际情况设置MIME类型
return blob
}
async function calculateImageHash (blob) {
const buffer = await blob.arrayBuffer()
const hashBuffer = await crypto.subtle.digest('SHA-256', buffer)
const hashArray = Array.from(new Uint8Array(hashBuffer))
const hashHex = hashArray
.map(byte => byte.toString(16).padStart(2, '0'))
.join('')
return hashHex
}
// 获取 rembg 模型
async function get_rembg_models () {
try {
const response = await fetch(`${get_url()}/mixlab/folder_paths`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
type: 'rembg'
})
})
const data = await response.json()
// console.log(data)
return data.names
} catch (error) {
console.error(error)
}
}
//自动抠图
async function run_rembg (model, base64) {
try {
const response = await fetch(`${get_url()}/mixlab/rembg`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
model,
base64
})
})
const data = await response.json()
// console.log(data)
return data.data
} catch (error) {
console.error(error)
}
}
function copyHtmlWithImagesToClipboard (data, cb) {
// 创建一个临时div元素
const tempDiv = document.createElement('div')
// 将HTML字符串赋值给div的innerHTML属性
tempDiv.innerHTML = data
// 获取div中的所有图像元素
const images = tempDiv.getElementsByTagName('img')
// 遍历图像元素,并将图像数据转换为Base64编码
for (let i = 0; i < images.length; i++) {
const image = images[i]
const canvas = document.createElement('canvas')
const context = canvas.getContext('2d')
// 设置canvas尺寸与图像尺寸相同
canvas.width = image.width
canvas.height = image.height
// 在canvas上绘制图像
context.drawImage(image, 0, 0)
// 将canvas转换为Base64编码
const imageData = canvas.toDataURL()
// 将Base64编码替换图像元素的src属性
image.src = imageData
}
let richText = tempDiv.innerHTML
// 创建一个新的Blob对象,并将富文本字符串作为数据传递进去
const blob = new Blob([richText], { type: 'text/html' })
// 创建一个ClipboardItem对象,并将Blob对象添加到其中
const clipboardItem = new ClipboardItem({ 'text/html': blob })
// 使用Clipboard API将内容复制到剪贴板
navigator.clipboard
.write([clipboardItem])
.then(() => {
console.log('富文本已成功复制到剪贴板')
tempDiv.remove()
if (cb) cb(true)
})
.catch(error => {
console.error('复制到剪贴板失败:', error)
tempDiv.remove()
if (cb) cb(false)
})
}
function copyImagesToClipboard (html, cb) {
const tempDiv = document.createElement('div')
tempDiv.innerHTML = html
const images = tempDiv.querySelectorAll('img')
const promises = Array.from(images).map(image => {
return new Promise(resolve => {
const img = new Image()
img.src = image.src
img.onload = () => {
const canvas = document.createElement('canvas')
const context = canvas.getContext('2d')
canvas.width = img.width
canvas.height = img.height
context.drawImage(img, 0, 0)
canvas.toBlob(blob => {
const clipboardItem = new ClipboardItem({ 'image/png': blob })
navigator.clipboard
.write([clipboardItem])
.then(() => {
resolve()
tempDiv.remove()
if (cb) cb(true)
})
.catch(error => {
reject(error)
tempDiv.remove()
if (cb) cb(false)
})
})
}
})
})
Promise.all([...promises])
.then(() => {
console.log('所有图片已成功复制到剪贴板')
if (cb) cb(true)
tempDiv.remove()
})
.catch(error => {
console.error('复制到剪贴板失败:', error)
if (cb) cb(false)
tempDiv.remove()
})
}
function copyTextToClipboard (html, cb) {
const tempDiv = document.createElement('div')
tempDiv.innerHTML = html
const text = tempDiv.innerText
const textData = new ClipboardItem({
'text/plain': new Blob([text], { type: 'text/plain' })
})
navigator.clipboard
.write([textData])
.then(() => {
console.log('所有文本已成功复制到剪贴板', text)
if (cb) cb(true)
tempDiv.remove()
})
.catch(error => {
console.error('复制到剪贴板失败:', error)
if (cb) cb(false)
tempDiv.remove()
})
}
// ComfyUI\web\extensions\core\dynamicPrompts.js
// 官方实现修改
// Allows for simple dynamic prompt replacement
// Inputs in the format {a|b} will have a random value of a or b chosen when the prompt is queued.
/*
* Strips C-style line and block comments from a string
*/
function dynamicPrompts (prompt) {
prompt = prompt.replace(/\/\*[\s\S]*?\*\/|\/\/.*/g, '')
while (
prompt.replace('\\{', '').includes('{') &&
prompt.replace('\\}', '').includes('}')
) {
const startIndex = prompt.replace('\\{', '00').indexOf('{')
const endIndex = prompt.replace('\\}', '00').indexOf('}')
const optionsString = prompt.substring(startIndex + 1, endIndex)
const options = optionsString.split('|')
const randomIndex = Math.floor(Math.random() * options.length)
const randomOption = options[randomIndex]
prompt =
prompt.substring(0, startIndex) +
randomOption +
prompt.substring(endIndex + 1)
}
return prompt
}
// 遍历所有组合,语法同 动态提示
function generateAllCombinations (prompt) {
prompt = prompt.replace(/\/\*[\s\S]*?\*\/|\/\/.*/g, '')
// Helper function to get all combinations
function getAllCombinations (parts) {
if (parts.length === 0) return ['']
const [firstPart, ...restParts] = parts
const restCombinations = getAllCombinations(restParts)
const allCombinations = []
firstPart.forEach(option => {
restCombinations.forEach(combination => {
allCombinations.push(option + combination)
})
})
return allCombinations
}
// Split prompt into static parts and dynamic parts
let parts = []
let startIndex = 0
while (
prompt.replace('\\{', '').includes('{') &&
prompt.replace('\\}', '').includes('}')
) {
startIndex = prompt.replace('\\{', '00').indexOf('{')
const endIndex = prompt.replace('\\}', '00').indexOf('}')
const staticPart = prompt.substring(0, startIndex)
const optionsString = prompt.substring(startIndex + 1, endIndex)
const options = optionsString.split('|')
parts.push([staticPart])
parts.push(options)
prompt = prompt.substring(endIndex + 1)
}
// Add the remaining static part
parts.push([prompt])
// Get all combinations
const combinations = getAllCombinations(parts)
return combinations
}
const _textNodes = [
'TextInput_',
'CLIPTextEncode',
'PromptSimplification',
'ChinesePrompt_Mix'
],
_loraNodes = ['CheckpointLoaderSimple', 'LoraLoader'],
_numberNodes = ['FloatSlider', 'IntNumber'],
_slideNodes = ['PromptSlide'],
_imageNodes = [
'LoadImage',
'VHS_LoadVideo',
'ImagesPrompt_',
'LoadImagesToBatch'
],
_colorNodes = ['Color'],
_audioNodes = ['LoadAndCombinedAudio_']
export default {
get_url,
get_my_app,
getAppInit,
getFilenameAndCategoryFromUrl,
success,
interrupt,
getQueue,
queuePromptNew,
savePromptResult,
uploadImage,
uploadMask,
run_rembg,
get_rembg_models,
parseImageToBase64,
createImage,
convertImageToBlackBasedOnAlpha,
blobToBase64,
base64ToBlob,
calculateImageHash,
copyHtmlWithImagesToClipboard,
copyImagesToClipboard,
copyTextToClipboard,
dynamicPrompts,
generateAllCombinations,
_textNodes,
_loraNodes,
_numberNodes,
_slideNodes,
_imageNodes,
_colorNodes,
_audioNodes
}
+203 -8
View File
@@ -1,5 +1,205 @@
import { app } from '../../../scripts/app.js'
// import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
async function getConfig () {
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
const res = await fetch(`${url}/mixlab`, {
method: 'POST'
})
return await res.json()
}
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
const getLocalData = key => {
let data = {}
try {
data = JSON.parse(localStorage.getItem(key)) || {}
} catch (error) {
return {}
}
return data
}
app.registerExtension({
name: 'Mixlab.GPT.ChatGPTOpenAI',
async getCustomWidgets (app) {
return {
KEY (node, inputName, inputData, app) {
// console.log('##inputData', inputData)
const widget = {
type: inputData[0], // the type, CHEESE
name: inputName, // the name, slice
size: [128, 32], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_api_key')
return data[node.id] || 'by Mixlab'
}
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
},
URL (node, inputName, inputData, app) {
// console.log('node', inputName, inputData[0])
const widget = {
type: inputData[0], // the type, CHEESE
name: inputName, // the name, slice
size: [128, 32], // a default size
draw (ctx, node, width, y) {
// a method to draw the widget (ctx is a CanvasRenderingContext2D)
},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_api_url')
return data[node.id] || 'https://api.openai.com/v1'
}
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'ChatGPTOpenAI') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const api_key = this.widgets.filter(w => w.name == 'api_key')[0]
const api_url = this.widgets.filter(w => w.name == 'api_url')[0]
console.log('ChatGPTOpenAI nodeData', this.widgets)
const widget = {
type: 'div',
name: 'chatgptdiv',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, api_key.y, node.size[1])
)
}
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
const inputDiv = (key, placeholder) => {
let div = document.createElement('div')
const ip = document.createElement('input')
ip.type = placeholder === 'Key' ? 'password' : 'text'
ip.className = `${'comfy-multiline-input'} ${placeholder}`
div.style = `display: flex;
align-items: center;
margin: 6px 8px;
margin-top: 0;`
ip.placeholder = placeholder
ip.value = placeholder
ip.style = `margin-left: 24px;
outline: none;
border: none;
padding: 4px;width: 100%;`
const label = document.createElement('label')
label.style = 'font-size: 10px;min-width:32px'
label.innerText = placeholder
div.appendChild(label)
div.appendChild(ip)
ip.addEventListener('change', () => {
let data = getLocalData(key)
data[this.id] = ip.value.trim()
localStorage.setItem(key, JSON.stringify(data))
console.log(this.id, key)
})
return div
}
let inputKey = inputDiv('_mixlab_api_key', 'Key')
let inputUrl = inputDiv('_mixlab_api_url', 'URL')
widget.div.appendChild(inputKey)
widget.div.appendChild(inputUrl)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
inputUrl.remove()
inputKey.remove()
widget.div.remove()
return onRemoved?.()
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
// Fires every time a node is constructed
// You can modify widgets/add handlers/etc here
if (node.type === 'ChatGPTOpenAI') {
let widget = node.widgets.filter(w => w.div)[0]
let apiKey = getLocalData('_mixlab_api_key'),
url = getLocalData('_mixlab_api_url')
let id = node.id
// console.log('ChatGPTOpenAI serialize_widgets', this)
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
widget.div.querySelector('.URL').value =
url[id] || 'https://api.openai.com/v1'
}
}
})
app.registerExtension({
name: 'Mixlab.GPT.ShowTextForGPT',
@@ -9,16 +209,13 @@ app.registerExtension({
text = text.filter(t => t && t?.trim())
if (this.widgets) {
// console.log('#ShowTextForGPT',this.widgets)
// const pos = this.widgets.findIndex(w => w.name === 'text')
for (let i = 0; i < this.widgets.length; i++) {
if (this.widgets[i].name == 'show_text')
this.widgets[i].onRemove?.()
if (this.widgets[i].name == 'show_text') this.widgets[i].onRemove?.()
}
this.widgets.length = 2
this.widgets.length = 1
}
// console.log('ShowTextForGPT',text)
for (let list of text) {
if (list) {
// console.log('#####', list)
@@ -31,8 +228,6 @@ app.registerExtension({
w.inputEl.readOnly = true
w.inputEl.style.opacity = 0.6
// w.inputEl.style.display='none'
try {
if (typeof list != 'string') {
let data = JSON.parse(list)
+17 -44
View File
@@ -675,17 +675,6 @@ const createInputImageForBatch = (base64, widget) => {
return im
}
// 添加新图片
const addBase64ToWidgetForLoadImagesToBatch = (
base64,
imagesWidget,
imagesDiv
) => {
imagesWidget.value.base64.push(base64)
let im = createInputImageForBatch(base64, imagesWidget)
imagesDiv.appendChild(im)
}
app.registerExtension({
name: 'Mixlab.Comfy.LoadImagesToBatch',
async getCustomWidgets (app) {
@@ -716,6 +705,7 @@ app.registerExtension({
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'LoadImagesToBatch') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
@@ -761,18 +751,13 @@ app.registerExtension({
base64 = await loadImageToCanvas(base64)
// console.log(base64)
if (!imagesWidget.value) imagesWidget.value = { base64: [] }
addBase64ToWidgetForLoadImagesToBatch(
base64,
imagesWidget,
imagesDiv
)
imagesWidget.value.base64.push(base64)
let im = createInputImageForBatch(base64, imagesWidget)
imagesDiv.appendChild(im)
}
reader.readAsDataURL(file)
})
// 如果是复制的,有数据 , 这个不生效,取不到数据, 需要在nodeCreated里获取
// console.log('#LoadImagesToBatch', imagesWidget.value?.base64)
const btn = document.createElement('button')
btn.innerText = 'Upload Image'
@@ -844,36 +829,18 @@ app.registerExtension({
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'LoadImagesToBatch') {
// await sleep(0)
let imagesWidget = node.widgets.filter(w => w.name === 'images')[0]
let imagePreview = node.widgets.filter(w => w.name == 'image_base64')[0]
// console.log('#LoadImagesToBatch', imagesWidget.value?.base64)
let imagesDiv = imagePreview.div.querySelector('.images_preview')
let pre = imagePreview.div.querySelector('.images_preview')
for (const d of imagesWidget.value?.base64 || []) {
let im = createInputImageForBatch(d, imagesWidget)
imagesDiv.appendChild(im)
pre.appendChild(im)
}
}
},
nodeCreated (node, app) {
//数据延迟??
setTimeout(() => {
// console.log('#LoadImagesToBatch', node.type)
if (node.type === 'LoadImagesToBatch') {
let imagesWidget = node.widgets.filter(w => w.name === 'images')[0]
let imagePreview = node.widgets.filter(w => w.name == 'image_base64')[0]
let imagesDiv = imagePreview?.div?.querySelector('.images_preview')
for (const d of imagesWidget.value?.base64 || []) {
let im = createInputImageForBatch(d, imagesWidget)
imagesDiv.appendChild(im)
}
}
}, 1000)
}
})
@@ -901,8 +868,8 @@ app.registerExtension({
const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated
? onNodeCreated.apply(this, arguments)
: undefined
? onNodeCreated.apply(this, arguments)
: undefined
this.size = [400, this.size[1]]
console.log('##onNodeCreated', this)
@@ -924,15 +891,20 @@ app.registerExtension({
this.addCustomWidget(widget)
this.serialize_widgets = true //需要保存参数
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
return r
return r
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
@@ -992,7 +964,7 @@ app.registerExtension({
label: 'After'
}
]
this.size = [this.size[0], 300]
this.size=[this.size[0],300]
}
}
},
@@ -1002,6 +974,7 @@ app.registerExtension({
// node.widgets[0].div.id = 'mix_comparingtowframes_' + node.id
// if (node.widgets_values && node.widgets_values[0]) {
// node.widgets[0].div.innerHTML = ''
// let slider = new juxtapose.JXSlider(
// '#mix_comparingtowframes_' + node.id,
// node.widgets_values,
+1 -1
View File
@@ -1267,7 +1267,7 @@ app.registerExtension({
})
widget.PictureInPicture = $el('button', {
innerText: 'Picture In Picture',
innerText: 'PictureInPicture',
style: {
display: 'pictureInPictureEnabled' in document ? 'block' : 'none',
cursor: 'pointer',
-295
View File
@@ -1,295 +0,0 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
import WaveSurfer from 'https://cdn.jsdelivr.net/npm/wavesurfer.js@7/dist/wavesurfer.esm.js'
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: '0',
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
//把文件转为url访问
const parseUrl = data => {
let { filename, subfolder, type, prompt } = data
return {
url: api.apiURL(
`/view?filename=${encodeURIComponent(
filename
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
),
prompt
}
}
const createWaveSurfer = (wavesurfer, id,url) => {
// Create an instance of WaveSurfer
if (wavesurfer) {
wavesurfer.destroy()
}
wavesurfer = WaveSurfer.create({
container: '#' + id,
waveColor: 'rgb(200, 0, 200)',
progressColor: 'rgb(100, 0, 100)',
// Set a bar width
barWidth: 10,
// Optionally, specify the spacing between bars
barGap: 2,
// And the bar radius
barRadius: 6,
url
})
wavesurfer._auto = true
// 监听播放结束事件,重新开始播放以实现循环播放
wavesurfer.on('finish', function () {
// console.log(wavesurfer)
if (wavesurfer._auto) wavesurfer.play()
})
wavesurfer.on('interaction', () => {
wavesurfer._auto = false
if (!wavesurfer.isPlaying()) wavesurfer.play()
})
// 获取当前播放时间的峰值
wavesurfer.on('audioprocess', () => {
if (wavesurfer.isPlaying()&&wavesurfer.getDecodedData()) {
const channelData = wavesurfer.getDecodedData().getChannelData(0);
const currentTime = wavesurfer.getCurrentTime()
// console.log(wavesurfer)
const sampleRate = wavesurfer.getDecodedData().sampleRate
// 定义要分析的时间窗口(例如1秒)
const windowSize = 1
const startSample = Math.floor(currentTime * sampleRate)
const endSample = Math.min(
startSample + windowSize * sampleRate,
channelData.length
)
let peak = 0
for (let i = startSample; i < endSample; i++) {
const value = Math.abs(channelData[i])
if (value > peak) {
peak = value
}
}
// console.log('Current Peak:', peak)
}
})
return wavesurfer
}
//更新gui
function updateWaveWidgetValue (widgets, id, url, prompt, wavesurfer) {
let widget = widgets.filter(w => w.name == 'AudioPlay')[0]
// 手动更新widget值
widget.value = [url, prompt]
if (widget.div) {
widget.div.querySelector('.wave').id = `AudioPlay_${id}`
}
wavesurfer = createWaveSurfer(wavesurfer, `AudioPlay_${id}`,url)
wavesurfer.on('ready', duration => {
console.log('Audio duration: ' + duration + ' seconds')
if (widget.div) {
widget.div.setAttribute('data-url', url)
widget.div.querySelector('.link').setAttribute('href', url)
widget.div.querySelector(
'.info'
).innerHTML = `<span style="font-size: 12px;
margin: 8px;">${duration.toFixed(
2
)} seconds</span> <br><span style="font-size: 14px;">${prompt||''}</span> <br>`
}
})
wavesurfer.load(url)
// console.log('updateWaveWidgetValue' ,url,wavesurfer)
return wavesurfer
}
app.registerExtension({
name: 'SoundLab.AudioPlay',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'AudioPlay') {
let that = this
// console.log('that', that)
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const widget = {
type: 'div',
name: 'AudioPlay',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1])
)
}
}
// console.log('AudioPlay nodeData', this)
widget.div = $el('div', {})
document.body.appendChild(widget.div)
// wave
const waveDiv = document.createElement('div')
waveDiv.className = 'wave'
waveDiv.style.minHeight = '172px'
widget.div.appendChild(waveDiv)
//prompt 相关信息展示
const infoDiv = document.createElement('div')
infoDiv.className = 'info'
infoDiv.style.marginBottom = '20px'
widget.div.appendChild(infoDiv)
// 按钮的区域
let btns = document.createElement('div')
btns.className = 'btns'
btns.style = `display: flex;
width: 100%;
justify-content: space-between;`
widget.div.appendChild(btns)
//play button
const playBtn = document.createElement('a')
playBtn.innerText = 'Play/Pause'
playBtn.style = `
display: flex;
padding: 4px 15px;
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;
color: var(--descrip-text);
text-decoration: none;
border-radius: 5px;
transition: background-color 0.3s ease 0s;
`
playBtn.addEventListener('click', e => {
e.preventDefault()
if (that[`wavesurfer_${this.id}`]) {
that[`wavesurfer_${this.id}`]?.playPause()
that[`wavesurfer_${this.id}`]._auto = true
}
})
btns.appendChild(playBtn)
const urlLink = document.createElement('a')
urlLink.className = 'link'
urlLink.innerText = 'URL'
urlLink.setAttribute('target', '_blank')
urlLink.style = `display: flex;
padding: 4px 15px;
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;
color: var(--descrip-text);
text-decoration: none;
border-radius: 5px;
transition: background-color 0.3s ease 0s;`
// urlLink.style.minHeight = '200px'
btns.appendChild(urlLink)
//todo 导出视频 that[`wavesurfer_${this.id}`].renderer.exportImage('image/png',1,'dataURL')
// https://github.com/diffusion-studio/ffmpeg-js
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
this.size = [this.size[0], 280]
this.serialize_widgets = true //需保存widget的值
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
const audio = message.audio
console.log('#onExecuted', `AudioPlay_${this.id}`, message,audio)
try {
let { url, prompt } = parseUrl(audio[0])
that[`wavesurfer_${this.id}`] = updateWaveWidgetValue(
this.widgets,
this.id,
url,
prompt,
that[`wavesurfer_${this.id}`]
)
that[`wavesurfer_${this.id}`]?.playPause()
} catch (error) {
console.log(error)
}
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'AudioPlay') {
let widget = node.widgets.filter(w => w.name == 'AudioPlay')[0]
if (widget.value) {
let [url, prompt] = widget.value
this[`wavesurfer_${node.id}`] = updateWaveWidgetValue(
node.widgets,
node.id,
url,
prompt,
this[`wavesurfer_${node.id}`]
)
}
console.log('#loadedGraphNode', node)
}
}
})
-323
View File
@@ -1,323 +0,0 @@
// touchdesigner的背景效果,把appinfo的输出,选择一张图片作为背景
window._bg_img = null
/**
* draws the back canvas (the one containing the background and the connections)
* @method drawBackCanvas
**/
LGraphCanvas.prototype.drawBackCanvas = function () {
var canvas = this.bgcanvas
if (
canvas.width != this.canvas.width ||
canvas.height != this.canvas.height
) {
canvas.width = this.canvas.width
canvas.height = this.canvas.height
}
if (!this.bgctx) {
this.bgctx = this.bgcanvas.getContext('2d')
}
var ctx = this.bgctx
if (ctx.start) {
ctx.start()
}
var viewport = this.viewport || [0, 0, ctx.canvas.width, ctx.canvas.height]
//clear
if (this.clear_background) {
ctx.clearRect(viewport[0], viewport[1], viewport[2], viewport[3])
}
//show subgraph stack header
if (this._graph_stack && this._graph_stack.length) {
ctx.save()
var parent_graph = this._graph_stack[this._graph_stack.length - 1]
var subgraph_node = this.graph._subgraph_node
ctx.strokeStyle = subgraph_node.bgcolor
ctx.lineWidth = 10
ctx.strokeRect(1, 1, canvas.width - 2, canvas.height - 2)
ctx.lineWidth = 1
ctx.font = '40px Arial'
ctx.textAlign = 'center'
ctx.fillStyle = subgraph_node.bgcolor || '#AAA'
var title = ''
for (var i = 1; i < this._graph_stack.length; ++i) {
title += this._graph_stack[i]._subgraph_node.getTitle() + ' >> '
}
ctx.fillText(title + subgraph_node.getTitle(), canvas.width * 0.5, 40)
ctx.restore()
}
var bg_already_painted = false
if (this.onRenderBackground) {
bg_already_painted = this.onRenderBackground(canvas, ctx)
}
//reset in case of error
if (!this.viewport) {
ctx.restore()
ctx.setTransform(1, 0, 0, 1, 0, 0)
}
this.visible_links.length = 0
if (this.graph) {
//apply transformations
ctx.save()
this.ds.toCanvasContext(ctx)
//render BG
if (
this.ds.scale < 1 &&
!bg_already_painted &&
this.clear_background_color
) {
ctx.fillStyle = this.clear_background_color
ctx.fillRect(
this.visible_area[0],
this.visible_area[1],
this.visible_area[2],
this.visible_area[3]
)
}
// 主要修改
if (this.background_image && this.ds.scale > 0.5 && !bg_already_painted) {
if (this.zoom_modify_alpha) {
//使得 alpha 越接近0时变化越缓慢。
let alpha = (1.0 - 0.5 / this.ds.scale) * this.editor_alpha
ctx.globalAlpha = Math.min(Math.max(0, Math.sqrt(alpha)), 1)
// console.log((1.0 - 0.5 / this.ds.scale) * this.editor_alpha)
} else {
ctx.globalAlpha = this.editor_alpha
}
ctx.imageSmoothingEnabled = ctx.imageSmoothingEnabled = false // ctx.mozImageSmoothingEnabled =
if (!this._bg_img || this._bg_img.name != this.background_image) {
this._bg_img = new Image()
this._bg_img.name = this.background_image
this._bg_img.src = this.background_image
var that = this
this._bg_img.onload = function () {
that.draw(true, true)
}
}
var pattern = null
if (this._pattern == null && this._bg_img.width > 0) {
pattern = ctx.createPattern(this._bg_img, 'repeat')
this._pattern_img = this._bg_img
this._pattern = pattern
} else {
pattern = this._pattern
}
if (pattern) {
ctx.fillStyle = pattern
ctx.fillRect(
this.visible_area[0],
this.visible_area[1],
this.visible_area[2],
this.visible_area[3]
)
ctx.fillStyle = 'transparent'
}
ctx.globalAlpha = 1.0
ctx.imageSmoothingEnabled = ctx.imageSmoothingEnabled = true //= ctx.mozImageSmoothingEnabled
}
//groups
if (this.graph._groups.length && !this.live_mode) {
this.drawGroups(canvas, ctx)
}
if (this.onDrawBackground) {
this.onDrawBackground(ctx, this.visible_area)
}
if (this.onBackgroundRender) {
//LEGACY
console.error(
'WARNING! onBackgroundRender deprecated, now is named onDrawBackground '
)
this.onBackgroundRender = null
}
//DEBUG: show clipping area
//ctx.fillStyle = "red";
//ctx.fillRect( this.visible_area[0] + 10, this.visible_area[1] + 10, this.visible_area[2] - 20, this.visible_area[3] - 20);
//bg
if (this.render_canvas_border) {
ctx.strokeStyle = '#235'
ctx.strokeRect(0, 0, canvas.width, canvas.height)
}
if (this.render_connections_shadows) {
ctx.shadowColor = '#000'
ctx.shadowOffsetX = 0
ctx.shadowOffsetY = 0
ctx.shadowBlur = 6
} else {
ctx.shadowColor = 'rgba(0,0,0,0)'
}
//draw connections
if (!this.live_mode) {
this.drawConnections(ctx)
}
ctx.shadowColor = 'rgba(0,0,0,0)'
//restore state
ctx.restore()
}
if (ctx.finish) {
ctx.finish()
}
this.dirty_bgcanvas = false
this.dirty_canvas = true //to force to repaint the front canvas with the bgcanvas
}
function imgToCanvasBase64 (img) {
const canvas = document.createElement('canvas')
const ctx = canvas.getContext('2d')
canvas.width = img.width
canvas.height = img.height
ctx.drawImage(img, 0, 0)
const base64 = canvas.toDataURL('image/png')
return base64
}
// 使用示例
function convertImageToBase64 (img) {
// const img = new Image()
// img.src = 'path/to/your/image.jpg' // 替换为你的图片路径
// console.log('convertImageToBase64',img)
try {
const base64 = imgToCanvasBase64(img)
return base64
} catch (error) {
console.error(error)
}
}
function getInputsAndOutputs () {
const outputs =
`PreviewImage,SaveImage,TransparentImage,VHS_VideoCombine,VideoCombine_Adv,Image Save,SaveImageAndMetadata_`.split(
','
)
let outputsId = []
for (let node of app.graph._nodes) {
if (outputs.includes(node.type)) {
outputsId.push(node.id)
}
}
return outputsId
}
function getRandomElement (arr) {
const randomIndex = Math.floor(Math.random() * arr.length)
return arr[randomIndex]
}
async function getBG () {
var outputs = []
for (let id of app.graph
.getNodeById(50)
.widgets.filter(w => w.name === 'output_ids')[0]
.value.split('\n')) {
if (getInputsAndOutputs().map(Number).includes(Number(id))) {
if (app.graph.getNodeById(id).imgs && app.graph.getNodeById(id).imgs[0]) {
let b = convertImageToBase64(app.graph.getNodeById(id).imgs[0])
// console.log(b)
outputs.push(b)
}
}
}
var BACKGROUND_IMAGE = getRandomElement(outputs),
CLEAR_BACKGROUND_COLOR = 'rgba(0,0,0,0.9)'
if (!window._bg_img) {
window._bg_img = app.canvas._bg_img.src
}
// let img=new Image();
// img.src=BACKGROUND_IMAGE;
//去掉透明度过度
// app.canvas.zoom_modify_alpha=false;
//整体透明度
app.canvas.editor_alpha = 1.1
// app.canvas._pattern=ctx.createPattern(img, "no-repeat");
app.canvas.updateBackground(BACKGROUND_IMAGE, CLEAR_BACKGROUND_COLOR)
app.canvas.draw(true, true)
}
class BgRunner {
constructor () {
this.intervalId = null
this.running = false
}
// 要运行的方法
bg () {
console.log('方法bg正在运行')
getBG()
}
// 启动bg方法每秒运行一次
start () {
if (!this.running) {
this.intervalId = setInterval(() => this.bg(), 1500)
this.running = true
}
}
// 停止bg方法的运行
stop () {
if (this.running) {
clearInterval(this.intervalId)
this.intervalId = null
this.running = false
if (window._bg_img) {
var BACKGROUND_IMAGE = window._bg_img,
CLEAR_BACKGROUND_COLOR = 'rgba(0,0,0,1)'
app.canvas.editor_alpha = 1
app.canvas.updateBackground(BACKGROUND_IMAGE, CLEAR_BACKGROUND_COLOR)
app.canvas.draw(true, true)
}
}
}
// 切换start和stop
toggle () {
if (this.running) {
this.stop()
} else {
this.start()
}
}
// 获取运行状态
isRunning () {
return this.running
}
}
// 示例用法
// const runner = new BgRunner();
// runner.start();
// setTimeout(() => runner.stop(), 5000);
export const td_bg = new BgRunner()
+39 -47
View File
@@ -100,7 +100,7 @@ async function start_llama (model = 'Phi-3-mini-4k-instruct-Q5_K_S.gguf') {
})
const data = await response.json()
if (data.llama_cpp_error||!data.port) {
if (data.llama_cpp_error) {
return
}
@@ -163,20 +163,17 @@ async function createMenu () {
// appsButton.onclick = () =>
appsButton.onclick = async () => {
// if (window._mixlab_llamacpp&&window._mixlab_llamacpp.model&&window._mixlab_llamacpp.model.length>0) {
// //显示运行的模型
// createModelsModal([
// window._mixlab_llamacpp.url,
// window._mixlab_llamacpp.model
// ])
// } else {
// // let ms = await get_llamafile_models()
// // ms = ms.filter(m => !m.match('-mmproj-'))
// // if (ms.length > 0) createModelsModal(ms)
// }
createModelsModal([
])
if (window._mixlab_llamacpp) {
//显示运行的模型
createModelsModal([
window._mixlab_llamacpp.url,
window._mixlab_llamacpp.model
])
} else {
let ms = await get_llamafile_models()
ms = ms.filter(m => !m.match('-mmproj-'))
if (ms.length > 0) createModelsModal(ms)
}
}
menu.append(appsButton)
}
@@ -803,11 +800,11 @@ async function fetchReadmeContent (url) {
async function startLLM (model) {
let res = await start_llama(model)
window._mixlab_llamacpp = res||{ model:[] }
window._mixlab_llamacpp = res
localStorage.setItem('_mixlab_llama_select', res?.model||'')
localStorage.setItem('_mixlab_llama_select', res.model)
if (document.body.querySelector('#mixlab_chatbot_by_llamacpp')&&window._mixlab_llamacpp?.url) {
if (document.body.querySelector('#mixlab_chatbot_by_llamacpp')&&window._mixlab_llamacpp.url) {
document.body
.querySelector('#mixlab_chatbot_by_llamacpp')
.setAttribute('title', window._mixlab_llamacpp.url)
@@ -935,16 +932,16 @@ function createModelsModal (models) {
const n_gpu_p = document.createElement('p')
n_gpu_p.innerText = 'n_gpu_layers'
const batchPageBtn = document.createElement('div')
batchPageBtn.style = `display: flex;
const n_gpu_div = document.createElement('div')
n_gpu_div.style = `display: flex;
justify-content: center;
align-items: center;
font-size: 12px;`
batchPageBtn.innerHTML=`<a href="${get_url()}/mixlab/app" target="_blank" style="color: var(--input-text);
background-color: var(--comfy-input-bg);">App</a>`
n_gpu_div.appendChild(n_gpu_p)
n_gpu_div.appendChild(n_gpu)
const title = document.createElement('p')
title.innerText = 'Mixlab Nodes'
title.innerText = 'Models'
title.style = `font-size: 18px;
margin-right: 8px;
margin-top: 0;`
@@ -956,9 +953,9 @@ function createModelsModal (models) {
font-size: 12px;
flex-direction: column; `
left_d.appendChild(title)
// title.appendChild(statusIcon)
// left_d.appendChild(linkIcon)
left_d.appendChild(batchPageBtn)
title.appendChild(statusIcon)
left_d.appendChild(linkIcon)
left_d.appendChild(n_gpu_div)
headTitleElement.appendChild(left_d)
// headTitleElement.appendChild(n_gpu_div)
@@ -1013,26 +1010,26 @@ function createModelsModal (models) {
var modalContent = document.createElement('div')
modalContent.classList.add('modal-content')
var inputForSystemPrompt = document.createElement('textarea')
inputForSystemPrompt.className = 'comfy-multiline-input'
inputForSystemPrompt.style = ` height: 260px;
var input = document.createElement('textarea')
input.className = 'comfy-multiline-input'
input.style = ` height: 260px;
width: 480px;
font-size: 16px;
padding: 18px;`
inputForSystemPrompt.value = localStorage.getItem('_mixlab_system_prompt')
input.value = localStorage.getItem('_mixlab_system_prompt')
inputForSystemPrompt.addEventListener('change', e => {
input.addEventListener('change', e => {
e.stopPropagation()
localStorage.setItem('_mixlab_system_prompt', inputForSystemPrompt.value)
localStorage.setItem('_mixlab_system_prompt', input.value)
})
inputForSystemPrompt.addEventListener('click', e => {
input.addEventListener('click', e => {
e.stopPropagation()
})
// modalContent.appendChild(inputForSystemPrompt)
modalContent.appendChild(input)
if (!window._mixlab_llamacpp||(window._mixlab_llamacpp?.model?.length==0)) {
if (!window._mixlab_llamacpp) {
for (const m of models) {
let d = document.createElement('div')
d.innerText = `${showTextByLanguage('Run', {
@@ -1043,10 +1040,10 @@ function createModelsModal (models) {
d.addEventListener('click', async e => {
e.stopPropagation()
div.remove()
// startLLM(m)
startLLM(m)
})
// modalContent.appendChild(d)
modalContent.appendChild(d)
}
}
modal.appendChild(modalContent)
@@ -1417,7 +1414,7 @@ app.registerExtension({
.setAttribute('title', res.url)
})
}else{
// startLLM('')
startLLM('')
}
LGraphCanvas.prototype.helpAboutNode = async function (node) {
@@ -1442,14 +1439,10 @@ app.registerExtension({
LGraphCanvas.prototype.fixTheNode = function (node) {
let new_node = LiteGraph.createNode(node.comfyClass)
console.log(node)
if(new_node){
new_node.pos = [node.pos[0], node.pos[1]]
app.canvas.graph.add(new_node, false)
copyNodeValues(node, new_node)
app.canvas.graph.remove(node)
}
new_node.pos = [node.pos[0], node.pos[1]]
app.canvas.graph.add(new_node, false)
copyNodeValues(node, new_node)
app.canvas.graph.remove(node)
}
smart_init()
@@ -1791,7 +1784,6 @@ app.registerExtension({
{
content: 'Help ♾️Mixlab', // with a name
callback: () => {
// console.log('#data',node)
LGraphCanvas.prototype.helpAboutNode(node)
} // and the callback
},
+26 -143
View File
@@ -1,5 +1,5 @@
import { app } from '../../../scripts/app.js'
import { $el } from '../../../scripts/ui.js'
import { $el } from '../../../scripts/ui.js'
const getLocalData = key => {
let data = {}
@@ -122,7 +122,7 @@ app.registerExtension({
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
console.log('Color nodeData', this.div)
// console.log('Color nodeData', this.widgets)
const widget = {
type: 'div',
@@ -273,19 +273,19 @@ app.registerExtension({
})
const min_max = node => {
if (node.widgets) {
if(node.widgets){
const min_value = node.widgets.filter(w => w.name === 'min_value')[0]
const max_value = node.widgets.filter(w => w.name === 'max_value')[0]
const number = node.widgets.filter(w => w.name === 'number')[0]
if (number) {
number.options.min = min_value.value
number.options.max = max_value.value
number.value = Math.min(number.options.max, number.value)
number.value = Math.max(number.options.min, number.value)
}
if (min_value)
min_value.callback = e => {
number.options.min = e
@@ -297,18 +297,22 @@ const min_max = node => {
number.value = e
}
}
}
app.registerExtension({
name: 'Mixlab.utils.FloatSlider',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'FloatSlider') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
const orig_nodeCreated = nodeType.prototype.onNodeCreated;
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
min_max(this)
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'FloatSlider') {
@@ -319,6 +323,7 @@ app.registerExtension({
app.registerExtension({
name: 'Mixlab.utils.IntNumber',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'IntNumber') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
@@ -326,6 +331,7 @@ app.registerExtension({
min_max(this)
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'IntNumber') {
@@ -334,145 +340,22 @@ app.registerExtension({
}
})
app.registerExtension({
name: 'Mixlab.utils.TESTNODE_',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'TESTNODE_') {
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
console.log('##', message)
}
const onExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments);
console.log('##',message)
};
}
}
})
app.registerExtension({
name: 'Mixlab.utils.KeyInput',
init () {},
async getCustomWidgets (app) {
return {
KEY (node, inputName, inputData, app) {
// console.log('##node', node)
const widget = {
type: inputData[0], // the type, CHEESE
name: inputName, // the name, slice
size: [128, 24], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_api_key')
return data[node.id] || 'by Mixlab'
}
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'KeyInput') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const widget = {
type: 'div',
name: 'input_key',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 24, node.size[1])
)
}
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
const inputDiv = (key, placeholder) => {
let div = document.createElement('div')
div.style = `
display: flex;
align-items: center;
margin: 6px 8px;
margin-top:0px;
height:44px;
width:220px;
`
const ip = document.createElement('input')
ip.type = 'password'
ip.className = `${'comfy-multiline-input'} ${placeholder}`
ip.placeholder = placeholder
// ip.value = placeholder
ip.style = `margin-left:8px;
outline: none;
border: none;
padding:12px;
width: 100%;
`
div.appendChild(ip)
ip.addEventListener('change', () => {
let data = getLocalData(key)
data[this.id] = ip.value.trim()
localStorage.setItem(key, JSON.stringify(data))
})
return div
}
let inputKey = inputDiv('_mixlab_api_key', 'Key')
widget.div.appendChild(inputKey)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
inputKey.remove()
widget.div.remove()
return onRemoved?.()
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'KeyInput') {
let widget = node.widgets.filter(w => w.div)[0]
let apiKey = getLocalData('_mixlab_api_key')
let id = node.id
if (widget.div.querySelector('.Key'))
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
}
},
nodeCreated (node, app) {
//数据延迟??
setTimeout(() => {
// console.log('#LoadImagesToBatch', node.type)
if (node.type === 'KeyInput') {
let widget = node.widgets.filter(w => w.div)[0]
let apiKey = getLocalData('_mixlab_api_key')
let id = node.id
if (widget.div.querySelector('.Key'))
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
}
}, 1000)
}
},
})
+49 -45
View File
@@ -6,6 +6,8 @@ import { $el } from '../../../scripts/ui.js'
// The code is based on ComfyUI-VideoHelperSuite modification.
function injectCSS (css) {
// 检查页面中是否已经存在具有相同内容的style标签
const existingStyle = document.querySelector('style')
@@ -238,7 +240,15 @@ app.registerExtension({
}
})
function offsetDOMWidget (widget, ctx, node, widgetWidth, widgetY, height) {
function offsetDOMWidget(
widget,
ctx,
node,
widgetWidth,
widgetY,
height
) {
const margin = 10
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
@@ -260,18 +270,18 @@ function offsetDOMWidget (widget, ctx, node, widgetWidth, widgetY, height) {
position: 'absolute',
background: !node.color ? '' : node.color,
color: !node.color ? '' : 'white',
zIndex: 5 //app.graph._nodes.indexOf(node),
zIndex: 5, //app.graph._nodes.indexOf(node),
})
}
export const hasWidgets = node => {
export const hasWidgets = (node) => {
if (!node.widgets || !node.widgets?.[Symbol.iterator]) {
return false
}
return true
}
export const cleanupNode = node => {
export const cleanupNode = (node) => {
if (!hasWidgets(node)) {
return
}
@@ -288,43 +298,43 @@ export const cleanupNode = node => {
}
}
const createPreviewElement = (name, val, format) => {
const [type] = format.split('/')
const CreatePreviewElement = (name, val, format) => {
const [type] = format.split('/')
const w = {
name,
type,
value: val,
draw: function (ctx, node, widgetWidth, widgetY, height) {
const [cw, ch] = this.computeSize(widgetWidth)
offsetDOMWidget(this, ctx, node, widgetWidth, widgetY, ch)
},
computeSize: function (_) {
const ratio = this.inputRatio || 1
const width = Math.max(220, this.parent.size[0])
return [width, width / ratio + 10]
},
onRemoved: function () {
if (this.inputEl) {
this.inputEl.remove()
}
name,
type,
value: val,
draw: function (ctx, node, widgetWidth, widgetY, height) {
const [cw, ch] = this.computeSize(widgetWidth)
offsetDOMWidget(this, ctx, node, widgetWidth, widgetY, ch)
},
computeSize: function (_) {
const ratio = this.inputRatio || 1
const width = Math.max(220, this.parent.size[0])
return [width, (width / ratio + 10)]
},
onRemoved: function () {
if (this.inputEl) {
this.inputEl.remove()
}
},
}
w.inputEl = document.createElement(type === 'video' ? 'video' : 'img')
w.inputEl.src = w.value
if (type === 'video') {
w.inputEl.setAttribute('type', 'video/webm');
w.inputEl.autoplay = true
w.inputEl.loop = true
w.inputEl.controls = false;
}
w.inputEl.onload = function () {
w.inputRatio = w.inputEl.naturalWidth / w.inputEl.naturalHeight
}
document.body.appendChild(w.inputEl)
return w
}
w.inputEl = document.createElement(type === 'video' ? 'video' : 'img')
w.inputEl.src = w.value
if (type === 'video' || format.match('.mp4')) {
w.inputEl.setAttribute('type', 'video/webm')
w.inputEl.autoplay = true
w.inputEl.loop = true
w.inputEl.controls = true
}
w.inputEl.onload = function () {
w.inputRatio = w.inputEl.naturalWidth / w.inputEl.naturalHeight
}
document.body.appendChild(w.inputEl)
return w
}
app.registerExtension({
name: 'Mixlab.Video.ImageListReplace',
@@ -459,17 +469,12 @@ app.registerExtension({
}
}
if (
nodeData?.name == 'VideoCombine_Adv' ||
nodeData?.name == 'CombineAudioVideo'
) {
if (nodeData?.name == 'VideoCombine_Adv') {
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
const prefix = 'vhs_gif_preview_'
const r = onExecuted ? onExecuted.apply(this, message) : undefined
if(!this.widgets) this.widgets=[]
if (this.widgets) {
const pos = this.widgets.findIndex(w => w.name === `${prefix}_0`)
if (pos !== -1) {
@@ -484,13 +489,12 @@ app.registerExtension({
'/view?' + new URLSearchParams(params).toString()
)
const w = this.addCustomWidget(
createPreviewElement(
CreatePreviewElement(
`${prefix}_${i}`,
previewUrl,
params.format || 'image/gif'
)
)
console.log(w)
w.parent = this
})
}
+525
View File
@@ -0,0 +1,525 @@
<!DOCTYPE html>
<!--
VoiceStreamAI Client Interface
Real-time audio transcription using self-hosted Whisper and WebSocket
Contributor:
- Alessandro Saccoia - alessandro.saccoia@gmail.com
-->
<html lang="en">
<head>
<meta charset="UTF-8">
<title>Audio Stream to WebSocket Server</title>
<style>
body {
font-family: Arial, sans-serif;
margin: 0;
padding: 0;
background: #f4f4f4;
text-align: center;
}
h1 {
color: #333;
}
.controls {
margin: 20px auto;
padding: 10px;
width: 80%;
display: flex;
justify-content: space-around;
align-items: center;
}
.control-group {
display: flex;
flex-direction: column;
align-items: center;
}
.controls input,
.controls button,
.controls select {
padding: 8px;
margin: 5px;
border: 1px solid #ddd;
border-radius: 5px;
font-size: 0.9em;
}
#transcription {
margin: 20px auto;
border: 1px solid #ddd;
padding: 10px;
width: 80%;
height: 150px;
overflow-y: auto;
background: white;
}
.label {
font-size: 0.9em;
color: #555;
margin-bottom: 5px;
}
button {
cursor: pointer;
}
.buffering-strategy-panel {
margin-top: 10px;
}
/* ... existing styles ... */
.hidden {
display: none;
}
</style>
<style>
body {
margin: 0;
padding: 0;
/* background-color: #333;
color:white */
}
#mic_container {
display: flex;
width: 100%;
align-items: center;
justify-content: center;
}
#mic {
border: 1px solid #ddd;
border-radius: 4px;
margin-top: 1rem;
width: 300px
}
#asistant {
width: 300px;
padding: 12px;
font-size: 12px;
overflow-y: scroll;
height: 300px;
}
</style>
</head>
<body>
<div id="mic_container">
<div id="mic"></div>
<div id="asistant"></div>
</div>
<h1>VAD + Whisper + </h1>
<button id="init_server">Server</button>
<div class="controls">
<div class="control-group">
<label class="label" for="websocketAddress">WebSocket Address:</label>
<input type="text" id="websocketAddress" value="ws://localhost:8725">
</div>
<div class="control-group">
<label class="label" for="bufferingStrategySelect" onchange="toggleBufferingStrategyPanel()">Buffering
Strategy:</label>
<select id="bufferingStrategySelect">
<option value="silence_at_end_of_chunk" selected>Silence at End of Chunk</option>
</select>
</div>
<div class="silence_at_end_of_chunk_options_panel">
<div class="control-group">
<label class="label" for="chunk_length_seconds">Chunk Length (s):</label>
<input type="number" id="chunk_length_seconds" value="3" min="1">
</div>
<div class="control-group">
<label class="label" for="chunk_offset_seconds">Silence at the End of Chunk (s):</label>
<input type="number" id="chunk_offset_seconds" value="0.1" min="0">
</div>
</div>
<div class="control-group">
<label class="label" for="languageSelect">Language:</label>
<select id="languageSelect">
<option value="multilingual">Multilingual</option>
<option value="english">English</option>
<option value="italian">Italian</option>
<option value="spanish">Spanish</option>
<option value="french">French</option>
<option value="german">German</option>
<option value="chinese">Chinese</option>
<option value="arabic">Arabic</option>
<option value="portuguese">Portuguese</option>
<option value="russian">Russian</option>
<option value="japanese">Japanese</option>
<option value="dutch">Dutch</option>
<option value="korean">Korean</option>
<option value="hindi">Hindi</option>
<option value="turkish">Turkish</option>
<option value="swedish">Swedish</option>
<option value="norwegian">Norwegian</option>
<option value="danish">Danish</option>
<option value="polish">Polish</option>
<option value="finnish">Finnish</option>
<option value="thai">Thai</option>
<option value="czech">Czech</option>
<option value="hungarian">Hungarian</option>
<option value="greek">Greek</option>
</select>
</div>
<button id="connectButton">Connect</button>
</div>
<button id="startButton" disabled>Start Streaming</button>
<button id="stopButton" disabled>Stop Streaming</button>
<div id="transcription"></div>
<br />
<div>WebSocket: <span id="webSocketStatus">Not Connected</span></div>
<div>Detected Language: <span id="detected_language">Undefined</span></div>
<div>Last Processing Time: <span id="processing_time">Undefined</span></div>
<script type="module">
// Record plugin
import WaveSurfer from 'https://cdn.jsdelivr.net/npm/wavesurfer.js@7/dist/wavesurfer.esm.js'
import RecordPlugin from 'https://cdn.jsdelivr.net/npm/wavesurfer.js/dist/plugins/record.esm.js'
// 可视化
let wavesurfer, record;
// ws服务
let websocket;
let context;
let processor;
let globalStream;
const websocket_uri = 'ws://localhost:8765';
const bufferSize = 4096;
let isRecording = false;
let chunk_length_seconds, chunk_offset_seconds, language;
// Record button
const startButton = document.getElementById('startButton'),
stopButton = document.getElementById('stopButton'),
connectButton = document.getElementById('connectButton'),
initServerButton = document.getElementById('init_server')
const createWaveSurfer = () => {
// Create an instance of WaveSurfer
if (wavesurfer) {
wavesurfer.destroy()
}
wavesurfer = WaveSurfer.create({
container: '#mic',
waveColor: 'rgb(200, 0, 200)',
progressColor: 'rgb(100, 0, 100)',
renderFunction: (channels, ctx) => {
const { width, height } = ctx.canvas
// console.log(width, height)
const scale = channels[0].length / width
const step = 20
ctx.translate(0, height / 2)
ctx.strokeStyle = ctx.fillStyle
ctx.beginPath()
for (let i = 0; i < width; i += step * 2) {
const index = Math.floor(i * scale)
const value = Math.abs(channels[0][index])
let x = i
let y = value * height * 1.2
ctx.moveTo(x, 0)
ctx.lineTo(x, y)
ctx.arc(x + step / 2, y, step / 2, Math.PI, 0, true)
ctx.lineTo(x + step, 0)
x = x + step
y = -y
ctx.moveTo(x, 0)
ctx.lineTo(x, y)
ctx.arc(x + step / 2, y, step / 2, Math.PI, 0, false)
ctx.lineTo(x + step, 0)
}
ctx.stroke()
ctx.closePath()
},
// // Set a bar width
// barWidth: 20,
// // Optionally, specify the spacing between bars
// barGap: 4,
// // And the bar radius
// barRadius: 2,
})
// Initialize the Record plugin
record = wavesurfer.registerPlugin(RecordPlugin.create({ scrollingWaveform: false, renderRecordedAudio: false }))
// Render recorded audio
// recButton.textContent = 'Record'
}
startButton.addEventListener('click', e => {
record.startRecording()
if (isRecording) return;
isRecording = true;
const AudioContext = window.AudioContext || window.webkitAudioContext;
context = new AudioContext();
navigator.mediaDevices.getUserMedia({ audio: true }).then(stream => {
globalStream = stream;
const input = context.createMediaStreamSource(stream);
processor = context.createScriptProcessor(bufferSize, 1, 1);
processor.onaudioprocess = e => processAudio(e);
input.connect(processor);
processor.connect(context.destination);
sendAudioConfig();
}).catch(error => console.error('Error accessing microphone', error));
// Disable start button and enable stop button
startButton.disabled = true;
stopButton.disabled = false;
});
stopButton.addEventListener('click', e => {
if (!isRecording) return;
isRecording = false;
if (globalStream) {
globalStream.getTracks().forEach(track => track.stop());
}
if (processor) {
processor.disconnect();
processor = null;
}
if (context) {
context.close().then(() => context = null);
}
startButton.disabled = false;
stopButton.disabled = true;
// 可视化部分
if (record.isRecording() || record.isPaused()) {
record.stopRecording()
return
}
})
connectButton.addEventListener('click', e => {
initWebSocket()
})
initServerButton.addEventListener('click', async e => {
const response = await fetch('/mixlab/start_live', {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
port: 8323,
model: ''
})
})
console.log(await response.json())
})
// Initialize on page load
window.onload = () => {
initWebSocket()
createWaveSurfer();
};
function initWebSocket() {
const websocketAddress = document.getElementById('websocketAddress').value;
chunk_length_seconds = document.getElementById('chunk_length_seconds').value;
chunk_offset_seconds = document.getElementById('chunk_offset_seconds').value;
const selectedLanguage = document.getElementById('languageSelect').value;
language = selectedLanguage !== 'multilingual' ? selectedLanguage : null;
if (!websocketAddress) {
console.log("WebSocket address is required.");
return;
}
if (websocket) websocket.close()
websocket = new WebSocket(websocketAddress);
websocket.onopen = () => {
console.log("WebSocket connection established");
document.getElementById("webSocketStatus").textContent = 'Connected';
startButton.disabled = false;
};
websocket.onclose = event => {
console.log("WebSocket connection closed", event);
document.getElementById("webSocketStatus").textContent = 'Not Connected';
stopButton.click()
startButton.disabled = true;
stopButton.disabled = true;
// setTimeout(()=>initWebSocket(),1000)
};
websocket.onmessage = event => {
console.log("Message from server:", event.data);
const transcript_data = JSON.parse(event.data);
if (transcript_data.status === 'chat_start') {
updateTranscription(transcript_data);
stopButton.click();
} else if (transcript_data.status === 'chat_end') {
let asistant = decodeURIComponent(transcript_data.asistant)
document.getElementById('asistant').innerText = asistant;
startButton.click();
}
};
websocket.onerror = () => {
// setTimeout(()=>initWebSocket(),1000)
}
}
function updateTranscription(transcript_data) {
const transcriptionDiv = document.getElementById('transcription');
const languageDiv = document.getElementById('detected_language');
if (transcript_data['words'] && transcript_data['words'].length > 0) {
// Append words with color based on their probability
transcript_data['words'].forEach(wordData => {
const span = document.createElement('span');
const probability = wordData['probability'];
span.textContent = wordData['word'] + ' ';
// Set the color based on the probability
if (probability > 0.9) {
span.style.color = 'green';
} else if (probability > 0.6) {
span.style.color = 'orange';
} else {
span.style.color = 'red';
}
transcriptionDiv.appendChild(span);
});
// Add a new line at the end
transcriptionDiv.appendChild(document.createElement('br'));
} else {
// Fallback to plain text
transcriptionDiv.textContent += transcript_data['text'] + '\n';
}
// Update the language information
if (transcript_data['language'] && transcript_data['language_probability']) {
languageDiv.textContent = transcript_data['language'] + ' (' + transcript_data['language_probability'].toFixed(2) + ')';
}
// Update the processing time, if available
const processingTimeDiv = document.getElementById('processing_time');
if (transcript_data['processing_time']) {
processingTimeDiv.textContent = 'Processing time: ' + transcript_data['processing_time'].toFixed(2) + ' seconds';
}
}
function sendAudioConfig() {
let selectedStrategy = document.getElementById('bufferingStrategySelect').value;
let processingArgs = {};
if (selectedStrategy === 'silence_at_end_of_chunk') {
processingArgs = {
chunk_length_seconds: parseFloat(document.getElementById('chunk_length_seconds').value),
chunk_offset_seconds: parseFloat(document.getElementById('chunk_offset_seconds').value)
};
}
const audioConfig = {
type: 'config',
data: {
sampleRate: context.sampleRate,
bufferSize: bufferSize,
channels: 1, // Assuming mono channel
language: language,
processing_strategy: selectedStrategy,
processing_args: processingArgs
}
};
websocket.send(JSON.stringify(audioConfig));
}
function downsampleBuffer(buffer, inputSampleRate, outputSampleRate) {
if (inputSampleRate === outputSampleRate) {
return buffer;
}
var sampleRateRatio = inputSampleRate / outputSampleRate;
var newLength = Math.round(buffer.length / sampleRateRatio);
var result = new Float32Array(newLength);
var offsetResult = 0;
var offsetBuffer = 0;
while (offsetResult < result.length) {
var nextOffsetBuffer = Math.round((offsetResult + 1) * sampleRateRatio);
var accum = 0, count = 0;
for (var i = offsetBuffer; i < nextOffsetBuffer && i < buffer.length; i++) {
accum += buffer[i];
count++;
}
result[offsetResult] = accum / count;
offsetResult++;
offsetBuffer = nextOffsetBuffer;
}
return result;
}
function processAudio(e) {
const inputSampleRate = context.sampleRate;
const outputSampleRate = 16000; // Target sample rate
const left = e.inputBuffer.getChannelData(0);
const downsampledBuffer = downsampleBuffer(left, inputSampleRate, outputSampleRate);
const audioData = convertFloat32ToInt16(downsampledBuffer);
if (websocket && websocket.readyState === WebSocket.OPEN) {
websocket.send(audioData);
}
}
function convertFloat32ToInt16(buffer) {
let l = buffer.length;
const buf = new Int16Array(l);
while (l--) {
buf[l] = Math.min(1, buffer[l]) * 0x7FFF;
}
return buf.buffer;
}
function toggleBufferingStrategyPanel() {
var selectedStrategy = document.getElementById('bufferingStrategySelect').value;
if (selectedStrategy === 'silence_at_end_of_chunk') {
var panel = document.getElementById('silence_at_end_of_chunk_options_panel');
panel.classList.remove('hidden');
} else {
var panel = document.getElementById('silence_at_end_of_chunk_options_panel');
panel.classList.add('hidden');
}
}
</script>
</body>
</html>
+439 -675
View File
File diff suppressed because it is too large Load Diff
-408
View File
@@ -1,408 +0,0 @@
{
"last_node_id": 21,
"last_link_id": 16,
"nodes": [
{
"id": 10,
"type": "ChatGPTOpenAI",
"pos": [
489,
689
],
"size": {
"0": 403.2580261230469,
"1": 309.2166442871094
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "api_key",
"type": "STRING",
"link": null,
"widget": {
"name": "api_key"
}
},
{
"name": "custom_model_name",
"type": "STRING",
"link": null,
"widget": {
"name": "custom_model_name"
}
},
{
"name": "custom_api_url",
"type": "STRING",
"link": 13,
"widget": {
"name": "custom_api_url"
},
"slot_index": 2
}
],
"outputs": [
{
"name": "text",
"type": "STRING",
"links": [
12
],
"shape": 3,
"slot_index": 0
},
{
"name": "messages",
"type": "STRING",
"links": null,
"shape": 3
},
{
"name": "session_history",
"type": "STRING",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "ChatGPTOpenAI"
},
"widgets_values": [
"hi",
"You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"gpt-3.5-turbo",
447210757728856,
"randomize",
1,
"openai",
"",
"",
""
]
},
{
"id": 3,
"type": "ShowTextForGPT",
"pos": [
982,
686
],
"size": {
"0": 400,
"1": 200
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "text",
"type": "STRING",
"link": 12,
"widget": {
"name": "text"
}
},
{
"name": "output_dir",
"type": "STRING",
"link": null,
"widget": {
"name": "output_dir"
}
}
],
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": null,
"shape": 6
}
],
"properties": {
"Node name for S&R": "ShowTextForGPT"
},
"widgets_values": [
"",
"",
" Hi there! What can I help you with?"
]
},
{
"id": 11,
"type": "SiliconflowLLM",
"pos": [
489,
318
],
"size": {
"0": 395.197998046875,
"1": 262
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "api_key",
"type": "STRING",
"link": 16,
"widget": {
"name": "api_key"
}
},
{
"name": "custom_model_name",
"type": "STRING",
"link": null,
"widget": {
"name": "custom_model_name"
}
}
],
"outputs": [
{
"name": "text",
"type": "STRING",
"links": [
15
],
"shape": 3,
"slot_index": 0
},
{
"name": "messages",
"type": "STRING",
"links": null,
"shape": 3
},
{
"name": "session_history",
"type": "STRING",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "SiliconflowLLM"
},
"widgets_values": [
"",
"",
"You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"Qwen/Qwen2-7B-Instruct",
593422808835285,
"randomize",
1,
""
]
},
{
"id": 17,
"type": "ShowTextForGPT",
"pos": [
975,
329
],
"size": {
"0": 400,
"1": 200
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "text",
"type": "STRING",
"link": 15,
"widget": {
"name": "text"
}
},
{
"name": "output_dir",
"type": "STRING",
"link": null,
"widget": {
"name": "output_dir"
}
}
],
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": null,
"shape": 6
}
],
"properties": {
"Node name for S&R": "ShowTextForGPT"
},
"widgets_values": [
"",
"",
"Hello! How can I assist you today?"
]
},
{
"id": 18,
"type": "KeyInput",
"pos": [
46,
319
],
"size": {
"0": 315,
"1": 70
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "key",
"type": "STRING",
"links": [
16
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "KeyInput"
},
"widgets_values": [
null,
null
]
},
{
"id": 12,
"type": "TextInput_",
"pos": [
31,
871
],
"size": [
407.9377612789413,
76
],
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": [
13
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "TextInput_"
},
"widgets_values": [
"http://127.0.0.1:8000/v1"
]
},
{
"id": 20,
"type": "Note",
"pos": [
35,
664
],
"size": [
350.04604707424306,
116.54209784249178
],
"flags": {},
"order": 2,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"api_key 填写对应平台的Key\ncustom model和api 根据需要自行填写\n\n如果不填写custome,则按照model和api_url选择的选项"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 21,
"type": "Note",
"pos": [
42,
438
],
"size": {
"0": 350.0460510253906,
"1": 116.54209899902344
},
"flags": {},
"order": 3,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"API key节点不会保存到workflow的json文件。\n\n::会保存到appinfo导出的app.json里\n\n\n注册https://cloud.siliconflow.cn/account/ak 领取免费的API"
],
"color": "#432",
"bgcolor": "#653"
}
],
"links": [
[
12,
10,
0,
3,
0,
"STRING"
],
[
13,
12,
0,
10,
2,
"STRING"
],
[
15,
11,
0,
17,
0,
"STRING"
],
[
16,
18,
0,
11,
0,
"STRING"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.9646149645000006,
"offset": [
170.81398913081276,
-128.0066534315481
]
}
},
"version": 0.4
}
File diff suppressed because it is too large Load Diff