Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6cb2b57463 | ||
|
|
be8ccc1dc4 | ||
|
|
228e5d9183 | ||
|
|
8afe6d0383 | ||
|
|
5f7190b08f | ||
|
|
a70a9b4bb1 | ||
|
|
b796e66890 | ||
|
|
f1a663779a | ||
|
|
b0aa972326 | ||
|
|
ef927a7ed1 | ||
|
|
aa8fc59051 | ||
|
|
ce62204392 | ||
|
|
837f28142d | ||
|
|
60c79c991d | ||
|
|
078aaeb679 | ||
|
|
d9edbd535e | ||
|
|
b4a61b21c3 | ||
|
|
bdc4193ffe | ||
|
|
74fdd6e396 | ||
|
|
b2479ebff2 | ||
|
|
ce2162c764 | ||
|
|
16ffd63c80 | ||
|
|
8faf68348d | ||
|
|
02dbc72856 | ||
|
|
da4dcf92dc | ||
|
|
49b750abcc | ||
|
|
4bb4122628 | ||
|
|
cee54f336e | ||
|
|
e95b3813cc | ||
|
|
6815cfb05e | ||
|
|
b6acbbce35 | ||
|
|
399e74877d | ||
|
|
61083e91a6 | ||
|
|
67b4ec3178 | ||
|
|
0fcb725a7a | ||
|
|
0dbdcdfdc7 | ||
|
|
e426d77353 | ||
|
|
bd15e29f17 | ||
|
|
b323d29567 | ||
|
|
0e54af3356 | ||
|
|
97f12f3bed | ||
|
|
5612047b97 | ||
|
|
2d147a3ae1 | ||
|
|
1a93c0f8e8 | ||
|
|
0a2b64881a | ||
|
|
42e7fe4d93 | ||
|
|
7ada28258c | ||
|
|
bd312afd00 | ||
|
|
d94a8af35b | ||
|
|
078fd10147 | ||
|
|
824e25d77c | ||
|
|
899b887e47 | ||
|
|
e58981d8a3 | ||
|
|
fc41d977a5 | ||
|
|
ab6210e667 | ||
|
|
f41805f053 | ||
|
|
baa809fcd6 | ||
|
|
a38d15e495 | ||
|
|
e97641372a | ||
|
|
9aecc2cb08 | ||
|
|
697667945e | ||
|
|
d908024577 | ||
|
|
a5a656d958 | ||
|
|
ddc3cf05dd | ||
|
|
66ad4b0abd | ||
|
|
a66023adc6 | ||
|
|
7277844128 | ||
|
|
6ef82b1d56 | ||
|
|
8ded4829f3 | ||
|
|
c4b6acb916 | ||
|
|
9beb81c303 | ||
|
|
f8dd4c6efa | ||
|
|
6ce5aa6a3a | ||
|
|
d2efa8a90a | ||
|
|
c6374063e9 | ||
|
|
0846013378 | ||
|
|
c141ba405f | ||
|
|
0320f13a9f | ||
|
|
8adc34be4d | ||
|
|
cb6810d3c1 | ||
|
|
d384f64abf | ||
|
|
ef7035f8ee | ||
|
|
bfcadde5c3 | ||
|
|
496ff41782 | ||
|
|
46f0be5484 | ||
|
|
f3db0131c1 | ||
|
|
83a8d47f51 | ||
|
|
fee0222910 | ||
|
|
1ed7b5511f | ||
|
|
b8f7c31537 | ||
|
|
164791c257 | ||
|
|
8f5e599928 | ||
|
|
7a7aaeb84d | ||
|
|
e2136ab2fc | ||
|
|
c75cb21946 | ||
|
|
bf95218c91 | ||
|
|
8cb4507a5f | ||
|
|
555890d1ba | ||
|
|
e4f54e83b6 | ||
|
|
692c4a709e | ||
|
|
cbd1961459 | ||
|
|
2e31a33ebf | ||
|
|
d16c6137d2 | ||
|
|
0416ab79ec | ||
|
|
fc9a1c62b9 | ||
|
|
5d4567b134 | ||
|
|
ae4a17d271 | ||
|
|
d110a08889 | ||
|
|
e0157293cb | ||
|
|
0d985b3b65 | ||
|
|
a65ade9fda | ||
|
|
874d6c8cb1 | ||
|
|
f70ba2afa3 | ||
|
|
e9f821e578 | ||
|
|
8e488d4b1d | ||
|
|
77201a457d | ||
|
|
076e3b1178 | ||
|
|
6b13fa64dc | ||
|
|
846671a890 | ||
|
|
05b3088b75 | ||
|
|
fe57286959 | ||
|
|
03645bbb33 | ||
|
|
93dba9a399 | ||
|
|
5627ea8073 | ||
|
|
7ba679c9ce | ||
|
|
c7a450e6ce | ||
|
|
beda5156bf | ||
|
|
76a9da7163 | ||
|
|
edd0303f59 | ||
|
|
be6f47a333 | ||
|
|
4cd6a072ca | ||
|
|
743a82efe9 |
@@ -1,19 +1,47 @@
|
||||

|
||||
|
||||
> 适配了最新版 comfyui 的 py3.11 ,torch 2.1.2+cu121
|
||||
> 适配了最新版 comfyui 的 py3.11 ,torch 2.3.1+cu121
|
||||
> [Mixlab nodes discord](https://discord.gg/cXs9vZSqeK)
|
||||
|
||||
商务合作请联系 389570357@qq.com
|
||||
For business cooperation, please contact email 389570357@qq.com
|
||||
|
||||

|
||||
|
||||
##### `最新`:
|
||||
|
||||
- 增加 SiliconflowLLM,可以使用由Siliconflow提供的免费LLM
|
||||
- 新增 SenseVoice
|
||||
|
||||
- 增加 Edit Mask,方便在生成的时候手动绘制 mask [workflow](./workflow/edit-mask-workflow.json)
|
||||
- [新增JS-SDK,方便直接在前端项目中使用comfyui](https://github.com/shadowcz007/comfyui-js-sdk)
|
||||
|
||||
- 新增API调用图像生成节点 TextToImage Siliconflow,可以直接调用Siliconflow提供的flux生成图像
|
||||
|
||||
- [增加 Her 的DEMO页面,和数字人对话](https://github.com/shadowcz007/ComfyUI-Backend-MixlabNodes/blob/main/workflow/her_demo_workflow.json)
|
||||
|
||||
- 右键菜单支持 text-to-text,方便对 prompt 词补全,支持云LLM或者是本地LLM。
|
||||
|
||||
- 增加 MiniCPM-V 2.6 int4
|
||||
|
||||
This is the int4 quantized version of MiniCPM-V 2.6.
|
||||
Running with int4 version would use lower GPU memory (about 7GB).
|
||||
|
||||
- 移动端适配、修改 app 模式的 Mask 编辑器
|
||||
|
||||
- 增加 p5.js 作为输入节点
|
||||
[workflow](./workflow/p5workflow.json)
|
||||
[workflow2](./workflow/p5-video-workflow.json)
|
||||
|
||||
- App 模式增加 batch prompt,批量提示词,可以把动态提示词批量组成后运行
|
||||
|
||||

|
||||
|
||||
- 增加 API Key Input 节点,用于管理 LLM 的 Key,同时优化 LLM 相关节点,为后续 agent 模式做准备
|
||||
|
||||
- 增加 SiliconflowLLM,可以使用由 Siliconflow 提供的免费 LLM
|
||||
|
||||
<!-- - 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)
|
||||
|
||||
@@ -24,7 +52,6 @@
|
||||

|
||||
 -->
|
||||
|
||||
|
||||
#### `相关插件推荐`
|
||||
|
||||
[comfyui-liveportrait](https://github.com/shadowcz007/comfyui-liveportrait)
|
||||
@@ -48,7 +75,8 @@
|
||||
- 发布为 app 的 workflow,可以在右键里再次编辑了
|
||||
- web app 可以设置分类,在 comfyui 右键菜单可以编辑更新 web app
|
||||
- 支持动态提示
|
||||
- 支持把输出显示到comfyui背景(TouchDesigner 风格)
|
||||
- 支持把输出显示到 comfyui 背景(TouchDesigner 风格)
|
||||
- 如果转为 web app 打开是空白的,注意检查下插件目录的名字需要是:comfyui-mixlab-nodes(如果是 zip 包下载会多了个-main 的后缀,需要去掉)
|
||||
|
||||

|
||||
|
||||
@@ -107,15 +135,20 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
|
||||
|
||||
[Voice + Real-time Face Swap Workflow](./workflow/语音+实时换脸workflow.json)
|
||||
|
||||
- Preview Audio
|
||||
|
||||
[text-to-audio](./workflow/text-to-audio-base-workflow.json)
|
||||
|
||||
### GPT
|
||||
|
||||
> 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
|
||||
> Support for calling multiple GPTs.Local LLM 、 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)
|
||||
|
||||
[workflow-5](./workflow/5-gpt-workflow.json)
|
||||
- SiliconflowLLM
|
||||
- ChatGPTOpenAI
|
||||
|
||||
最新:ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
|
||||
<!-- 最新:ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
|
||||
|
||||
Model download,move to :`models/llamafile/`
|
||||
|
||||
@@ -143,7 +176,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
|
||||
|
||||
@@ -172,7 +205,6 @@ pip install llama-cpp-python \
|
||||
|
||||
> 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"
|
||||
|
||||
|
||||

|
||||
|
||||

|
||||
@@ -206,9 +238,16 @@ pip install llama-cpp-python \
|
||||
|
||||
#### TextImage
|
||||
|
||||
> [下载字体](https://drxie.github.io/OSFCC/)放到 ```custom_nodes/comfyui-mixlab-nodes/assets/fonts```
|
||||
> [下载字体](https://drxie.github.io/OSFCC/)放到 `custom_nodes/comfyui-mixlab-nodes/assets/fonts`
|
||||
|
||||
#### MiniCPM-VQA Simple
|
||||
|
||||
This is the int4 quantized version of MiniCPM-V 2.6.
|
||||
Running with int4 version would use lower GPU memory (about 7GB).
|
||||
|
||||
[模型](https://huggingface.co/openbmb/MiniCPM-V-2_6-int4)
|
||||
|
||||

|
||||
|
||||
### Style
|
||||
|
||||
@@ -232,6 +271,8 @@ pip install llama-cpp-python \
|
||||
|
||||
### Other Nodes
|
||||
|
||||
- 增加 Edit Mask,方便在生成的时候手动绘制 mask [workflow](./workflow/edit-mask-workflow.json)
|
||||
|
||||

|
||||

|
||||
|
||||
@@ -247,27 +288,45 @@ Add edges to an image.
|
||||
|
||||

|
||||
|
||||
> LaMaInpainting
|
||||
> LaMaInpainting(需要手动安装)
|
||||
|
||||
- simple-lama-inpainting 里的 pillow 造成冲突,暂时从依赖里移除,如果有安装 simple-lama-inpainting ,节点会自动添加,没有,则不会自动添加。
|
||||
|
||||
from [simple-lama-inpainting](https://github.com/enesmsahin/simple-lama-inpainting)
|
||||
|
||||
- [问题汇总](https://github.com/shadowcz007/comfyui-mixlab-nodes/issues/294)
|
||||
|
||||
> rembgNode
|
||||
|
||||
"briarmbg","u2net","u2netp","u2net_human_seg","u2net_cloth_seg","silueta","isnet-general-use","isnet-anime"
|
||||
|
||||
**_ briarmbg _** model was developed by BRlA Al and can be used as an open-source model for non-commercial purposes
|
||||
|
||||
### Improvement
|
||||
### Enhancement
|
||||
|
||||
- Add "help" option to the context menu for each node.
|
||||
- Add "Nodes Map" option to the global context menu.
|
||||
- Direct "Help" option accessible through node context menu.
|
||||
|
||||
An improvement has been made to directly redirect to GitHub to search for missing nodes when loading the graph.
|
||||
- "Nodes Map" feature added to global context menu.
|
||||
|
||||
- An improvement has been made to directly redirect to GitHub to search for missing nodes when loading the graph.
|
||||
|
||||
*** If not needed, you can comment out ```app.showMissingNodesError``` in the ```ui_mixlab.js``` file.
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
|
||||
- Right-click shortcut
|
||||
|
||||
右键菜单支持 text-to-text,方便对 prompt 词补全,支持云LLM或者是本地LLM。
|
||||
The right-click menu supports text-to-text conversion, facilitating prompt word completion, and supports cloud LLMs or local LLMs.
|
||||
|
||||
Local LLM API example:```http://localhost:1234/v1```
|
||||
|
||||

|
||||
|
||||
|
||||
### Models
|
||||
|
||||
- [Download TripoSR](https://huggingface.co/stabilityai/TripoSR/blob/main/model.ckpt) and place it in `models/triposr`
|
||||
|
||||
|
After Width: | Height: | Size: 537 KiB |
|
After Width: | Height: | Size: 340 KiB |
|
After Width: | Height: | Size: 29 KiB |
|
After Width: | Height: | Size: 366 KiB |
|
After Width: | Height: | Size: 2.2 MiB |
|
Before Width: | Height: | Size: 35 KiB After Width: | Height: | Size: 94 KiB |
@@ -11,9 +11,9 @@ if exist "%python_exec%" (
|
||||
%python_exec% -s -m pip install "%%i" -i https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
)
|
||||
|
||||
%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 --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
|
||||
|
||||
%python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
|
||||
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
|
||||
|
||||
|
||||
) else (
|
||||
|
||||
@@ -6,14 +6,79 @@ import folder_paths
|
||||
import hashlib
|
||||
import codecs,sys
|
||||
import importlib.util
|
||||
import subprocess
|
||||
import requests
|
||||
from PIL import Image
|
||||
from io import BytesIO
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
python = sys.executable
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def is_installed(package):
|
||||
# 从文本中提取json
|
||||
def extract_json_strings(text):
|
||||
json_strings = []
|
||||
brace_level = 0
|
||||
json_str = ''
|
||||
in_json = False
|
||||
|
||||
for char in text:
|
||||
if char == '{':
|
||||
brace_level += 1
|
||||
in_json = True
|
||||
if in_json:
|
||||
json_str += char
|
||||
if char == '}':
|
||||
brace_level -= 1
|
||||
if in_json and brace_level == 0:
|
||||
json_strings.append(json_str)
|
||||
json_str = ''
|
||||
in_json = False
|
||||
|
||||
return json_strings[0] if len(json_strings)>0 else "{}"
|
||||
|
||||
|
||||
def is_installed(package, package_overwrite=None,auto_install=True):
|
||||
is_has=False
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
is_has=spec is not None
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return spec is not None
|
||||
pass
|
||||
|
||||
package = package_overwrite or package
|
||||
|
||||
if spec is None:
|
||||
if auto_install==True:
|
||||
print(f"Installing {package}...")
|
||||
# 清华源 -i https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
command = f'"{python}" -m pip install {package}'
|
||||
|
||||
result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=True, env=os.environ)
|
||||
|
||||
is_has=True
|
||||
|
||||
if result.returncode != 0:
|
||||
print(f"Couldn't install\nCommand: {command}\nError code: {result.returncode}")
|
||||
is_has=False
|
||||
else:
|
||||
print(package+'## OK')
|
||||
|
||||
return is_has
|
||||
|
||||
|
||||
|
||||
# def is_installed(package):
|
||||
# try:
|
||||
# spec = importlib.util.find_spec(package)
|
||||
# except ModuleNotFoundError:
|
||||
# return False
|
||||
# return spec is not None
|
||||
|
||||
|
||||
def get_unique_hash(string):
|
||||
@@ -59,24 +124,8 @@ def openai_client(key,url):
|
||||
return client
|
||||
|
||||
def ZhipuAI_client(key):
|
||||
|
||||
try:
|
||||
if is_installed('zhipuai')==False:
|
||||
import subprocess
|
||||
|
||||
# 安装
|
||||
print('#pip install zhipuai')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'zhipuai'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0:
|
||||
print("#install success")
|
||||
from zhipuai import ZhipuAI
|
||||
else:
|
||||
print("#install error")
|
||||
|
||||
else:
|
||||
if is_installed('zhipuai')==True:
|
||||
from zhipuai import ZhipuAI
|
||||
except:
|
||||
print("#install zhipuai error")
|
||||
@@ -97,72 +146,75 @@ 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=get_llama_models()
|
||||
# llama_modes_list=[]
|
||||
|
||||
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 ):
|
||||
if is_installed('json_repair'):
|
||||
from json_repair import repair_json
|
||||
|
||||
|
||||
def chat(client, model_name,messages,max_tokens=4096,temperature=0.6 ):
|
||||
print('#chat',model_name,messages)
|
||||
try_count = 0
|
||||
while True:
|
||||
@@ -171,7 +223,9 @@ def chat(client, model_name,messages ):
|
||||
if hasattr(client, "chat"):
|
||||
response = client.chat.completions.create(
|
||||
model=model_name,
|
||||
messages=messages
|
||||
messages=messages,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature
|
||||
)
|
||||
else:
|
||||
# 是llama的
|
||||
@@ -206,6 +260,36 @@ 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()
|
||||
@@ -215,8 +299,9 @@ class ChatGPTNode:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
model_list=llama_modes_list+[
|
||||
"gpt-3.5-turbo",
|
||||
|
||||
model_list=[
|
||||
"gpt-3.5-turbo",
|
||||
"gpt-3.5-turbo-16k",
|
||||
"gpt-4o",
|
||||
"gpt-4o-2024-05-13",
|
||||
@@ -242,25 +327,32 @@ class ChatGPTNode:
|
||||
"01-ai/Yi-1.5-9B-Chat-16K",
|
||||
"meta-llama/Meta-Llama-3.1-8B-Instruct"
|
||||
]
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
|
||||
"api_url":("URL", {"default": "", "multiline": True,"dynamicPrompts": False}),
|
||||
# "api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
|
||||
# "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}),
|
||||
"api_url":(list(llm_apis_dict.keys()),
|
||||
{"default": list(llm_apis_dict.keys())[0]}),
|
||||
},
|
||||
"hidden": {
|
||||
"unique_id": "UNIQUE_ID",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO",
|
||||
},
|
||||
"optional":{
|
||||
"api_key":("STRING", {"forceInput": True,}),
|
||||
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
|
||||
"custom_api_url":("STRING", {"forceInput": True,}), #适合自定义model
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING","STRING","STRING",)
|
||||
@@ -272,12 +364,29 @@ class ChatGPTNode:
|
||||
|
||||
|
||||
def generate_contextual_text(self,
|
||||
api_key,
|
||||
api_url,
|
||||
# api_key,
|
||||
prompt,
|
||||
system_content,
|
||||
model,
|
||||
seed,context_size,unique_id = None, extra_pnginfo=None):
|
||||
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"
|
||||
|
||||
# print(api_key!='',api_url,prompt,system_content,model,seed)
|
||||
# 可以选择保留会话历史以维持上下文记忆
|
||||
# 或者在此处清除会话历史 self.session_history.clear()
|
||||
@@ -290,7 +399,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)
|
||||
@@ -299,9 +408,9 @@ 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)
|
||||
@@ -351,14 +460,15 @@ class SiliconflowFreeNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
model_list= [
|
||||
"Qwen/Qwen2-7B-Instruct",
|
||||
"Qwen/Qwen2.5-7B-Instruct",
|
||||
"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":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
|
||||
"api_key":("STRING", {"forceInput": True,}),
|
||||
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
|
||||
"system_content": ("STRING",
|
||||
{
|
||||
@@ -369,11 +479,11 @@ class SiliconflowFreeNode:
|
||||
{"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}),
|
||||
"max_tokens":("INT", {"default": 512, "min": 512, "max":200000, "step": 1}),
|
||||
},
|
||||
"hidden": {
|
||||
"unique_id": "UNIQUE_ID",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO",
|
||||
},
|
||||
"optional":{
|
||||
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING","STRING","STRING",)
|
||||
@@ -385,12 +495,18 @@ class SiliconflowFreeNode:
|
||||
|
||||
|
||||
def generate_contextual_text(self,
|
||||
api_key,
|
||||
prompt,
|
||||
system_content,
|
||||
api_key,
|
||||
prompt,
|
||||
system_content,
|
||||
model,
|
||||
seed,context_size,unique_id = None, extra_pnginfo=None):
|
||||
|
||||
seed,
|
||||
context_size,
|
||||
max_tokens,
|
||||
custom_model_name=None):
|
||||
|
||||
if custom_model_name!=None:
|
||||
model=custom_model_name
|
||||
|
||||
api_url="https://api.siliconflow.cn/v1"
|
||||
|
||||
# 把系统信息和初始信息添加到会话历史中
|
||||
@@ -418,7 +534,7 @@ class SiliconflowFreeNode:
|
||||
|
||||
messages=[{"role": "system", "content": self.system_content}]+session_history+[{"role": "user", "content": prompt}]
|
||||
|
||||
response_content = chat(client,model,messages)
|
||||
response_content = chat(client,model,messages,max_tokens)
|
||||
|
||||
self.session_history=self.session_history+[{"role": "user", "content": prompt}]+[{'role':'assistant',"content":response_content}]
|
||||
|
||||
@@ -426,6 +542,82 @@ class SiliconflowFreeNode:
|
||||
|
||||
|
||||
|
||||
class SiliconflowTextToImageNode:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
model_list= [
|
||||
"black-forest-labs/FLUX.1-schnell",
|
||||
]
|
||||
return {
|
||||
"required": {
|
||||
"api_key":("STRING", {"forceInput": True,}),
|
||||
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
|
||||
"width": ("INT", {"default": 512, "min": 512, "max": 4096, "step": 8}),
|
||||
"height": ("INT", {"default": 512, "min": 512, "max": 4096, "step": 8}),
|
||||
"model": ( model_list,
|
||||
{"default": model_list[0]}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
|
||||
},
|
||||
"optional":{
|
||||
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("image",)
|
||||
FUNCTION = "generate_contextual_text"
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
|
||||
def generate_contextual_text(self,
|
||||
api_key,
|
||||
prompt,
|
||||
width,
|
||||
height,
|
||||
model,
|
||||
seed,
|
||||
custom_model_name=None):
|
||||
|
||||
if custom_model_name!=None:
|
||||
model=custom_model_name
|
||||
|
||||
url=f"https://api.siliconflow.cn/v1/{model}/text-to-image"
|
||||
|
||||
headers = {
|
||||
"Accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {api_key}"
|
||||
}
|
||||
post_data = {
|
||||
"prompt":prompt,
|
||||
"image_size": f'{width}x{height}',
|
||||
}
|
||||
|
||||
empty_img= pil2tensor(Image.new('RGB', (1, 1), color='white'))
|
||||
|
||||
try:
|
||||
response = requests.post(url, headers=headers, data=json.dumps(post_data))
|
||||
response_data = response.json()
|
||||
|
||||
if response_data.get('code') == 20021:
|
||||
return (empty_img,)
|
||||
|
||||
image_url = response_data['images'][0]['url']
|
||||
|
||||
# Fetch the image using the image URL and read it with PIL
|
||||
image_response = requests.get(image_url)
|
||||
image = Image.open(BytesIO(image_response.content))
|
||||
|
||||
image=pil2tensor(image)
|
||||
return (image,)
|
||||
except Exception as error:
|
||||
print(error)
|
||||
return (empty_img,)
|
||||
|
||||
|
||||
|
||||
class ShowTextForGPT:
|
||||
@classmethod
|
||||
@@ -587,3 +779,41 @@ class TextSplitByDelimiter:
|
||||
arr= arr[start_index:start_index + max_count * (skip_every+1):(skip_every+1)]
|
||||
|
||||
return (arr,)
|
||||
|
||||
|
||||
class JsonRepair:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"json_string":("STRING", {"forceInput": True,}),
|
||||
"key":("STRING", {"multiline": False,"dynamicPrompts": False,"default": ""}),
|
||||
}
|
||||
}
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
RETURN_TYPES = ("STRING","STRING",)
|
||||
RETURN_NAMES = ("json_string","value",)
|
||||
FUNCTION = "run"
|
||||
# OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (False,False,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/GPT"
|
||||
|
||||
def run(self, json_string,key=""):
|
||||
|
||||
json_string=extract_json_strings(json_string)
|
||||
# print(json_string)
|
||||
good_json_string = repair_json(json_string)
|
||||
|
||||
# 将 JSON 字符串解析为 Python 对象
|
||||
data = json.loads(good_json_string)
|
||||
|
||||
v=""
|
||||
if key!="" and (key in data):
|
||||
v=data[key]
|
||||
|
||||
# 将 Python 对象转换回 JSON 字符串,确保中文字符不被转义
|
||||
json_str_with_chinese = json.dumps(data, ensure_ascii=False)
|
||||
|
||||
return (json_str_with_chinese,v,)
|
||||
@@ -0,0 +1,250 @@
|
||||
# 修改自 https://github.com/AnyaCoder/ComfyUI-fish-speech/
|
||||
|
||||
import torch,os
|
||||
from pathlib import Path
|
||||
from .fish_speech.llama_utils import load_model as load_llama_model
|
||||
from .fish_speech.vqgan_utils import load_model as load_vqgan_model
|
||||
from .fish_speech.vqgan_utils import audio2prompt, semantic2audio
|
||||
from .fish_speech.llama_utils import prompt2semantic
|
||||
|
||||
import folder_paths
|
||||
|
||||
def get_checkpoints_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('fish_speech')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "fish_speech")
|
||||
|
||||
current_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
configs_dir=os.path.join(current_directory,"fish_speech","configs")
|
||||
|
||||
CKPTS_FOLDER = Path(get_checkpoints_path())
|
||||
|
||||
CONFIGS_FOLDER = Path(configs_dir)
|
||||
|
||||
|
||||
class LoadVQGAN:
|
||||
def __init__(self):
|
||||
self.vqgan = None
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"config": ([str(c.relative_to(CONFIGS_FOLDER)) for c in CONFIGS_FOLDER.glob("*vq*.yaml")], {"default": "firefly_gan_vq.yaml"}),
|
||||
"model": ([str(p.relative_to(CKPTS_FOLDER)) for p in CKPTS_FOLDER.glob("*vq*.pth")], ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, model):
|
||||
return ""
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(s, model):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("VQGAN", )
|
||||
RETURN_NAMES = ("vqgan", )
|
||||
|
||||
FUNCTION = "load_vqgan"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def load_vqgan(self, config, model, device):
|
||||
config = config.rsplit(".", 1)[0]
|
||||
model = str(CKPTS_FOLDER / model)
|
||||
if self.vqgan is None:
|
||||
self.vqgan = load_vqgan_model(config,model, device=device)
|
||||
return (self.vqgan, )
|
||||
|
||||
|
||||
|
||||
class AudioToPrompt:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"vqgan": ("VQGAN", ),
|
||||
"audio": ("AUDIO", ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("AUDIO", "NUMPY")
|
||||
RETURN_NAMES = ("restored_audio", "prompt_tokens")
|
||||
|
||||
FUNCTION = "encode"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def encode(self, vqgan, audio, device):
|
||||
return audio2prompt(vqgan, audio, device)
|
||||
|
||||
|
||||
|
||||
class Prompt2Semantic:
|
||||
|
||||
def __init__(self):
|
||||
self.llama = None
|
||||
self.decode_func = None
|
||||
pass
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
"prompt_text": ("STRING", {"multiline": True}),
|
||||
"prompt_tokens": ("NUMPY", ),
|
||||
"max_new_tokens": ("INT", {
|
||||
"default": 1024,
|
||||
"min": 0,
|
||||
"max": 2048,
|
||||
"step": 8,
|
||||
"display": "number",
|
||||
}),
|
||||
"top_p": ("FLOAT", {
|
||||
"default": 0.7,
|
||||
"min": 0.6,
|
||||
"max": 0.9,
|
||||
"step": 0.01,
|
||||
"display": "number",
|
||||
}),
|
||||
"repetition_penalty": ("FLOAT", {
|
||||
"default": 1.2,
|
||||
"min": 1.0,
|
||||
"max": 1.5,
|
||||
"step": 0.01,
|
||||
"display": "number",
|
||||
}),
|
||||
"temperature": ("FLOAT", {
|
||||
"default": 0.7,
|
||||
"min": 0.6,
|
||||
"max": 0.9,
|
||||
"step": 0.01,
|
||||
"display": "number",
|
||||
}),
|
||||
|
||||
"seed": ("INT", {
|
||||
"default": 42,
|
||||
"min": 0,
|
||||
"max": 4294967295,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
}),
|
||||
"iterative_prompt": (["yes", "no"], {"default": "yes"}),
|
||||
"chunk_length": ("INT", {
|
||||
"default": 100,
|
||||
"min": 0,
|
||||
"max": 500,
|
||||
"step": 8,
|
||||
"display": "number",
|
||||
}),
|
||||
|
||||
"compile": (["yes", "no"], {"default": "no"}),
|
||||
"precision": (["bf16", "half"], {"default": "bf16"}),
|
||||
|
||||
# "decode_func": ("DECODE_FUNC", ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NUMPY", )
|
||||
RETURN_NAMES = ("codes", )
|
||||
|
||||
FUNCTION = "decode"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def decode(
|
||||
self,
|
||||
|
||||
text: str,
|
||||
prompt_text: str,
|
||||
prompt_tokens,
|
||||
max_new_tokens: int,
|
||||
top_p: float,
|
||||
repetition_penalty: float,
|
||||
temperature: float,
|
||||
|
||||
seed: int,
|
||||
iterative_prompt: str,
|
||||
chunk_length: int,
|
||||
|
||||
compile: str,
|
||||
precision,
|
||||
device: str,
|
||||
):
|
||||
|
||||
model = get_checkpoints_path()
|
||||
precision = torch.bfloat16 if precision == "bf16" else torch.half
|
||||
compile=True if compile == "yes" else False
|
||||
if self.llama is None or self.decode_func is None:
|
||||
self.llama, self.decode_func = load_llama_model(model, device, precision, compile)
|
||||
|
||||
|
||||
return prompt2semantic(
|
||||
self.llama,
|
||||
self.decode_func,
|
||||
text,
|
||||
[prompt_text,],
|
||||
[prompt_tokens,],
|
||||
max_new_tokens,
|
||||
top_p,
|
||||
repetition_penalty,
|
||||
temperature,
|
||||
device,
|
||||
compile=True if compile == "yes" else False,
|
||||
seed=seed,
|
||||
iterative_prompt=True if iterative_prompt == "yes" else False,
|
||||
chunk_length=chunk_length,
|
||||
)
|
||||
|
||||
|
||||
|
||||
class Semantic2Audio:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"vqgan": ("VQGAN", ),
|
||||
"codes": ("NUMPY", ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO", )
|
||||
RETURN_NAMES = ("generated_audio", )
|
||||
|
||||
FUNCTION = "generate"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def generate(self, vqgan, codes, device):
|
||||
return semantic2audio(vqgan, codes, device)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -8,17 +8,31 @@ from PIL.PngImagePlugin import PngInfo
|
||||
import base64,os,random
|
||||
from io import BytesIO
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
import json,io
|
||||
import comfy.utils
|
||||
from comfy.cli_args import args
|
||||
import cv2
|
||||
import string
|
||||
import string,re
|
||||
import math,glob
|
||||
from .Watcher import FolderWatcher
|
||||
|
||||
from itertools import product
|
||||
|
||||
|
||||
# 文件名排序
|
||||
def sort_by_filename(items):
|
||||
def extract_parts(filename):
|
||||
# 使用正则表达式将文件名拆分为数字和非数字部分
|
||||
parts = re.split(r'(\d+)', filename)
|
||||
# 将数字部分转换为整数以便正确排序,同时保留非数字部分
|
||||
parts = [int(part) if part.isdigit() else part for part in parts]
|
||||
return parts
|
||||
|
||||
# 按照 file_name 的拆分部分进行排序
|
||||
sorted_items = sorted(items, key=lambda x: extract_parts(x['file_name']))
|
||||
return sorted_items
|
||||
|
||||
# 将PIL图片转换为OpenCV格式
|
||||
def pil_to_opencv(image):
|
||||
open_cv_image = cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
|
||||
@@ -43,14 +57,13 @@ def get_files_with_extension(directory, 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)
|
||||
scale = max(scale, 1 / bwidth)
|
||||
scale = max(scale, 1 / bheight)
|
||||
|
||||
def determine_scale_option(width, height):
|
||||
return 'height' if height > width else 'width'
|
||||
@@ -69,9 +82,9 @@ def composite_images(foreground, background, mask, is_multiply_blend=False, posi
|
||||
else:
|
||||
scale_option = determine_scale_option(width, height)
|
||||
if scale_option == 'height':
|
||||
scale = int(bheight * scale) / height
|
||||
scale = bheight * scale / height
|
||||
else:
|
||||
scale = int(bwidth * scale) / width
|
||||
scale = bwidth * scale / width
|
||||
|
||||
new_width = int(width * scale)
|
||||
new_height = int(height * scale)
|
||||
@@ -109,22 +122,13 @@ def composite_images(foreground, background, mask, is_multiply_blend=False, posi
|
||||
"mask": mask
|
||||
}
|
||||
|
||||
layer_image = layer['image']
|
||||
layer_mask = layer['mask']
|
||||
# Resize the foreground image with antialiasing
|
||||
layer_image = layer['image'].resize((layer['width'], layer['height']), Image.ANTIALIAS)
|
||||
layer_mask = layer['mask'].resize((layer['width'], layer['height']), Image.ANTIALIAS)
|
||||
|
||||
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.paste(layer_image, (layer['x'], layer['y']), layer_mask)
|
||||
|
||||
bg_image = bg_image.convert('RGB')
|
||||
|
||||
return bg_image
|
||||
return bg_image.convert('RGB')
|
||||
|
||||
|
||||
|
||||
@@ -491,6 +495,53 @@ def load_image(fp,white_bg=False):
|
||||
|
||||
return images
|
||||
|
||||
|
||||
# 读取图片数据,转成tensor
|
||||
def load_image_to_tensor( image):
|
||||
image_path = folder_paths.get_annotated_filepath(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:
|
||||
mask = torch.zeros((64,64), 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 (output_image, output_mask)
|
||||
|
||||
|
||||
def load_image_and_mask_from_url(url, timeout=10):
|
||||
# Load the image from the URL
|
||||
response = requests.get(url, timeout=timeout)
|
||||
@@ -731,8 +782,7 @@ def areaToMask(x,y,w,h,image):
|
||||
# return bg_image
|
||||
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
# ps的正片叠底
|
||||
# 可以基于https://www.cnblogs.com/jsxyhelu/p/16947810.html ,用gpt写python代码
|
||||
@@ -912,9 +962,68 @@ 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):
|
||||
# Split text into lines based on line breaks
|
||||
lines = text.split("\n")
|
||||
|
||||
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,
|
||||
max_characters_per_line=48,
|
||||
fixed_width=None):
|
||||
|
||||
def split_text(text, max_chars, fixed_width=False):
|
||||
lines = []
|
||||
current_line = ""
|
||||
current_length = 0
|
||||
|
||||
for char in text:
|
||||
if char == '\n':
|
||||
lines.append(current_line)
|
||||
current_line = ""
|
||||
current_length = 0
|
||||
elif '\u4e00' <= char <= '\u9fff': # Chinese character
|
||||
if current_length + 1 <= max_chars:
|
||||
current_line += char
|
||||
current_length += 1
|
||||
else:
|
||||
lines.append(current_line)
|
||||
current_line = char
|
||||
current_length = 1
|
||||
else: # English character or other
|
||||
if char == ' ':
|
||||
space_length = 1
|
||||
else:
|
||||
space_length = 1
|
||||
if current_length + space_length <= max_chars:
|
||||
current_line += char
|
||||
current_length += space_length
|
||||
else:
|
||||
lines.append(current_line)
|
||||
current_line = char
|
||||
current_length = space_length
|
||||
|
||||
if current_line:
|
||||
lines.append(current_line)
|
||||
|
||||
# Pad lines to max_chars if fixed_width is provided
|
||||
if fixed_width:
|
||||
lines = [line.ljust(max_chars) for line in lines]
|
||||
|
||||
# If there's only one line and fixed_width is True, pad it
|
||||
if fixed_width and len(lines) == 1:
|
||||
lines[0] = lines[0].ljust(max_chars)
|
||||
|
||||
return lines
|
||||
|
||||
# lines = text.split("\n")
|
||||
# Split text into lines based on max_characters_per_line
|
||||
lines = split_text(text, max_characters_per_line,fixed_width)
|
||||
|
||||
# Load font
|
||||
font = ImageFont.truetype(font_path, font_size)
|
||||
@@ -943,7 +1052,6 @@ def generate_text_image(text, font_path, font_size, text_color, vertical=True, s
|
||||
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
|
||||
else:
|
||||
for line in lines:
|
||||
@@ -955,10 +1063,8 @@ def generate_text_image(text, font_path, font_size, text_color, vertical=True, s
|
||||
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
|
||||
|
||||
# 3. Create image with calculated width and height
|
||||
@@ -971,10 +1077,10 @@ def generate_text_image(text, font_path, font_size, text_color, vertical=True, s
|
||||
for char in 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
|
||||
@@ -988,10 +1094,18 @@ def generate_text_image(text, font_path, font_size, text_color, vertical=True, s
|
||||
|
||||
image = image.convert('RGB')
|
||||
|
||||
# 5. Scale the image if fixed_width is specified
|
||||
if fixed_width and fixed_width < max_width:
|
||||
scaling_factor = fixed_width / max_width
|
||||
new_height = int(max_height * scaling_factor)
|
||||
image = image.resize((fixed_width, new_height), Image.ANTIALIAS)
|
||||
alpha_image = alpha_image.resize((fixed_width, new_height), Image.ANTIALIAS)
|
||||
|
||||
return (image, alpha_image)
|
||||
|
||||
|
||||
|
||||
|
||||
def base64_to_image(base64_string):
|
||||
# 去除前缀
|
||||
prefix, base64_data = base64_string.split(",", 1)
|
||||
@@ -1332,7 +1446,7 @@ class LoadImagesFromPath:
|
||||
},
|
||||
"optional":{
|
||||
"white_bg": (["disable","enable"],),
|
||||
"newest_files": (["enable", "disable"],),
|
||||
"sort_by": (["file_name", "newest"],),#根据文件名来排序,还是按照最新创建时间
|
||||
"index_variable":("INT", {
|
||||
"default": 0,
|
||||
"min": -1, #Minimum value
|
||||
@@ -1342,13 +1456,13 @@ class LoadImagesFromPath:
|
||||
}),
|
||||
"watcher":(["disable","enable"],),
|
||||
"result": ("WATCHER",),#为了激活本节点运行
|
||||
"prompt": ("PROMPT",),
|
||||
# "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"prompt": ("PROMPT",),
|
||||
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('IMAGE','MASK','STRING','STRING',)
|
||||
RETURN_NAMES = ("IMAGE","MASK","prompt_for_FloatingVideo","filepaths",)
|
||||
RETURN_NAMES = ("image list","MASK","prompt_for_FloatingVideo","filepaths",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
@@ -1361,7 +1475,7 @@ class LoadImagesFromPath:
|
||||
watcher_folder=None
|
||||
|
||||
# 运行的函数
|
||||
def run(self,file_path,white_bg,newest_files,index_variable,watcher,result,prompt):
|
||||
def run(self,file_path,white_bg,sort_by,index_variable,watcher,result,prompt,seed=1):
|
||||
global watcher_folder
|
||||
# print('###监听:',watcher_folder,watcher,file_path,result)
|
||||
|
||||
@@ -1384,19 +1498,23 @@ class LoadImagesFromPath:
|
||||
# 当开启了监听,则取最新的,第一个文件
|
||||
if watcher=='enable':
|
||||
index_variable=0
|
||||
newest_files='enable'
|
||||
sort_by='newest'
|
||||
|
||||
# 排序
|
||||
sorted_files = sorted(images, key=lambda x: os.path.getmtime(x['file_path']), reverse=(newest_files=='enable'))
|
||||
if sort_by=='newest':
|
||||
sorted_files = sorted(images, key=lambda x: os.path.getmtime(x['file_path']), reverse=True)
|
||||
elif sort_by=='file_name':
|
||||
# 根据文件名排序
|
||||
sorted_files = sort_by_filename(images)
|
||||
|
||||
imgs=[]
|
||||
masks=[]
|
||||
file_names=[]
|
||||
file_paths=[]
|
||||
|
||||
for im in sorted_files:
|
||||
imgs.append(im['image'])
|
||||
masks.append(im['mask'])
|
||||
file_names.append(im['file_name'])
|
||||
file_paths.append(im['file_path'])
|
||||
|
||||
# print('index_variable',index_variable)
|
||||
|
||||
@@ -1404,12 +1522,13 @@ class LoadImagesFromPath:
|
||||
if index_variable!=-1:
|
||||
imgs=[imgs[index_variable]] if index_variable < len(imgs) else None
|
||||
masks=[masks[index_variable]] if index_variable < len(masks) else None
|
||||
file_names=[file_names[index_variable]] if index_variable < len(file_names) else None
|
||||
file_paths=[file_paths[index_variable]] if index_variable < len(file_paths) else None
|
||||
except Exception as e:
|
||||
print("发生了一个未知的错误:", str(e))
|
||||
|
||||
# print('#prompt::::',prompt)
|
||||
return {"ui": {"seed": [1]}, "result":(imgs,masks,prompt,file_names,)}
|
||||
# return {"ui": {"seed": [1]}, "result":(imgs,masks,prompt,file_names,)}
|
||||
return (imgs,masks,prompt,file_paths,)
|
||||
|
||||
|
||||
# TODO 扩大选区的功能,重新输出mask
|
||||
@@ -1515,35 +1634,49 @@ class TextImage:
|
||||
"font": (get_files_with_extension(FONT_PATH,['.ttf','.otf']),),#后缀为 ttf
|
||||
"font_size": ("INT",{
|
||||
"default":100,
|
||||
"min": 100, #Minimum value
|
||||
"max": 1000, #Maximum value
|
||||
"min": 1, #Minimum value
|
||||
"max": 10000000, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"spacing": ("INT",{
|
||||
"default":12,
|
||||
"min": -200, #Minimum value
|
||||
"max": 200, #Maximum value
|
||||
"min": -2000000000, #Minimum value
|
||||
"max": 2000000000, #Maximum value
|
||||
"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
|
||||
"min": -2000000000, #Minimum value
|
||||
"max": 2000000000, #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
|
||||
"max": 2000000000, #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},),
|
||||
"max_characters_per_line": ("INT",{
|
||||
"default":44,
|
||||
"min": 1, #Minimum value
|
||||
"max": 2000000000, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"fixed_width":("INT",{
|
||||
"default":0,
|
||||
"min": 0, #Minimum value
|
||||
"max": 2000000000, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -1557,14 +1690,20 @@ 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,font_size,spacing,line_spacing,padding,text_color,vertical,stroke,max_characters_per_line,fixed_width):
|
||||
|
||||
font_path=os.path.join(FONT_PATH,font)
|
||||
|
||||
if text=="":
|
||||
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)
|
||||
# max_characters_per_line 英文字按照空格计算1个,中文按照字数计算
|
||||
if fixed_width==0:
|
||||
fixed_width=None
|
||||
img,mask=generate_text_image(text,font_path,font_size,text_color,vertical,stroke,(0, 0, 0),1,
|
||||
spacing,line_spacing,padding,max_characters_per_line,
|
||||
fixed_width
|
||||
)
|
||||
|
||||
img=pil2tensor(img)
|
||||
mask=pil2tensor(mask)
|
||||
@@ -1598,7 +1737,7 @@ class LoadImagesFromURL:
|
||||
|
||||
def run(self,url,seed=0):
|
||||
global urls_image
|
||||
print(urls_image)
|
||||
# print(urls_image)
|
||||
def filter_http_urls(urls):
|
||||
filtered_urls = []
|
||||
for url in urls.split('\n'):
|
||||
@@ -1687,28 +1826,51 @@ class Image3D:
|
||||
def run(self,upload,material=None):
|
||||
# print('material',material)
|
||||
# print(upload )
|
||||
image = base64_to_image(upload['image'])
|
||||
|
||||
mat=None
|
||||
if 'material' in upload and upload['material']:
|
||||
mat=base64_to_image(upload['material'])
|
||||
mat=mat.convert('RGB')
|
||||
mat=pil2tensor(mat)
|
||||
# 截取的系列角度截图
|
||||
images=upload['images'] if "images" in upload else []
|
||||
|
||||
mask = image.split()[3]
|
||||
image=image.convert('RGB')
|
||||
ims=[]
|
||||
for im in images:
|
||||
if 'type' in im and (not f"[{im['type']}]" in im['name']):
|
||||
im['name']=im['name']+" "+f"[{im['type']}]"
|
||||
output_image, output_mask = load_image_to_tensor(im['name'])
|
||||
ims.append(output_image)
|
||||
|
||||
mask=mask.convert('L')
|
||||
|
||||
|
||||
mask=None
|
||||
bg_image=None
|
||||
if 'bg_image' in upload and upload['bg_image']:
|
||||
bg_image = base64_to_image(upload['bg_image'])
|
||||
bg_image=bg_image.convert('RGB')
|
||||
bg_image=pil2tensor(bg_image)
|
||||
mat=None
|
||||
|
||||
# 如果没有系列截图
|
||||
if len(ims)==0:
|
||||
# 这个是3d模型当前截图
|
||||
image = base64_to_image(upload['image'])
|
||||
|
||||
|
||||
if 'material' in upload and upload['material']:
|
||||
mat=base64_to_image(upload['material'])
|
||||
mat=mat.convert('RGB')
|
||||
mat=pil2tensor(mat)
|
||||
|
||||
mask = image.split()[3]
|
||||
image=image.convert('RGB')
|
||||
|
||||
mask=mask.convert('L')
|
||||
|
||||
|
||||
if 'bg_image' in upload and upload['bg_image']:
|
||||
bg_image = base64_to_image(upload['bg_image'])
|
||||
bg_image=bg_image.convert('RGB')
|
||||
bg_image=pil2tensor(bg_image)
|
||||
|
||||
|
||||
mask=pil2tensor(mask)
|
||||
image=pil2tensor(image)
|
||||
mask=pil2tensor(mask)
|
||||
image=pil2tensor(image)
|
||||
else:
|
||||
|
||||
image = torch.cat(ims, dim=0)
|
||||
|
||||
|
||||
m=[]
|
||||
if not material is None:
|
||||
@@ -1829,12 +1991,9 @@ class CompositeImages:
|
||||
|
||||
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)
|
||||
@@ -2731,14 +2890,14 @@ class ResizeImage:
|
||||
"default": 512,
|
||||
"min": 1, #Minimum value
|
||||
"max": 8192, #Maximum value
|
||||
"step": 8, #Slider's step
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"height": ("INT",{
|
||||
"default": 512,
|
||||
"min": 1, #Minimum value
|
||||
"max": 8192, #Maximum value
|
||||
"step": 8, #Slider's step
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"scale_option": (["width","height",'overall','center'],),
|
||||
@@ -2754,7 +2913,7 @@ class ResizeImage:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE","IMAGE","STRING","MASK",)
|
||||
RETURN_NAMES = ("image","average_image","average_hex","mask",)
|
||||
RETURN_NAMES = ("image list","average_image","average_hex","mask",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
@@ -3223,3 +3382,99 @@ class ImageListToBatch_:
|
||||
out = torch.cat(out, dim=0)
|
||||
|
||||
return (out,)
|
||||
|
||||
|
||||
# https://github.com/gokayfem/ComfyUI-Depth-Visualization?tab=readme-ov-file
|
||||
class DepthViewer_:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"depth_map": ("IMAGE",),
|
||||
},
|
||||
"optional":{
|
||||
"frames":("IMAGEBASE64",),
|
||||
},
|
||||
}
|
||||
|
||||
def __init__(self):
|
||||
self.saved_reference = []
|
||||
self.saved_depth = []
|
||||
|
||||
self.full_output_folder,self.filename,self.counter, self.subfolder, self.filename_prefix = folder_paths.get_save_image_path(
|
||||
"imagesave",
|
||||
folder_paths.get_output_directory())
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("frames",)
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/3D"
|
||||
def run(self, image, depth_map,frames=None):
|
||||
self.saved_reference.clear()
|
||||
self.saved_depth.clear()
|
||||
image = image[0].detach().cpu().numpy()
|
||||
depth = depth_map[0].detach().cpu().numpy()
|
||||
|
||||
image = Image.fromarray(np.clip(255. * image, 0, 255).astype(np.uint8)).convert('RGB')
|
||||
depth = Image.fromarray(np.clip(255. * depth, 0, 255).astype(np.uint8))
|
||||
|
||||
return self.display([image], [depth],frames)
|
||||
|
||||
def display(self, reference_image, depth_map,frames):
|
||||
for (batch_number, (single_image, single_depth)) in enumerate(zip(reference_image, depth_map)):
|
||||
filename_with_batch_num = self.filename.replace("%batch_num%", str(batch_number))
|
||||
|
||||
image_file = f"{filename_with_batch_num}_{self.counter:05}_reference.png"
|
||||
single_image.save(os.path.join(self.full_output_folder, image_file))
|
||||
|
||||
depth_file = f"{filename_with_batch_num}_{self.counter:05}_depth.png"
|
||||
single_depth.save(os.path.join(self.full_output_folder, depth_file))
|
||||
|
||||
self.saved_reference.append({
|
||||
"filename": image_file,
|
||||
"subfolder": self.subfolder,
|
||||
"type": "output"
|
||||
})
|
||||
|
||||
self.saved_depth.append({
|
||||
"filename": depth_file,
|
||||
"subfolder": self.subfolder,
|
||||
"type": "output"
|
||||
})
|
||||
self.counter += 1
|
||||
|
||||
|
||||
ims=[]
|
||||
image1 = Image.new('RGB', (512, 512), color='black')
|
||||
image1=pil2tensor(image1)
|
||||
|
||||
# print('frames',frames)
|
||||
if frames!=None and "images" in frames:
|
||||
|
||||
for im in frames['images']:
|
||||
# print(im)
|
||||
if 'type' in im and (not f"[{im['type']}]" in im['name']):
|
||||
im['name']=im['name']+" "+f"[{im['type']}]"
|
||||
|
||||
try:
|
||||
output_image, output_mask = load_image_to_tensor(im['name'])
|
||||
ims.append(output_image)
|
||||
except:
|
||||
print("no")
|
||||
|
||||
|
||||
if len(ims)>0:
|
||||
image1 = ims[0]
|
||||
for image2 in ims[1:]:
|
||||
if image1.shape[1:] != image2.shape[1:]:
|
||||
image2 = comfy.utils.common_upscale(image2.movedim(-1, 1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1, -1)
|
||||
image1 = torch.cat((image1, image2), dim=0)
|
||||
|
||||
return {"ui": {"reference_image": self.saved_reference, "depth_map": self.saved_depth}, "result": (image1,)}
|
||||
@@ -0,0 +1,127 @@
|
||||
# Referenced some code:https://github.com/IuvenisSapiens/ComfyUI_MiniCPM-V-2_6-int4
|
||||
|
||||
import os
|
||||
import torch
|
||||
import folder_paths
|
||||
from transformers import AutoTokenizer, AutoModel
|
||||
from torchvision.transforms.v2 import ToPILImage
|
||||
# from decord import VideoReader, cpu # pip install decord
|
||||
# from PIL import Image
|
||||
|
||||
def get_model_path(n=""):
|
||||
try:
|
||||
return folder_paths.get_folder_paths(n)[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, n)
|
||||
|
||||
|
||||
class MiniCPM_VQA_Simple:
|
||||
def __init__(self):
|
||||
self.model_checkpoint = None
|
||||
self.tokenizer = None
|
||||
self.model = None
|
||||
self.device = (
|
||||
torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
)
|
||||
self.bf16_support = (
|
||||
torch.cuda.is_available()
|
||||
and torch.cuda.get_device_capability(self.device)[0] >= 8
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"text": ("STRING", {"default": "", "multiline": True}),
|
||||
"seed": ("INT", {"default": -1}), # add seed parameter, default is -1
|
||||
"temperature": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.7,
|
||||
},
|
||||
),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
|
||||
def inference(
|
||||
self,
|
||||
images,
|
||||
text,
|
||||
seed, # add seed parameter, default is -1
|
||||
temperature,
|
||||
keep_model_loaded,
|
||||
):
|
||||
if seed != -1:
|
||||
torch.manual_seed(seed)
|
||||
model_id = "openbmb/MiniCPM-V-2_6-int4"
|
||||
|
||||
self.model_checkpoint = os.path.join( get_model_path("prompt_generator"), os.path.basename(model_id))
|
||||
|
||||
if not os.path.exists(self.model_checkpoint):
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
snapshot_download(
|
||||
repo_id=model_id,
|
||||
local_dir=self.model_checkpoint,
|
||||
local_dir_use_symlinks=False,
|
||||
endpoint='https://hf-mirror.com'
|
||||
)
|
||||
|
||||
if self.tokenizer is None:
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||
self.model_checkpoint,
|
||||
trust_remote_code=True,
|
||||
low_cpu_mem_usage=True,
|
||||
)
|
||||
|
||||
if self.model is None:
|
||||
self.model = AutoModel.from_pretrained(
|
||||
self.model_checkpoint,
|
||||
trust_remote_code=True,
|
||||
low_cpu_mem_usage=True,
|
||||
attn_implementation="sdpa",
|
||||
torch_dtype=torch.bfloat16 if self.bf16_support else torch.float16,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
images = images.permute([0, 3, 1, 2])
|
||||
images = [ToPILImage()(img).convert("RGB") for img in images]
|
||||
msgs = [{"role": "user", "content": images + [text]}]
|
||||
|
||||
params = {"use_image_id": False, }
|
||||
|
||||
# offload model to CPU
|
||||
# self.model = self.model.to(torch.device("cpu"))
|
||||
# self.model.eval()
|
||||
|
||||
result = self.model.chat(
|
||||
image=None,
|
||||
msgs=msgs,
|
||||
tokenizer=self.tokenizer,
|
||||
sampling=True,
|
||||
# top_k=top_k,
|
||||
# top_p=top_p,
|
||||
temperature=temperature,
|
||||
# repetition_penalty=repetition_penalty,
|
||||
# max_new_tokens=max_new_tokens,
|
||||
**params,
|
||||
)
|
||||
# offload model to GPU
|
||||
# self.model = self.model.to(torch.device("cpu"))
|
||||
# self.model.eval()
|
||||
if not keep_model_loaded:
|
||||
del self.tokenizer # release tokenizer memory
|
||||
del self.model # release model memory
|
||||
self.tokenizer = None # set tokenizer to None
|
||||
self.model = None # set model to None
|
||||
torch.cuda.empty_cache() # release GPU memory
|
||||
torch.cuda.ipc_collect()
|
||||
# print(result)
|
||||
return (result,)
|
||||
@@ -0,0 +1,104 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image,ImageSequence,ImageOps
|
||||
import base64
|
||||
import io
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
def load_image_to_tensor( image):
|
||||
image_path = folder_paths.get_annotated_filepath(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:
|
||||
mask = torch.zeros((64,64), 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 (output_image, output_mask)
|
||||
|
||||
|
||||
|
||||
class P5Input:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"frames":("IMAGEBASE64",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("frames",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self, frames):
|
||||
ims=[]
|
||||
for im in frames['images']:
|
||||
# print(im)
|
||||
if 'type' in im and (not f"[{im['type']}]" in im['name']):
|
||||
im['name']=im['name']+" "+f"[{im['type']}]"
|
||||
|
||||
output_image, output_mask = load_image_to_tensor(im['name'])
|
||||
ims.append(output_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:]:
|
||||
image2 = comfy.utils.common_upscale(image2.movedim(-1, 1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1, -1)
|
||||
image1 = torch.cat((image1, image2), dim=0)
|
||||
|
||||
# 用于节点提示:p5节点提示有多少帧
|
||||
return {"ui": {"_info": [len(frames['images'])]}, "result": (image1,)}
|
||||
@@ -18,7 +18,13 @@ import json
|
||||
# req = request.Request("http://127.0.0.1:8188/prompt", data=data)
|
||||
# request.urlopen(req)
|
||||
|
||||
embeddings_path=os.path.join(folder_paths.models_dir, "embeddings")
|
||||
def get_model_path(n=""):
|
||||
try:
|
||||
return folder_paths.get_folder_paths(n)[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, n)
|
||||
|
||||
embeddings_path=get_model_path("embeddings")
|
||||
|
||||
def get_files_with_extension(directory, extension):
|
||||
|
||||
@@ -181,7 +187,8 @@ class PromptImage:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("json_str",)
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@@ -196,12 +203,19 @@ class PromptImage:
|
||||
filename_prefix="mixlab_"
|
||||
filename_prefix += self.prefix_append
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
|
||||
filename_prefix, self.output_dir, images[0].shape[1], images[0].shape[0])
|
||||
filename_prefix,self.output_dir, images[0].shape[1], images[0].shape[0])
|
||||
|
||||
full_output_folder=os.path.join(full_output_folder,'PromptImage')
|
||||
subfolder='PromptImage'
|
||||
|
||||
results = list()
|
||||
|
||||
save_to_image=save_to_image[0]=='enable'
|
||||
|
||||
#保存到本地的json文件,记录图片和prompt的对应关系
|
||||
output_images=[]
|
||||
output_prompt=[]
|
||||
|
||||
for index in range(len(images)):
|
||||
res=[]
|
||||
imgs=images[index]
|
||||
@@ -209,24 +223,36 @@ class PromptImage:
|
||||
for image in imgs:
|
||||
img=tensor2pil(image)
|
||||
|
||||
prompt_text=prompts[index]
|
||||
|
||||
metadata = None
|
||||
if save_to_image:
|
||||
metadata = PngInfo()
|
||||
prompt_text=prompts[index]
|
||||
if prompt_text is not None:
|
||||
metadata.add_text("prompt_text", prompt_text)
|
||||
|
||||
file = f"{filename}_{index}_{counter:05}_.png"
|
||||
img.save(os.path.join(full_output_folder, file), pnginfo=metadata, compress_level=self.compress_level)
|
||||
fp=os.path.join(full_output_folder,file)
|
||||
img.save(fp, pnginfo=metadata, compress_level=self.compress_level)
|
||||
res.append({
|
||||
"filename": file,
|
||||
"subfolder": subfolder,
|
||||
"type": self.type
|
||||
})
|
||||
output_images.append(fp)
|
||||
output_prompt.append(prompt_text)
|
||||
counter += 1
|
||||
results.append(res)
|
||||
|
||||
return { "ui": { "_images": results,"prompts":prompts } }
|
||||
|
||||
# if save_to_image:
|
||||
# # 保存为本地文件
|
||||
# with open(os.path.join(full_output_folder,'PromptImage.json'), 'w') as file:
|
||||
# json.dump(output_dict, file, ensure_ascii=False, indent=4)
|
||||
|
||||
return { "ui": { "_images": results,"prompts":prompts },"result":(json.dumps({
|
||||
"images":output_images,
|
||||
"prompts":output_prompt
|
||||
}),) }
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,215 @@
|
||||
# -*- coding:utf-8 -*-
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
|
||||
from huggingface_hub import snapshot_download
|
||||
import torch,re
|
||||
from sensevoice.onnx.sense_voice_ort_session import SenseVoiceInferenceSession
|
||||
from sensevoice.utils.frontend import WavFrontend
|
||||
from sensevoice.utils.fsmn_vad import FSMNVad
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
|
||||
languages = {"auto": 0, "zh": 3, "en": 4, "yue": 7, "ja": 11, "ko": 12, "nospeech": 13}
|
||||
|
||||
# 设置环境变量
|
||||
os.environ['HF_ENDPOINT'] = 'https://hf-mirror.com'
|
||||
|
||||
#
|
||||
def get_model_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('sense_voice')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "sense_voice")
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
|
||||
any_type = AnyType("*")
|
||||
|
||||
# 字幕
|
||||
def format_to_srt(channel_id, start_time_ms, end_time_ms, asr_result):
|
||||
start_time = start_time_ms / 1000
|
||||
end_time = end_time_ms / 1000
|
||||
|
||||
def format_time(seconds):
|
||||
hours = int(seconds // 3600)
|
||||
minutes = int((seconds % 3600) // 60)
|
||||
seconds = seconds % 60
|
||||
milliseconds = int((seconds - int(seconds)) * 1000)
|
||||
return f"{hours:02}:{minutes:02}:{int(seconds):02},{milliseconds:03}"
|
||||
|
||||
start_time_str = format_time(start_time)
|
||||
end_time_str = format_time(end_time)
|
||||
|
||||
pattern = r"<\|(.+?)\|><\|(.+?)\|><\|(.+?)\|><\|(.+?)\|>(.+)"
|
||||
match = re.match(pattern,asr_result)
|
||||
lang, emotion, audio_type, itn, text = match.groups()
|
||||
# 😊 表示高兴,😡 表示愤怒,😔 表示悲伤。对于音频事件,🎼 表示音乐,😀 表示笑声,👏 表示掌声
|
||||
|
||||
srt_content = f"1\n{start_time_str} --> {end_time_str}\n{text}\n"
|
||||
|
||||
logging.info(f"[Channel {channel_id}] [{start_time}s - {end_time}s] [{lang}] [{emotion}] [{audio_type}] [{itn}] {text}")
|
||||
|
||||
return lang, emotion, audio_type, itn,srt_content,start_time,end_time,text
|
||||
|
||||
|
||||
class SenseVoiceProcessor:
|
||||
def __init__(self, download_model_path, device, num_threads, use_int8):
|
||||
|
||||
if not os.path.exists(download_model_path):
|
||||
logging.info(
|
||||
"Downloading model from huggingface hub from https://huggingface.co/lovemefan/SenseVoice-onnx"
|
||||
)
|
||||
logging.info(
|
||||
"You can speed up with `export HF_ENDPOINT=https://hf-mirror.com`"
|
||||
)
|
||||
snapshot_download(
|
||||
repo_id="lovemefan/SenseVoice-onnx", local_dir=download_model_path
|
||||
)
|
||||
|
||||
self.download_model_path = download_model_path
|
||||
self.device = device
|
||||
self.num_threads = num_threads
|
||||
self.use_int8 = use_int8
|
||||
self.front = WavFrontend(os.path.join(download_model_path, "am.mvn"))
|
||||
self.model = SenseVoiceInferenceSession(
|
||||
os.path.join(download_model_path, "embedding.npy"),
|
||||
os.path.join(
|
||||
download_model_path,
|
||||
"sense-voice-encoder-int8.onnx"
|
||||
if use_int8
|
||||
else "sense-voice-encoder.onnx",
|
||||
),
|
||||
os.path.join(download_model_path, "chn_jpn_yue_eng_ko_spectok.bpe.model"),
|
||||
device,
|
||||
num_threads,
|
||||
)
|
||||
self.vad = FSMNVad(download_model_path)
|
||||
|
||||
def process_audio(self, waveform, _sample_rate, language, use_itn):
|
||||
|
||||
start = time.time()
|
||||
pbar = comfy.utils.ProgressBar(waveform.shape[1]) # 进度条
|
||||
|
||||
results = []
|
||||
|
||||
for channel_id, channel_data in enumerate(waveform.T):
|
||||
segments = self.vad.segments_offline(channel_data)
|
||||
|
||||
for part in segments:
|
||||
audio_feats = self.front.get_features(channel_data[part[0] * 16 : part[1] * 16])
|
||||
asr_result = self.model(
|
||||
audio_feats[None, ...],
|
||||
language=languages[language],
|
||||
use_itn=use_itn,
|
||||
)
|
||||
|
||||
lang, emotion, audio_type, itn,srt_content,start_time,end_time,text=format_to_srt(
|
||||
channel_id,
|
||||
part[0] ,
|
||||
part[1],
|
||||
asr_result)
|
||||
|
||||
results.append({
|
||||
"language":lang,
|
||||
"emotion":emotion,
|
||||
"audio_type":audio_type,
|
||||
"itn":itn,
|
||||
"srt_content":srt_content,
|
||||
"start_time":start_time,
|
||||
"end_time":end_time,
|
||||
"text":text
|
||||
})
|
||||
|
||||
self.vad.vad.all_reset_detection()
|
||||
pbar.update(1) # 更新进度条
|
||||
|
||||
decoding_time = time.time() - start
|
||||
logging.info(f"Decoder audio takes {decoding_time} seconds")
|
||||
logging.info(f"The RTF is {decoding_time/(waveform.shape[1] * len(waveform) / _sample_rate)}.")
|
||||
return results
|
||||
|
||||
|
||||
class SenseVoiceNode:
|
||||
|
||||
def __init__(self):
|
||||
self.processor = None
|
||||
self.download_model_path=get_model_path()
|
||||
self.device="cpu"
|
||||
self.num_threads = 4
|
||||
self.use_int8 = True
|
||||
self.language='auto'
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {"required": {
|
||||
"audio": ("AUDIO", ),
|
||||
"device": ( ['auto','cpu'], {"default": 'auto'}),
|
||||
"language": (list(languages.keys()), {"default": 'auto'}),# 不能直接写 languages.keys(),json.dumps会报错
|
||||
"num_threads":("INT",{
|
||||
"default":4,
|
||||
"min": 1, #Minimum value
|
||||
"max": 32, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
},),
|
||||
"use_int8":("BOOLEAN", {"default": True},),
|
||||
"use_itn":("BOOLEAN", {"default": True},),
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ("result",)
|
||||
|
||||
def run(self,audio,device,language,num_threads,use_int8,use_itn ):
|
||||
|
||||
if device!=self.device:
|
||||
self.device=device
|
||||
self.processor=None
|
||||
if language!=self.language:
|
||||
self.language=language
|
||||
self.processor=None
|
||||
if num_threads!=self.num_threads:
|
||||
self.num_threads=num_threads
|
||||
self.processor=None
|
||||
if use_int8!=self.use_int8:
|
||||
self.use_int8=use_int8
|
||||
self.processor=None
|
||||
|
||||
if device=='auto' and torch.cuda.is_available():
|
||||
self.device='cuda'
|
||||
|
||||
# num_threads=4
|
||||
# use_int8=True
|
||||
|
||||
if self.processor==None:
|
||||
self.processor = SenseVoiceProcessor(self.download_model_path,
|
||||
self.device,
|
||||
self.num_threads,
|
||||
self.use_int8)
|
||||
|
||||
if 'waveform' in audio and 'sample_rate' in audio:
|
||||
waveform = audio['waveform']
|
||||
# print("Original shape:", waveform.shape) # 打印原始形状
|
||||
if waveform.ndim == 3 and waveform.shape[0] == 1: # 检查是否为三维且 batch_size 为 1
|
||||
waveform = waveform.squeeze(0) # 移除 batch_size 维度
|
||||
waveform_numpy = waveform.numpy().transpose(1, 0) # 转换为 (num_samples, num_channels)
|
||||
else:
|
||||
raise ValueError("Unexpected waveform dimensions")
|
||||
|
||||
_sample_rate = audio['sample_rate']
|
||||
|
||||
results=self.processor.process_audio(waveform_numpy, _sample_rate, language, use_itn)
|
||||
|
||||
|
||||
return (results,)
|
||||
@@ -280,7 +280,7 @@ class ChinesePrompt:
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 0xffffffffffffffff}),
|
||||
|
||||
},
|
||||
|
||||
@@ -331,13 +331,15 @@ class ChinesePrompt:
|
||||
|
||||
for t in texts:
|
||||
if t:
|
||||
# translated_text = translated_word = translate(zh_en_tokenizer,zh_en_model,str(t))
|
||||
parser = Lark(grammar, start="start", parser="lalr", transformer=ChinesePromptTranslate())
|
||||
# print('t',t)
|
||||
result = parser.parse(t).children
|
||||
# print('en_result',result)
|
||||
# en_text=translate(zh_en_tokenizer,zh_en_model,text_without_syntax)
|
||||
en_texts.append(result[0])
|
||||
try:
|
||||
result = parser.parse(t).children
|
||||
en_texts.append(result[0])
|
||||
except:
|
||||
print(f"Error parsing '{t}'")
|
||||
t = translate(str(t))
|
||||
en_texts.append(t)
|
||||
|
||||
|
||||
zh_en_model.to('cpu')
|
||||
print("test en_text",en_texts)
|
||||
@@ -384,7 +386,7 @@ class PromptGenerate:
|
||||
|
||||
"optional":{
|
||||
"multiple": (["off","on"],),
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 0xffffffffffffffff}),
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
@@ -82,13 +82,13 @@ def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
|
||||
def create_temp_file(image):
|
||||
def create_temp_file(image,counter=1):
|
||||
output_dir = folder_paths.get_temp_directory()
|
||||
|
||||
(
|
||||
full_output_folder,
|
||||
filename,
|
||||
counter,
|
||||
_,
|
||||
subfolder,
|
||||
_,
|
||||
) = folder_paths.get_save_image_path('tmp', output_dir)
|
||||
@@ -181,6 +181,28 @@ 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
|
||||
@@ -579,6 +601,7 @@ class AppInfo:
|
||||
"link":("STRING",{"multiline": False,"default": "https://","dynamicPrompts": False}),
|
||||
"category":("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
|
||||
"auto_save": (["enable","disable"],),
|
||||
"idle_animation": ("BOOLEAN", {"default": False},),
|
||||
}
|
||||
|
||||
}
|
||||
@@ -594,14 +617,21 @@ 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,idle_animation):
|
||||
name=name[0]
|
||||
|
||||
idle_animation=idle_animation[0]
|
||||
|
||||
im=None
|
||||
im=[]
|
||||
if image:
|
||||
im=image[0][0]
|
||||
#TODO batch 的方式需要处理
|
||||
im=create_temp_file(im)
|
||||
images=[image]
|
||||
# batch 的方式需要处理
|
||||
images=flatten_list(images)
|
||||
# img=image[0][0]
|
||||
print('AppInfo_image',len(images))
|
||||
for i in range(len(images)):
|
||||
img=images[i]
|
||||
im.append(create_temp_file(img,i+1)[0])
|
||||
# image [img,] img[batch,w,h,a] 列表里面是batch,
|
||||
|
||||
input_ids=input_ids[0]
|
||||
@@ -614,7 +644,50 @@ class AppInfo:
|
||||
|
||||
# id=get_json_hash([name,im,input_ids,output_ids,description,version])
|
||||
|
||||
return {"ui": {"json": [name,im,input_ids,output_ids,description,version,share_prefix,link,category]}, "result": ()}
|
||||
return {"ui": {"json": [name,im,input_ids,output_ids,description,version,share_prefix,link,category,idle_animation]}, "result": ()}
|
||||
|
||||
|
||||
|
||||
class CreateJsonNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"key": ("STRING",{"multiline": False,"default": "data","dynamicPrompts": False}),
|
||||
"value":(any_type,),
|
||||
"save":("BOOLEAN", {"default": True},),
|
||||
},
|
||||
"optional":{
|
||||
"json_str":("STRING", {"forceInput": True,}),
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("json_str",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Output"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = False
|
||||
# OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self,key,value,save,json_str=None):
|
||||
data={}
|
||||
|
||||
data[key]=value
|
||||
|
||||
if json_str:
|
||||
json_obj = json.loads(json_str)
|
||||
data.update(json_obj)
|
||||
|
||||
if save:
|
||||
# 保存为本地文件
|
||||
with open(os.path.join(folder_paths.get_output_directory(),'data.json'), 'w') as file:
|
||||
json.dump(data, file, ensure_ascii=False, indent=4)
|
||||
|
||||
return (json.dumps(data),)
|
||||
|
||||
|
||||
|
||||
@@ -795,7 +868,7 @@ class TESTNODE_:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"ANY":(any_type,),
|
||||
"ANY":(any_type,),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -810,6 +883,9 @@ class TESTNODE_:
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self,ANY):
|
||||
|
||||
print('#TESTNODE_',len(ANY))
|
||||
|
||||
print(type(ANY))
|
||||
try:
|
||||
print(ANY[0].shape)
|
||||
|
||||
@@ -22,7 +22,15 @@ import base64
|
||||
|
||||
import mimetypes
|
||||
|
||||
|
||||
# 使用递归的方法将嵌套的列表展平为一维列表
|
||||
def flatten_list(nested_list):
|
||||
flat_list = []
|
||||
for item in nested_list:
|
||||
if isinstance(item, list):
|
||||
flat_list.extend(flatten_list(item))
|
||||
else:
|
||||
flat_list.append(item)
|
||||
return flat_list
|
||||
|
||||
def get_frames(frame_count, frames, revert=False):
|
||||
if not revert:
|
||||
@@ -913,7 +921,7 @@ class GenerateFramesByCount:
|
||||
def r(self, frames, frame_count, revert):
|
||||
|
||||
image_list = [frames[i:i + 1, ...] for i in range(frames.shape[0])]
|
||||
|
||||
print('#image_list',len(image_list),frame_count)
|
||||
image_list=get_frames(frame_count,image_list,revert)
|
||||
|
||||
images = torch.cat(image_list, dim=0)
|
||||
@@ -928,7 +936,6 @@ class scenesNode_:
|
||||
return {"required": {
|
||||
"scenes_video": ('SCENE_VIDEO',),
|
||||
"index": ("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
|
||||
},}
|
||||
|
||||
RETURN_TYPES = ('IMAGE','INT',)
|
||||
@@ -940,14 +947,15 @@ class scenesNode_:
|
||||
INPUT_IS_LIST = True
|
||||
|
||||
def load_video_cv_fallback(self, video, frame_load_cap, skip_first_frames):
|
||||
# print('#video',video)
|
||||
|
||||
images = []
|
||||
total_frame_count = 0
|
||||
video_cap = cv2.VideoCapture(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)
|
||||
|
||||
@@ -986,11 +994,13 @@ class scenesNode_:
|
||||
finally:
|
||||
video_cap.release()
|
||||
|
||||
print("total_frame_count",total_frame_count)
|
||||
images = torch.cat(images, dim=0)
|
||||
|
||||
return (images, frames_added,)
|
||||
|
||||
def run(self, scenes_video,index):
|
||||
scenes_video=flatten_list(scenes_video)
|
||||
print('#scenes_video',index,scenes_video)
|
||||
index=index[0]
|
||||
if len(scenes_video) > index:
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
import itertools
|
||||
import re
|
||||
|
||||
LANGUAGE_UNICODE_RANGE_MAP = {
|
||||
"ZH": [(0x4E00, 0x9FFF)],
|
||||
"JP": [(0x4E00, 0x9FFF), (0x3040, 0x309F), (0x30A0, 0x30FF), (0x31F0, 0x31FF)],
|
||||
"EN": [(0x0000, 0x007F)],
|
||||
}
|
||||
|
||||
SYMBOLS_MAPPING = {
|
||||
":": ",",
|
||||
";": ",",
|
||||
",": ",",
|
||||
"。": ".",
|
||||
"!": "!",
|
||||
"?": "?",
|
||||
"\n": ".",
|
||||
"·": ",",
|
||||
"、": ",",
|
||||
"...": "…",
|
||||
"“": "'",
|
||||
"”": "'",
|
||||
"‘": "'",
|
||||
"’": "'",
|
||||
"(": "'",
|
||||
")": "'",
|
||||
"(": "'",
|
||||
")": "'",
|
||||
"《": "'",
|
||||
"》": "'",
|
||||
"【": "'",
|
||||
"】": "'",
|
||||
"[": "'",
|
||||
"]": "'",
|
||||
"—": "-",
|
||||
"~": "-",
|
||||
"~": "-",
|
||||
"・": "-",
|
||||
"「": "'",
|
||||
"」": "'",
|
||||
";": ",",
|
||||
":": ",",
|
||||
}
|
||||
|
||||
REPLACE_SYMBOL_REGEX = re.compile(
|
||||
"|".join(re.escape(p) for p in SYMBOLS_MAPPING.keys())
|
||||
)
|
||||
ALL_KNOWN_UTF8_RANGE = list(
|
||||
itertools.chain.from_iterable(LANGUAGE_UNICODE_RANGE_MAP.values())
|
||||
)
|
||||
REMOVE_UNKNOWN_SYMBOL_REGEX = re.compile(
|
||||
"[^"
|
||||
+ "".join(
|
||||
f"{re.escape(chr(start))}-{re.escape(chr(end))}"
|
||||
for start, end in ALL_KNOWN_UTF8_RANGE
|
||||
)
|
||||
+ "]"
|
||||
)
|
||||
|
||||
|
||||
def clean_text(text):
|
||||
# Clean the text
|
||||
text = text.strip()
|
||||
|
||||
# Replace all chinese symbols with their english counterparts
|
||||
text = REPLACE_SYMBOL_REGEX.sub(lambda x: SYMBOLS_MAPPING[x.group()], text)
|
||||
text = REMOVE_UNKNOWN_SYMBOL_REGEX.sub("", text)
|
||||
|
||||
return text
|
||||
@@ -0,0 +1,87 @@
|
||||
# Base configuration for training a model
|
||||
paths:
|
||||
run_dir: results/${project}
|
||||
ckpt_dir: ${paths.run_dir}/checkpoints
|
||||
|
||||
hydra:
|
||||
run:
|
||||
dir: ${paths.run_dir}
|
||||
|
||||
# Lightning Trainer
|
||||
trainer:
|
||||
_target_: lightning.pytorch.trainer.Trainer
|
||||
|
||||
default_root_dir: ${paths.run_dir}
|
||||
accelerator: gpu
|
||||
num_nodes: 1
|
||||
devices: auto
|
||||
strategy:
|
||||
_target_: lightning.pytorch.strategies.DDPStrategy
|
||||
process_group_backend: nccl # This should be override when training on windows
|
||||
|
||||
precision: bf16-mixed
|
||||
|
||||
# disable validation by epoch end
|
||||
check_val_every_n_epoch: null
|
||||
val_check_interval: 5000
|
||||
max_steps: 100_000
|
||||
|
||||
# Use torch.backends.cudnn.benchmark to speed up training
|
||||
benchmark: true
|
||||
|
||||
# Callbacks
|
||||
callbacks:
|
||||
model_checkpoint:
|
||||
_target_: lightning.pytorch.callbacks.ModelCheckpoint
|
||||
dirpath: ${paths.ckpt_dir}
|
||||
filename: "step_{step:09d}"
|
||||
save_last: false # additionally always save an exact copy of the last checkpoint to a file last.ckpt
|
||||
save_top_k: 5 # save 5 latest checkpoints
|
||||
monitor: step # use step to monitor checkpoints
|
||||
mode: max # save the latest checkpoint with the highest global_step
|
||||
every_n_epochs: null # don't save checkpoints by epoch end
|
||||
every_n_train_steps: 5000 # save checkpoints every 5000 steps
|
||||
auto_insert_metric_name: false
|
||||
|
||||
model_summary:
|
||||
_target_: lightning.pytorch.callbacks.ModelSummary
|
||||
max_depth: 2 # the maximum depth of layer nesting that the summary will include
|
||||
|
||||
learning_rate_monitor:
|
||||
_target_: lightning.pytorch.callbacks.LearningRateMonitor
|
||||
logging_interval: step
|
||||
log_momentum: false
|
||||
|
||||
grad_norm_monitor:
|
||||
_target_: fish_speech.callbacks.GradNormMonitor
|
||||
norm_type: 2
|
||||
logging_interval: step
|
||||
|
||||
# Logger
|
||||
logger:
|
||||
tensorboard:
|
||||
_target_: lightning.pytorch.loggers.tensorboard.TensorBoardLogger
|
||||
save_dir: "${paths.run_dir}/tensorboard/"
|
||||
name: null
|
||||
log_graph: false
|
||||
default_hp_metric: true
|
||||
prefix: ""
|
||||
|
||||
# wandb:
|
||||
# _target_: lightning.pytorch.loggers.wandb.WandbLogger
|
||||
# # name: "" # name of the run (normally generated by wandb)
|
||||
# save_dir: "${paths.run_dir}"
|
||||
# offline: False
|
||||
# id: null # pass correct id to resume experiment!
|
||||
# anonymous: null # enable anonymous logging
|
||||
# project: "fish-speech"
|
||||
# log_model: False # upload lightning ckpts
|
||||
# prefix: "" # a string to put at the beginning of metric keys
|
||||
# # entity: "" # set to name of your wandb team
|
||||
# group: ""
|
||||
# tags: ["vq", "hq", "finetune"]
|
||||
# job_type: ""
|
||||
|
||||
# Loop
|
||||
train: true
|
||||
test: false
|
||||
@@ -0,0 +1,33 @@
|
||||
_target_: fish_speech.models.vqgan.modules.firefly.FireflyArchitecture
|
||||
spec_transform:
|
||||
_target_: fish_speech.utils.spectrogram.LogMelSpectrogram
|
||||
sample_rate: 44100
|
||||
n_mels: 160
|
||||
n_fft: 2048
|
||||
hop_length: 512
|
||||
win_length: 2048
|
||||
backbone:
|
||||
_target_: fish_speech.models.vqgan.modules.firefly.ConvNeXtEncoder
|
||||
input_channels: 160
|
||||
depths: [3, 3, 9, 3]
|
||||
dims: [128, 256, 384, 512]
|
||||
drop_path_rate: 0.2
|
||||
kernel_size: 7
|
||||
head:
|
||||
_target_: fish_speech.models.vqgan.modules.firefly.HiFiGANGenerator
|
||||
hop_length: 512
|
||||
upsample_rates: [8, 8, 2, 2, 2] # aka. strides
|
||||
upsample_kernel_sizes: [16, 16, 4, 4, 4]
|
||||
resblock_kernel_sizes: [3, 7, 11]
|
||||
resblock_dilation_sizes: [[1, 3, 5], [1, 3, 5], [1, 3, 5]]
|
||||
num_mels: 512
|
||||
upsample_initial_channel: 512
|
||||
pre_conv_kernel_size: 13
|
||||
post_conv_kernel_size: 13
|
||||
quantizer:
|
||||
_target_: fish_speech.models.vqgan.modules.fsq.DownsampleFiniteScalarQuantize
|
||||
input_dim: 512
|
||||
n_groups: 8
|
||||
n_codebooks: 1
|
||||
levels: [8, 5, 5, 5]
|
||||
downsample_factor: [2, 2]
|
||||
@@ -0,0 +1,4 @@
|
||||
_target_: fish_speech.models.text2semantic.lora.LoraConfig
|
||||
r: 8
|
||||
lora_alpha: 16
|
||||
lora_dropout: 0.01
|
||||
@@ -0,0 +1,83 @@
|
||||
defaults:
|
||||
- base
|
||||
- _self_
|
||||
|
||||
project: text2semantic_finetune_dual_ar
|
||||
max_length: 4096
|
||||
pretrained_ckpt_path: checkpoints/fish-speech-1.4
|
||||
|
||||
# Lightning Trainer
|
||||
trainer:
|
||||
accumulate_grad_batches: 1
|
||||
gradient_clip_val: 1.0
|
||||
gradient_clip_algorithm: "norm"
|
||||
max_steps: 1000
|
||||
precision: bf16-true
|
||||
limit_val_batches: 10
|
||||
val_check_interval: 100
|
||||
|
||||
# Dataset Configuration
|
||||
tokenizer:
|
||||
_target_: transformers.AutoTokenizer.from_pretrained
|
||||
pretrained_model_name_or_path: ${pretrained_ckpt_path}
|
||||
|
||||
# Dataset Configuration
|
||||
train_dataset:
|
||||
_target_: fish_speech.datasets.semantic.AutoTextSemanticInstructionDataset
|
||||
proto_files:
|
||||
- data/protos
|
||||
tokenizer: ${tokenizer}
|
||||
causal: true
|
||||
max_length: ${max_length}
|
||||
use_speaker: false
|
||||
interactive_prob: 0.7
|
||||
|
||||
val_dataset:
|
||||
_target_: fish_speech.datasets.semantic.AutoTextSemanticInstructionDataset
|
||||
proto_files:
|
||||
- data/protos
|
||||
tokenizer: ${tokenizer}
|
||||
causal: true
|
||||
max_length: ${max_length}
|
||||
use_speaker: false
|
||||
interactive_prob: 0.7
|
||||
|
||||
data:
|
||||
_target_: fish_speech.datasets.semantic.SemanticDataModule
|
||||
train_dataset: ${train_dataset}
|
||||
val_dataset: ${val_dataset}
|
||||
num_workers: 4
|
||||
batch_size: 8
|
||||
tokenizer: ${tokenizer}
|
||||
max_length: ${max_length}
|
||||
|
||||
# Model Configuration
|
||||
model:
|
||||
_target_: fish_speech.models.text2semantic.lit_module.TextToSemantic
|
||||
model:
|
||||
_target_: fish_speech.models.text2semantic.llama.BaseTransformer.from_pretrained
|
||||
path: ${pretrained_ckpt_path}
|
||||
load_weights: true
|
||||
max_length: ${max_length}
|
||||
lora_config: null
|
||||
|
||||
optimizer:
|
||||
_target_: torch.optim.AdamW
|
||||
_partial_: true
|
||||
lr: 1e-4
|
||||
weight_decay: 0
|
||||
betas: [0.9, 0.95]
|
||||
eps: 1e-5
|
||||
|
||||
lr_scheduler:
|
||||
_target_: torch.optim.lr_scheduler.LambdaLR
|
||||
_partial_: true
|
||||
lr_lambda:
|
||||
_target_: fish_speech.scheduler.get_constant_schedule_with_warmup_lr_lambda
|
||||
_partial_: true
|
||||
num_warmup_steps: 10
|
||||
|
||||
# Callbacks
|
||||
callbacks:
|
||||
model_checkpoint:
|
||||
every_n_train_steps: ${trainer.val_check_interval}
|
||||
@@ -0,0 +1,2 @@
|
||||
SEMANTIC_TOKEN = "<|semantic|>"
|
||||
CODEBOOK_PAD_TOKEN_ID = 0
|
||||
@@ -0,0 +1,53 @@
|
||||
import bisect
|
||||
import random
|
||||
from typing import Iterable
|
||||
|
||||
from torch.utils.data import Dataset, IterableDataset
|
||||
|
||||
|
||||
class ConcatRepeatDataset(Dataset):
|
||||
datasets: list[Dataset]
|
||||
cumulative_sizes: list[int]
|
||||
repeats: list[int]
|
||||
|
||||
@staticmethod
|
||||
def cumsum(sequence, repeats):
|
||||
r, s = [], 0
|
||||
for dataset, repeat in zip(sequence, repeats):
|
||||
l = len(dataset) * repeat
|
||||
r.append(l + s)
|
||||
s += l
|
||||
return r
|
||||
|
||||
def __init__(self, datasets: Iterable[Dataset], repeats: list[int]):
|
||||
super().__init__()
|
||||
|
||||
self.datasets = list(datasets)
|
||||
self.repeats = repeats
|
||||
|
||||
assert len(self.datasets) > 0, "datasets should not be an empty iterable"
|
||||
assert len(self.datasets) == len(
|
||||
repeats
|
||||
), "datasets and repeats should have the same length"
|
||||
|
||||
for d in self.datasets:
|
||||
assert not isinstance(
|
||||
d, IterableDataset
|
||||
), "ConcatRepeatDataset does not support IterableDataset"
|
||||
|
||||
self.cumulative_sizes = self.cumsum(self.datasets, self.repeats)
|
||||
|
||||
def __len__(self):
|
||||
return self.cumulative_sizes[-1]
|
||||
|
||||
def __getitem__(self, idx):
|
||||
dataset_idx = bisect.bisect_right(self.cumulative_sizes, idx)
|
||||
|
||||
if dataset_idx == 0:
|
||||
sample_idx = idx
|
||||
else:
|
||||
sample_idx = idx - self.cumulative_sizes[dataset_idx - 1]
|
||||
|
||||
dataset = self.datasets[dataset_idx]
|
||||
|
||||
return dataset[sample_idx % len(dataset)]
|
||||
@@ -0,0 +1,24 @@
|
||||
syntax = "proto3";
|
||||
|
||||
package text_data;
|
||||
|
||||
message Semantics {
|
||||
repeated uint32 values = 1;
|
||||
}
|
||||
|
||||
message Sentence {
|
||||
repeated string texts = 1;
|
||||
repeated Semantics semantics = 3;
|
||||
}
|
||||
|
||||
message TextData {
|
||||
string source = 1;
|
||||
string name = 2;
|
||||
repeated Sentence sentences = 4;
|
||||
}
|
||||
|
||||
message SampledData {
|
||||
string source = 1;
|
||||
string name = 2;
|
||||
repeated Sentence samples = 3;
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Generated by the protocol buffer compiler. DO NOT EDIT!
|
||||
# source: text-data.proto
|
||||
# Protobuf Python Version: 4.25.1
|
||||
"""Generated protocol buffer code."""
|
||||
from google.protobuf import descriptor as _descriptor
|
||||
from google.protobuf import descriptor_pool as _descriptor_pool
|
||||
from google.protobuf import symbol_database as _symbol_database
|
||||
from google.protobuf.internal import builder as _builder
|
||||
|
||||
# @@protoc_insertion_point(imports)
|
||||
|
||||
_sym_db = _symbol_database.Default()
|
||||
|
||||
|
||||
DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(
|
||||
b'\n\x0ftext-data.proto\x12\ttext_data"\x1b\n\tSemantics\x12\x0e\n\x06values\x18\x01 \x03(\r"B\n\x08Sentence\x12\r\n\x05texts\x18\x01 \x03(\t\x12\'\n\tsemantics\x18\x03 \x03(\x0b\x32\x14.text_data.Semantics"P\n\x08TextData\x12\x0e\n\x06source\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12&\n\tsentences\x18\x04 \x03(\x0b\x32\x13.text_data.Sentence"Q\n\x0bSampledData\x12\x0e\n\x06source\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12$\n\x07samples\x18\x03 \x03(\x0b\x32\x13.text_data.Sentenceb\x06proto3'
|
||||
)
|
||||
|
||||
_globals = globals()
|
||||
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
|
||||
_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, "text_data_pb2", _globals)
|
||||
if _descriptor._USE_C_DESCRIPTORS == False:
|
||||
DESCRIPTOR._options = None
|
||||
_globals["_SEMANTICS"]._serialized_start = 30
|
||||
_globals["_SEMANTICS"]._serialized_end = 57
|
||||
_globals["_SENTENCE"]._serialized_start = 59
|
||||
_globals["_SENTENCE"]._serialized_end = 125
|
||||
_globals["_TEXTDATA"]._serialized_start = 127
|
||||
_globals["_TEXTDATA"]._serialized_end = 207
|
||||
_globals["_SAMPLEDDATA"]._serialized_start = 209
|
||||
_globals["_SAMPLEDDATA"]._serialized_end = 290
|
||||
# @@protoc_insertion_point(module_scope)
|
||||
@@ -0,0 +1,36 @@
|
||||
import struct
|
||||
|
||||
from .text_data_pb2 import TextData
|
||||
|
||||
|
||||
def read_pb_stream(f):
|
||||
while True:
|
||||
buf = f.read(4)
|
||||
if len(buf) == 0:
|
||||
break
|
||||
size = struct.unpack("I", buf)[0]
|
||||
buf = f.read(size)
|
||||
text_data = TextData()
|
||||
text_data.ParseFromString(buf)
|
||||
yield text_data
|
||||
|
||||
|
||||
def write_pb_stream(f, text_data):
|
||||
buf = text_data.SerializeToString()
|
||||
f.write(struct.pack("I", len(buf)))
|
||||
f.write(buf)
|
||||
|
||||
|
||||
def pack_pb_stream(text_data):
|
||||
buf = text_data.SerializeToString()
|
||||
return struct.pack("I", len(buf)) + buf
|
||||
|
||||
|
||||
def split_pb_stream(f):
|
||||
while True:
|
||||
head = f.read(4)
|
||||
if len(head) == 0:
|
||||
break
|
||||
size = struct.unpack("I", head)[0]
|
||||
buf = f.read(size)
|
||||
yield head + buf
|
||||
@@ -0,0 +1,496 @@
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain
|
||||
from pathlib import Path
|
||||
from random import Random
|
||||
from typing import Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from datasets.download.streaming_download_manager import xopen
|
||||
from huggingface_hub import HfApi
|
||||
from lightning import LightningDataModule
|
||||
from torch.distributed import get_rank, get_world_size, is_initialized
|
||||
from torch.utils.data import DataLoader, IterableDataset, get_worker_info
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fish_speech.conversation import CODEBOOK_PAD_TOKEN_ID
|
||||
from fish_speech.datasets.protos.text_data_pb2 import SampledData
|
||||
from fish_speech.datasets.protos.text_data_stream import read_pb_stream
|
||||
from fish_speech.text.clean import clean_text
|
||||
from fish_speech.utils import RankedLogger
|
||||
from fish_speech.utils.braceexpand import braceexpand
|
||||
|
||||
log = RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
def split_by_rank_worker(files):
|
||||
# We need to know the total number of devices
|
||||
# to split the data properly
|
||||
|
||||
total_devices = 1
|
||||
if is_initialized():
|
||||
total_devices = get_world_size()
|
||||
|
||||
worker_info = get_worker_info()
|
||||
if worker_info is not None:
|
||||
total_devices *= worker_info.num_workers
|
||||
|
||||
if len(files) < total_devices:
|
||||
# Repeat the files N times to match the number of devices
|
||||
files = files * (total_devices // len(files) + 1)
|
||||
|
||||
# DDP
|
||||
if is_initialized():
|
||||
files = files[get_rank() :: get_world_size()]
|
||||
|
||||
# Split by worker
|
||||
if worker_info is not None:
|
||||
files = files[worker_info.id :: worker_info.num_workers]
|
||||
|
||||
return files
|
||||
|
||||
|
||||
class AutoTextSemanticInstructionDataset(IterableDataset):
|
||||
"""
|
||||
Auto Augment Dataset by Speaker
|
||||
|
||||
1. Random concatenate multiple sentences from the same speaker to form a longer sentence
|
||||
2. Automatically normalize the text
|
||||
|
||||
For interactive mode, we use the following format (multiple sequences):
|
||||
<s> [INST] [SPK: speaker] text [/INST] ... [INST] text [/INST] </s>
|
||||
|
||||
For non-interactive mode, we use the following format (one long sequence):
|
||||
<s> [INST] text [/INST] ... </s>
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
proto_files: list[str],
|
||||
seed: int = 42,
|
||||
interactive_prob: float = 0.5,
|
||||
max_length: int = 1024,
|
||||
tokenizer: AutoTokenizer = None,
|
||||
use_speaker: bool | float = True,
|
||||
causal: bool = True,
|
||||
num_codebooks: Optional[int] = None,
|
||||
skip_text_prob: float = 0.0,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
proto_files: proto buf files if using local data
|
||||
seed: random seed
|
||||
interactive_prob: probability to use interactive mode
|
||||
max_length: max length of the text
|
||||
tokenizer: tokenizer
|
||||
use_speaker: include speaker information in the prompt
|
||||
causal: use causal sampling when using local data, disable will lead to random sampling
|
||||
num_codebooks: number of codebooks, if None, it will be automatically detected
|
||||
skip_text_prob: probability to skip the text (audio only), this only applies to interactive mode
|
||||
"""
|
||||
|
||||
super().__init__()
|
||||
|
||||
assert 0 <= interactive_prob <= 1, "interactive_prob must be in [0, 1]"
|
||||
|
||||
self.seed = seed
|
||||
self.max_length = max_length
|
||||
self.tokenizer = tokenizer
|
||||
self.interactive_prob = interactive_prob
|
||||
self.use_speaker = use_speaker
|
||||
self.proto_files = proto_files
|
||||
self.causal = causal
|
||||
self.num_codebooks = num_codebooks
|
||||
self.skip_text_prob = skip_text_prob
|
||||
|
||||
self.semantic_token_id = self.tokenizer.convert_tokens_to_ids("<|semantic|>")
|
||||
self.groups = None
|
||||
|
||||
def init_mock_data_server(self):
|
||||
if self.groups is not None:
|
||||
return
|
||||
|
||||
# Expand the proto files
|
||||
expanded_proto_files = []
|
||||
for filename in self.proto_files:
|
||||
for i in braceexpand(filename):
|
||||
i = Path(i)
|
||||
if i.is_file():
|
||||
expanded_proto_files.append(i)
|
||||
elif i.is_dir():
|
||||
expanded_proto_files.extend(i.rglob("*.proto"))
|
||||
expanded_proto_files.extend(i.rglob("*.protos"))
|
||||
else:
|
||||
raise ValueError(f"{i} is not a file or directory")
|
||||
|
||||
expanded_proto_files = sorted(expanded_proto_files)
|
||||
Random(self.seed).shuffle(expanded_proto_files)
|
||||
|
||||
self.groups = []
|
||||
shard_proto_files = split_by_rank_worker(expanded_proto_files)
|
||||
log.info(
|
||||
f"Reading {len(shard_proto_files)} / {len(expanded_proto_files)} files"
|
||||
)
|
||||
|
||||
count = 0
|
||||
for filename in shard_proto_files:
|
||||
with open(filename, "rb") as f:
|
||||
for text_data in read_pb_stream(f):
|
||||
self.groups.append(text_data)
|
||||
count += 1
|
||||
|
||||
log.info(f"Read total {count} groups of data")
|
||||
|
||||
# Shuffle the lines
|
||||
Random(self.seed).shuffle(self.groups)
|
||||
self.group_weights = [len(i.sentences) for i in self.groups]
|
||||
|
||||
def __iter__(self):
|
||||
while True:
|
||||
yield self.augment()
|
||||
|
||||
def tokenize_sentence(self, sentence: str):
|
||||
sentence = clean_text(sentence)
|
||||
tokens = self.tokenizer.encode(
|
||||
f"{sentence}",
|
||||
max_length=10**6,
|
||||
add_special_tokens=False,
|
||||
truncation=False,
|
||||
)
|
||||
return sentence, len(tokens)
|
||||
|
||||
def sample_data(self):
|
||||
if self.groups is None:
|
||||
self.init_mock_data_server()
|
||||
|
||||
# Shuffle unique lines, estimate that each sample is at least 20 tokens
|
||||
num_samples = self.max_length // 20
|
||||
|
||||
# choice group based on their number of samples
|
||||
group = random.choices(self.groups, weights=self.group_weights, k=1)[0]
|
||||
|
||||
if self.causal:
|
||||
# Sample in order
|
||||
if num_samples >= len(group.sentences):
|
||||
samples = group.sentences
|
||||
else:
|
||||
begin = random.randint(0, len(group.sentences) - num_samples)
|
||||
samples = group.sentences[begin : begin + num_samples]
|
||||
else:
|
||||
samples = random.choices(
|
||||
group.sentences, k=min(num_samples, len(group.sentences))
|
||||
)
|
||||
|
||||
return SampledData(
|
||||
source=group.source,
|
||||
name=group.name,
|
||||
samples=samples,
|
||||
)
|
||||
|
||||
def augment(self):
|
||||
final_text, final_semantic = [], []
|
||||
response = self.sample_data()
|
||||
if len(response.samples) == 0:
|
||||
# Invalid group
|
||||
return None
|
||||
|
||||
samples = list(response.samples)
|
||||
idx = 0
|
||||
use_interactive = random.random() < self.interactive_prob
|
||||
|
||||
if use_interactive is False:
|
||||
# Random sample based on speaker using a truncated normal distribution
|
||||
a = torch.tensor([0], dtype=torch.float32)
|
||||
torch.nn.init.trunc_normal_(
|
||||
a,
|
||||
mean=self.max_length // 2,
|
||||
std=self.max_length // 4,
|
||||
a=10,
|
||||
b=self.max_length,
|
||||
)
|
||||
remaining_tokens = a.long().item() - 4
|
||||
else:
|
||||
remaining_tokens = self.max_length
|
||||
|
||||
# Use speaker
|
||||
if isinstance(self.use_speaker, float):
|
||||
use_speaker = random.random() < self.use_speaker
|
||||
else:
|
||||
use_speaker = self.use_speaker
|
||||
|
||||
all_tokens, all_labels = [], []
|
||||
while remaining_tokens > 0 and len(samples) > 0:
|
||||
sentence = samples.pop(0)
|
||||
|
||||
text = random.choice(sentence.texts)
|
||||
text, length = self.tokenize_sentence(text)
|
||||
remaining_tokens -= length + len(sentence.semantics[0].values)
|
||||
|
||||
if use_interactive is False:
|
||||
final_text.append(text)
|
||||
final_semantic.append(sentence.semantics)
|
||||
else:
|
||||
# For interactive mode, we only apply speaker for the first sentence
|
||||
# [INST] [SPK: speaker] text [/INST] ... [INST] text [/INST]
|
||||
tokens, labels = self.pack_sentences(
|
||||
sentences=[text],
|
||||
semantics=[sentence.semantics],
|
||||
speaker=response.name if use_speaker else None,
|
||||
skip_text=random.random() < self.skip_text_prob,
|
||||
)
|
||||
|
||||
all_tokens.append(tokens)
|
||||
all_labels.append(labels)
|
||||
|
||||
idx += 1
|
||||
|
||||
if use_interactive is False:
|
||||
tokens, labels = self.pack_sentences(
|
||||
final_text,
|
||||
semantics=final_semantic,
|
||||
speaker=response.name if use_speaker else None,
|
||||
)
|
||||
all_tokens.append(tokens)
|
||||
all_labels.append(labels)
|
||||
|
||||
tokens = torch.cat(all_tokens, dim=1)
|
||||
labels = torch.cat(all_labels, dim=1)
|
||||
|
||||
# Verify that the length is correct
|
||||
assert tokens.size(1) == labels.size(1), f"{tokens.size(1)} != {labels.size(1)}"
|
||||
|
||||
data = {"tokens": tokens, "labels": labels}
|
||||
|
||||
return data
|
||||
|
||||
def pack_sentences(
|
||||
self,
|
||||
sentences: list[str],
|
||||
semantics: list,
|
||||
speaker: Optional[str] = None,
|
||||
skip_text: bool = False,
|
||||
):
|
||||
if speaker is None:
|
||||
speaker = "assistant"
|
||||
|
||||
cated_sentences = " ".join(sentences)
|
||||
if skip_text:
|
||||
cated_sentences = "<|skip_text|>"
|
||||
|
||||
final_text = "<|im_start|>user\n" + cated_sentences + "<|im_end|>"
|
||||
final_text = final_text + f"<|im_start|>{speaker}\n"
|
||||
|
||||
encoded = self.tokenizer.encode(
|
||||
final_text,
|
||||
add_special_tokens=False,
|
||||
truncation=False,
|
||||
max_length=10**6,
|
||||
)
|
||||
semantic_length = sum([len(i[0].values) for i in semantics])
|
||||
prompt_length = len(encoded)
|
||||
num_codebooks = (
|
||||
len(semantics[0]) if self.num_codebooks is None else self.num_codebooks
|
||||
)
|
||||
|
||||
# Pack the tokens and semantics (add <s> and </s> to semantic tokens)
|
||||
tokens = (
|
||||
encoded
|
||||
+ [self.semantic_token_id] * semantic_length
|
||||
+ self.tokenizer.convert_tokens_to_ids(["<|im_end|>"])
|
||||
)
|
||||
|
||||
# Codebook bos/padding: 0, eos: 1
|
||||
codes = [[CODEBOOK_PAD_TOKEN_ID] * prompt_length for _ in range(num_codebooks)]
|
||||
for segment in semantics:
|
||||
for book_idx, book in zip(range(num_codebooks), segment):
|
||||
for j in book.values:
|
||||
codes[book_idx].append(int(j) + 1)
|
||||
|
||||
for book in codes:
|
||||
book.extend([CODEBOOK_PAD_TOKEN_ID] * 1)
|
||||
|
||||
tokens = [tokens] + codes
|
||||
|
||||
tokens = torch.tensor(tokens, dtype=torch.long)
|
||||
labels = tokens.clone()
|
||||
|
||||
if skip_text:
|
||||
# If text is not provided, the sentence is used for condition only, all labels are -100
|
||||
torch.fill_(labels, -100)
|
||||
return tokens, labels
|
||||
|
||||
# Mask out the <s> tokens for semantic, predict semantic tokens only
|
||||
# Since we don't mask out the input tokens, the language modeling still works
|
||||
labels[1:, :prompt_length] = -100
|
||||
|
||||
tokens = tokens[:, :-1]
|
||||
labels = labels[:, 1:]
|
||||
|
||||
# Verify the padding is correct, and the last token is eos
|
||||
assert (tokens[1:, :prompt_length] == CODEBOOK_PAD_TOKEN_ID).all()
|
||||
assert (labels[1:, -1:] == CODEBOOK_PAD_TOKEN_ID).all()
|
||||
|
||||
return tokens, labels
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextDataCollator:
|
||||
tokenizer: AutoTokenizer
|
||||
max_length: int = 1024
|
||||
|
||||
def __call__(self, examples):
|
||||
if "negative_tokens" in examples:
|
||||
positive_examples = []
|
||||
negative_examples = []
|
||||
|
||||
for i in examples:
|
||||
positive_examples.append(
|
||||
{
|
||||
"tokens": i["tokens"],
|
||||
"labels": i["labels"],
|
||||
}
|
||||
)
|
||||
negative_examples.append(
|
||||
{
|
||||
"tokens": i["negative_tokens"],
|
||||
"labels": i["negative_labels"],
|
||||
}
|
||||
)
|
||||
|
||||
examples = positive_examples + negative_examples
|
||||
|
||||
return self.batchify(examples)
|
||||
|
||||
def batchify(self, examples, tokens_key="tokens", labels_key="labels"):
|
||||
tokens, attention_masks, labels = [], [], []
|
||||
|
||||
# Calculate the max length
|
||||
max_tokens_length = 0
|
||||
for example in examples:
|
||||
max_tokens_length = max(max_tokens_length, example[tokens_key].size(1))
|
||||
max_tokens_length = min(max_tokens_length, self.max_length)
|
||||
|
||||
for example in examples:
|
||||
_tokens = example[tokens_key][:, :max_tokens_length]
|
||||
_labels = example[labels_key][:, :max_tokens_length]
|
||||
_attention_mask = torch.ones((max_tokens_length,), dtype=torch.bool)
|
||||
tokens_length = _tokens.size(1)
|
||||
_attention_mask[:tokens_length] = False
|
||||
|
||||
assert tokens_length == _labels.size(
|
||||
1
|
||||
), f"{tokens_length} != {_labels.size(1)}"
|
||||
|
||||
if tokens_length < max_tokens_length:
|
||||
_tokens = F.pad(
|
||||
_tokens,
|
||||
(0, max_tokens_length - tokens_length),
|
||||
value=self.tokenizer.eos_token_id,
|
||||
)
|
||||
_tokens[1:, tokens_length:] = CODEBOOK_PAD_TOKEN_ID
|
||||
_labels = F.pad(
|
||||
_labels, (0, max_tokens_length - _labels.size(1)), value=-100
|
||||
)
|
||||
|
||||
tokens.append(_tokens)
|
||||
attention_masks.append(_attention_mask)
|
||||
labels.append(_labels)
|
||||
|
||||
tokens = torch.stack(tokens, dim=0)
|
||||
attention_masks = torch.stack(attention_masks, dim=0)
|
||||
labels = torch.stack(labels, dim=0)
|
||||
|
||||
return {
|
||||
"inputs": tokens,
|
||||
"attention_masks": attention_masks,
|
||||
"labels": labels,
|
||||
}
|
||||
|
||||
|
||||
class InterleaveDataset(IterableDataset):
|
||||
def __init__(
|
||||
self,
|
||||
datasets: list[IterableDataset],
|
||||
probabilities: list[float],
|
||||
seed: int = 42,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.datasets = datasets
|
||||
self.probabilities = probabilities
|
||||
self.seed = seed
|
||||
|
||||
def __iter__(self):
|
||||
rng = np.random.default_rng(self.seed)
|
||||
dataset_iterators = [iter(dataset) for dataset in self.datasets]
|
||||
|
||||
while True:
|
||||
# Random choice one
|
||||
dataset_idx = rng.choice(len(self.datasets), p=self.probabilities)
|
||||
dataset_iterator = dataset_iterators[dataset_idx]
|
||||
|
||||
try:
|
||||
yield next(dataset_iterator)
|
||||
except StopIteration:
|
||||
# Exhausted, create a new iterator
|
||||
dataset_iterators[dataset_idx] = iter(self.datasets[dataset_idx])
|
||||
yield next(dataset_iterators[dataset_idx])
|
||||
|
||||
|
||||
class SemanticDataModule(LightningDataModule):
|
||||
def __init__(
|
||||
self,
|
||||
train_dataset: Union[AutoTextSemanticInstructionDataset, InterleaveDataset],
|
||||
val_dataset: Union[AutoTextSemanticInstructionDataset, InterleaveDataset],
|
||||
batch_size: int = 32,
|
||||
tokenizer: AutoTokenizer = None,
|
||||
max_length: int = 1024,
|
||||
num_workers: int = 4,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.train_dataset = train_dataset
|
||||
self.val_dataset = val_dataset
|
||||
self.batch_size = batch_size
|
||||
self.tokenizer = tokenizer
|
||||
self.max_length = max_length
|
||||
self.num_workers = num_workers
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=self.batch_size,
|
||||
collate_fn=TextDataCollator(self.tokenizer, self.max_length),
|
||||
num_workers=self.num_workers,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
def val_dataloader(self):
|
||||
return DataLoader(
|
||||
self.val_dataset,
|
||||
batch_size=self.batch_size,
|
||||
collate_fn=TextDataCollator(self.tokenizer, self.max_length),
|
||||
num_workers=self.num_workers,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from tqdm import tqdm
|
||||
|
||||
ds = AutoTextSemanticInstructionDataset(
|
||||
["data/protos"],
|
||||
tokenizer=AutoTokenizer.from_pretrained("fishaudio/fish-speech-1"),
|
||||
use_speaker=False,
|
||||
interactive_prob=1.0,
|
||||
skip_text_prob=0.5,
|
||||
)
|
||||
|
||||
for i in ds:
|
||||
print(ds.tokenizer.decode(i["tokens"][0], skip_special_tokens=False))
|
||||
# i["labels"][0][i["labels"][0] == -100] = 0
|
||||
# print(ds.tokenizer.decode(i["labels"][0], skip_special_tokens=False))
|
||||
break
|
||||
@@ -0,0 +1,147 @@
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
from lightning import LightningDataModule
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
|
||||
from fish_speech.utils import RankedLogger
|
||||
|
||||
logger = RankedLogger(__name__, rank_zero_only=False)
|
||||
|
||||
|
||||
class VQGANDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
filelist: str,
|
||||
sample_rate: int = 32000,
|
||||
hop_length: int = 640,
|
||||
slice_frames: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
filelist = Path(filelist)
|
||||
root = filelist.parent
|
||||
|
||||
self.files = [
|
||||
root / line.strip()
|
||||
for line in filelist.read_text(encoding="utf-8").splitlines()
|
||||
if line.strip()
|
||||
]
|
||||
self.sample_rate = sample_rate
|
||||
self.hop_length = hop_length
|
||||
self.slice_frames = slice_frames
|
||||
|
||||
def __len__(self):
|
||||
return len(self.files)
|
||||
|
||||
def get_item(self, idx):
|
||||
file = self.files[idx]
|
||||
|
||||
audio, _ = librosa.load(file, sr=self.sample_rate, mono=True)
|
||||
|
||||
# Slice audio and features
|
||||
if (
|
||||
self.slice_frames is not None
|
||||
and audio.shape[0] > self.slice_frames * self.hop_length
|
||||
):
|
||||
start = np.random.randint(
|
||||
0, audio.shape[0] - self.slice_frames * self.hop_length
|
||||
)
|
||||
audio = audio[start : start + self.slice_frames * self.hop_length]
|
||||
|
||||
if len(audio) == 0:
|
||||
return None
|
||||
|
||||
max_value = np.abs(audio).max()
|
||||
if max_value > 1.0:
|
||||
audio = audio / max_value
|
||||
|
||||
return {
|
||||
"audio": torch.from_numpy(audio),
|
||||
}
|
||||
|
||||
def __getitem__(self, idx):
|
||||
try:
|
||||
return self.get_item(idx)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
logger.error(f"Error loading {self.files[idx]}: {e}")
|
||||
return None
|
||||
|
||||
|
||||
@dataclass
|
||||
class VQGANCollator:
|
||||
def __call__(self, batch):
|
||||
batch = [x for x in batch if x is not None]
|
||||
|
||||
audio_lengths = torch.tensor([len(x["audio"]) for x in batch])
|
||||
audio_maxlen = audio_lengths.max()
|
||||
|
||||
# Rounds up to nearest multiple of 2 (audio_lengths)
|
||||
audios = []
|
||||
for x in batch:
|
||||
audios.append(
|
||||
torch.nn.functional.pad(x["audio"], (0, audio_maxlen - len(x["audio"])))
|
||||
)
|
||||
|
||||
return {
|
||||
"audios": torch.stack(audios),
|
||||
"audio_lengths": audio_lengths,
|
||||
}
|
||||
|
||||
|
||||
class VQGANDataModule(LightningDataModule):
|
||||
def __init__(
|
||||
self,
|
||||
train_dataset: VQGANDataset,
|
||||
val_dataset: VQGANDataset,
|
||||
batch_size: int = 32,
|
||||
num_workers: int = 4,
|
||||
val_batch_size: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.train_dataset = train_dataset
|
||||
self.val_dataset = val_dataset
|
||||
self.batch_size = batch_size
|
||||
self.val_batch_size = val_batch_size or batch_size
|
||||
self.num_workers = num_workers
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=self.batch_size,
|
||||
collate_fn=VQGANCollator(),
|
||||
num_workers=self.num_workers,
|
||||
shuffle=True,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
def val_dataloader(self):
|
||||
return DataLoader(
|
||||
self.val_dataset,
|
||||
batch_size=self.val_batch_size,
|
||||
collate_fn=VQGANCollator(),
|
||||
num_workers=self.num_workers,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = VQGANDataset("data/LibriTTS_R/vq_train_filelist.txt")
|
||||
dataloader = DataLoader(
|
||||
dataset, batch_size=4, shuffle=False, collate_fn=VQGANCollator()
|
||||
)
|
||||
|
||||
for batch in dataloader:
|
||||
print(batch["audios"].shape)
|
||||
print(batch["features"].shape)
|
||||
print(batch["audio_lengths"])
|
||||
print(batch["feature_lengths"])
|
||||
break
|
||||
@@ -0,0 +1,104 @@
|
||||
|
||||
import torch
|
||||
from .models.text2semantic.llama import BaseTransformer, NaiveTransformer, DualARTransformer
|
||||
from .tools.llama.generate import decode_one_token_ar, decode_one_token_naive, generate_long
|
||||
import numpy as np
|
||||
import time
|
||||
from typing import Union
|
||||
from loguru import logger
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
def load_model(checkpoint_path, device, precision, compile=False):
|
||||
model: Union[NaiveTransformer, DualARTransformer] = BaseTransformer.from_pretrained(
|
||||
checkpoint_path, load_weights=True
|
||||
)
|
||||
|
||||
model = model.to(device=device, dtype=precision)
|
||||
logger.info(f"Restored model from checkpoint")
|
||||
|
||||
if isinstance(model, DualARTransformer):
|
||||
decode_one_token = decode_one_token_ar
|
||||
logger.info("Using DualARTransformer")
|
||||
else:
|
||||
decode_one_token = decode_one_token_naive
|
||||
logger.info("Using NaiveTransformer")
|
||||
|
||||
if compile:
|
||||
logger.info("Compiling function...")
|
||||
decode_one_token = torch.compile(
|
||||
decode_one_token, mode="reduce-overhead", fullgraph=True
|
||||
)
|
||||
|
||||
return model.eval(), decode_one_token
|
||||
|
||||
|
||||
def prompt2semantic(
|
||||
model: DualARTransformer,
|
||||
decode_one_token: callable,
|
||||
text: str,
|
||||
prompt_text: Optional[list[str]],
|
||||
prompt_tokens: Optional[list[np.ndarray]],
|
||||
max_new_tokens: int,
|
||||
top_p: float,
|
||||
repetition_penalty: float,
|
||||
temperature: float,
|
||||
device: str,
|
||||
compile: bool,
|
||||
seed: int,
|
||||
iterative_prompt: bool,
|
||||
chunk_length: int,
|
||||
):
|
||||
|
||||
if prompt_text is not None and len(prompt_text) != len(prompt_tokens):
|
||||
raise ValueError(
|
||||
f"Number of prompt text ({len(prompt_text)}) and prompt tokens ({len(prompt_tokens)}) should be the same"
|
||||
)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
if prompt_tokens is not None:
|
||||
prompt_tokens = [torch.from_numpy(pt).to(device) for pt in prompt_tokens]
|
||||
|
||||
torch.manual_seed(seed)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
generator = generate_long(
|
||||
model=model,
|
||||
device=device,
|
||||
decode_one_token=decode_one_token,
|
||||
text=text,
|
||||
num_samples=1,
|
||||
max_new_tokens=max_new_tokens,
|
||||
top_p=top_p,
|
||||
repetition_penalty=repetition_penalty,
|
||||
temperature=temperature,
|
||||
compile=compile,
|
||||
iterative_prompt=iterative_prompt,
|
||||
chunk_length=chunk_length,
|
||||
prompt_text=prompt_text,
|
||||
prompt_tokens=prompt_tokens,
|
||||
)
|
||||
|
||||
idx = 0
|
||||
all_codes = []
|
||||
codes = []
|
||||
|
||||
for response in generator:
|
||||
if response.action == "sample":
|
||||
codes.append(response.codes)
|
||||
logger.info(f"Sampled text: {response.text}")
|
||||
elif response.action == "next":
|
||||
if codes:
|
||||
all_codes.append(torch.cat(codes, dim=1).cpu().numpy())
|
||||
logger.info(f"Saved codes to codes_{idx}.npy")
|
||||
logger.info(f"Next sample")
|
||||
codes = []
|
||||
idx += 1
|
||||
else:
|
||||
logger.error(f"Error: {response}")
|
||||
|
||||
return all_codes
|
||||
@@ -0,0 +1,202 @@
|
||||
from typing import Any, Optional
|
||||
|
||||
import lightning as L
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from lightning.pytorch.utilities.types import OptimizerLRScheduler
|
||||
|
||||
import fish_speech.utils as utils
|
||||
from fish_speech.conversation import CODEBOOK_PAD_TOKEN_ID
|
||||
from fish_speech.models.text2semantic.llama import NaiveTransformer
|
||||
|
||||
log = utils.RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
class TextToSemantic(L.LightningModule):
|
||||
def __init__(
|
||||
self,
|
||||
model: NaiveTransformer,
|
||||
optimizer: Any,
|
||||
lr_scheduler: Any,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.model = model
|
||||
self.optimizer_builder = optimizer
|
||||
self.lr_scheduler_builder = lr_scheduler
|
||||
|
||||
def forward(self, x):
|
||||
return self.model(x)
|
||||
|
||||
def on_save_checkpoint(self, checkpoint):
|
||||
# Save only LoRA parameters
|
||||
state_dict = checkpoint["state_dict"]
|
||||
use_lora = any("lora" in name for name in state_dict.keys())
|
||||
if not use_lora:
|
||||
return
|
||||
|
||||
for name in list(state_dict.keys()):
|
||||
if "lora" not in name:
|
||||
state_dict.pop(name)
|
||||
|
||||
def configure_optimizers(self) -> OptimizerLRScheduler:
|
||||
# Get weight decay parameters
|
||||
weight_decay_parameters, other_parameters = [], []
|
||||
for name, param in self.named_parameters():
|
||||
if ".bias" in name or "norm.weight" in name or ".embeddings." in name:
|
||||
other_parameters.append(param)
|
||||
else:
|
||||
weight_decay_parameters.append(param)
|
||||
|
||||
optimizer = self.optimizer_builder(
|
||||
[
|
||||
{"params": weight_decay_parameters},
|
||||
{"params": other_parameters, "weight_decay": 0.0},
|
||||
]
|
||||
)
|
||||
|
||||
# Print the parameters and their weight decay
|
||||
for i in optimizer.param_groups:
|
||||
log.info(
|
||||
f"Set weight decay: {i['weight_decay']} for {len(i['params'])} parameters"
|
||||
)
|
||||
|
||||
lr_scheduler = self.lr_scheduler_builder(optimizer)
|
||||
|
||||
return {
|
||||
"optimizer": optimizer,
|
||||
"lr_scheduler": {
|
||||
"scheduler": lr_scheduler,
|
||||
"interval": "step",
|
||||
},
|
||||
}
|
||||
|
||||
# Copied from https://github.com/eric-mitchell/direct-preference-optimization/blob/main/trainers.py#L90
|
||||
def get_batch_logps(
|
||||
self,
|
||||
logits: torch.FloatTensor,
|
||||
labels: torch.LongTensor,
|
||||
average_log_prob: bool = False,
|
||||
) -> torch.FloatTensor:
|
||||
"""Compute the log probabilities of the given labels under the given logits.
|
||||
|
||||
Args:
|
||||
logits: Logits of the model (unnormalized). Shape: (batch_size, sequence_length, codebook_size, vocab_size)
|
||||
labels: Labels for which to compute the log probabilities. Label tokens with a value of -100 are ignored. Shape: (batch_size, sequence_length, codebook_size)
|
||||
average_log_prob: If True, return the average log probability per (non-masked) token. Otherwise, return the sum of the log probabilities of the (non-masked) tokens.
|
||||
|
||||
Returns:
|
||||
A tensor of shape (batch_size,) containing the average/sum log probabilities of the given labels under the given logits.
|
||||
"""
|
||||
assert logits.shape[:-1] == labels.shape
|
||||
|
||||
labels = labels.clone()
|
||||
loss_mask = labels != -100
|
||||
|
||||
# dummy token; we'll ignore the losses on these tokens later
|
||||
labels[labels == -100] = 0
|
||||
|
||||
per_token_logps = torch.gather(
|
||||
logits.log_softmax(-1), dim=-1, index=labels.unsqueeze(-1)
|
||||
).squeeze(-1)
|
||||
|
||||
if average_log_prob:
|
||||
return (per_token_logps * loss_mask).sum(-1) / loss_mask.sum(-1)
|
||||
else:
|
||||
return (per_token_logps * loss_mask).sum(-1)
|
||||
|
||||
def _step(self, batch, batch_idx, stage: str):
|
||||
is_train = stage == "train"
|
||||
|
||||
if is_train:
|
||||
# Key part to make lora work
|
||||
# Otherwise the parameters are merged, which lead to incorrect gradients
|
||||
self.model.train()
|
||||
|
||||
# Do positive and negative samples in the same batch to speed up training
|
||||
labels = batch["labels"]
|
||||
outputs = self.model(
|
||||
inp=batch["inputs"],
|
||||
key_padding_mask=batch["attention_masks"],
|
||||
)
|
||||
token_logits = outputs.token_logits
|
||||
codebook_logits = outputs.codebook_logits
|
||||
|
||||
# Generate labels
|
||||
base_loss = F.cross_entropy(
|
||||
token_logits.view(-1, token_logits.size(-1)),
|
||||
labels[:, 0].reshape(-1),
|
||||
ignore_index=-100,
|
||||
)
|
||||
|
||||
codebook_labels = labels[:, 1 : 1 + self.model.config.num_codebooks].mT
|
||||
semantic_loss = F.cross_entropy(
|
||||
codebook_logits.view(-1, codebook_logits.size(-1)),
|
||||
codebook_labels.reshape(-1),
|
||||
ignore_index=-100,
|
||||
)
|
||||
|
||||
loss = base_loss + semantic_loss
|
||||
|
||||
self.log(
|
||||
f"{stage}/loss",
|
||||
loss,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=True,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
self.log(
|
||||
f"{stage}/base_loss",
|
||||
base_loss,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=False,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
self.log(
|
||||
f"{stage}/semantic_loss",
|
||||
semantic_loss,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=False,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
# Top-5 accuracy
|
||||
accuracy = self.get_accuracy(codebook_logits, codebook_labels)
|
||||
self.log(
|
||||
f"{stage}/top_5_accuracy",
|
||||
accuracy,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=True,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
return loss
|
||||
|
||||
def get_accuracy(self, logits, labels):
|
||||
mask = (labels != -100) & (labels != CODEBOOK_PAD_TOKEN_ID)
|
||||
if mask.sum() == 0:
|
||||
return torch.tensor(0.0, device=logits.device)
|
||||
|
||||
_, indices = logits.topk(5, dim=-1)
|
||||
correct = indices.eq(labels.unsqueeze(-1))
|
||||
correct[~mask] = 0
|
||||
correct = correct.sum()
|
||||
accuracy = correct / mask.sum()
|
||||
|
||||
return accuracy
|
||||
|
||||
def training_step(self, batch, batch_idx):
|
||||
return self._step(batch, batch_idx, "train")
|
||||
|
||||
def validation_step(self, batch, batch_idx):
|
||||
return self._step(batch, batch_idx, "val")
|
||||
@@ -0,0 +1,779 @@
|
||||
import json
|
||||
import math
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from loguru import logger
|
||||
from torch import Tensor
|
||||
from torch.nn import functional as F
|
||||
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fish_speech.conversation import SEMANTIC_TOKEN
|
||||
from fish_speech.utils import RankedLogger
|
||||
|
||||
from .lora import LoraConfig, setup_lora
|
||||
|
||||
log = RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
def find_multiple(n: int, k: int) -> int:
|
||||
if n % k == 0:
|
||||
return n
|
||||
return n + k - (n % k)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseModelArgs:
|
||||
model_type: str = "base"
|
||||
|
||||
vocab_size: int = 32000
|
||||
n_layer: int = 32
|
||||
n_head: int = 32
|
||||
dim: int = 4096
|
||||
intermediate_size: int = None
|
||||
n_local_heads: int = -1
|
||||
head_dim: int = 64
|
||||
rope_base: float = 10000
|
||||
norm_eps: float = 1e-5
|
||||
max_seq_len: int = 2048
|
||||
dropout: float = 0.0
|
||||
tie_word_embeddings: bool = True
|
||||
attention_qkv_bias: bool = False
|
||||
|
||||
# Codebook configs
|
||||
codebook_size: int = 160
|
||||
num_codebooks: int = 4
|
||||
|
||||
# Gradient checkpointing
|
||||
use_gradient_checkpointing: bool = True
|
||||
|
||||
# Initialize the model
|
||||
initializer_range: float = 0.02
|
||||
|
||||
def __post_init__(self):
|
||||
if self.n_local_heads == -1:
|
||||
self.n_local_heads = self.n_head
|
||||
if self.intermediate_size is None:
|
||||
hidden_dim = 4 * self.dim
|
||||
n_hidden = int(2 * hidden_dim / 3)
|
||||
self.intermediate_size = find_multiple(n_hidden, 256)
|
||||
self.head_dim = self.dim // self.n_head
|
||||
|
||||
@staticmethod
|
||||
def from_pretrained(path: str):
|
||||
path = Path(path)
|
||||
|
||||
if path.is_dir():
|
||||
path = path / "config.json"
|
||||
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
|
||||
match data["model_type"]:
|
||||
case "naive":
|
||||
cls = NaiveModelArgs
|
||||
case "dual_ar":
|
||||
cls = DualARModelArgs
|
||||
case _:
|
||||
raise ValueError(f"Unknown model type: {data['model_type']}")
|
||||
|
||||
return cls(**data)
|
||||
|
||||
def save(self, path: str):
|
||||
with open(path, "w") as f:
|
||||
json.dump(self.__dict__, f, indent=4, sort_keys=True, ensure_ascii=False)
|
||||
|
||||
|
||||
@dataclass
|
||||
class NaiveModelArgs(BaseModelArgs):
|
||||
model_type: str = "naive"
|
||||
|
||||
|
||||
@dataclass
|
||||
class DualARModelArgs(BaseModelArgs):
|
||||
model_type: str = "dual_ar"
|
||||
n_fast_layer: int = 4
|
||||
|
||||
|
||||
class KVCache(nn.Module):
|
||||
def __init__(
|
||||
self, max_batch_size, max_seq_len, n_heads, head_dim, dtype=torch.bfloat16
|
||||
):
|
||||
super().__init__()
|
||||
cache_shape = (max_batch_size, n_heads, max_seq_len, head_dim)
|
||||
self.register_buffer("k_cache", torch.zeros(cache_shape, dtype=dtype))
|
||||
self.register_buffer("v_cache", torch.zeros(cache_shape, dtype=dtype))
|
||||
|
||||
def update(self, input_pos, k_val, v_val):
|
||||
# input_pos: [S], k_val: [B, H, S, D]
|
||||
assert input_pos.shape[0] == k_val.shape[2]
|
||||
|
||||
k_out = self.k_cache
|
||||
v_out = self.v_cache
|
||||
k_out[:, :, input_pos] = k_val
|
||||
v_out[:, :, input_pos] = v_val
|
||||
|
||||
return k_out, v_out
|
||||
|
||||
|
||||
@dataclass
|
||||
class TransformerForwardResult:
|
||||
token_logits: Tensor
|
||||
codebook_logits: Tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseTransformerForwardResult:
|
||||
logits: Tensor
|
||||
hidden_states: Tensor
|
||||
|
||||
|
||||
class BaseTransformer(nn.Module):
|
||||
def __init__(
|
||||
self, config: BaseModelArgs, tokenizer: AutoTokenizer, init_weights: bool = True
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.semantic_token_id = tokenizer.convert_tokens_to_ids(SEMANTIC_TOKEN)
|
||||
|
||||
# Slow transformer
|
||||
self.embeddings = nn.Embedding(
|
||||
config.vocab_size,
|
||||
config.dim,
|
||||
)
|
||||
self.codebook_embeddings = nn.Embedding(
|
||||
config.codebook_size * config.num_codebooks,
|
||||
config.dim,
|
||||
)
|
||||
self.layers = nn.ModuleList(
|
||||
TransformerBlock(config, use_sdpa=True) for _ in range(config.n_layer)
|
||||
)
|
||||
self.norm = RMSNorm(config.dim, eps=config.norm_eps)
|
||||
|
||||
if self.config.tie_word_embeddings is False:
|
||||
self.output = nn.Linear(
|
||||
config.dim,
|
||||
config.vocab_size,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.register_buffer(
|
||||
"freqs_cis",
|
||||
precompute_freqs_cis(
|
||||
config.max_seq_len,
|
||||
config.dim // config.n_head,
|
||||
config.rope_base,
|
||||
),
|
||||
persistent=False,
|
||||
)
|
||||
self.register_buffer(
|
||||
"causal_mask",
|
||||
torch.tril(
|
||||
torch.ones(
|
||||
config.max_seq_len,
|
||||
config.max_seq_len,
|
||||
dtype=torch.bool,
|
||||
)
|
||||
),
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
# For kv cache
|
||||
self.max_batch_size = -1
|
||||
self.max_seq_len = -1
|
||||
|
||||
if init_weights:
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def setup_caches(
|
||||
self, max_batch_size: int, max_seq_len: int, dtype: torch.dtype = torch.bfloat16
|
||||
):
|
||||
if self.max_seq_len >= max_seq_len and self.max_batch_size >= max_batch_size:
|
||||
return
|
||||
|
||||
head_dim = self.config.dim // self.config.n_head
|
||||
max_seq_len = find_multiple(max_seq_len, 8)
|
||||
self.max_seq_len = max_seq_len
|
||||
self.max_batch_size = max_batch_size
|
||||
|
||||
for b in self.layers:
|
||||
b.attention.kv_cache = KVCache(
|
||||
max_batch_size,
|
||||
max_seq_len,
|
||||
self.config.n_local_heads,
|
||||
head_dim,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
def embed(self, x: Tensor) -> Tensor:
|
||||
vocab_embeds = [self.embeddings(x[:, 0])]
|
||||
for i in range(self.config.num_codebooks):
|
||||
emb = self.codebook_embeddings(x[:, i + 1] + i * self.config.codebook_size)
|
||||
emb[x[:, 0] != self.semantic_token_id] = 0
|
||||
vocab_embeds.append(emb)
|
||||
|
||||
x = torch.stack(vocab_embeds, dim=3)
|
||||
x = x.sum(dim=3)
|
||||
|
||||
return x
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inp: Tensor,
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
) -> BaseTransformerForwardResult:
|
||||
seq_len = inp.size(2)
|
||||
|
||||
# Here we want to merge the embeddings of the codebooks
|
||||
x = self.embed(inp)
|
||||
|
||||
freqs_cis = self.freqs_cis[:seq_len]
|
||||
|
||||
# Not that the causal mask here follows the definition of scaled_dot_product_attention
|
||||
# That is, FALSE means masked out
|
||||
# To maintain consistency, key_padding_mask use TRUE to mask out
|
||||
mask = None
|
||||
if key_padding_mask is not None:
|
||||
mask = self.causal_mask[None, None, :seq_len, :seq_len] # (B, N, Q, K)
|
||||
mask = mask & key_padding_mask[:, None, None, :].logical_not()
|
||||
|
||||
for layer in self.layers:
|
||||
if self.config.use_gradient_checkpointing and self.training:
|
||||
x = checkpoint(layer, x, freqs_cis, mask, use_reentrant=True)
|
||||
else:
|
||||
x = layer(x, freqs_cis, mask)
|
||||
|
||||
# We got slow_out here
|
||||
slow_out = self.norm(x)
|
||||
|
||||
if self.config.tie_word_embeddings:
|
||||
token_logits = F.linear(slow_out, self.embeddings.weight)
|
||||
else:
|
||||
token_logits = self.output(slow_out)
|
||||
|
||||
return BaseTransformerForwardResult(
|
||||
logits=token_logits,
|
||||
hidden_states=x,
|
||||
)
|
||||
|
||||
def forward_generate(
|
||||
self,
|
||||
x: Tensor,
|
||||
input_pos: Optional[Tensor] = None,
|
||||
return_all: bool = False,
|
||||
) -> BaseTransformerForwardResult:
|
||||
# This is used for generation, optimized for torch compile
|
||||
assert (
|
||||
self.max_seq_len != -1 and self.max_batch_size != -1
|
||||
), "Please call setup_caches before forward_generate"
|
||||
|
||||
x = self.embed(x)
|
||||
|
||||
mask = self.causal_mask[
|
||||
None, None, input_pos, : self.max_seq_len
|
||||
] # (B, N, Q, K)
|
||||
freqs_cis = self.freqs_cis[input_pos]
|
||||
|
||||
for layer in self.layers:
|
||||
x = layer(x, freqs_cis, mask, input_pos=input_pos)
|
||||
|
||||
# If prefill, we only calculate the logits of last token
|
||||
if x.size(1) > 1 and not return_all:
|
||||
x = x[:, -1:]
|
||||
|
||||
# We got slow_out here
|
||||
slow_out = self.norm(x)
|
||||
|
||||
if self.config.tie_word_embeddings:
|
||||
token_logits = F.linear(slow_out, self.embeddings.weight)
|
||||
else:
|
||||
token_logits = self.output(slow_out)
|
||||
|
||||
return BaseTransformerForwardResult(
|
||||
logits=token_logits,
|
||||
hidden_states=x,
|
||||
)
|
||||
|
||||
def _init_weights(self, module):
|
||||
std = self.config.initializer_range
|
||||
if isinstance(module, nn.Linear):
|
||||
module.weight.data.normal_(mean=0.0, std=std)
|
||||
if module.bias is not None:
|
||||
module.bias.data.zero_()
|
||||
elif isinstance(module, nn.Embedding):
|
||||
module.weight.data.normal_(mean=0.0, std=std)
|
||||
if module.padding_idx is not None:
|
||||
module.weight.data[module.padding_idx].zero_()
|
||||
|
||||
@staticmethod
|
||||
def from_pretrained(
|
||||
path: str,
|
||||
load_weights: bool = False,
|
||||
max_length: int | None = None,
|
||||
lora_config: LoraConfig | None = None,
|
||||
rope_base: int | None = None,
|
||||
) -> "BaseTransformer":
|
||||
config = BaseModelArgs.from_pretrained(str(path))
|
||||
if max_length is not None:
|
||||
config.max_seq_len = max_length
|
||||
log.info(f"Override max_seq_len to {max_length}")
|
||||
|
||||
if rope_base is not None:
|
||||
config.rope_base = rope_base
|
||||
log.info(f"Override rope_base to {rope_base}")
|
||||
|
||||
match config.model_type:
|
||||
case "naive":
|
||||
model_cls = NaiveTransformer
|
||||
case "dual_ar":
|
||||
model_cls = DualARTransformer
|
||||
case _:
|
||||
raise ValueError(f"Unknown model type: {config.model_type}")
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(str(path))
|
||||
log.info(f"Loading model from {path}, config: {config}")
|
||||
model = model_cls(config, tokenizer=tokenizer)
|
||||
|
||||
if lora_config is not None:
|
||||
setup_lora(model, lora_config)
|
||||
log.info(f"LoRA setup: {lora_config}")
|
||||
|
||||
if load_weights is False:
|
||||
log.info("Randomly initialized model")
|
||||
else:
|
||||
|
||||
if "int8" in str(Path(path)):
|
||||
logger.info("Using int8 weight-only quantization!")
|
||||
from tools.llama.quantize import WeightOnlyInt8QuantHandler
|
||||
|
||||
simple_quantizer = WeightOnlyInt8QuantHandler(model)
|
||||
model = simple_quantizer.convert_for_runtime()
|
||||
|
||||
if "int4" in str(Path(path)):
|
||||
logger.info("Using int4 quantization!")
|
||||
path_comps = path.name.split("-")
|
||||
assert path_comps[-2].startswith("g")
|
||||
groupsize = int(path_comps[-2][1:])
|
||||
from tools.llama.quantize import WeightOnlyInt4QuantHandler
|
||||
|
||||
simple_quantizer = WeightOnlyInt4QuantHandler(model, groupsize)
|
||||
model = simple_quantizer.convert_for_runtime()
|
||||
|
||||
weights = torch.load(
|
||||
Path(path) / "model.pth", map_location="cpu", mmap=True
|
||||
)
|
||||
|
||||
if "state_dict" in weights:
|
||||
logger.warning(
|
||||
"Using a TextToSemantic LightningModule checkpoint, "
|
||||
"please make sure it is a full model, not a LoRA model."
|
||||
)
|
||||
weights = weights["state_dict"]
|
||||
|
||||
if next(iter(weights.keys())).startswith("model."):
|
||||
logger.info(
|
||||
f"Remove prefix 'model.' created by TextToSemantic LightningModule from keys"
|
||||
)
|
||||
new_weights = OrderedDict()
|
||||
for k, v in weights.items():
|
||||
new_weights[k.replace("model.", "")] = v
|
||||
weights = new_weights
|
||||
|
||||
# Verify the name and shape of parameters since strict=False in load_state_dict.
|
||||
for k, v in model.named_parameters():
|
||||
if k not in weights:
|
||||
logger.warning(f"No weight for {k}")
|
||||
elif v.shape != weights[k].shape:
|
||||
logger.warning(
|
||||
f"Shape mismatch for {k}: {v.shape} vs {weights[k].shape}"
|
||||
)
|
||||
|
||||
err = model.load_state_dict(weights, strict=False, assign=True)
|
||||
log.info(f"Loaded weights with error: {err}")
|
||||
|
||||
return model
|
||||
|
||||
def save_pretrained(self, path: str, drop_lora: bool = False):
|
||||
path = Path(path)
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
self.config.save(path / "config.json")
|
||||
state_dict = self.state_dict()
|
||||
|
||||
if drop_lora:
|
||||
for key in list(state_dict.keys()):
|
||||
if "lora" not in key:
|
||||
continue
|
||||
|
||||
state_dict.pop(key)
|
||||
log.info(f"Drop LoRA parameter: {key}")
|
||||
|
||||
torch.save(state_dict, path / "model.pth")
|
||||
self.tokenizer.save_pretrained(path)
|
||||
|
||||
|
||||
class NaiveTransformer(BaseTransformer):
|
||||
def __init__(self, config: NaiveModelArgs, tokenizer: AutoTokenizer) -> None:
|
||||
super().__init__(config, init_weights=False, tokenizer=tokenizer)
|
||||
|
||||
self.codebook_norm = RMSNorm(config.dim, eps=config.norm_eps)
|
||||
self.codebook_output = nn.Linear(
|
||||
config.dim,
|
||||
config.codebook_size * config.num_codebooks,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def decode(self, result: BaseTransformerForwardResult) -> TransformerForwardResult:
|
||||
token_logits = result.logits
|
||||
x = result.hidden_states
|
||||
|
||||
# Codebook
|
||||
codebook_logits = self.codebook_output(self.codebook_norm(x))
|
||||
codebook_logits = rearrange(
|
||||
codebook_logits, "b n (c d) -> b n c d", c=self.config.num_codebooks
|
||||
)
|
||||
|
||||
return TransformerForwardResult(
|
||||
token_logits=token_logits,
|
||||
codebook_logits=codebook_logits,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inp: Tensor,
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
) -> TransformerForwardResult:
|
||||
result = super().forward(
|
||||
inp=inp,
|
||||
key_padding_mask=key_padding_mask,
|
||||
)
|
||||
return self.decode(result)
|
||||
|
||||
def forward_generate(
|
||||
self, x: Tensor, input_pos: Optional[Tensor] = None
|
||||
) -> TransformerForwardResult:
|
||||
result = super().forward_generate(x, input_pos)
|
||||
return self.decode(result)
|
||||
|
||||
|
||||
class DualARTransformer(BaseTransformer):
|
||||
def __init__(self, config: NaiveModelArgs, tokenizer: AutoTokenizer) -> None:
|
||||
super().__init__(config, init_weights=False, tokenizer=tokenizer)
|
||||
|
||||
# Fast transformer
|
||||
self.fast_embeddings = nn.Embedding(config.codebook_size, config.dim)
|
||||
|
||||
# The equivalent bs is so large that sdpa doesn't work
|
||||
self.fast_layers = nn.ModuleList(
|
||||
TransformerBlock(config, use_sdpa=False) for _ in range(config.n_fast_layer)
|
||||
)
|
||||
self.fast_norm = RMSNorm(config.dim, eps=config.norm_eps)
|
||||
self.fast_output = nn.Linear(
|
||||
config.dim,
|
||||
config.codebook_size,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def setup_caches(
|
||||
self, max_batch_size: int, max_seq_len: int, dtype: torch.dtype = torch.bfloat16
|
||||
):
|
||||
super().setup_caches(max_batch_size, max_seq_len, dtype)
|
||||
|
||||
head_dim = self.config.dim // self.config.n_head
|
||||
|
||||
# Fast transformer
|
||||
# The max seq len here is the number of codebooks
|
||||
for b in self.fast_layers:
|
||||
b.attention.kv_cache = KVCache(
|
||||
max_batch_size,
|
||||
self.config.num_codebooks,
|
||||
self.config.n_local_heads,
|
||||
head_dim,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inp: Tensor,
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
) -> TransformerForwardResult:
|
||||
parent_result = super().forward(inp, key_padding_mask)
|
||||
token_logits = parent_result.logits
|
||||
x = parent_result.hidden_states
|
||||
|
||||
# Fast transformer
|
||||
fast_seq_len = self.config.num_codebooks
|
||||
fast_mask = self.causal_mask[
|
||||
None, None, :fast_seq_len, :fast_seq_len
|
||||
] # (B, N, Q, K)
|
||||
fast_freqs_cis = self.freqs_cis[:fast_seq_len]
|
||||
|
||||
# Drop the last token and rotate left
|
||||
codebooks = inp[:, 1:-1, 1:]
|
||||
codebooks = F.pad(codebooks, (0, 1), value=0)
|
||||
codebook_embeddings = self.fast_embeddings(codebooks)
|
||||
x = torch.cat([x[:, None], codebook_embeddings], dim=1)
|
||||
b, s = x.size(0), x.size(2)
|
||||
x = rearrange(x, "b n s d -> (b s) n d") # flatten the batch and seq_len
|
||||
|
||||
# Remove padded part
|
||||
codebooks = rearrange(codebooks, "b n s -> (b s) n")
|
||||
codebook_mask = (codebooks == 0).all(dim=-1)
|
||||
|
||||
if torch.all(codebook_mask):
|
||||
# If all codebooks are padded, we keep first 8 to make sure the model runs
|
||||
codebook_mask[:8] = False
|
||||
|
||||
x_bs, x_len = x.size(0), x.size(1)
|
||||
x = x[~codebook_mask]
|
||||
|
||||
for layer in self.fast_layers:
|
||||
if self.config.use_gradient_checkpointing and self.training:
|
||||
x = checkpoint(layer, x, fast_freqs_cis, fast_mask, use_reentrant=True)
|
||||
else:
|
||||
x = layer(x, fast_freqs_cis, fast_mask)
|
||||
|
||||
# unflatten the batch and num_codebooks
|
||||
fast_out = self.fast_norm(x)
|
||||
codebook_logits = self.fast_output(fast_out)
|
||||
|
||||
# Re-pad the codebook_logits
|
||||
buffer = torch.zeros(
|
||||
x_bs,
|
||||
x_len,
|
||||
codebook_logits.size(-1),
|
||||
device=codebook_logits.device,
|
||||
dtype=codebook_logits.dtype,
|
||||
)
|
||||
buffer[~codebook_mask] = codebook_logits
|
||||
codebook_logits = buffer
|
||||
|
||||
assert codebook_logits.shape[1] == self.config.num_codebooks
|
||||
codebook_logits = rearrange(
|
||||
codebook_logits,
|
||||
"(b s) n d -> b s n d",
|
||||
b=b,
|
||||
s=s,
|
||||
n=self.config.num_codebooks,
|
||||
)
|
||||
|
||||
return TransformerForwardResult(
|
||||
token_logits=token_logits,
|
||||
codebook_logits=codebook_logits,
|
||||
)
|
||||
|
||||
def forward_generate_fast(
|
||||
self, x: Tensor, input_pos: Optional[Tensor] = None
|
||||
) -> Tensor:
|
||||
# Fast transformer
|
||||
x = x.view(1, 1, -1)
|
||||
|
||||
fast_mask = self.causal_mask[
|
||||
None, None, input_pos, : self.config.num_codebooks
|
||||
] # (B, N, Q, K)
|
||||
fast_freqs_cis = self.freqs_cis[input_pos]
|
||||
|
||||
for layer in self.fast_layers:
|
||||
x = layer(x, fast_freqs_cis, fast_mask, input_pos=input_pos)
|
||||
|
||||
# unflatten the batch and num_codebooks
|
||||
fast_out = self.fast_norm(x) # only take the last token
|
||||
codebook_logits = self.fast_output(fast_out)
|
||||
|
||||
return codebook_logits
|
||||
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
def __init__(self, config: BaseModelArgs, use_sdpa: bool = True) -> None:
|
||||
super().__init__()
|
||||
self.attention = Attention(config, use_sdpa=use_sdpa)
|
||||
self.feed_forward = FeedForward(config)
|
||||
self.ffn_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.attention_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
|
||||
def forward(
|
||||
self, x: Tensor, freqs_cis: Tensor, mask: Tensor, input_pos: Tensor = None
|
||||
) -> Tensor:
|
||||
h = x + self.attention(self.attention_norm(x), freqs_cis, mask, input_pos)
|
||||
out = h + self.feed_forward(self.ffn_norm(h))
|
||||
return out
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(self, config: BaseModelArgs, use_sdpa: bool = True):
|
||||
super().__init__()
|
||||
assert config.dim % config.n_head == 0
|
||||
|
||||
total_head_dim = (config.n_head + 2 * config.n_local_heads) * config.head_dim
|
||||
# key, query, value projections for all heads, but in a batch
|
||||
self.wqkv = nn.Linear(
|
||||
config.dim, total_head_dim, bias=config.attention_qkv_bias
|
||||
)
|
||||
self.wo = nn.Linear(config.dim, config.dim, bias=False)
|
||||
self.kv_cache = None
|
||||
|
||||
self.dropout = config.dropout
|
||||
self.n_head = config.n_head
|
||||
self.head_dim = config.head_dim
|
||||
self.n_local_heads = config.n_local_heads
|
||||
self.dim = config.dim
|
||||
self.use_sdpa = use_sdpa
|
||||
self._register_load_state_dict_pre_hook(self.load_hook)
|
||||
|
||||
def load_hook(self, state_dict, prefix, *args):
|
||||
if prefix + "wq.weight" in state_dict:
|
||||
wq = state_dict.pop(prefix + "wq.weight")
|
||||
wk = state_dict.pop(prefix + "wk.weight")
|
||||
wv = state_dict.pop(prefix + "wv.weight")
|
||||
state_dict[prefix + "wqkv.weight"] = torch.cat([wq, wk, wv])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
freqs_cis: Tensor,
|
||||
mask: Tensor,
|
||||
input_pos: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
bsz, seqlen, _ = x.shape
|
||||
|
||||
kv_size = self.n_local_heads * self.head_dim
|
||||
q, k, v = self.wqkv(x).split([self.dim, kv_size, kv_size], dim=-1)
|
||||
|
||||
q = q.view(bsz, seqlen, self.n_head, self.head_dim)
|
||||
k = k.view(bsz, seqlen, self.n_local_heads, self.head_dim)
|
||||
v = v.view(bsz, seqlen, self.n_local_heads, self.head_dim)
|
||||
|
||||
q = apply_rotary_emb(q, freqs_cis)
|
||||
k = apply_rotary_emb(k, freqs_cis)
|
||||
|
||||
q, k, v = map(lambda x: x.transpose(1, 2), (q, k, v))
|
||||
|
||||
if self.kv_cache is not None:
|
||||
k, v = self.kv_cache.update(input_pos, k, v)
|
||||
|
||||
k = k.repeat_interleave(self.n_head // self.n_local_heads, dim=1)
|
||||
v = v.repeat_interleave(self.n_head // self.n_local_heads, dim=1)
|
||||
|
||||
if self.use_sdpa:
|
||||
if mask is None:
|
||||
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
|
||||
y = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=self.dropout if self.training else 0.0,
|
||||
is_causal=True,
|
||||
# No third party attn_mask here to use flash_attention
|
||||
)
|
||||
else:
|
||||
y = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=mask,
|
||||
dropout_p=self.dropout if self.training else 0.0,
|
||||
)
|
||||
else:
|
||||
y = self.eq_scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=mask,
|
||||
dropout_p=self.dropout if self.training else 0.0,
|
||||
)
|
||||
|
||||
y = y.transpose(1, 2).contiguous().view(bsz, seqlen, self.dim)
|
||||
|
||||
return self.wo(y)
|
||||
|
||||
def eq_scaled_dot_product_attention(
|
||||
self,
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=None,
|
||||
dropout_p=0.0,
|
||||
) -> torch.Tensor:
|
||||
# This is a standard scaled dot product attention
|
||||
# It's low efficient, but it doesn't raise cuda error
|
||||
|
||||
L, S = query.size(-2), key.size(-2)
|
||||
scale_factor = 1 / math.sqrt(query.size(-1))
|
||||
attn_bias = torch.zeros(1, 1, L, S, dtype=query.dtype, device=query.device)
|
||||
|
||||
if attn_mask is not None:
|
||||
if attn_mask.dtype == torch.bool:
|
||||
attn_bias.masked_fill_(attn_mask.logical_not(), float("-inf"))
|
||||
else:
|
||||
attn_bias += attn_mask
|
||||
|
||||
attn_weight = query @ key.transpose(-2, -1) * scale_factor
|
||||
attn_weight += attn_bias
|
||||
attn_weight = torch.softmax(attn_weight, dim=-1)
|
||||
attn_weight = torch.dropout(attn_weight, dropout_p, train=True)
|
||||
|
||||
return attn_weight @ value
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, config: BaseModelArgs) -> None:
|
||||
super().__init__()
|
||||
self.w1 = nn.Linear(config.dim, config.intermediate_size, bias=False)
|
||||
self.w3 = nn.Linear(config.dim, config.intermediate_size, bias=False)
|
||||
self.w2 = nn.Linear(config.intermediate_size, config.dim, bias=False)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim: int, eps: float = 1e-5):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(torch.mean(x * x, dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
output = self._norm(x.float()).type_as(x)
|
||||
return output * self.weight
|
||||
|
||||
|
||||
def precompute_freqs_cis(seq_len: int, n_elem: int, base: int = 10000) -> Tensor:
|
||||
freqs = 1.0 / (
|
||||
base ** (torch.arange(0, n_elem, 2)[: (n_elem // 2)].float() / n_elem)
|
||||
)
|
||||
t = torch.arange(seq_len, device=freqs.device)
|
||||
freqs = torch.outer(t, freqs)
|
||||
freqs_cis = torch.polar(torch.ones_like(freqs), freqs)
|
||||
cache = torch.stack([freqs_cis.real, freqs_cis.imag], dim=-1)
|
||||
return cache.to(dtype=torch.bfloat16)
|
||||
|
||||
|
||||
def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
|
||||
xshaped = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||
freqs_cis = freqs_cis.view(1, xshaped.size(1), 1, xshaped.size(3), 2)
|
||||
x_out2 = torch.stack(
|
||||
[
|
||||
xshaped[..., 0] * freqs_cis[..., 0] - xshaped[..., 1] * freqs_cis[..., 1],
|
||||
xshaped[..., 1] * freqs_cis[..., 0] + xshaped[..., 0] * freqs_cis[..., 1],
|
||||
],
|
||||
-1,
|
||||
)
|
||||
|
||||
x_out2 = x_out2.flatten(3)
|
||||
return x_out2.type_as(x)
|
||||
@@ -0,0 +1,92 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
import loralib as lora
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoraConfig:
|
||||
r: int
|
||||
lora_alpha: float
|
||||
lora_dropout: float = 0.0
|
||||
|
||||
|
||||
def setup_lora(model, lora_config):
|
||||
# Replace the embedding layer with a LoRA layer
|
||||
model.embeddings = lora.Embedding(
|
||||
num_embeddings=model.embeddings.num_embeddings,
|
||||
embedding_dim=model.embeddings.embedding_dim,
|
||||
padding_idx=model.embeddings.padding_idx,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
)
|
||||
|
||||
model.codebook_embeddings = lora.Embedding(
|
||||
num_embeddings=model.codebook_embeddings.num_embeddings,
|
||||
embedding_dim=model.codebook_embeddings.embedding_dim,
|
||||
padding_idx=model.codebook_embeddings.padding_idx,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
)
|
||||
|
||||
# Replace output layer with a LoRA layer
|
||||
linears = [(model, "output")]
|
||||
|
||||
# Replace all linear layers with LoRA layers
|
||||
for layer in model.layers:
|
||||
linears.extend([(layer.attention, "wqkv"), (layer.attention, "wo")])
|
||||
linears.extend(
|
||||
[
|
||||
(layer.feed_forward, "w1"),
|
||||
(layer.feed_forward, "w2"),
|
||||
(layer.feed_forward, "w3"),
|
||||
]
|
||||
)
|
||||
|
||||
if hasattr(model, "fast_layers"):
|
||||
model.fast_embeddings = lora.Embedding(
|
||||
num_embeddings=model.fast_embeddings.num_embeddings,
|
||||
embedding_dim=model.fast_embeddings.embedding_dim,
|
||||
padding_idx=model.fast_embeddings.padding_idx,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
)
|
||||
|
||||
# Dual-AR model
|
||||
linears.append((model, "fast_output"))
|
||||
|
||||
for layer in model.fast_layers:
|
||||
linears.extend([(layer.attention, "wqkv"), (layer.attention, "wo")])
|
||||
linears.extend(
|
||||
[
|
||||
(layer.feed_forward, "w1"),
|
||||
(layer.feed_forward, "w2"),
|
||||
(layer.feed_forward, "w3"),
|
||||
]
|
||||
)
|
||||
|
||||
for module, layer in linears:
|
||||
updated_linear = lora.Linear(
|
||||
in_features=getattr(module, layer).in_features,
|
||||
out_features=getattr(module, layer).out_features,
|
||||
bias=getattr(module, layer).bias,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
lora_dropout=lora_config.lora_dropout,
|
||||
)
|
||||
setattr(module, layer, updated_linear)
|
||||
|
||||
# Mark only the LoRA layers as trainable
|
||||
lora.mark_only_lora_as_trainable(model, bias="none")
|
||||
|
||||
|
||||
def get_merged_state_dict(model):
|
||||
# This line will merge the state dict of the model and the LoRA parameters
|
||||
model.eval()
|
||||
|
||||
# Then we need to remove the LoRA parameters from the state dict
|
||||
state_dict = model.state_dict()
|
||||
for name in list(state_dict.keys()):
|
||||
if "lora" in name:
|
||||
state_dict.pop(name)
|
||||
|
||||
return state_dict
|
||||
@@ -0,0 +1,596 @@
|
||||
import math
|
||||
from functools import partial
|
||||
from math import prod
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from torch.nn.utils.parametrizations import weight_norm
|
||||
from torch.nn.utils.parametrize import remove_parametrizations
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
|
||||
def sequence_mask(length, max_length=None):
|
||||
if max_length is None:
|
||||
max_length = length.max()
|
||||
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
|
||||
return x.unsqueeze(0) < length.unsqueeze(1)
|
||||
|
||||
|
||||
def init_weights(m, mean=0.0, std=0.01):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv1D") != -1:
|
||||
m.weight.data.normal_(mean, std)
|
||||
|
||||
|
||||
def get_padding(kernel_size, dilation=1):
|
||||
return (kernel_size * dilation - dilation) // 2
|
||||
|
||||
|
||||
def unpad1d(x: torch.Tensor, paddings: tuple[int, int]):
|
||||
"""Remove padding from x, handling properly zero padding. Only for 1d!"""
|
||||
padding_left, padding_right = paddings
|
||||
assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
|
||||
assert (padding_left + padding_right) <= x.shape[-1]
|
||||
end = x.shape[-1] - padding_right
|
||||
return x[..., padding_left:end]
|
||||
|
||||
|
||||
def get_extra_padding_for_conv1d(
|
||||
x: torch.Tensor, kernel_size: int, stride: int, padding_total: int = 0
|
||||
) -> int:
|
||||
"""See `pad_for_conv1d`."""
|
||||
length = x.shape[-1]
|
||||
n_frames = (length - kernel_size + padding_total) / stride + 1
|
||||
ideal_length = (math.ceil(n_frames) - 1) * stride + (kernel_size - padding_total)
|
||||
return ideal_length - length
|
||||
|
||||
|
||||
def pad1d(
|
||||
x: torch.Tensor,
|
||||
paddings: tuple[int, int],
|
||||
mode: str = "zeros",
|
||||
value: float = 0.0,
|
||||
):
|
||||
"""Tiny wrapper around F.pad, just to allow for reflect padding on small input.
|
||||
If this is the case, we insert extra 0 padding to the right
|
||||
before the reflection happen.
|
||||
"""
|
||||
length = x.shape[-1]
|
||||
padding_left, padding_right = paddings
|
||||
assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
|
||||
if mode == "reflect":
|
||||
max_pad = max(padding_left, padding_right)
|
||||
extra_pad = 0
|
||||
if length <= max_pad:
|
||||
extra_pad = max_pad - length + 1
|
||||
x = F.pad(x, (0, extra_pad))
|
||||
padded = F.pad(x, paddings, mode, value)
|
||||
end = padded.shape[-1] - extra_pad
|
||||
return padded[..., :end]
|
||||
else:
|
||||
return F.pad(x, paddings, mode, value)
|
||||
|
||||
|
||||
class FishConvNet(nn.Module):
|
||||
def __init__(
|
||||
self, in_channels, out_channels, kernel_size, dilation=1, stride=1, groups=1
|
||||
):
|
||||
super(FishConvNet, self).__init__()
|
||||
self.conv = nn.Conv1d(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=stride,
|
||||
dilation=dilation,
|
||||
groups=groups,
|
||||
)
|
||||
self.stride = stride
|
||||
self.kernel_size = (kernel_size - 1) * dilation + 1
|
||||
self.dilation = dilation
|
||||
|
||||
def forward(self, x):
|
||||
pad = self.kernel_size - self.stride
|
||||
extra_padding = get_extra_padding_for_conv1d(
|
||||
x, self.kernel_size, self.stride, pad
|
||||
)
|
||||
x = pad1d(x, (pad, extra_padding), mode="constant", value=0)
|
||||
return self.conv(x).contiguous()
|
||||
|
||||
def weight_norm(self, name="weight", dim=0):
|
||||
self.conv = weight_norm(self.conv, name=name, dim=dim)
|
||||
return self
|
||||
|
||||
def remove_weight_norm(self):
|
||||
self.conv = remove_parametrizations(self.conv)
|
||||
return self
|
||||
|
||||
|
||||
class FishTransConvNet(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel_size, dilation=1, stride=1):
|
||||
super(FishTransConvNet, self).__init__()
|
||||
self.conv = nn.ConvTranspose1d(
|
||||
in_channels, out_channels, kernel_size, stride=stride, dilation=dilation
|
||||
)
|
||||
self.stride = stride
|
||||
self.kernel_size = kernel_size
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
pad = self.kernel_size - self.stride
|
||||
padding_right = math.ceil(pad)
|
||||
padding_left = pad - padding_right
|
||||
x = unpad1d(x, (padding_left, padding_right))
|
||||
return x.contiguous()
|
||||
|
||||
def weight_norm(self, name="weight", dim=0):
|
||||
self.conv = weight_norm(self.conv, name=name, dim=dim)
|
||||
return self
|
||||
|
||||
def remove_weight_norm(self):
|
||||
self.conv = remove_parametrizations(self.conv)
|
||||
return self
|
||||
|
||||
|
||||
class ResBlock1(torch.nn.Module):
|
||||
def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)):
|
||||
super().__init__()
|
||||
|
||||
self.convs1 = nn.ModuleList(
|
||||
[
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[0]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[1]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[2]
|
||||
).weight_norm(),
|
||||
]
|
||||
)
|
||||
self.convs1.apply(init_weights)
|
||||
|
||||
self.convs2 = nn.ModuleList(
|
||||
[
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[0]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[1]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[2]
|
||||
).weight_norm(),
|
||||
]
|
||||
)
|
||||
self.convs2.apply(init_weights)
|
||||
|
||||
def forward(self, x):
|
||||
for c1, c2 in zip(self.convs1, self.convs2):
|
||||
xt = F.silu(x)
|
||||
xt = c1(xt)
|
||||
xt = F.silu(xt)
|
||||
xt = c2(xt)
|
||||
x = xt + x
|
||||
return x
|
||||
|
||||
def remove_parametrizations(self):
|
||||
for conv in self.convs1:
|
||||
remove_parametrizations(conv, tensor_name="weight")
|
||||
for conv in self.convs2:
|
||||
remove_parametrizations(conv, tensor_name="weight")
|
||||
|
||||
|
||||
class ParallelBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
channels: int,
|
||||
kernel_sizes: tuple[int] = (3, 7, 11),
|
||||
dilation_sizes: tuple[tuple[int]] = ((1, 3, 5), (1, 3, 5), (1, 3, 5)),
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
assert len(kernel_sizes) == len(dilation_sizes)
|
||||
|
||||
self.blocks = nn.ModuleList()
|
||||
for k, d in zip(kernel_sizes, dilation_sizes):
|
||||
self.blocks.append(ResBlock1(channels, k, d))
|
||||
|
||||
def forward(self, x):
|
||||
return torch.stack([block(x) for block in self.blocks], dim=0).mean(dim=0)
|
||||
|
||||
def remove_parametrizations(self):
|
||||
for block in self.blocks:
|
||||
block.remove_parametrizations()
|
||||
|
||||
|
||||
class HiFiGANGenerator(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
hop_length: int = 512,
|
||||
upsample_rates: tuple[int] = (8, 8, 2, 2, 2),
|
||||
upsample_kernel_sizes: tuple[int] = (16, 16, 8, 2, 2),
|
||||
resblock_kernel_sizes: tuple[int] = (3, 7, 11),
|
||||
resblock_dilation_sizes: tuple[tuple[int]] = ((1, 3, 5), (1, 3, 5), (1, 3, 5)),
|
||||
num_mels: int = 128,
|
||||
upsample_initial_channel: int = 512,
|
||||
pre_conv_kernel_size: int = 7,
|
||||
post_conv_kernel_size: int = 7,
|
||||
post_activation: Callable = partial(nn.SiLU, inplace=True),
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
assert (
|
||||
prod(upsample_rates) == hop_length
|
||||
), f"hop_length must be {prod(upsample_rates)}"
|
||||
|
||||
self.conv_pre = FishConvNet(
|
||||
num_mels,
|
||||
upsample_initial_channel,
|
||||
pre_conv_kernel_size,
|
||||
stride=1,
|
||||
).weight_norm()
|
||||
|
||||
self.num_upsamples = len(upsample_rates)
|
||||
self.num_kernels = len(resblock_kernel_sizes)
|
||||
|
||||
self.noise_convs = nn.ModuleList()
|
||||
self.ups = nn.ModuleList()
|
||||
|
||||
for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
|
||||
self.ups.append(
|
||||
FishTransConvNet(
|
||||
upsample_initial_channel // (2**i),
|
||||
upsample_initial_channel // (2 ** (i + 1)),
|
||||
k,
|
||||
stride=u,
|
||||
).weight_norm()
|
||||
)
|
||||
|
||||
self.resblocks = nn.ModuleList()
|
||||
for i in range(len(self.ups)):
|
||||
ch = upsample_initial_channel // (2 ** (i + 1))
|
||||
self.resblocks.append(
|
||||
ParallelBlock(ch, resblock_kernel_sizes, resblock_dilation_sizes)
|
||||
)
|
||||
|
||||
self.activation_post = post_activation()
|
||||
self.conv_post = FishConvNet(
|
||||
ch, 1, post_conv_kernel_size, stride=1
|
||||
).weight_norm()
|
||||
self.ups.apply(init_weights)
|
||||
self.conv_post.apply(init_weights)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_pre(x)
|
||||
|
||||
for i in range(self.num_upsamples):
|
||||
x = F.silu(x, inplace=True)
|
||||
x = self.ups[i](x)
|
||||
|
||||
if self.training and self.checkpointing:
|
||||
x = checkpoint(
|
||||
self.resblocks[i],
|
||||
x,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
x = self.resblocks[i](x)
|
||||
|
||||
x = self.activation_post(x)
|
||||
x = self.conv_post(x)
|
||||
x = torch.tanh(x)
|
||||
|
||||
return x
|
||||
|
||||
def remove_parametrizations(self):
|
||||
for up in self.ups:
|
||||
remove_parametrizations(up, tensor_name="weight")
|
||||
for block in self.resblocks:
|
||||
block.remove_parametrizations()
|
||||
remove_parametrizations(self.conv_pre, tensor_name="weight")
|
||||
remove_parametrizations(self.conv_post, tensor_name="weight")
|
||||
|
||||
|
||||
# DropPath copied from timm library
|
||||
def drop_path(
|
||||
x, drop_prob: float = 0.0, training: bool = False, scale_by_keep: bool = True
|
||||
):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||
|
||||
This is the same as the DropConnect impl I created for EfficientNet, etc networks, however,
|
||||
the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
|
||||
See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for
|
||||
changing the layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use
|
||||
'survival rate' as the argument.
|
||||
|
||||
""" # noqa: E501
|
||||
|
||||
if drop_prob == 0.0 or not training:
|
||||
return x
|
||||
keep_prob = 1 - drop_prob
|
||||
shape = (x.shape[0],) + (1,) * (
|
||||
x.ndim - 1
|
||||
) # work with diff dim tensors, not just 2D ConvNets
|
||||
random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
|
||||
if keep_prob > 0.0 and scale_by_keep:
|
||||
random_tensor.div_(keep_prob)
|
||||
return x * random_tensor
|
||||
|
||||
|
||||
class DropPath(nn.Module):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).""" # noqa: E501
|
||||
|
||||
def __init__(self, drop_prob: float = 0.0, scale_by_keep: bool = True):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
self.scale_by_keep = scale_by_keep
|
||||
|
||||
def forward(self, x):
|
||||
return drop_path(x, self.drop_prob, self.training, self.scale_by_keep)
|
||||
|
||||
def extra_repr(self):
|
||||
return f"drop_prob={round(self.drop_prob,3):0.3f}"
|
||||
|
||||
|
||||
class LayerNorm(nn.Module):
|
||||
r"""LayerNorm that supports two data formats: channels_last (default) or channels_first.
|
||||
The ordering of the dimensions in the inputs. channels_last corresponds to inputs with
|
||||
shape (batch_size, height, width, channels) while channels_first corresponds to inputs
|
||||
with shape (batch_size, channels, height, width).
|
||||
""" # noqa: E501
|
||||
|
||||
def __init__(self, normalized_shape, eps=1e-6, data_format="channels_last"):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(normalized_shape))
|
||||
self.bias = nn.Parameter(torch.zeros(normalized_shape))
|
||||
self.eps = eps
|
||||
self.data_format = data_format
|
||||
if self.data_format not in ["channels_last", "channels_first"]:
|
||||
raise NotImplementedError
|
||||
self.normalized_shape = (normalized_shape,)
|
||||
|
||||
def forward(self, x):
|
||||
if self.data_format == "channels_last":
|
||||
return F.layer_norm(
|
||||
x, self.normalized_shape, self.weight, self.bias, self.eps
|
||||
)
|
||||
elif self.data_format == "channels_first":
|
||||
u = x.mean(1, keepdim=True)
|
||||
s = (x - u).pow(2).mean(1, keepdim=True)
|
||||
x = (x - u) / torch.sqrt(s + self.eps)
|
||||
x = self.weight[:, None] * x + self.bias[:, None]
|
||||
return x
|
||||
|
||||
|
||||
# ConvNeXt Block copied from https://github.com/fishaudio/fish-diffusion/blob/main/fish_diffusion/modules/convnext.py
|
||||
class ConvNeXtBlock(nn.Module):
|
||||
r"""ConvNeXt Block. There are two equivalent implementations:
|
||||
(1) DwConv -> LayerNorm (channels_first) -> 1x1 Conv -> GELU -> 1x1 Conv; all in (N, C, H, W)
|
||||
(2) DwConv -> Permute to (N, H, W, C); LayerNorm (channels_last) -> Linear -> GELU -> Linear; Permute back
|
||||
We use (2) as we find it slightly faster in PyTorch
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
drop_path (float): Stochastic depth rate. Default: 0.0
|
||||
layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.0.
|
||||
kernel_size (int): Kernel size for depthwise conv. Default: 7.
|
||||
dilation (int): Dilation for depthwise conv. Default: 1.
|
||||
""" # noqa: E501
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
drop_path: float = 0.0,
|
||||
layer_scale_init_value: float = 1e-6,
|
||||
mlp_ratio: float = 4.0,
|
||||
kernel_size: int = 7,
|
||||
dilation: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.dwconv = FishConvNet(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size=kernel_size,
|
||||
# padding=int(dilation * (kernel_size - 1) / 2),
|
||||
groups=dim,
|
||||
) # depthwise conv
|
||||
self.norm = LayerNorm(dim, eps=1e-6)
|
||||
self.pwconv1 = nn.Linear(
|
||||
dim, int(mlp_ratio * dim)
|
||||
) # pointwise/1x1 convs, implemented with linear layers
|
||||
self.act = nn.GELU()
|
||||
self.pwconv2 = nn.Linear(int(mlp_ratio * dim), dim)
|
||||
self.gamma = (
|
||||
nn.Parameter(layer_scale_init_value * torch.ones((dim)), requires_grad=True)
|
||||
if layer_scale_init_value > 0
|
||||
else None
|
||||
)
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
|
||||
|
||||
def forward(self, x, apply_residual: bool = True):
|
||||
input = x
|
||||
|
||||
x = self.dwconv(x)
|
||||
x = x.permute(0, 2, 1) # (N, C, L) -> (N, L, C)
|
||||
x = self.norm(x)
|
||||
x = self.pwconv1(x)
|
||||
x = self.act(x)
|
||||
x = self.pwconv2(x)
|
||||
|
||||
if self.gamma is not None:
|
||||
x = self.gamma * x
|
||||
|
||||
x = x.permute(0, 2, 1) # (N, L, C) -> (N, C, L)
|
||||
x = self.drop_path(x)
|
||||
|
||||
if apply_residual:
|
||||
x = input + x
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class ConvNeXtEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_channels: int = 3,
|
||||
depths: list[int] = [3, 3, 9, 3],
|
||||
dims: list[int] = [96, 192, 384, 768],
|
||||
drop_path_rate: float = 0.0,
|
||||
layer_scale_init_value: float = 1e-6,
|
||||
kernel_size: int = 7,
|
||||
):
|
||||
super().__init__()
|
||||
assert len(depths) == len(dims)
|
||||
|
||||
self.downsample_layers = nn.ModuleList()
|
||||
stem = nn.Sequential(
|
||||
FishConvNet(
|
||||
input_channels,
|
||||
dims[0],
|
||||
kernel_size=7,
|
||||
# padding=3,
|
||||
# padding_mode="replicate",
|
||||
# padding_mode="zeros",
|
||||
),
|
||||
LayerNorm(dims[0], eps=1e-6, data_format="channels_first"),
|
||||
)
|
||||
self.downsample_layers.append(stem)
|
||||
|
||||
for i in range(len(depths) - 1):
|
||||
mid_layer = nn.Sequential(
|
||||
LayerNorm(dims[i], eps=1e-6, data_format="channels_first"),
|
||||
nn.Conv1d(dims[i], dims[i + 1], kernel_size=1),
|
||||
)
|
||||
self.downsample_layers.append(mid_layer)
|
||||
|
||||
self.stages = nn.ModuleList()
|
||||
dp_rates = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]
|
||||
|
||||
cur = 0
|
||||
for i in range(len(depths)):
|
||||
stage = nn.Sequential(
|
||||
*[
|
||||
ConvNeXtBlock(
|
||||
dim=dims[i],
|
||||
drop_path=dp_rates[cur + j],
|
||||
layer_scale_init_value=layer_scale_init_value,
|
||||
kernel_size=kernel_size,
|
||||
)
|
||||
for j in range(depths[i])
|
||||
]
|
||||
)
|
||||
self.stages.append(stage)
|
||||
cur += depths[i]
|
||||
|
||||
self.norm = LayerNorm(dims[-1], eps=1e-6, data_format="channels_first")
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, (nn.Conv1d, nn.Linear)):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
for i in range(len(self.downsample_layers)):
|
||||
x = self.downsample_layers[i](x)
|
||||
x = self.stages[i](x)
|
||||
|
||||
return self.norm(x)
|
||||
|
||||
|
||||
class FireflyArchitecture(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
backbone: nn.Module,
|
||||
head: nn.Module,
|
||||
quantizer: nn.Module,
|
||||
spec_transform: nn.Module,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.backbone = backbone
|
||||
self.head = head
|
||||
self.quantizer = quantizer
|
||||
self.spec_transform = spec_transform
|
||||
self.downsample_factor = math.prod(self.quantizer.downsample_factor)
|
||||
|
||||
def forward(self, x: torch.Tensor, template=None, mask=None) -> torch.Tensor:
|
||||
if self.spec_transform is not None:
|
||||
x = self.spec_transform(x)
|
||||
|
||||
x = self.backbone(x)
|
||||
if mask is not None:
|
||||
x = x * mask
|
||||
|
||||
if self.quantizer is not None:
|
||||
vq_result = self.quantizer(x)
|
||||
x = vq_result.z
|
||||
|
||||
if mask is not None:
|
||||
x = x * mask
|
||||
|
||||
x = self.head(x, template=template)
|
||||
|
||||
if x.ndim == 2:
|
||||
x = x[:, None, :]
|
||||
|
||||
if self.vq is not None:
|
||||
return x, vq_result
|
||||
|
||||
return x
|
||||
|
||||
def encode(self, audios, audio_lengths):
|
||||
audios = audios.float()
|
||||
|
||||
mels = self.spec_transform(audios)
|
||||
mel_lengths = audio_lengths // self.spec_transform.hop_length
|
||||
mel_masks = sequence_mask(mel_lengths, mels.shape[2])
|
||||
mel_masks_float_conv = mel_masks[:, None, :].float()
|
||||
mels = mels * mel_masks_float_conv
|
||||
|
||||
# Encode
|
||||
encoded_features = self.backbone(mels) * mel_masks_float_conv
|
||||
feature_lengths = mel_lengths // self.downsample_factor
|
||||
|
||||
return self.quantizer.encode(encoded_features), feature_lengths
|
||||
|
||||
def decode(self, indices, feature_lengths) -> torch.Tensor:
|
||||
mel_masks = sequence_mask(
|
||||
feature_lengths * self.downsample_factor,
|
||||
indices.shape[2] * self.downsample_factor,
|
||||
)
|
||||
mel_masks_float_conv = mel_masks[:, None, :].float()
|
||||
audio_lengths = (
|
||||
feature_lengths * self.downsample_factor * self.spec_transform.hop_length
|
||||
)
|
||||
|
||||
audio_masks = sequence_mask(
|
||||
audio_lengths,
|
||||
indices.shape[2] * self.downsample_factor * self.spec_transform.hop_length,
|
||||
)
|
||||
audio_masks_float_conv = audio_masks[:, None, :].float()
|
||||
|
||||
z = self.quantizer.decode(indices) * mel_masks_float_conv
|
||||
x = self.head(z) * audio_masks_float_conv
|
||||
|
||||
return x, audio_lengths
|
||||
|
||||
def remove_parametrizations(self):
|
||||
if hasattr(self.backbone, "remove_parametrizations"):
|
||||
self.backbone.remove_parametrizations()
|
||||
|
||||
if hasattr(self.head, "remove_parametrizations"):
|
||||
self.head.remove_parametrizations()
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
@@ -0,0 +1,116 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from vector_quantize_pytorch import GroupedResidualFSQ
|
||||
|
||||
from .firefly import ConvNeXtBlock, FishConvNet, FishTransConvNet
|
||||
|
||||
|
||||
@dataclass
|
||||
class FSQResult:
|
||||
z: torch.Tensor
|
||||
codes: torch.Tensor
|
||||
latents: torch.Tensor
|
||||
|
||||
|
||||
class DownsampleFiniteScalarQuantize(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int = 512,
|
||||
n_codebooks: int = 9,
|
||||
n_groups: int = 1,
|
||||
levels: tuple[int] = (8, 5, 5, 5), # Approximate 2**10
|
||||
downsample_factor: tuple[int] = (2, 2),
|
||||
downsample_dims: tuple[int] | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if downsample_dims is None:
|
||||
downsample_dims = [input_dim for _ in range(len(downsample_factor))]
|
||||
|
||||
all_dims = (input_dim,) + tuple(downsample_dims)
|
||||
|
||||
self.residual_fsq = GroupedResidualFSQ(
|
||||
dim=all_dims[-1],
|
||||
levels=levels,
|
||||
num_quantizers=n_codebooks,
|
||||
groups=n_groups,
|
||||
)
|
||||
|
||||
self.downsample_factor = downsample_factor
|
||||
self.downsample_dims = downsample_dims
|
||||
|
||||
self.downsample = nn.Sequential(
|
||||
*[
|
||||
nn.Sequential(
|
||||
FishConvNet(
|
||||
all_dims[idx],
|
||||
all_dims[idx + 1],
|
||||
kernel_size=factor,
|
||||
stride=factor,
|
||||
),
|
||||
ConvNeXtBlock(dim=all_dims[idx + 1]),
|
||||
)
|
||||
for idx, factor in enumerate(downsample_factor)
|
||||
]
|
||||
)
|
||||
|
||||
self.upsample = nn.Sequential(
|
||||
*[
|
||||
nn.Sequential(
|
||||
FishTransConvNet(
|
||||
all_dims[idx + 1],
|
||||
all_dims[idx],
|
||||
kernel_size=factor,
|
||||
stride=factor,
|
||||
),
|
||||
ConvNeXtBlock(dim=all_dims[idx]),
|
||||
)
|
||||
for idx, factor in reversed(list(enumerate(downsample_factor)))
|
||||
]
|
||||
)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, (nn.Conv1d, nn.Linear)):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, z) -> FSQResult:
|
||||
original_shape = z.shape
|
||||
z = self.downsample(z)
|
||||
quantized, indices = self.residual_fsq(z.mT)
|
||||
result = FSQResult(
|
||||
z=quantized.mT,
|
||||
codes=indices.mT,
|
||||
latents=z,
|
||||
)
|
||||
result.z = self.upsample(result.z)
|
||||
|
||||
# Pad or crop z to match original shape
|
||||
diff = original_shape[-1] - result.z.shape[-1]
|
||||
left = diff // 2
|
||||
right = diff - left
|
||||
|
||||
if diff > 0:
|
||||
result.z = F.pad(result.z, (left, right))
|
||||
elif diff < 0:
|
||||
result.z = result.z[..., left:-right]
|
||||
|
||||
return result
|
||||
|
||||
def encode(self, z):
|
||||
z = self.downsample(z)
|
||||
_, indices = self.residual_fsq(z.mT)
|
||||
indices = rearrange(indices, "g b l r -> b (g r) l")
|
||||
return indices
|
||||
|
||||
def decode(self, indices: torch.Tensor):
|
||||
indices = rearrange(indices, "b (g r) l -> g b l r", g=self.residual_fsq.groups)
|
||||
z_q = self.residual_fsq.get_output_from_indices(indices)
|
||||
z_q = self.upsample(z_q.mT)
|
||||
return z_q
|
||||
@@ -0,0 +1,94 @@
|
||||
import matplotlib
|
||||
import torch
|
||||
from matplotlib import pyplot as plt
|
||||
|
||||
matplotlib.use("Agg")
|
||||
|
||||
|
||||
def convert_pad_shape(pad_shape):
|
||||
l = pad_shape[::-1]
|
||||
pad_shape = [item for sublist in l for item in sublist]
|
||||
return pad_shape
|
||||
|
||||
|
||||
def sequence_mask(length, max_length=None):
|
||||
if max_length is None:
|
||||
max_length = length.max()
|
||||
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
|
||||
return x.unsqueeze(0) < length.unsqueeze(1)
|
||||
|
||||
|
||||
def init_weights(m, mean=0.0, std=0.01):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv") != -1:
|
||||
m.weight.data.normal_(mean, std)
|
||||
|
||||
|
||||
def get_padding(kernel_size, dilation=1):
|
||||
return int((kernel_size * dilation - dilation) / 2)
|
||||
|
||||
|
||||
def plot_mel(data, titles=None):
|
||||
fig, axes = plt.subplots(len(data), 1, squeeze=False)
|
||||
|
||||
if titles is None:
|
||||
titles = [None for i in range(len(data))]
|
||||
|
||||
plt.tight_layout()
|
||||
|
||||
for i in range(len(data)):
|
||||
mel = data[i]
|
||||
|
||||
if isinstance(mel, torch.Tensor):
|
||||
mel = mel.float().detach().cpu().numpy()
|
||||
|
||||
axes[i][0].imshow(mel, origin="lower")
|
||||
axes[i][0].set_aspect(2.5, adjustable="box")
|
||||
axes[i][0].set_ylim(0, mel.shape[0])
|
||||
axes[i][0].set_title(titles[i], fontsize="medium")
|
||||
axes[i][0].tick_params(labelsize="x-small", left=False, labelleft=False)
|
||||
axes[i][0].set_anchor("W")
|
||||
|
||||
return fig
|
||||
|
||||
|
||||
def slice_segments(x, ids_str, segment_size=4):
|
||||
ret = torch.zeros_like(x[:, :, :segment_size])
|
||||
for i in range(x.size(0)):
|
||||
idx_str = ids_str[i]
|
||||
idx_end = idx_str + segment_size
|
||||
ret[i] = x[i, :, idx_str:idx_end]
|
||||
|
||||
return ret
|
||||
|
||||
|
||||
def rand_slice_segments(x, x_lengths=None, segment_size=4):
|
||||
b, d, t = x.size()
|
||||
if x_lengths is None:
|
||||
x_lengths = t
|
||||
ids_str_max = torch.clamp(x_lengths - segment_size + 1, min=0)
|
||||
ids_str = (torch.rand([b], device=x.device) * ids_str_max).to(dtype=torch.long)
|
||||
ret = slice_segments(x, ids_str, segment_size)
|
||||
return ret, ids_str
|
||||
|
||||
|
||||
@torch.jit.script
|
||||
def fused_add_tanh_sigmoid_multiply(in_act, n_channels):
|
||||
n_channels_int = n_channels[0]
|
||||
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
||||
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
|
||||
acts = t_act * s_act
|
||||
|
||||
return acts
|
||||
|
||||
|
||||
def avg_with_mask(x, mask):
|
||||
assert mask.dtype == torch.float, "Mask should be float"
|
||||
|
||||
if mask.ndim == 2:
|
||||
mask = mask.unsqueeze(1)
|
||||
|
||||
if mask.shape[1] == 1:
|
||||
mask = mask.expand_as(x)
|
||||
|
||||
return (x * mask).sum() / mask.sum()
|
||||
@@ -0,0 +1,130 @@
|
||||
import re
|
||||
import string
|
||||
|
||||
from .clean import clean_text
|
||||
|
||||
|
||||
def utf_8_len(text):
|
||||
return len(text.encode("utf-8"))
|
||||
|
||||
|
||||
def break_text(texts, length, splits: set):
|
||||
for text in texts:
|
||||
if utf_8_len(text) <= length:
|
||||
yield text
|
||||
continue
|
||||
|
||||
curr = ""
|
||||
for char in text:
|
||||
curr += char
|
||||
|
||||
if char in splits:
|
||||
yield curr
|
||||
curr = ""
|
||||
|
||||
if curr:
|
||||
yield curr
|
||||
|
||||
|
||||
def break_text_by_length(texts, length):
|
||||
for text in texts:
|
||||
if utf_8_len(text) <= length:
|
||||
yield text
|
||||
continue
|
||||
|
||||
curr = ""
|
||||
for char in text:
|
||||
curr += char
|
||||
|
||||
if utf_8_len(curr) >= length:
|
||||
yield curr
|
||||
curr = ""
|
||||
|
||||
if curr:
|
||||
yield curr
|
||||
|
||||
|
||||
def add_cleaned(curr, segments):
|
||||
curr = curr.strip()
|
||||
if curr and not all(c.isspace() or c in string.punctuation for c in curr):
|
||||
segments.append(curr)
|
||||
|
||||
|
||||
def protect_float(text):
|
||||
# Turns 3.14 into <3_f_14> to prevent splitting
|
||||
return re.sub(r"(\d+)\.(\d+)", r"<\1_f_\2>", text)
|
||||
|
||||
|
||||
def unprotect_float(text):
|
||||
# Turns <3_f_14> into 3.14
|
||||
return re.sub(r"<(\d+)_f_(\d+)>", r"\1.\2", text)
|
||||
|
||||
|
||||
def split_text(text, length):
|
||||
text = clean_text(text)
|
||||
|
||||
# Break the text into pieces with following rules:
|
||||
# 1. Split the text at ".", "!", "?" if text is NOT a float
|
||||
# 2. If the text is longer than length, split at ","
|
||||
# 3. If the text is still longer than length, split at " "
|
||||
# 4. If the text is still longer than length, split at any character to length
|
||||
|
||||
texts = [text]
|
||||
texts = map(protect_float, texts)
|
||||
texts = break_text(texts, length, {".", "!", "?"})
|
||||
texts = map(unprotect_float, texts)
|
||||
texts = break_text(texts, length, {","})
|
||||
texts = break_text(texts, length, {" "})
|
||||
texts = list(break_text_by_length(texts, length))
|
||||
|
||||
# Then, merge the texts into segments with length <= length
|
||||
segments = []
|
||||
curr = ""
|
||||
|
||||
for text in texts:
|
||||
if utf_8_len(curr) + utf_8_len(text) <= length:
|
||||
curr += text
|
||||
else:
|
||||
add_cleaned(curr, segments)
|
||||
curr = text
|
||||
|
||||
if curr:
|
||||
add_cleaned(curr, segments)
|
||||
|
||||
return segments
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Test the split_text function
|
||||
|
||||
text = "This is a test sentence. This is another test sentence. And a third one."
|
||||
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence.",
|
||||
"This is another test sentence. And a third one.",
|
||||
]
|
||||
assert split_text("a,aaaaaa3.14", 10) == ["a,", "aaaaaa3.14"]
|
||||
assert split_text(" ", 10) == []
|
||||
assert split_text("a", 10) == ["a"]
|
||||
|
||||
text = "This is a test sentence with only commas, and no dots, and no exclamation marks, and no question marks, and no newlines."
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence with only commas,",
|
||||
"and no dots, and no exclamation marks,",
|
||||
"and no question marks, and no newlines.",
|
||||
]
|
||||
|
||||
text = "This is a test sentence This is a test sentence This is a test sentence. This is a test sentence, This is a test sentence, This is a test sentence."
|
||||
# First half split at " ", second half split at ","
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence This is a test sentence",
|
||||
"This is a test sentence. This is a test sentence,",
|
||||
"This is a test sentence, This is a test sentence.",
|
||||
]
|
||||
|
||||
text = "这是一段很长的中文文本,而且没有句号,也没有感叹号,也没有问号,也没有换行符。"
|
||||
assert split_text(text, 50) == [
|
||||
"这是一段很长的中文文本,",
|
||||
"而且没有句号,也没有感叹号,",
|
||||
"也没有问号,也没有换行符.",
|
||||
]
|
||||
@@ -0,0 +1,4 @@
|
||||
from .clean import clean_text
|
||||
from .spliter import split_text
|
||||
|
||||
__all__ = ["clean_text", "split_text"]
|
||||
@@ -0,0 +1,114 @@
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
|
||||
# C extensions
|
||||
*.so
|
||||
|
||||
# Distribution / packaging
|
||||
.Python
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
MANIFEST
|
||||
|
||||
# PyInstaller
|
||||
# Usually these files are written by a python script from a template
|
||||
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||
*.manifest
|
||||
*.spec
|
||||
|
||||
# Installer logs
|
||||
pip-log.txt
|
||||
pip-delete-this-directory.txt
|
||||
|
||||
# Unit test / coverage reports
|
||||
htmlcov/
|
||||
.tox/
|
||||
.coverage
|
||||
.coverage.*
|
||||
.cache
|
||||
nosetests.xml
|
||||
coverage.xml
|
||||
*.cover
|
||||
.hypothesis/
|
||||
.pytest_cache/
|
||||
|
||||
# Translations
|
||||
*.mo
|
||||
*.pot
|
||||
|
||||
# Django stuff:
|
||||
*.log
|
||||
local_settings.py
|
||||
db.sqlite3
|
||||
|
||||
# Flask stuff:
|
||||
instance/
|
||||
.webassets-cache
|
||||
|
||||
# Scrapy stuff:
|
||||
.scrapy
|
||||
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
|
||||
# PyBuilder
|
||||
target/
|
||||
|
||||
# Jupyter Notebook
|
||||
.ipynb_checkpoints
|
||||
|
||||
# pyenv
|
||||
.python-version
|
||||
|
||||
# celery beat schedule file
|
||||
celerybeat-schedule
|
||||
|
||||
# SageMath parsed files
|
||||
*.sage.py
|
||||
|
||||
# Environments
|
||||
.env
|
||||
.venv
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
env.bak/
|
||||
venv.bak/
|
||||
|
||||
# Spyder project settings
|
||||
.spyderproject
|
||||
.spyproject
|
||||
|
||||
# Rope project settings
|
||||
.ropeproject
|
||||
|
||||
# mkdocs documentation
|
||||
/site
|
||||
|
||||
# mypy
|
||||
.mypy_cache/
|
||||
|
||||
# JetBrains PyCharm
|
||||
.idea
|
||||
|
||||
# Customize
|
||||
references
|
||||
url.txt
|
||||
|
||||
# Git
|
||||
.git
|
||||
@@ -0,0 +1,36 @@
|
||||
# This account is no longer in use, see [Atomicoo](https://github.com/atomicoo) for my latest works.
|
||||
|
||||
# Chn Text Norm
|
||||
|
||||
this is a repository for chinese text normalization (no longer maintained).
|
||||
|
||||
## Quick Start ##
|
||||
|
||||
### Git Clone Repo ###
|
||||
|
||||
git clone this repo to the root directory of your project which need to use it.
|
||||
|
||||
cd /path/to/proj
|
||||
git clone https://github.com/Joee1995/chn-text-norm.git
|
||||
|
||||
after that, your doc tree should be:
|
||||
```
|
||||
proj # root of your project
|
||||
|--- chn_text_norm # this chn-text-norm tool
|
||||
|--- text.py
|
||||
|--- ...
|
||||
|--- text_normalize.py # your text normalization code
|
||||
|--- ...
|
||||
```
|
||||
|
||||
### How to Use ? ###
|
||||
|
||||
# text_normalize.py
|
||||
from chn_text_norm.text import *
|
||||
|
||||
raw_text = 'your raw text'
|
||||
text = Text(raw_text=raw_text).normalize()
|
||||
|
||||
### How to add quantums ###
|
||||
|
||||
打开test.py,然后你就知道怎么做了。
|
||||
@@ -0,0 +1,172 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""基本类
|
||||
中文字符类
|
||||
中文数字/数位类
|
||||
中文数字类
|
||||
中文数位类
|
||||
中文数字系统类
|
||||
中文数学符号类
|
||||
*中文其他符号类
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-02"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_constant import NUMBERING_TYPES
|
||||
|
||||
|
||||
class ChineseChar(object):
|
||||
"""
|
||||
中文字符
|
||||
每个字符对应简体和繁体,
|
||||
e.g. 简体 = '负', 繁体 = '負'
|
||||
转换时可转换为简体或繁体
|
||||
"""
|
||||
|
||||
def __init__(self, simplified, traditional):
|
||||
self.simplified = simplified
|
||||
self.traditional = traditional
|
||||
self.__repr__ = self.__str__
|
||||
|
||||
def __str__(self):
|
||||
return self.simplified or self.traditional or None
|
||||
|
||||
def __repr__(self):
|
||||
return self.__str__()
|
||||
|
||||
|
||||
class ChineseNumberUnit(ChineseChar):
|
||||
"""
|
||||
中文数字/数位字符
|
||||
每个字符除繁简体外还有一个额外的大写字符
|
||||
e.g. '陆' 和 '陸'
|
||||
"""
|
||||
|
||||
def __init__(self, power, simplified, traditional, big_s, big_t):
|
||||
super(ChineseNumberUnit, self).__init__(simplified, traditional)
|
||||
self.power = power
|
||||
self.big_s = big_s
|
||||
self.big_t = big_t
|
||||
|
||||
def __str__(self):
|
||||
return "10^{}".format(self.power)
|
||||
|
||||
@classmethod
|
||||
def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
|
||||
|
||||
if small_unit:
|
||||
return ChineseNumberUnit(
|
||||
power=index + 1,
|
||||
simplified=value[0],
|
||||
traditional=value[1],
|
||||
big_s=value[1],
|
||||
big_t=value[1],
|
||||
)
|
||||
elif numbering_type == NUMBERING_TYPES[0]:
|
||||
return ChineseNumberUnit(
|
||||
power=index + 8,
|
||||
simplified=value[0],
|
||||
traditional=value[1],
|
||||
big_s=value[0],
|
||||
big_t=value[1],
|
||||
)
|
||||
elif numbering_type == NUMBERING_TYPES[1]:
|
||||
return ChineseNumberUnit(
|
||||
power=(index + 2) * 4,
|
||||
simplified=value[0],
|
||||
traditional=value[1],
|
||||
big_s=value[0],
|
||||
big_t=value[1],
|
||||
)
|
||||
elif numbering_type == NUMBERING_TYPES[2]:
|
||||
return ChineseNumberUnit(
|
||||
power=pow(2, index + 3),
|
||||
simplified=value[0],
|
||||
traditional=value[1],
|
||||
big_s=value[0],
|
||||
big_t=value[1],
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Counting type should be in {0} ({1} provided).".format(
|
||||
NUMBERING_TYPES, numbering_type
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class ChineseNumberDigit(ChineseChar):
|
||||
"""
|
||||
中文数字字符
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None
|
||||
):
|
||||
super(ChineseNumberDigit, self).__init__(simplified, traditional)
|
||||
self.value = value
|
||||
self.big_s = big_s
|
||||
self.big_t = big_t
|
||||
self.alt_s = alt_s
|
||||
self.alt_t = alt_t
|
||||
|
||||
def __str__(self):
|
||||
return str(self.value)
|
||||
|
||||
@classmethod
|
||||
def create(cls, i, v):
|
||||
return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
|
||||
|
||||
|
||||
class ChineseMath(ChineseChar):
|
||||
"""
|
||||
中文数位字符
|
||||
"""
|
||||
|
||||
def __init__(self, simplified, traditional, symbol, expression=None):
|
||||
super(ChineseMath, self).__init__(simplified, traditional)
|
||||
self.symbol = symbol
|
||||
self.expression = expression
|
||||
self.big_s = simplified
|
||||
self.big_t = traditional
|
||||
|
||||
|
||||
CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
|
||||
|
||||
|
||||
class NumberSystem(object):
|
||||
"""
|
||||
中文数字系统
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class MathSymbol(object):
|
||||
"""
|
||||
用于中文数字系统的数学符号 (繁/简体), e.g.
|
||||
positive = ['正', '正']
|
||||
negative = ['负', '負']
|
||||
point = ['点', '點']
|
||||
"""
|
||||
|
||||
def __init__(self, positive, negative, point):
|
||||
self.positive = positive
|
||||
self.negative = negative
|
||||
self.point = point
|
||||
|
||||
def __iter__(self):
|
||||
for v in self.__dict__.values():
|
||||
yield v
|
||||
|
||||
|
||||
# class OtherSymbol(object):
|
||||
# """
|
||||
# 其他符号
|
||||
# """
|
||||
#
|
||||
# def __init__(self, sil):
|
||||
# self.sil = sil
|
||||
#
|
||||
# def __iter__(self):
|
||||
# for v in self.__dict__.values():
|
||||
# yield v
|
||||
@@ -0,0 +1,30 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""基本常量
|
||||
中文数字/数位/符号字符常量
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-02"
|
||||
|
||||
CHINESE_DIGIS = "零一二三四五六七八九"
|
||||
BIG_CHINESE_DIGIS_SIMPLIFIED = "零壹贰叁肆伍陆柒捌玖"
|
||||
BIG_CHINESE_DIGIS_TRADITIONAL = "零壹貳參肆伍陸柒捌玖"
|
||||
SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = "十百千万"
|
||||
SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = "拾佰仟萬"
|
||||
LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = "亿兆京垓秭穰沟涧正载"
|
||||
LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = "億兆京垓秭穰溝澗正載"
|
||||
SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = "十百千万"
|
||||
SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = "拾佰仟萬"
|
||||
|
||||
ZERO_ALT = "〇"
|
||||
ONE_ALT = "幺"
|
||||
TWO_ALTS = ["两", "兩"]
|
||||
|
||||
POSITIVE = ["正", "正"]
|
||||
NEGATIVE = ["负", "負"]
|
||||
POINT = ["点", "點"]
|
||||
# PLUS = [u'加', u'加']
|
||||
# SIL = [u'杠', u'槓']
|
||||
|
||||
# 中文数字系统类型
|
||||
NUMBERING_TYPES = ["low", "mid", "high"]
|
||||
@@ -0,0 +1,342 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""基本方法
|
||||
创建中文数字系统 方法
|
||||
中文字符串 <=> 数字串 方法
|
||||
数字串 <=> 中文字符串 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-02"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_class import *
|
||||
from fish_speech.text.chn_text_norm.basic_constant import *
|
||||
|
||||
|
||||
def create_system(numbering_type=NUMBERING_TYPES[1]):
|
||||
"""
|
||||
根据数字系统类型返回创建相应的数字系统,默认为 mid
|
||||
NUMBERING_TYPES = ['low', 'mid', 'high']: 中文数字系统类型
|
||||
low: '兆' = '亿' * '十' = $10^{9}$, '京' = '兆' * '十', etc.
|
||||
mid: '兆' = '亿' * '万' = $10^{12}$, '京' = '兆' * '万', etc.
|
||||
high: '兆' = '亿' * '亿' = $10^{16}$, '京' = '兆' * '兆', etc.
|
||||
返回对应的数字系统
|
||||
"""
|
||||
|
||||
# chinese number units of '亿' and larger
|
||||
all_larger_units = zip(
|
||||
LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED,
|
||||
LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL,
|
||||
)
|
||||
larger_units = [
|
||||
CNU.create(i, v, numbering_type, False) for i, v in enumerate(all_larger_units)
|
||||
]
|
||||
# chinese number units of '十, 百, 千, 万'
|
||||
all_smaller_units = zip(
|
||||
SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED,
|
||||
SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL,
|
||||
)
|
||||
smaller_units = [
|
||||
CNU.create(i, v, small_unit=True) for i, v in enumerate(all_smaller_units)
|
||||
]
|
||||
# digis
|
||||
chinese_digis = zip(
|
||||
CHINESE_DIGIS,
|
||||
CHINESE_DIGIS,
|
||||
BIG_CHINESE_DIGIS_SIMPLIFIED,
|
||||
BIG_CHINESE_DIGIS_TRADITIONAL,
|
||||
)
|
||||
digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
|
||||
digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
|
||||
digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
|
||||
digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
|
||||
|
||||
# symbols
|
||||
positive_cn = CM(POSITIVE[0], POSITIVE[1], "+", lambda x: x)
|
||||
negative_cn = CM(NEGATIVE[0], NEGATIVE[1], "-", lambda x: -x)
|
||||
point_cn = CM(POINT[0], POINT[1], ".", lambda x, y: float(str(x) + "." + str(y)))
|
||||
# sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
|
||||
system = NumberSystem()
|
||||
system.units = smaller_units + larger_units
|
||||
system.digits = digits
|
||||
system.math = MathSymbol(positive_cn, negative_cn, point_cn)
|
||||
# system.symbols = OtherSymbol(sil_cn)
|
||||
return system
|
||||
|
||||
|
||||
def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
|
||||
|
||||
def get_symbol(char, system):
|
||||
for u in system.units:
|
||||
if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
|
||||
return u
|
||||
for d in system.digits:
|
||||
if char in [
|
||||
d.traditional,
|
||||
d.simplified,
|
||||
d.big_s,
|
||||
d.big_t,
|
||||
d.alt_s,
|
||||
d.alt_t,
|
||||
]:
|
||||
return d
|
||||
for m in system.math:
|
||||
if char in [m.traditional, m.simplified]:
|
||||
return m
|
||||
|
||||
def string2symbols(chinese_string, system):
|
||||
int_string, dec_string = chinese_string, ""
|
||||
for p in [system.math.point.simplified, system.math.point.traditional]:
|
||||
if p in chinese_string:
|
||||
int_string, dec_string = chinese_string.split(p)
|
||||
break
|
||||
return [get_symbol(c, system) for c in int_string], [
|
||||
get_symbol(c, system) for c in dec_string
|
||||
]
|
||||
|
||||
def correct_symbols(integer_symbols, system):
|
||||
"""
|
||||
一百八 to 一百八十
|
||||
一亿一千三百万 to 一亿 一千万 三百万
|
||||
"""
|
||||
|
||||
if integer_symbols and isinstance(integer_symbols[0], CNU):
|
||||
if integer_symbols[0].power == 1:
|
||||
integer_symbols = [system.digits[1]] + integer_symbols
|
||||
|
||||
if len(integer_symbols) > 1:
|
||||
if isinstance(integer_symbols[-1], CND) and isinstance(
|
||||
integer_symbols[-2], CNU
|
||||
):
|
||||
integer_symbols.append(
|
||||
CNU(integer_symbols[-2].power - 1, None, None, None, None)
|
||||
)
|
||||
|
||||
result = []
|
||||
unit_count = 0
|
||||
for s in integer_symbols:
|
||||
if isinstance(s, CND):
|
||||
result.append(s)
|
||||
unit_count = 0
|
||||
elif isinstance(s, CNU):
|
||||
current_unit = CNU(s.power, None, None, None, None)
|
||||
unit_count += 1
|
||||
|
||||
if unit_count == 1:
|
||||
result.append(current_unit)
|
||||
elif unit_count > 1:
|
||||
for i in range(len(result)):
|
||||
if (
|
||||
isinstance(result[-i - 1], CNU)
|
||||
and result[-i - 1].power < current_unit.power
|
||||
):
|
||||
result[-i - 1] = CNU(
|
||||
result[-i - 1].power + current_unit.power,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
return result
|
||||
|
||||
def compute_value(integer_symbols):
|
||||
"""
|
||||
Compute the value.
|
||||
When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
|
||||
e.g. '两千万' = 2000 * 10000 not 2000 + 10000
|
||||
"""
|
||||
value = [0]
|
||||
last_power = 0
|
||||
for s in integer_symbols:
|
||||
if isinstance(s, CND):
|
||||
value[-1] = s.value
|
||||
elif isinstance(s, CNU):
|
||||
value[-1] *= pow(10, s.power)
|
||||
if s.power > last_power:
|
||||
value[:-1] = list(map(lambda v: v * pow(10, s.power), value[:-1]))
|
||||
last_power = s.power
|
||||
value.append(0)
|
||||
return sum(value)
|
||||
|
||||
system = create_system(numbering_type)
|
||||
int_part, dec_part = string2symbols(chinese_string, system)
|
||||
int_part = correct_symbols(int_part, system)
|
||||
int_str = str(compute_value(int_part))
|
||||
dec_str = "".join([str(d.value) for d in dec_part])
|
||||
if dec_part:
|
||||
return "{0}.{1}".format(int_str, dec_str)
|
||||
else:
|
||||
return int_str
|
||||
|
||||
|
||||
def num2chn(
|
||||
number_string,
|
||||
numbering_type=NUMBERING_TYPES[1],
|
||||
big=False,
|
||||
traditional=False,
|
||||
alt_zero=False,
|
||||
alt_one=False,
|
||||
alt_two=True,
|
||||
use_zeros=True,
|
||||
use_units=True,
|
||||
):
|
||||
|
||||
def get_value(value_string, use_zeros=True):
|
||||
|
||||
striped_string = value_string.lstrip("0")
|
||||
|
||||
# record nothing if all zeros
|
||||
if not striped_string:
|
||||
return []
|
||||
|
||||
# record one digits
|
||||
elif len(striped_string) == 1:
|
||||
if use_zeros and len(value_string) != len(striped_string):
|
||||
return [system.digits[0], system.digits[int(striped_string)]]
|
||||
else:
|
||||
return [system.digits[int(striped_string)]]
|
||||
|
||||
# recursively record multiple digits
|
||||
else:
|
||||
result_unit = next(
|
||||
u for u in reversed(system.units) if u.power < len(striped_string)
|
||||
)
|
||||
result_string = value_string[: -result_unit.power]
|
||||
return (
|
||||
get_value(result_string)
|
||||
+ [result_unit]
|
||||
+ get_value(striped_string[-result_unit.power :])
|
||||
)
|
||||
|
||||
system = create_system(numbering_type)
|
||||
|
||||
int_dec = number_string.split(".")
|
||||
if len(int_dec) == 1:
|
||||
int_string = int_dec[0]
|
||||
dec_string = ""
|
||||
elif len(int_dec) == 2:
|
||||
int_string = int_dec[0]
|
||||
dec_string = int_dec[1]
|
||||
else:
|
||||
raise ValueError(
|
||||
"invalid input num string with more than one dot: {}".format(number_string)
|
||||
)
|
||||
|
||||
if use_units and len(int_string) > 1:
|
||||
result_symbols = get_value(int_string)
|
||||
else:
|
||||
result_symbols = [system.digits[int(c)] for c in int_string]
|
||||
dec_symbols = [system.digits[int(c)] for c in dec_string]
|
||||
if dec_string:
|
||||
result_symbols += [system.math.point] + dec_symbols
|
||||
|
||||
if alt_two:
|
||||
liang = CND(
|
||||
2,
|
||||
system.digits[2].alt_s,
|
||||
system.digits[2].alt_t,
|
||||
system.digits[2].big_s,
|
||||
system.digits[2].big_t,
|
||||
)
|
||||
for i, v in enumerate(result_symbols):
|
||||
if isinstance(v, CND) and v.value == 2:
|
||||
next_symbol = (
|
||||
result_symbols[i + 1] if i < len(result_symbols) - 1 else None
|
||||
)
|
||||
previous_symbol = result_symbols[i - 1] if i > 0 else None
|
||||
if isinstance(next_symbol, CNU) and isinstance(
|
||||
previous_symbol, (CNU, type(None))
|
||||
):
|
||||
if next_symbol.power != 1 and (
|
||||
(previous_symbol is None) or (previous_symbol.power != 1)
|
||||
):
|
||||
result_symbols[i] = liang
|
||||
|
||||
# if big is True, '两' will not be used and `alt_two` has no impact on output
|
||||
if big:
|
||||
attr_name = "big_"
|
||||
if traditional:
|
||||
attr_name += "t"
|
||||
else:
|
||||
attr_name += "s"
|
||||
else:
|
||||
if traditional:
|
||||
attr_name = "traditional"
|
||||
else:
|
||||
attr_name = "simplified"
|
||||
|
||||
result = "".join([getattr(s, attr_name) for s in result_symbols])
|
||||
|
||||
# if not use_zeros:
|
||||
# result = result.strip(getattr(system.digits[0], attr_name))
|
||||
|
||||
if alt_zero:
|
||||
result = result.replace(
|
||||
getattr(system.digits[0], attr_name), system.digits[0].alt_s
|
||||
)
|
||||
|
||||
if alt_one:
|
||||
result = result.replace(
|
||||
getattr(system.digits[1], attr_name), system.digits[1].alt_s
|
||||
)
|
||||
|
||||
for i, p in enumerate(POINT):
|
||||
if result.startswith(p):
|
||||
return CHINESE_DIGIS[0] + result
|
||||
|
||||
# ^10, 11, .., 19
|
||||
if (
|
||||
len(result) >= 2
|
||||
and result[1]
|
||||
in [
|
||||
SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
|
||||
SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0],
|
||||
]
|
||||
and result[0]
|
||||
in [
|
||||
CHINESE_DIGIS[1],
|
||||
BIG_CHINESE_DIGIS_SIMPLIFIED[1],
|
||||
BIG_CHINESE_DIGIS_TRADITIONAL[1],
|
||||
]
|
||||
):
|
||||
result = result[1:]
|
||||
|
||||
return result
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
all_chinese_number_string = (
|
||||
CHINESE_DIGIS
|
||||
+ BIG_CHINESE_DIGIS_SIMPLIFIED
|
||||
+ BIG_CHINESE_DIGIS_TRADITIONAL
|
||||
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED
|
||||
+ LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL
|
||||
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED
|
||||
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL
|
||||
+ ZERO_ALT
|
||||
+ ONE_ALT
|
||||
+ "".join(TWO_ALTS + POSITIVE + NEGATIVE + POINT)
|
||||
)
|
||||
|
||||
print("num:", chn2num("一万零四百零三点八零五"))
|
||||
print("num:", chn2num("一亿六点三"))
|
||||
print("num:", chn2num("一亿零六点三"))
|
||||
print("num:", chn2num("两千零一亿六点三"))
|
||||
# print('num:', chn2num('一零零八六'))
|
||||
print("txt:", num2chn("10260.03", alt_zero=True))
|
||||
print("txt:", num2chn("20037.090", numbering_type="low", traditional=True))
|
||||
print("txt:", num2chn("100860001.77", numbering_type="high", big=True))
|
||||
print(
|
||||
"txt:",
|
||||
num2chn(
|
||||
"059523810880",
|
||||
alt_one=True,
|
||||
alt_two=False,
|
||||
use_lzeros=True,
|
||||
use_rzeros=True,
|
||||
use_units=False,
|
||||
),
|
||||
)
|
||||
|
||||
print(all_chinese_number_string)
|
||||
@@ -0,0 +1,32 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""CARDINAL类 (包含小数DECIMAL类)
|
||||
纯数 <=> 中文字符串 方法
|
||||
中文字符串 <=> 纯数 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class Cardinal:
|
||||
"""
|
||||
CARDINAL类
|
||||
"""
|
||||
|
||||
def __init__(self, cardinal=None, chntext=None):
|
||||
self.cardinal = cardinal
|
||||
self.chntext = chntext
|
||||
|
||||
def chntext2cardinal(self):
|
||||
return chn2num(self.chntext)
|
||||
|
||||
def cardinal2chntext(self):
|
||||
return num2chn(self.cardinal)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Cardinal(cardinal="21357.230").cardinal2chntext())
|
||||
@@ -0,0 +1,75 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""DATE类
|
||||
日期 <=> 中文字符串 方法
|
||||
中文字符串 <=> 日期 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-07"
|
||||
|
||||
from fish_speech.text.chn_text_norm.cardinal import Cardinal
|
||||
from fish_speech.text.chn_text_norm.digit import Digit
|
||||
|
||||
|
||||
class Date:
|
||||
"""
|
||||
DATE类
|
||||
"""
|
||||
|
||||
def __init__(self, date=None, chntext=None):
|
||||
self.date = date
|
||||
self.chntext = chntext
|
||||
|
||||
# def chntext2date(self):
|
||||
# chntext = self.chntext
|
||||
# try:
|
||||
# year, other = chntext.strip().split('年', maxsplit=1)
|
||||
# year = Digit(chntext=year).digit2chntext() + '年'
|
||||
# except ValueError:
|
||||
# other = chntext
|
||||
# year = ''
|
||||
# if other:
|
||||
# try:
|
||||
# month, day = other.strip().split('月', maxsplit=1)
|
||||
# month = Cardinal(chntext=month).chntext2cardinal() + '月'
|
||||
# except ValueError:
|
||||
# day = chntext
|
||||
# month = ''
|
||||
# if day:
|
||||
# day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
|
||||
# else:
|
||||
# month = ''
|
||||
# day = ''
|
||||
# date = year + month + day
|
||||
# self.date = date
|
||||
# return self.date
|
||||
|
||||
def date2chntext(self):
|
||||
date = self.date
|
||||
try:
|
||||
year, other = date.strip().split("年", maxsplit=1)
|
||||
year = Digit(digit=year).digit2chntext() + "年"
|
||||
except ValueError:
|
||||
other = date
|
||||
year = ""
|
||||
if other:
|
||||
try:
|
||||
month, day = other.strip().split("月", maxsplit=1)
|
||||
month = Cardinal(cardinal=month).cardinal2chntext() + "月"
|
||||
except ValueError:
|
||||
day = date
|
||||
month = ""
|
||||
if day:
|
||||
day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
|
||||
else:
|
||||
month = ""
|
||||
day = ""
|
||||
chntext = year + month + day
|
||||
self.chntext = chntext
|
||||
return self.chntext
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试
|
||||
print(Date(date="09年3月16日").date2chntext())
|
||||
@@ -0,0 +1,32 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""DIGIT类
|
||||
数字串 <=> 中文字符串 方法
|
||||
中文字符串 <=> 数字串 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class Digit:
|
||||
"""
|
||||
DIGIT类
|
||||
"""
|
||||
|
||||
def __init__(self, digit=None, chntext=None):
|
||||
self.digit = digit
|
||||
self.chntext = chntext
|
||||
|
||||
# def chntext2digit(self):
|
||||
# return chn2num(self.chntext)
|
||||
|
||||
def digit2chntext(self):
|
||||
return num2chn(self.digit, alt_two=False, use_units=False)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Digit(digit="2016").digit2chntext())
|
||||
@@ -0,0 +1,35 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""FRACTION类
|
||||
分数 <=> 中文字符串 方法
|
||||
中文字符串 <=> 分数 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class Fraction:
|
||||
"""
|
||||
FRACTION类
|
||||
"""
|
||||
|
||||
def __init__(self, fraction=None, chntext=None):
|
||||
self.fraction = fraction
|
||||
self.chntext = chntext
|
||||
|
||||
def chntext2fraction(self):
|
||||
denominator, numerator = self.chntext.split("分之")
|
||||
return chn2num(numerator) + "/" + chn2num(denominator)
|
||||
|
||||
def fraction2chntext(self):
|
||||
numerator, denominator = self.fraction.split("/")
|
||||
return num2chn(denominator) + "分之" + num2chn(numerator)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Fraction(fraction="2135/7230").fraction2chntext())
|
||||
print(Fraction(chntext="五百八十一分之三百六十九").chntext2fraction())
|
||||
@@ -0,0 +1,43 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""MONEY类
|
||||
金钱 <=> 中文字符串 方法
|
||||
中文字符串 <=> 金钱 方法
|
||||
"""
|
||||
import re
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-08"
|
||||
|
||||
from fish_speech.text.chn_text_norm.cardinal import Cardinal
|
||||
|
||||
|
||||
class Money:
|
||||
"""
|
||||
MONEY类
|
||||
"""
|
||||
|
||||
def __init__(self, money=None, chntext=None):
|
||||
self.money = money
|
||||
self.chntext = chntext
|
||||
|
||||
# def chntext2money(self):
|
||||
# return self.money
|
||||
|
||||
def money2chntext(self):
|
||||
money = self.money
|
||||
pattern = re.compile(r"(\d+(\.\d+)?)")
|
||||
matchers = pattern.findall(money)
|
||||
if matchers:
|
||||
for matcher in matchers:
|
||||
money = money.replace(
|
||||
matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext()
|
||||
)
|
||||
self.chntext = money
|
||||
return self.chntext
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试
|
||||
print(Money(money="21.5万元").money2chntext())
|
||||
print(Money(money="230块5毛").money2chntext())
|
||||
@@ -0,0 +1,33 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""PERCENTAGE类
|
||||
百分数 <=> 中文字符串 方法
|
||||
中文字符串 <=> 百分数 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-06"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class Percentage:
|
||||
"""
|
||||
PERCENTAGE类
|
||||
"""
|
||||
|
||||
def __init__(self, percentage=None, chntext=None):
|
||||
self.percentage = percentage
|
||||
self.chntext = chntext
|
||||
|
||||
def chntext2percentage(self):
|
||||
return chn2num(self.chntext.strip().strip("百分之")) + "%"
|
||||
|
||||
def percentage2chntext(self):
|
||||
return "百分之" + num2chn(self.percentage.strip().strip("%"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Percentage(chntext="百分之五十六点零三").chntext2percentage())
|
||||
print(Percentage(percentage="65.3%").percentage2chntext())
|
||||
@@ -0,0 +1,51 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""TELEPHONE类
|
||||
电话号码 <=> 中文字符串 方法
|
||||
中文字符串 <=> 电话号码 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class TelePhone:
|
||||
"""
|
||||
TELEPHONE类
|
||||
"""
|
||||
|
||||
def __init__(self, telephone=None, raw_chntext=None, chntext=None):
|
||||
self.telephone = telephone
|
||||
self.raw_chntext = raw_chntext
|
||||
self.chntext = chntext
|
||||
|
||||
# def chntext2telephone(self):
|
||||
# sil_parts = self.raw_chntext.split('<SIL>')
|
||||
# self.telephone = '-'.join([
|
||||
# str(chn2num(p)) for p in sil_parts
|
||||
# ])
|
||||
# return self.telephone
|
||||
|
||||
def telephone2chntext(self, fixed=False):
|
||||
|
||||
if fixed:
|
||||
sil_parts = self.telephone.split("-")
|
||||
self.raw_chntext = "<SIL>".join(
|
||||
[num2chn(part, alt_two=False, use_units=False) for part in sil_parts]
|
||||
)
|
||||
self.chntext = self.raw_chntext.replace("<SIL>", "")
|
||||
else:
|
||||
sp_parts = self.telephone.strip("+").split()
|
||||
self.raw_chntext = "<SP>".join(
|
||||
[num2chn(part, alt_two=False, use_units=False) for part in sp_parts]
|
||||
)
|
||||
self.chntext = self.raw_chntext.replace("<SP>", "")
|
||||
return self.chntext
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(TelePhone(telephone="0595-23980880").telephone2chntext())
|
||||
# print(TelePhone(raw_chntext='零五九五杠二三八六五零九八').chntext2telephone())
|
||||
@@ -0,0 +1,177 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
TEXT类
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
import re
|
||||
|
||||
from fish_speech.text.chn_text_norm.cardinal import Cardinal
|
||||
from fish_speech.text.chn_text_norm.date import Date
|
||||
from fish_speech.text.chn_text_norm.digit import Digit
|
||||
from fish_speech.text.chn_text_norm.fraction import Fraction
|
||||
from fish_speech.text.chn_text_norm.money import Money
|
||||
from fish_speech.text.chn_text_norm.percentage import Percentage
|
||||
from fish_speech.text.chn_text_norm.telephone import TelePhone
|
||||
|
||||
CURRENCY_NAMES = (
|
||||
"(人民币|美元|日元|英镑|欧元|马克|法郎|加拿大元|澳元|港币|先令|芬兰马克|爱尔兰镑|"
|
||||
"里拉|荷兰盾|埃斯库多|比塞塔|印尼盾|林吉特|新西兰元|比索|卢布|新加坡元|韩元|泰铢)"
|
||||
)
|
||||
CURRENCY_UNITS = "((亿|千万|百万|万|千|百)|(亿|千万|百万|万|千|百|)元|(亿|千万|百万|万|千|百|)块|角|毛|分)"
|
||||
COM_QUANTIFIERS = (
|
||||
"(匹|张|座|回|场|尾|条|个|首|阙|阵|网|炮|顶|丘|棵|只|支|袭|辆|挑|担|颗|壳|窠|曲|墙|群|腔|"
|
||||
"砣|座|客|贯|扎|捆|刀|令|打|手|罗|坡|山|岭|江|溪|钟|队|单|双|对|出|口|头|脚|板|跳|枝|件|贴|"
|
||||
"针|线|管|名|位|身|堂|课|本|页|家|户|层|丝|毫|厘|分|钱|两|斤|担|铢|石|钧|锱|忽|(千|毫|微)克|"
|
||||
"毫|厘|分|寸|尺|丈|里|寻|常|铺|程|(千|分|厘|毫|微)米|撮|勺|合|升|斗|石|盘|碗|碟|叠|桶|笼|盆|"
|
||||
"盒|杯|钟|斛|锅|簋|篮|盘|桶|罐|瓶|壶|卮|盏|箩|箱|煲|啖|袋|钵|年|月|日|季|刻|时|周|天|秒|分|旬|"
|
||||
"纪|岁|世|更|夜|春|夏|秋|冬|代|伏|辈|丸|泡|粒|颗|幢|堆|条|根|支|道|面|片|张|颗|块|人|抽)"
|
||||
)
|
||||
|
||||
|
||||
class Text:
|
||||
"""
|
||||
Text类
|
||||
"""
|
||||
|
||||
def __init__(self, raw_text, norm_text=None):
|
||||
self.raw_text = "^" + raw_text + "$"
|
||||
self.norm_text = norm_text
|
||||
|
||||
def _particular(self):
|
||||
text = self.norm_text
|
||||
pattern = re.compile(r"(([a-zA-Z]+)二([a-zA-Z]+))")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('particular')
|
||||
for matcher in matchers:
|
||||
text = text.replace(matcher[0], matcher[1] + "2" + matcher[2], 1)
|
||||
self.norm_text = text
|
||||
return self.norm_text
|
||||
|
||||
def normalize(self):
|
||||
text = self.raw_text
|
||||
|
||||
# 规范化日期
|
||||
pattern = re.compile(
|
||||
r"\D+((([089]\d|(19|20)\d{2})年)?(\d{1,2}月(\d{1,2}[日号])?)?)"
|
||||
)
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('date')
|
||||
for matcher in matchers:
|
||||
text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
|
||||
|
||||
# 规范化金钱
|
||||
pattern = re.compile(
|
||||
r"\D+((\d+(\.\d+)?)[多余几]?"
|
||||
+ CURRENCY_UNITS
|
||||
+ "(\d"
|
||||
+ CURRENCY_UNITS
|
||||
+ "?)?)"
|
||||
)
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('money')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0], Money(money=matcher[0]).money2chntext(), 1
|
||||
)
|
||||
|
||||
# 规范化固话/手机号码
|
||||
# 手机
|
||||
# http://www.jihaoba.com/news/show/13680
|
||||
# 移动:139、138、137、136、135、134、159、158、157、150、151、152、188、187、182、183、184、178、198
|
||||
# 联通:130、131、132、156、155、186、185、176
|
||||
# 电信:133、153、189、180、181、177
|
||||
pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('telephone')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1
|
||||
)
|
||||
# 固话
|
||||
pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('fixed telephone')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0],
|
||||
TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True),
|
||||
1,
|
||||
)
|
||||
|
||||
# 规范化分数
|
||||
pattern = re.compile(r"(\d+/\d+)")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('fraction')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher, Fraction(fraction=matcher).fraction2chntext(), 1
|
||||
)
|
||||
|
||||
# 规范化百分数
|
||||
text = text.replace("%", "%")
|
||||
pattern = re.compile(r"(\d+(\.\d+)?%)")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('percentage')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0],
|
||||
Percentage(percentage=matcher[0]).percentage2chntext(),
|
||||
1,
|
||||
)
|
||||
|
||||
# 规范化纯数+量词
|
||||
pattern = re.compile(r"(\d+(\.\d+)?)[多余几]?" + COM_QUANTIFIERS)
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('cardinal+quantifier')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1
|
||||
)
|
||||
|
||||
# 规范化数字编号
|
||||
pattern = re.compile(r"(\d{4,32})")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('digit')
|
||||
for matcher in matchers:
|
||||
text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
|
||||
|
||||
# 规范化纯数
|
||||
pattern = re.compile(r"(\d+(\.\d+)?)")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('cardinal')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1
|
||||
)
|
||||
|
||||
self.norm_text = text
|
||||
self._particular()
|
||||
|
||||
return self.norm_text.lstrip("^").rstrip("$")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Text(raw_text="固话:0595-23865596或23880880。").normalize())
|
||||
print(Text(raw_text="手机:+86 19859213959或15659451527。").normalize())
|
||||
print(Text(raw_text="分数:32477/76391。").normalize())
|
||||
print(Text(raw_text="百分数:80.03%。").normalize())
|
||||
print(Text(raw_text="编号:31520181154418。").normalize())
|
||||
print(Text(raw_text="纯数:2983.07克或12345.60米。").normalize())
|
||||
print(Text(raw_text="日期:1999年2月20日或09年3月15号。").normalize())
|
||||
print(Text(raw_text="金钱:12块5,34.5元,20.1万").normalize())
|
||||
print(Text(raw_text="特殊:O2O或B2C。").normalize())
|
||||
@@ -0,0 +1,31 @@
|
||||
import re
|
||||
|
||||
SYMBOLS_MAPPING = {
|
||||
"“": "'",
|
||||
"”": "'",
|
||||
"‘": "'",
|
||||
"’": "'",
|
||||
"【": "",
|
||||
"】": "",
|
||||
"[": "",
|
||||
"]": "",
|
||||
"(": "",
|
||||
")": "",
|
||||
"(": "",
|
||||
")": "",
|
||||
"・": "·",
|
||||
}
|
||||
|
||||
REPLACE_SYMBOL_REGEX = re.compile(
|
||||
"|".join(re.escape(p) for p in SYMBOLS_MAPPING.keys())
|
||||
)
|
||||
|
||||
|
||||
def clean_text(text):
|
||||
# Clean the text
|
||||
text = text.strip()
|
||||
|
||||
# Replace all chinese symbols with their english counterparts
|
||||
text = REPLACE_SYMBOL_REGEX.sub(lambda x: SYMBOLS_MAPPING[x.group()], text)
|
||||
|
||||
return text
|
||||
@@ -0,0 +1,130 @@
|
||||
import re
|
||||
import string
|
||||
|
||||
from fish_speech.text.clean import clean_text
|
||||
|
||||
|
||||
def utf_8_len(text):
|
||||
return len(text.encode("utf-8"))
|
||||
|
||||
|
||||
def break_text(texts, length, splits: set):
|
||||
for text in texts:
|
||||
if utf_8_len(text) <= length:
|
||||
yield text
|
||||
continue
|
||||
|
||||
curr = ""
|
||||
for char in text:
|
||||
curr += char
|
||||
|
||||
if char in splits:
|
||||
yield curr
|
||||
curr = ""
|
||||
|
||||
if curr:
|
||||
yield curr
|
||||
|
||||
|
||||
def break_text_by_length(texts, length):
|
||||
for text in texts:
|
||||
if utf_8_len(text) <= length:
|
||||
yield text
|
||||
continue
|
||||
|
||||
curr = ""
|
||||
for char in text:
|
||||
curr += char
|
||||
|
||||
if utf_8_len(curr) >= length:
|
||||
yield curr
|
||||
curr = ""
|
||||
|
||||
if curr:
|
||||
yield curr
|
||||
|
||||
|
||||
def add_cleaned(curr, segments):
|
||||
curr = curr.strip()
|
||||
if curr and not all(c.isspace() or c in string.punctuation for c in curr):
|
||||
segments.append(curr)
|
||||
|
||||
|
||||
def protect_float(text):
|
||||
# Turns 3.14 into <3_f_14> to prevent splitting
|
||||
return re.sub(r"(\d+)\.(\d+)", r"<\1_f_\2>", text)
|
||||
|
||||
|
||||
def unprotect_float(text):
|
||||
# Turns <3_f_14> into 3.14
|
||||
return re.sub(r"<(\d+)_f_(\d+)>", r"\1.\2", text)
|
||||
|
||||
|
||||
def split_text(text, length):
|
||||
text = clean_text(text)
|
||||
|
||||
# Break the text into pieces with following rules:
|
||||
# 1. Split the text at ".", "!", "?" if text is NOT a float
|
||||
# 2. If the text is longer than length, split at ","
|
||||
# 3. If the text is still longer than length, split at " "
|
||||
# 4. If the text is still longer than length, split at any character to length
|
||||
|
||||
texts = [text]
|
||||
texts = map(protect_float, texts)
|
||||
texts = break_text(texts, length, {".", "!", "?", "。", "!", "?"})
|
||||
texts = map(unprotect_float, texts)
|
||||
texts = break_text(texts, length, {",", ","})
|
||||
texts = break_text(texts, length, {" "})
|
||||
texts = list(break_text_by_length(texts, length))
|
||||
|
||||
# Then, merge the texts into segments with length <= length
|
||||
segments = []
|
||||
curr = ""
|
||||
|
||||
for text in texts:
|
||||
if utf_8_len(curr) + utf_8_len(text) <= length:
|
||||
curr += text
|
||||
else:
|
||||
add_cleaned(curr, segments)
|
||||
curr = text
|
||||
|
||||
if curr:
|
||||
add_cleaned(curr, segments)
|
||||
|
||||
return segments
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Test the split_text function
|
||||
|
||||
text = "This is a test sentence. This is another test sentence. And a third one."
|
||||
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence.",
|
||||
"This is another test sentence. And a third one.",
|
||||
]
|
||||
assert split_text("a,aaaaaa3.14", 10) == ["a,", "aaaaaa3.14"]
|
||||
assert split_text(" ", 10) == []
|
||||
assert split_text("a", 10) == ["a"]
|
||||
|
||||
text = "This is a test sentence with only commas, and no dots, and no exclamation marks, and no question marks, and no newlines."
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence with only commas,",
|
||||
"and no dots, and no exclamation marks,",
|
||||
"and no question marks, and no newlines.",
|
||||
]
|
||||
|
||||
text = "This is a test sentence This is a test sentence This is a test sentence. This is a test sentence, This is a test sentence, This is a test sentence."
|
||||
# First half split at " ", second half split at ","
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence This is a test sentence",
|
||||
"This is a test sentence. This is a test sentence,",
|
||||
"This is a test sentence, This is a test sentence.",
|
||||
]
|
||||
|
||||
text = "这是一段很长的中文文本,而且没有句号,也没有感叹号,也没有问号,也没有换行符。"
|
||||
assert split_text(text, 50) == [
|
||||
"这是一段很长的中文文本,",
|
||||
"而且没有句号,也没有感叹号,",
|
||||
"也没有问号,也没有换行符.",
|
||||
]
|
||||
@@ -0,0 +1,169 @@
|
||||
import itertools
|
||||
import os
|
||||
import re
|
||||
from collections import defaultdict
|
||||
from functools import partial
|
||||
from multiprocessing import Pool
|
||||
from pathlib import Path
|
||||
|
||||
import click
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
|
||||
from fish_speech.datasets.protos.text_data_pb2 import Semantics, Sentence, TextData
|
||||
from fish_speech.datasets.protos.text_data_stream import pack_pb_stream
|
||||
from tools.file import load_filelist
|
||||
|
||||
# To avoid CPU overload
|
||||
os.environ["MKL_NUM_THREADS"] = "1"
|
||||
os.environ["OMP_NUM_THREADS"] = "1"
|
||||
|
||||
|
||||
def task_generator_folder(root: Path, text_extension: str):
|
||||
files = list(tqdm(Path(root).rglob("*.npy"), desc=f"Loading {root}"))
|
||||
files = sorted(files)
|
||||
|
||||
grouped_files = defaultdict(list)
|
||||
for file in tqdm(files, desc=f"Grouping {root}"):
|
||||
p = str(file.parent)
|
||||
speaker = file.parent.name
|
||||
|
||||
try:
|
||||
if isinstance(text_extension, str):
|
||||
texts = [file.with_suffix(text_extension).read_text(encoding="utf-8")]
|
||||
else:
|
||||
texts = [
|
||||
file.with_suffix(ext).read_text(encoding="utf-8")
|
||||
for ext in text_extension
|
||||
]
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to read text {file}: {e}")
|
||||
continue
|
||||
|
||||
grouped_files[p].append((speaker, file, texts))
|
||||
|
||||
logger.info(
|
||||
f"Found {len(grouped_files)} groups in {root}, {list(grouped_files.keys())[:5]}..."
|
||||
)
|
||||
|
||||
for i in grouped_files.values():
|
||||
subset = [(f, t) for _, f, t in i]
|
||||
yield i[0][0], subset, "folder"
|
||||
|
||||
|
||||
def task_generator_filelist(filelist):
|
||||
grouped_files = defaultdict(list)
|
||||
for filename, speaker, _, text in load_filelist(filelist):
|
||||
grouped_files[speaker].append((Path(filename), [text]))
|
||||
|
||||
logger.info(f"Found {len(grouped_files)} groups in {filelist}")
|
||||
for speaker, values in grouped_files.items():
|
||||
yield speaker, values, "filelist"
|
||||
|
||||
|
||||
def run_task(task):
|
||||
name, subset, source = task
|
||||
|
||||
# Parse the files
|
||||
sentences = []
|
||||
for file, texts in subset:
|
||||
np_file = file.with_suffix(".npy")
|
||||
if np_file.exists() is False:
|
||||
logger.warning(f"Can't find {np_file}")
|
||||
continue
|
||||
|
||||
new_texts = []
|
||||
|
||||
for text in texts:
|
||||
# Simple cleaning: replace { xxx } and < xxx > with space
|
||||
text = re.sub(r"\{.*?\}", " ", text)
|
||||
text = re.sub(r"<.*?>", " ", text)
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
new_texts.append(text)
|
||||
|
||||
try:
|
||||
semantics = np.load(np_file)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to parse {file}: {e}")
|
||||
continue
|
||||
|
||||
if isinstance(semantics, np.ndarray):
|
||||
semantics = semantics.tolist()
|
||||
|
||||
sentences.append(
|
||||
Sentence(
|
||||
texts=new_texts,
|
||||
semantics=[Semantics(values=s) for s in semantics],
|
||||
)
|
||||
)
|
||||
|
||||
# Pack the sentences
|
||||
return pack_pb_stream(
|
||||
TextData(
|
||||
source=source,
|
||||
name=name,
|
||||
sentences=sentences,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@click.command()
|
||||
@click.option(
|
||||
"--input",
|
||||
type=click.Path(path_type=Path),
|
||||
required=True,
|
||||
help="A folder containing the dataset or a filelist",
|
||||
multiple=True,
|
||||
)
|
||||
@click.option(
|
||||
"--output", type=click.Path(path_type=Path), default="data/quantized-dataset-ft"
|
||||
)
|
||||
@click.option("--num-workers", type=int, default=16)
|
||||
@click.option("--text-extension", type=str, default=[".txt"], multiple=True)
|
||||
@click.option(
|
||||
"--shard-size", type=int, default=10, help="The maximum size of each shard in mb"
|
||||
)
|
||||
def main(input, output, num_workers, text_extension, shard_size):
|
||||
generator_fns = []
|
||||
|
||||
for f in input:
|
||||
assert f.exists(), f"{f} not found"
|
||||
|
||||
if f.is_dir():
|
||||
generator_fn = task_generator_folder(f, text_extension)
|
||||
else:
|
||||
generator_fn = task_generator_filelist(f)
|
||||
|
||||
generator_fns.append(generator_fn)
|
||||
|
||||
generator_fn = itertools.chain(*generator_fns)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
dataset_fp = None
|
||||
tar_idx = 0
|
||||
written_size = 0
|
||||
|
||||
with Pool(num_workers) as p:
|
||||
for result in tqdm(p.imap_unordered(run_task, generator_fn)):
|
||||
if dataset_fp is None:
|
||||
dataset_fp = open(Path(output) / f"{tar_idx:08d}.protos", "wb")
|
||||
|
||||
dataset_fp.write(result)
|
||||
written_size += len(result)
|
||||
|
||||
if written_size > shard_size * 1024 * 1024:
|
||||
logger.info(f"Finished writing {tar_idx} shards to {output}")
|
||||
dataset_fp.close()
|
||||
dataset_fp = None
|
||||
written_size = 0
|
||||
tar_idx += 1
|
||||
|
||||
if dataset_fp is not None:
|
||||
dataset_fp.close()
|
||||
|
||||
logger.info(f"Finished writing {tar_idx + 1} shards to {output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,171 @@
|
||||
import pyrootutils
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from matplotlib import pyplot as plt
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
# register eval resolver and root
|
||||
pyrootutils.setup_root(__file__, indicator=".project-root", pythonpath=True)
|
||||
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from fish_speech.datasets.semantic import AutoAugTextDataset, TextDataCollator
|
||||
from tools.llama.generate import load_model
|
||||
|
||||
|
||||
def smooth(
|
||||
scalars: list[float], weight: float
|
||||
) -> list[float]: # Weight between 0 and 1
|
||||
last = scalars[0] # First value in the plot (first timestep)
|
||||
smoothed = list()
|
||||
for point in scalars:
|
||||
smoothed_val = last * weight + (1 - weight) * point # Calculate smoothed value
|
||||
smoothed.append(smoothed_val) # Save it
|
||||
last = smoothed_val # Anchor the last smoothed value
|
||||
|
||||
return smoothed
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def analyze_one_model(loader, config, weight, max_length):
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model = load_model(
|
||||
config,
|
||||
weight,
|
||||
device,
|
||||
torch.bfloat16,
|
||||
max_length,
|
||||
compile=False,
|
||||
)[0]
|
||||
|
||||
current_step = 0
|
||||
model.eval()
|
||||
|
||||
semantic_loss_sum = torch.zeros(
|
||||
max_length,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
counter = torch.zeros(
|
||||
max_length,
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
)
|
||||
|
||||
for batch in loader:
|
||||
batch = {k: v.to(device) for k, v in batch.items()}
|
||||
|
||||
labels = batch["labels"]
|
||||
outputs = model(
|
||||
inp=batch["inputs"],
|
||||
key_padding_mask=batch["attention_masks"],
|
||||
)
|
||||
|
||||
token_logits = outputs.token_logits
|
||||
codebook_logits = outputs.codebook_logits
|
||||
|
||||
# Generate labels
|
||||
base_loss = F.cross_entropy(
|
||||
token_logits.reshape(-1, token_logits.size(-1)),
|
||||
labels[:, 0].reshape(-1),
|
||||
ignore_index=-100,
|
||||
reduction="none",
|
||||
)
|
||||
|
||||
codebook_labels = labels[:, 1 : 1 + model.config.num_codebooks].mT
|
||||
semantic_loss = F.cross_entropy(
|
||||
codebook_logits.reshape(-1, codebook_logits.size(-1)),
|
||||
codebook_labels.reshape(-1),
|
||||
ignore_index=-100,
|
||||
reduction="none",
|
||||
)
|
||||
|
||||
base_loss = base_loss.reshape(labels[:, 0].shape)
|
||||
semantic_loss = semantic_loss.reshape(codebook_labels.shape)
|
||||
|
||||
semantic_loss_frame = semantic_loss.mean(-1)
|
||||
pad_pos = codebook_labels.sum(-1) == -100 * model.config.num_codebooks
|
||||
|
||||
for loss_sample, pad in zip(semantic_loss_frame, pad_pos):
|
||||
semantic_loss_sum[~pad] += loss_sample[~pad]
|
||||
counter[~pad] += 1
|
||||
|
||||
current_step += 1
|
||||
if current_step == 10:
|
||||
break
|
||||
|
||||
semantic_loss = semantic_loss.cpu()
|
||||
counter = counter.cpu()
|
||||
xs, ys = [], []
|
||||
|
||||
for i, (loss, count) in enumerate(zip(semantic_loss_sum, counter)):
|
||||
if count > 0:
|
||||
xs.append(i)
|
||||
ys.append((loss / count).item()) # for better loss visualization
|
||||
|
||||
smoothed_ys = smooth(ys, 0.95)
|
||||
|
||||
# Unload model
|
||||
del model
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return xs, ys, smoothed_ys
|
||||
|
||||
|
||||
def main():
|
||||
tokenizer = AutoTokenizer.from_pretrained("fishaudio/fish-speech-1")
|
||||
max_length = 4096
|
||||
|
||||
ds = AutoAugTextDataset(
|
||||
["data/protos/sft/云天河"],
|
||||
tokenizer=tokenizer,
|
||||
use_speaker=False,
|
||||
interactive_prob=1.0,
|
||||
max_length=max_length,
|
||||
)
|
||||
|
||||
loader = DataLoader(
|
||||
ds,
|
||||
batch_size=8,
|
||||
collate_fn=TextDataCollator(tokenizer, max_length=max_length),
|
||||
num_workers=0,
|
||||
shuffle=False,
|
||||
)
|
||||
|
||||
plt.figure(figsize=(10, 5), dpi=200)
|
||||
|
||||
plt.xlabel("Frame")
|
||||
plt.ylabel("Loss")
|
||||
plt.yscale("log")
|
||||
plt.title("Semantic Loss")
|
||||
plt.grid(which="both", axis="both")
|
||||
plt.xlim(0, max_length)
|
||||
|
||||
tests = [
|
||||
(
|
||||
"pertrain-medium",
|
||||
"dual_ar_2_codebook_medium",
|
||||
"checkpoints/text2semantic-pretrain-medium-2k-v1.pth",
|
||||
),
|
||||
(
|
||||
"sft-medium",
|
||||
"dual_ar_2_codebook_medium",
|
||||
"checkpoints/text2semantic-sft-medium-v1.1-4k.pth",
|
||||
),
|
||||
(
|
||||
"sft-large",
|
||||
"dual_ar_2_codebook_large",
|
||||
"checkpoints/text2semantic-sft-large-v1.1-4k.pth",
|
||||
),
|
||||
]
|
||||
|
||||
for name, config, weight in tests:
|
||||
xs, _, smoothed_ys = analyze_one_model(loader, config, weight, max_length)
|
||||
plt.plot(xs, smoothed_ys, label=name)
|
||||
|
||||
plt.legend()
|
||||
plt.savefig("semantic_loss.png")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,699 @@
|
||||
import os
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Literal, Optional, Tuple, Union
|
||||
|
||||
import click
|
||||
import hydra
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch._dynamo.config
|
||||
import torch._inductor.config
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
|
||||
from fish_speech.conversation import CODEBOOK_PAD_TOKEN_ID
|
||||
from fish_speech.clean import clean_text
|
||||
from fish_speech.spliter import split_text
|
||||
|
||||
|
||||
import comfy.utils
|
||||
|
||||
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
torch._inductor.config.coordinate_descent_tuning = True
|
||||
torch._inductor.config.triton.unique_kernel_names = True
|
||||
|
||||
if hasattr(torch._inductor.config, "fx_graph_cache"):
|
||||
# Experimental feature to reduce compilation times, will be on by default in future
|
||||
torch._inductor.config.fx_graph_cache = True
|
||||
|
||||
|
||||
from ...models.text2semantic.llama import BaseTransformer, DualARTransformer, NaiveTransformer
|
||||
|
||||
|
||||
def multinomial_sample_one_no_sync(
|
||||
probs_sort,
|
||||
): # Does multinomial sampling without a cuda synchronization
|
||||
q = torch.empty_like(probs_sort).exponential_(1)
|
||||
return torch.argmax(probs_sort / q, dim=-1, keepdim=True).to(dtype=torch.int)
|
||||
|
||||
|
||||
def logits_to_probs(
|
||||
logits,
|
||||
previous_tokens: Optional[torch.Tensor] = None,
|
||||
temperature: torch.Tensor = 1.0,
|
||||
top_p: torch.Tensor = 1.0,
|
||||
repetition_penalty: torch.Tensor = 1.0,
|
||||
) -> torch.Tensor:
|
||||
# Apply repetition penalty
|
||||
if previous_tokens is not None:
|
||||
previous_tokens = previous_tokens.long()
|
||||
score = torch.gather(logits, dim=0, index=previous_tokens)
|
||||
score = torch.where(
|
||||
score < 0, score * repetition_penalty, score / repetition_penalty
|
||||
)
|
||||
logits.scatter_(dim=0, index=previous_tokens, src=score)
|
||||
|
||||
# Apply top-p sampling
|
||||
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
|
||||
cum_probs = torch.cumsum(torch.nn.functional.softmax(sorted_logits, dim=-1), dim=-1)
|
||||
sorted_indices_to_remove = cum_probs > top_p
|
||||
sorted_indices_to_remove[0] = False # keep at least one option
|
||||
indices_to_remove = sorted_indices_to_remove.scatter(
|
||||
dim=0, index=sorted_indices, src=sorted_indices_to_remove
|
||||
)
|
||||
logits = logits.masked_fill(indices_to_remove, -float("Inf"))
|
||||
|
||||
logits = logits / max(temperature, 1e-5)
|
||||
|
||||
probs = torch.nn.functional.softmax(logits, dim=-1)
|
||||
return probs
|
||||
|
||||
|
||||
def sample(
|
||||
logits,
|
||||
previous_tokens: Optional[torch.Tensor] = None,
|
||||
**sampling_kwargs,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
probs = logits_to_probs(
|
||||
logits=logits[0, -1], previous_tokens=previous_tokens, **sampling_kwargs
|
||||
)
|
||||
idx_next = multinomial_sample_one_no_sync(probs)
|
||||
return idx_next, probs
|
||||
|
||||
|
||||
def decode_one_token_ar(
|
||||
model: DualARTransformer,
|
||||
x: torch.Tensor,
|
||||
input_pos: torch.Tensor,
|
||||
previous_tokens: torch.Tensor = None,
|
||||
**sampling_kwargs,
|
||||
) -> torch.Tensor:
|
||||
x = model.forward_generate(x, input_pos)
|
||||
codebooks = [
|
||||
sample(
|
||||
x.logits,
|
||||
previous_tokens=(
|
||||
previous_tokens[0] if previous_tokens is not None else None
|
||||
), # Disable repetition penalty for the token codebook
|
||||
**sampling_kwargs,
|
||||
)[0]
|
||||
]
|
||||
x = x.hidden_states
|
||||
|
||||
# Cleanup the cache
|
||||
for layer in model.fast_layers:
|
||||
layer.attention.kv_cache.k_cache.fill_(0)
|
||||
layer.attention.kv_cache.v_cache.fill_(0)
|
||||
|
||||
for codebook_idx in range(model.config.num_codebooks):
|
||||
input_pos = torch.tensor([codebook_idx], device=x.device, dtype=torch.long)
|
||||
logits = model.forward_generate_fast(x, input_pos)
|
||||
a = sample(
|
||||
logits,
|
||||
previous_tokens=(
|
||||
previous_tokens[codebook_idx + 1]
|
||||
if previous_tokens is not None
|
||||
else None
|
||||
),
|
||||
**sampling_kwargs,
|
||||
)[0]
|
||||
x = model.fast_embeddings(a)
|
||||
codebooks.append(a)
|
||||
|
||||
return torch.stack(codebooks, dim=0)
|
||||
|
||||
|
||||
def decode_one_token_naive(
|
||||
model: NaiveTransformer,
|
||||
x: torch.Tensor,
|
||||
input_pos: torch.Tensor,
|
||||
previous_tokens: torch.Tensor = None,
|
||||
**sampling_kwargs,
|
||||
) -> torch.Tensor:
|
||||
x = model.forward_generate(x, input_pos)
|
||||
|
||||
codebooks = [
|
||||
sample(
|
||||
x.token_logits,
|
||||
previous_tokens=None, # Disable repetition penalty for the token codebook
|
||||
**sampling_kwargs,
|
||||
)[0]
|
||||
]
|
||||
|
||||
for i in range(model.config.num_codebooks):
|
||||
codebooks.append(
|
||||
sample(
|
||||
x.codebook_logits[:, :, i],
|
||||
previous_tokens=(
|
||||
previous_tokens[i + 1] if previous_tokens is not None else None
|
||||
),
|
||||
**sampling_kwargs,
|
||||
)[0]
|
||||
)
|
||||
|
||||
return torch.stack(codebooks, dim=0)
|
||||
|
||||
|
||||
def decode_n_tokens(
|
||||
model: NaiveTransformer,
|
||||
cur_token: torch.Tensor,
|
||||
input_pos: torch.Tensor,
|
||||
num_new_tokens: int,
|
||||
im_end_id: int = 4,
|
||||
decode_one_token=decode_one_token_naive,
|
||||
**sampling_kwargs,
|
||||
):
|
||||
previous_tokens = torch.zeros(
|
||||
(model.config.num_codebooks + 1, model.config.max_seq_len),
|
||||
dtype=torch.int,
|
||||
device=cur_token.device,
|
||||
)
|
||||
|
||||
for i in tqdm(range(num_new_tokens)):
|
||||
# We need to get windowed repeat penalty
|
||||
win_size = 16
|
||||
if i < win_size:
|
||||
window = previous_tokens[:, :win_size]
|
||||
else:
|
||||
window = previous_tokens[:, i - win_size : i]
|
||||
|
||||
with torch.backends.cuda.sdp_kernel(
|
||||
enable_flash=False, enable_mem_efficient=False, enable_math=True
|
||||
): # Actually better for Inductor to codegen attention here
|
||||
next_token = decode_one_token(
|
||||
model=model,
|
||||
x=cur_token,
|
||||
input_pos=input_pos,
|
||||
previous_tokens=window,
|
||||
**sampling_kwargs,
|
||||
)
|
||||
|
||||
input_pos += 1
|
||||
cur_token = next_token.view(1, model.config.num_codebooks + 1, -1)
|
||||
previous_tokens[:, i : i + 1] = next_token.view(
|
||||
model.config.num_codebooks + 1, -1
|
||||
)
|
||||
|
||||
if cur_token[0, 0, -1] == im_end_id:
|
||||
break
|
||||
|
||||
return previous_tokens[:, : i + 1]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.inference_mode()
|
||||
def generate(
|
||||
*,
|
||||
model: NaiveTransformer,
|
||||
prompt: torch.Tensor,
|
||||
max_new_tokens: int,
|
||||
im_end_id: int = 4,
|
||||
decode_one_token=decode_one_token_naive,
|
||||
**sampling_kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Takes a conditioning sequence (prompt) as input and continues to generate as many tokens as requested.
|
||||
"""
|
||||
|
||||
# create an empty tensor of the expected final shape and fill in the current tokens
|
||||
T = prompt.size(1)
|
||||
|
||||
if max_new_tokens:
|
||||
if T + max_new_tokens > model.config.max_seq_len:
|
||||
max_new_tokens = model.config.max_seq_len - T
|
||||
logger.info(f"Truncating max_new_tokens to {max_new_tokens}")
|
||||
|
||||
T_new = T + max_new_tokens
|
||||
else:
|
||||
T_new = model.config.max_seq_len
|
||||
max_new_tokens = T_new - T
|
||||
|
||||
device, dtype = prompt.device, prompt.dtype
|
||||
with torch.device(device):
|
||||
model.setup_caches(
|
||||
max_batch_size=1, max_seq_len=T_new, dtype=next(model.parameters()).dtype
|
||||
)
|
||||
|
||||
codebook_dim = 1 + model.config.num_codebooks
|
||||
# create an empty tensor of the expected final shape and fill in the current tokens
|
||||
empty = torch.empty((codebook_dim, T_new), dtype=dtype, device=device)
|
||||
empty[:, :T] = prompt
|
||||
seq = empty
|
||||
input_pos = torch.arange(0, T, device=device)
|
||||
|
||||
# Use non-accelerated version for now, to avoid compilation overhead
|
||||
prefill_decode = (
|
||||
decode_one_token_naive
|
||||
if isinstance(model, NaiveTransformer)
|
||||
else decode_one_token_ar
|
||||
)
|
||||
|
||||
next_token = prefill_decode(
|
||||
model, prompt.view(1, codebook_dim, -1), input_pos, **sampling_kwargs
|
||||
)
|
||||
seq[:, T : T + 1] = next_token
|
||||
|
||||
input_pos = torch.tensor([T], device=device, dtype=torch.int)
|
||||
x = decode_n_tokens(
|
||||
model,
|
||||
next_token.view(1, codebook_dim, -1),
|
||||
input_pos,
|
||||
max_new_tokens - 1,
|
||||
im_end_id=im_end_id,
|
||||
decode_one_token=decode_one_token,
|
||||
**sampling_kwargs,
|
||||
)
|
||||
# x = torch.cat(generated_tokens, dim=1)
|
||||
seq = seq[:, : T + 1 + x.size(1)]
|
||||
seq[:, T + 1 :] = x
|
||||
|
||||
return seq
|
||||
|
||||
|
||||
def encode_tokens(
|
||||
tokenizer,
|
||||
string,
|
||||
device="cuda",
|
||||
prompt_tokens=None,
|
||||
num_codebooks=4,
|
||||
):
|
||||
string = clean_text(string)
|
||||
string = f"<|im_start|>user\n{string}<|im_end|><|im_start|>assistant\n"
|
||||
|
||||
new_tokens = tokenizer.encode(
|
||||
string,
|
||||
add_special_tokens=False,
|
||||
max_length=10**6,
|
||||
truncation=False,
|
||||
)
|
||||
tokens = torch.tensor([new_tokens], dtype=torch.int, device=device)
|
||||
|
||||
# Codebooks
|
||||
zeros = (
|
||||
torch.ones((num_codebooks, tokens.size(1)), dtype=torch.int, device=device)
|
||||
* CODEBOOK_PAD_TOKEN_ID
|
||||
)
|
||||
prompt = torch.cat((tokens, zeros), dim=0)
|
||||
|
||||
if prompt_tokens is None:
|
||||
return prompt
|
||||
|
||||
# Get prompt tokens
|
||||
if prompt_tokens.ndim == 3:
|
||||
assert (
|
||||
prompt_tokens.shape[0] == 1
|
||||
), f"3 dim prompt tokens should have shape (1, num_codebooks, seq_len)"
|
||||
prompt_tokens = prompt_tokens[0]
|
||||
|
||||
assert prompt_tokens.ndim == 2
|
||||
data = prompt_tokens + 1
|
||||
|
||||
if prompt_tokens.shape[0] > num_codebooks:
|
||||
logger.warning(
|
||||
f"Prompt tokens shape {prompt_tokens.shape} is larger than num_codebooks {num_codebooks}, getting first {num_codebooks} codebooks"
|
||||
)
|
||||
data = data[:num_codebooks]
|
||||
|
||||
# Add pad token for each codebook
|
||||
data = torch.cat(
|
||||
(data, torch.zeros((data.size(0), 1), dtype=torch.int, device=device)),
|
||||
dim=1,
|
||||
)
|
||||
|
||||
# Since 1.0, we use <|semantic|>
|
||||
s0_token_id = tokenizer.convert_tokens_to_ids("<|semantic|>")
|
||||
end_token_id = tokenizer.convert_tokens_to_ids("<|im_end|>")
|
||||
main_token_ids = (
|
||||
torch.ones((1, data.size(1)), dtype=torch.int, device=device) * s0_token_id
|
||||
)
|
||||
main_token_ids[0, -1] = end_token_id
|
||||
|
||||
data = torch.cat((main_token_ids, data), dim=0)
|
||||
prompt = torch.cat((prompt, data), dim=1)
|
||||
|
||||
return prompt
|
||||
|
||||
|
||||
def load_model(checkpoint_path, device, precision, compile=False):
|
||||
model: Union[NaiveTransformer, DualARTransformer] = BaseTransformer.from_pretrained(
|
||||
checkpoint_path, load_weights=True
|
||||
)
|
||||
|
||||
model = model.to(device=device, dtype=precision)
|
||||
logger.info(f"Restored model from checkpoint")
|
||||
|
||||
if isinstance(model, DualARTransformer):
|
||||
decode_one_token = decode_one_token_ar
|
||||
logger.info("Using DualARTransformer")
|
||||
else:
|
||||
decode_one_token = decode_one_token_naive
|
||||
logger.info("Using NaiveTransformer")
|
||||
|
||||
if compile:
|
||||
logger.info("Compiling function...")
|
||||
decode_one_token = torch.compile(
|
||||
decode_one_token, mode="reduce-overhead", fullgraph=True
|
||||
)
|
||||
|
||||
return model.eval(), decode_one_token
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerateResponse:
|
||||
action: Literal["sample", "next"]
|
||||
codes: Optional[torch.Tensor] = None
|
||||
text: Optional[str] = None
|
||||
|
||||
|
||||
def generate_long(
|
||||
*,
|
||||
model,
|
||||
device: str | torch.device,
|
||||
decode_one_token: callable,
|
||||
text: str,
|
||||
num_samples: int = 1,
|
||||
max_new_tokens: int = 0,
|
||||
top_p: int = 0.7,
|
||||
repetition_penalty: float = 1.5,
|
||||
temperature: float = 0.7,
|
||||
compile: bool = False,
|
||||
iterative_prompt: bool = True,
|
||||
max_length: int = 2048,
|
||||
chunk_length: int = 150,
|
||||
prompt_text: Optional[str | list[str]] = None,
|
||||
prompt_tokens: Optional[torch.Tensor | list[torch.Tensor]] = None,
|
||||
):
|
||||
assert 0 < top_p <= 1, "top_p must be in (0, 1]"
|
||||
assert 0 < repetition_penalty < 2, "repetition_penalty must be in (0, 2)"
|
||||
assert 0 < temperature < 2, "temperature must be in (0, 2)"
|
||||
|
||||
use_prompt = prompt_text is not None and prompt_tokens is not None
|
||||
if use_prompt and isinstance(prompt_text, str):
|
||||
prompt_text = [prompt_text]
|
||||
prompt_tokens = [prompt_tokens]
|
||||
|
||||
assert use_prompt is False or len(prompt_text) == len(
|
||||
prompt_tokens
|
||||
), "Prompt text and tokens must have the same length"
|
||||
|
||||
model_size = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
tokenizer = model.tokenizer
|
||||
im_end_id = tokenizer.convert_tokens_to_ids("<|im_end|>")
|
||||
|
||||
encoded = []
|
||||
texts = split_text(text, chunk_length) if iterative_prompt else [text]
|
||||
encoded_prompts = []
|
||||
|
||||
if use_prompt:
|
||||
for idx, (t, c) in enumerate(zip(prompt_text, prompt_tokens)):
|
||||
encoded_prompts.append(
|
||||
encode_tokens(
|
||||
tokenizer,
|
||||
string=t,
|
||||
device=device,
|
||||
prompt_tokens=c,
|
||||
num_codebooks=model.config.num_codebooks,
|
||||
)
|
||||
)
|
||||
|
||||
for idx, text in enumerate(texts):
|
||||
encoded.append(
|
||||
encode_tokens(
|
||||
tokenizer,
|
||||
string=text,
|
||||
device=device,
|
||||
num_codebooks=model.config.num_codebooks,
|
||||
)
|
||||
)
|
||||
logger.info(f"Encoded text: {text}")
|
||||
|
||||
# Move temperature, top_p, repetition_penalty to device
|
||||
# This is important so that changing params doesn't trigger recompile
|
||||
temperature = torch.tensor(temperature, device=device, dtype=torch.float)
|
||||
top_p = torch.tensor(top_p, device=device, dtype=torch.float)
|
||||
repetition_penalty = torch.tensor(
|
||||
repetition_penalty, device=device, dtype=torch.float
|
||||
)
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(num_samples*len(encoded))
|
||||
|
||||
for sample_idx in range(num_samples):
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
global_encoded = []
|
||||
seg_idx = 0
|
||||
|
||||
while seg_idx < len(encoded):
|
||||
logger.info(
|
||||
f"Generating sentence {seg_idx + 1}/{len(encoded)} of sample {sample_idx + 1}/{num_samples}"
|
||||
)
|
||||
pbar.update(1)
|
||||
seg = encoded[seg_idx]
|
||||
global_encoded.append(seg)
|
||||
|
||||
lengths = reversed([seg.size(1) for seg in global_encoded])
|
||||
|
||||
# Pick last 2000 tokens
|
||||
count = 0
|
||||
for i, length in enumerate(lengths):
|
||||
count += length
|
||||
if count + length > max_length - 1024 - sum(
|
||||
t.shape[1] for t in encoded_prompts
|
||||
):
|
||||
break
|
||||
|
||||
if i != 0 and i % 2 == 0:
|
||||
i -= 1
|
||||
|
||||
# Rotate the list, always make sure first segment is included to avoid drift
|
||||
if i < len(global_encoded) - 2:
|
||||
partial_encoded = global_encoded[:2] + global_encoded[-i:]
|
||||
else:
|
||||
partial_encoded = global_encoded
|
||||
|
||||
if use_prompt:
|
||||
partial_encoded = encoded_prompts + partial_encoded
|
||||
|
||||
cat_encoded = torch.cat(partial_encoded, dim=1)
|
||||
prompt_length = cat_encoded.size(1)
|
||||
|
||||
t0 = time.perf_counter()
|
||||
y = generate(
|
||||
model=model,
|
||||
prompt=cat_encoded,
|
||||
max_new_tokens=max_new_tokens,
|
||||
im_end_id=im_end_id,
|
||||
decode_one_token=decode_one_token,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
repetition_penalty=repetition_penalty,
|
||||
)
|
||||
|
||||
if sample_idx == 0 and seg_idx == 0 and compile:
|
||||
logger.info(f"Compilation time: {time.perf_counter() - t0:.2f} seconds")
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
t = time.perf_counter() - t0
|
||||
|
||||
tokens_generated = y.size(1) - prompt_length
|
||||
tokens_sec = tokens_generated / t
|
||||
logger.info(
|
||||
f"Generated {tokens_generated} tokens in {t:.02f} seconds, {tokens_sec:.02f} tokens/sec"
|
||||
)
|
||||
logger.info(
|
||||
f"Bandwidth achieved: {model_size * tokens_sec / 1e9:.02f} GB/s"
|
||||
)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
logger.info(
|
||||
f"GPU Memory used: {torch.cuda.max_memory_reserved() / 1e9:.02f} GB"
|
||||
)
|
||||
|
||||
# Put the generated tokens
|
||||
# since there is <im_end> and <eos> tokens, we remove last 2 tokens
|
||||
codes = y[1:, prompt_length:-1].clone()
|
||||
codes = codes - 1
|
||||
assert (codes >= 0).all(), f"Negative code found"
|
||||
|
||||
decoded = y[:, prompt_length:-1].clone()
|
||||
# But for global encoding, we should keep the <im_end> token
|
||||
|
||||
global_encoded.append(decoded)
|
||||
assert (codes >= 0).all(), f"Negative code found: {codes}"
|
||||
yield GenerateResponse(action="sample", codes=codes, text=texts[seg_idx])
|
||||
seg_idx += 1
|
||||
|
||||
# This indicates the end of the current sample
|
||||
yield GenerateResponse(action="next")
|
||||
|
||||
|
||||
@dataclass
|
||||
class WrappedGenerateResponse:
|
||||
status: Literal["success", "error"]
|
||||
response: Optional[GenerateResponse | Exception] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerateRequest:
|
||||
request: dict
|
||||
response_queue: queue.Queue
|
||||
|
||||
|
||||
def launch_thread_safe_queue(
|
||||
checkpoint_path,
|
||||
device,
|
||||
precision,
|
||||
compile: bool = False,
|
||||
):
|
||||
input_queue = queue.Queue()
|
||||
init_event = threading.Event()
|
||||
|
||||
def worker():
|
||||
model, decode_one_token = load_model(
|
||||
checkpoint_path, device, precision, compile=compile
|
||||
)
|
||||
init_event.set()
|
||||
|
||||
while True:
|
||||
item: GenerateRequest | None = input_queue.get()
|
||||
if item is None:
|
||||
break
|
||||
|
||||
kwargs = item.request
|
||||
response_queue = item.response_queue
|
||||
|
||||
try:
|
||||
for chunk in generate_long(
|
||||
model=model, decode_one_token=decode_one_token, **kwargs
|
||||
):
|
||||
response_queue.put(
|
||||
WrappedGenerateResponse(status="success", response=chunk)
|
||||
)
|
||||
except Exception as e:
|
||||
response_queue.put(WrappedGenerateResponse(status="error", response=e))
|
||||
|
||||
threading.Thread(target=worker, daemon=True).start()
|
||||
init_event.wait()
|
||||
|
||||
return input_queue
|
||||
|
||||
|
||||
@click.command()
|
||||
@click.option(
|
||||
"--text",
|
||||
type=str,
|
||||
default="你说的对, 但是原神是一款由米哈游自主研发的开放世界手游.",
|
||||
)
|
||||
@click.option("--prompt-text", type=str, default=None, multiple=True)
|
||||
@click.option(
|
||||
"--prompt-tokens",
|
||||
type=click.Path(path_type=Path, exists=True),
|
||||
default=None,
|
||||
multiple=True,
|
||||
)
|
||||
@click.option("--num-samples", type=int, default=1)
|
||||
@click.option("--max-new-tokens", type=int, default=0)
|
||||
@click.option("--top-p", type=float, default=0.7)
|
||||
@click.option("--repetition-penalty", type=float, default=1.2)
|
||||
@click.option("--temperature", type=float, default=0.7)
|
||||
@click.option(
|
||||
"--checkpoint-path",
|
||||
type=click.Path(path_type=Path, exists=True),
|
||||
default="checkpoints/fish-speech-1.2-sft",
|
||||
)
|
||||
@click.option("--device", type=str, default="cuda")
|
||||
@click.option("--compile/--no-compile", default=False)
|
||||
@click.option("--seed", type=int, default=42)
|
||||
@click.option("--half/--no-half", default=False)
|
||||
@click.option("--iterative-prompt/--no-iterative-prompt", default=True)
|
||||
@click.option("--chunk-length", type=int, default=100)
|
||||
def main(
|
||||
text: str,
|
||||
prompt_text: Optional[list[str]],
|
||||
prompt_tokens: Optional[list[Path]],
|
||||
num_samples: int,
|
||||
max_new_tokens: int,
|
||||
top_p: int,
|
||||
repetition_penalty: float,
|
||||
temperature: float,
|
||||
checkpoint_path: Path,
|
||||
device: str,
|
||||
compile: bool,
|
||||
seed: int,
|
||||
half: bool,
|
||||
iterative_prompt: bool,
|
||||
chunk_length: int,
|
||||
) -> None:
|
||||
|
||||
precision = torch.half if half else torch.bfloat16
|
||||
|
||||
if prompt_text is not None and len(prompt_text) != len(prompt_tokens):
|
||||
raise ValueError(
|
||||
f"Number of prompt text ({len(prompt_text)}) and prompt tokens ({len(prompt_tokens)}) should be the same"
|
||||
)
|
||||
|
||||
logger.info("Loading model ...")
|
||||
t0 = time.time()
|
||||
model, decode_one_token = load_model(
|
||||
checkpoint_path, device, precision, compile=compile
|
||||
)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
logger.info(f"Time to load model: {time.time() - t0:.02f} seconds")
|
||||
|
||||
if prompt_tokens is not None:
|
||||
prompt_tokens = [torch.from_numpy(np.load(p)).to(device) for p in prompt_tokens]
|
||||
|
||||
torch.manual_seed(seed)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
generator = generate_long(
|
||||
model=model,
|
||||
device=device,
|
||||
decode_one_token=decode_one_token,
|
||||
text=text,
|
||||
num_samples=num_samples,
|
||||
max_new_tokens=max_new_tokens,
|
||||
top_p=top_p,
|
||||
repetition_penalty=repetition_penalty,
|
||||
temperature=temperature,
|
||||
compile=compile,
|
||||
iterative_prompt=iterative_prompt,
|
||||
chunk_length=chunk_length,
|
||||
prompt_text=prompt_text,
|
||||
prompt_tokens=prompt_tokens,
|
||||
)
|
||||
|
||||
idx = 0
|
||||
codes = []
|
||||
|
||||
for response in generator:
|
||||
if response.action == "sample":
|
||||
codes.append(response.codes)
|
||||
logger.info(f"Sampled text: {response.text}")
|
||||
elif response.action == "next":
|
||||
if codes:
|
||||
np.save(f"codes_{idx}.npy", torch.cat(codes, dim=1).cpu().numpy())
|
||||
logger.info(f"Saved codes to codes_{idx}.npy")
|
||||
logger.info(f"Next sample")
|
||||
codes = []
|
||||
idx += 1
|
||||
else:
|
||||
logger.error(f"Error: {response}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,95 @@
|
||||
import shutil
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
|
||||
import click
|
||||
import hydra
|
||||
import torch
|
||||
from hydra import compose, initialize
|
||||
from hydra.utils import instantiate
|
||||
from loguru import logger
|
||||
|
||||
from fish_speech.models.text2semantic.llama import BaseTransformer
|
||||
from fish_speech.models.text2semantic.lora import get_merged_state_dict
|
||||
|
||||
|
||||
@click.command()
|
||||
@click.option("--lora-config", type=str, default="r_8_alpha_16")
|
||||
@click.option("--base-weight", type=str, default="checkpoints/fish-speech-1.4")
|
||||
@click.option("--lora-weight", type=str, required=True)
|
||||
@click.option("--output", type=str, required=True)
|
||||
def merge(lora_config, base_weight, lora_weight, output):
|
||||
output = Path(output)
|
||||
logger.info(
|
||||
f"Merging {base_weight} and {lora_weight} into {output} with {lora_config}"
|
||||
)
|
||||
|
||||
with initialize(version_base="1.3", config_path="../../fish_speech/configs/lora"):
|
||||
cfg = compose(config_name=lora_config)
|
||||
|
||||
lora_config = instantiate(cfg)
|
||||
logger.info(f"Loaded lora model with config {lora_config}")
|
||||
|
||||
llama_model = BaseTransformer.from_pretrained(
|
||||
path=base_weight,
|
||||
load_weights=True,
|
||||
lora_config=lora_config,
|
||||
)
|
||||
logger.info(f"Loaded llama model")
|
||||
|
||||
llama_state_dict = llama_model.state_dict()
|
||||
llama_state_dict = {k: v for k, v in llama_state_dict.items() if "lora" not in k}
|
||||
llama_state_dict_copy = deepcopy(llama_state_dict)
|
||||
lora_state_dict = torch.load(lora_weight, map_location="cpu")
|
||||
|
||||
if "state_dict" in llama_state_dict:
|
||||
llama_state_dict = llama_state_dict["state_dict"]
|
||||
|
||||
if "state_dict" in lora_state_dict:
|
||||
lora_state_dict = lora_state_dict["state_dict"]
|
||||
|
||||
# remove prefix model.
|
||||
if any(k.startswith("model.") for k in llama_state_dict.keys()):
|
||||
llama_state_dict = {
|
||||
k.replace("model.", ""): v
|
||||
for k, v in llama_state_dict.items()
|
||||
if k.startswith("model.")
|
||||
}
|
||||
if any(k.startswith("model.") for k in lora_state_dict.keys()):
|
||||
lora_state_dict = {
|
||||
k.replace("model.", ""): v
|
||||
for k, v in lora_state_dict.items()
|
||||
if k.startswith("model.")
|
||||
}
|
||||
|
||||
logger.info(f"Found {len(llama_state_dict)} keys in llama model")
|
||||
logger.info(f"Found {len(lora_state_dict)} keys in lora model")
|
||||
|
||||
merged_state_dict = llama_state_dict | lora_state_dict
|
||||
llama_model.load_state_dict(merged_state_dict, strict=True)
|
||||
logger.info(f"Merged model loaded")
|
||||
|
||||
# Trigger eval mode to merge lora
|
||||
llama_model.eval()
|
||||
llama_model.save_pretrained(output, drop_lora=True)
|
||||
logger.info(f"Saved merged model to {output}, validating")
|
||||
|
||||
new_state_dict = torch.load(output / "model.pth", map_location="cpu")
|
||||
original_keys = set(llama_state_dict_copy.keys())
|
||||
merged_keys = set(new_state_dict.keys())
|
||||
|
||||
assert original_keys == merged_keys, "Keys should be same"
|
||||
|
||||
for key in original_keys:
|
||||
diff_l1 = (new_state_dict[key] - llama_state_dict_copy[key]).abs().sum().item()
|
||||
if diff_l1 != 0:
|
||||
break
|
||||
else:
|
||||
logger.error("Merged model is same as the original model")
|
||||
exit(1)
|
||||
|
||||
logger.info("Merged model is different from the original model, check passed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
merge()
|
||||
@@ -0,0 +1,497 @@
|
||||
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
||||
# All rights reserved.
|
||||
import datetime
|
||||
import shutil
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import click
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fish_speech.models.text2semantic.llama import find_multiple
|
||||
from tools.llama.generate import load_model
|
||||
|
||||
##### Quantization Primitives ######
|
||||
|
||||
|
||||
def dynamically_quantize_per_channel(x, quant_min, quant_max, target_dtype):
|
||||
# assumes symmetric quantization
|
||||
# assumes axis == 0
|
||||
# assumes dense memory format
|
||||
# TODO(future): relax ^ as needed
|
||||
|
||||
# default setup for affine quantization of activations
|
||||
eps = torch.finfo(torch.float32).eps
|
||||
|
||||
# get min and max
|
||||
min_val, max_val = torch.aminmax(x, dim=1)
|
||||
|
||||
# calculate scales and zero_points based on min and max
|
||||
# reference: https://fburl.com/code/srbiybme
|
||||
min_val_neg = torch.min(min_val, torch.zeros_like(min_val))
|
||||
max_val_pos = torch.max(max_val, torch.zeros_like(max_val))
|
||||
device = min_val_neg.device
|
||||
|
||||
# reference: https://fburl.com/code/4wll53rk
|
||||
max_val_pos = torch.max(-min_val_neg, max_val_pos)
|
||||
scales = max_val_pos / (float(quant_max - quant_min) / 2)
|
||||
# ensure scales is the same dtype as the original tensor
|
||||
scales = torch.clamp(scales, min=eps).to(x.dtype)
|
||||
zero_points = torch.zeros(min_val_neg.size(), dtype=torch.int64, device=device)
|
||||
|
||||
# quantize based on qmin/qmax/scales/zp
|
||||
# reference: https://www.internalfb.com/code/fbsource/[8edc275012b1]/fbcode/caffe2/torch/ao/quantization/fx/_decomposed.py?lines=63
|
||||
x_div = x / scales.unsqueeze(-1)
|
||||
x_round = torch.round(x_div)
|
||||
x_zp = x_round + zero_points.unsqueeze(-1)
|
||||
quant = torch.clamp(x_zp, quant_min, quant_max).to(target_dtype)
|
||||
|
||||
return quant, scales, zero_points
|
||||
|
||||
|
||||
def get_group_qparams(w, n_bit=4, groupsize=128):
|
||||
# needed for GPTQ with padding
|
||||
if groupsize > w.shape[-1]:
|
||||
groupsize = w.shape[-1]
|
||||
assert groupsize > 1
|
||||
assert w.shape[-1] % groupsize == 0
|
||||
assert w.dim() == 2
|
||||
|
||||
to_quant = w.reshape(-1, groupsize)
|
||||
assert torch.isnan(to_quant).sum() == 0
|
||||
|
||||
max_val = to_quant.amax(dim=1, keepdim=True)
|
||||
min_val = to_quant.amin(dim=1, keepdim=True)
|
||||
max_int = 2**n_bit - 1
|
||||
scales = (max_val - min_val).clamp(min=1e-6) / max_int
|
||||
zeros = min_val + scales * (2 ** (n_bit - 1))
|
||||
return scales.to(torch.bfloat16).reshape(w.shape[0], -1), zeros.to(
|
||||
torch.bfloat16
|
||||
).reshape(w.shape[0], -1)
|
||||
|
||||
|
||||
def pack_scales_and_zeros(scales, zeros):
|
||||
assert scales.shape == zeros.shape
|
||||
assert scales.dtype == torch.bfloat16
|
||||
assert zeros.dtype == torch.bfloat16
|
||||
return (
|
||||
torch.cat(
|
||||
[
|
||||
scales.reshape(scales.size(0), scales.size(1), 1),
|
||||
zeros.reshape(zeros.size(0), zeros.size(1), 1),
|
||||
],
|
||||
2,
|
||||
)
|
||||
.transpose(0, 1)
|
||||
.contiguous()
|
||||
)
|
||||
|
||||
|
||||
def unpack_scales_and_zeros(scales_and_zeros):
|
||||
assert len(scales_and_zeros.shape) == 3 and scales_and_zeros.shape[2] == 2
|
||||
assert scales_and_zeros.dtype == torch.float
|
||||
return torch.split(scales_and_zeros.transpose(0, 1), 1, 2)
|
||||
|
||||
|
||||
def group_quantize_tensor_from_qparams(w, scales, zeros, n_bit=4, groupsize=128):
|
||||
assert groupsize > 1
|
||||
# needed for GPTQ single column quantize
|
||||
if groupsize > w.shape[-1] and scales.shape[-1] == 1:
|
||||
groupsize = w.shape[-1]
|
||||
|
||||
assert w.shape[-1] % groupsize == 0
|
||||
assert w.dim() == 2
|
||||
|
||||
to_quant = w.reshape(-1, groupsize)
|
||||
assert torch.isnan(to_quant).sum() == 0
|
||||
|
||||
scales = scales.reshape(-1, 1)
|
||||
zeros = zeros.reshape(-1, 1)
|
||||
min_val = zeros - scales * (2 ** (n_bit - 1))
|
||||
max_int = 2**n_bit - 1
|
||||
min_int = 0
|
||||
w_int32 = (
|
||||
to_quant.sub(min_val)
|
||||
.div(scales)
|
||||
.round()
|
||||
.clamp_(min_int, max_int)
|
||||
.to(torch.int32)
|
||||
.reshape_as(w)
|
||||
)
|
||||
|
||||
return w_int32
|
||||
|
||||
|
||||
def group_quantize_tensor(w, n_bit=4, groupsize=128):
|
||||
scales, zeros = get_group_qparams(w, n_bit, groupsize)
|
||||
w_int32 = group_quantize_tensor_from_qparams(w, scales, zeros, n_bit, groupsize)
|
||||
scales_and_zeros = pack_scales_and_zeros(scales, zeros)
|
||||
return w_int32, scales_and_zeros
|
||||
|
||||
|
||||
def group_dequantize_tensor_from_qparams(
|
||||
w_int32, scales, zeros, n_bit=4, groupsize=128
|
||||
):
|
||||
assert groupsize > 1
|
||||
# needed for GPTQ single column dequantize
|
||||
if groupsize > w_int32.shape[-1] and scales.shape[-1] == 1:
|
||||
groupsize = w_int32.shape[-1]
|
||||
assert w_int32.shape[-1] % groupsize == 0
|
||||
assert w_int32.dim() == 2
|
||||
|
||||
w_int32_grouped = w_int32.reshape(-1, groupsize)
|
||||
scales = scales.reshape(-1, 1)
|
||||
zeros = zeros.reshape(-1, 1)
|
||||
|
||||
w_dq = (
|
||||
w_int32_grouped.sub(2 ** (n_bit - 1)).mul(scales).add(zeros).reshape_as(w_int32)
|
||||
)
|
||||
return w_dq
|
||||
|
||||
|
||||
def group_dequantize_tensor(w_int32, scales_and_zeros, n_bit=4, groupsize=128):
|
||||
scales, zeros = unpack_scales_and_zeros(scales_and_zeros)
|
||||
return group_dequantize_tensor_from_qparams(
|
||||
w_int32, scales, zeros, n_bit, groupsize
|
||||
)
|
||||
|
||||
|
||||
class QuantHandler:
|
||||
def __init__(self, mod):
|
||||
self.mod = mod
|
||||
|
||||
def create_quantized_state_dict(self) -> "StateDict":
|
||||
pass
|
||||
|
||||
def convert_for_runtime(self) -> "nn.Module":
|
||||
pass
|
||||
|
||||
|
||||
##### Weight-only int8 per-channel quantized code ######
|
||||
|
||||
|
||||
def replace_linear_weight_only_int8_per_channel(module):
|
||||
for name, child in module.named_children():
|
||||
if isinstance(child, nn.Linear):
|
||||
setattr(
|
||||
module,
|
||||
name,
|
||||
WeightOnlyInt8Linear(child.in_features, child.out_features),
|
||||
)
|
||||
else:
|
||||
replace_linear_weight_only_int8_per_channel(child)
|
||||
|
||||
|
||||
class WeightOnlyInt8QuantHandler:
|
||||
def __init__(self, mod):
|
||||
self.mod = mod
|
||||
|
||||
@torch.no_grad()
|
||||
def create_quantized_state_dict(self):
|
||||
cur_state_dict = self.mod.state_dict()
|
||||
for fqn, mod in self.mod.named_modules():
|
||||
if isinstance(mod, torch.nn.Linear):
|
||||
int8_weight, scales, _ = dynamically_quantize_per_channel(
|
||||
mod.weight.float(), -128, 127, torch.int8
|
||||
)
|
||||
cur_state_dict[f"{fqn}.weight"] = int8_weight
|
||||
cur_state_dict[f"{fqn}.scales"] = scales.to(mod.weight.dtype)
|
||||
|
||||
return cur_state_dict
|
||||
|
||||
def convert_for_runtime(self):
|
||||
replace_linear_weight_only_int8_per_channel(self.mod)
|
||||
return self.mod
|
||||
|
||||
|
||||
class WeightOnlyInt8Linear(torch.nn.Module):
|
||||
__constants__ = ["in_features", "out_features"]
|
||||
in_features: int
|
||||
out_features: int
|
||||
weight: torch.Tensor
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
bias: bool = True,
|
||||
device=None,
|
||||
dtype=None,
|
||||
) -> None:
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
self.register_buffer(
|
||||
"weight", torch.empty((out_features, in_features), dtype=torch.int8)
|
||||
)
|
||||
self.register_buffer("scales", torch.ones(out_features, dtype=torch.bfloat16))
|
||||
|
||||
def forward(self, input: torch.Tensor) -> torch.Tensor:
|
||||
return F.linear(input, self.weight.to(dtype=input.dtype)) * self.scales
|
||||
|
||||
|
||||
##### weight only int4 per channel groupwise quantized code ######
|
||||
|
||||
|
||||
def prepare_int4_weight_and_scales_and_zeros(weight_bf16, groupsize, inner_k_tiles):
|
||||
weight_int32, scales_and_zeros = group_quantize_tensor(
|
||||
weight_bf16, n_bit=4, groupsize=groupsize
|
||||
)
|
||||
weight_int4pack = torch.ops.aten._convert_weight_to_int4pack(
|
||||
weight_int32, inner_k_tiles
|
||||
)
|
||||
return weight_int4pack, scales_and_zeros
|
||||
|
||||
|
||||
def linear_forward_int4(x, weight_int4pack, scales_and_zeros, out_features, groupsize):
|
||||
origin_x_size = x.size()
|
||||
x = x.reshape(-1, origin_x_size[-1])
|
||||
c = torch.ops.aten._weight_int4pack_mm(
|
||||
x, weight_int4pack, groupsize, scales_and_zeros
|
||||
)
|
||||
new_shape = origin_x_size[:-1] + (out_features,)
|
||||
c = c.reshape(new_shape)
|
||||
return c
|
||||
|
||||
|
||||
def _check_linear_int4_k(k, groupsize=1, inner_k_tiles=1):
|
||||
return k % groupsize == 0 and k % (inner_k_tiles * 16) == 0
|
||||
|
||||
|
||||
def replace_linear_int4(module, groupsize, inner_k_tiles, padding):
|
||||
for name, child in module.named_children():
|
||||
if isinstance(child, nn.Linear):
|
||||
if _check_linear_int4_k(child.in_features, groupsize, inner_k_tiles):
|
||||
setattr(
|
||||
module,
|
||||
name,
|
||||
WeightOnlyInt4Linear(
|
||||
child.in_features,
|
||||
child.out_features,
|
||||
bias=False,
|
||||
groupsize=groupsize,
|
||||
inner_k_tiles=inner_k_tiles,
|
||||
padding=False,
|
||||
),
|
||||
)
|
||||
elif padding:
|
||||
setattr(
|
||||
module,
|
||||
name,
|
||||
WeightOnlyInt4Linear(
|
||||
child.in_features,
|
||||
child.out_features,
|
||||
bias=False,
|
||||
groupsize=groupsize,
|
||||
inner_k_tiles=inner_k_tiles,
|
||||
padding=True,
|
||||
),
|
||||
)
|
||||
else:
|
||||
replace_linear_int4(child, groupsize, inner_k_tiles, padding)
|
||||
|
||||
|
||||
class WeightOnlyInt4QuantHandler:
|
||||
def __init__(self, mod, groupsize=128, inner_k_tiles=8, padding=True):
|
||||
self.mod = mod
|
||||
self.groupsize = groupsize
|
||||
self.inner_k_tiles = inner_k_tiles
|
||||
self.padding = padding
|
||||
assert groupsize in [32, 64, 128, 256]
|
||||
assert inner_k_tiles in [2, 4, 8]
|
||||
|
||||
@torch.no_grad()
|
||||
def create_quantized_state_dict(self):
|
||||
cur_state_dict = self.mod.state_dict()
|
||||
for fqn, mod in self.mod.named_modules():
|
||||
if isinstance(mod, torch.nn.Linear):
|
||||
assert not mod.bias
|
||||
out_features = mod.out_features
|
||||
in_features = mod.in_features
|
||||
assert out_features % 8 == 0, "require out_features % 8 == 0"
|
||||
print(f"linear: {fqn}, in={in_features}, out={out_features}")
|
||||
|
||||
weight = mod.weight.data
|
||||
if not _check_linear_int4_k(
|
||||
in_features, self.groupsize, self.inner_k_tiles
|
||||
):
|
||||
if self.padding:
|
||||
import torch.nn.functional as F
|
||||
|
||||
print(
|
||||
f"warning: {fqn} is padded to satisfy in_features % 1024 == 0"
|
||||
)
|
||||
padded_in_features = find_multiple(in_features, 1024)
|
||||
weight = F.pad(
|
||||
weight, pad=(0, padded_in_features - in_features)
|
||||
)
|
||||
else:
|
||||
print(
|
||||
f"warning: {fqn} is skipped, int4 requires that in_features is 32, 64, or is divisible by 1024, "
|
||||
+ "and that groupsize and inner_k_tiles*16 evenly divide into it"
|
||||
)
|
||||
continue
|
||||
(
|
||||
weight_int4pack,
|
||||
scales_and_zeros,
|
||||
) = prepare_int4_weight_and_scales_and_zeros(
|
||||
weight.to(torch.bfloat16).to("cuda"),
|
||||
self.groupsize,
|
||||
self.inner_k_tiles,
|
||||
)
|
||||
cur_state_dict[f"{fqn}.weight"] = weight_int4pack.to("cpu")
|
||||
cur_state_dict[f"{fqn}.scales_and_zeros"] = scales_and_zeros.to("cpu")
|
||||
|
||||
return cur_state_dict
|
||||
|
||||
def convert_for_runtime(self):
|
||||
replace_linear_int4(self.mod, self.groupsize, self.inner_k_tiles, self.padding)
|
||||
return self.mod
|
||||
|
||||
|
||||
class WeightOnlyInt4Linear(torch.nn.Module):
|
||||
__constants__ = ["in_features", "out_features"]
|
||||
in_features: int
|
||||
out_features: int
|
||||
weight: torch.Tensor
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
bias=True,
|
||||
device=None,
|
||||
dtype=None,
|
||||
groupsize: int = 128,
|
||||
inner_k_tiles: int = 8,
|
||||
padding: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.padding = padding
|
||||
if padding:
|
||||
self.origin_in_features = in_features
|
||||
in_features = find_multiple(in_features, 1024)
|
||||
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
assert not bias, "require bias=False"
|
||||
self.groupsize = groupsize
|
||||
self.inner_k_tiles = inner_k_tiles
|
||||
|
||||
assert out_features % 8 == 0, "require out_features % 8 == 0"
|
||||
assert (
|
||||
in_features % (inner_k_tiles * 16) == 0
|
||||
), "require in_features % (innerKTiles * 16) == 0"
|
||||
self.register_buffer(
|
||||
"weight",
|
||||
torch.empty(
|
||||
(
|
||||
out_features // 8,
|
||||
in_features // (inner_k_tiles * 16),
|
||||
32,
|
||||
inner_k_tiles // 2,
|
||||
),
|
||||
dtype=torch.int32,
|
||||
),
|
||||
)
|
||||
self.register_buffer(
|
||||
"scales_and_zeros",
|
||||
torch.empty(
|
||||
(in_features // groupsize, out_features, 2), dtype=torch.bfloat16
|
||||
),
|
||||
)
|
||||
|
||||
def forward(self, input: torch.Tensor) -> torch.Tensor:
|
||||
input = input.to(torch.bfloat16)
|
||||
if self.padding:
|
||||
import torch.nn.functional as F
|
||||
|
||||
input = F.pad(input, pad=(0, self.in_features - self.origin_in_features))
|
||||
return linear_forward_int4(
|
||||
input, self.weight, self.scales_and_zeros, self.out_features, self.groupsize
|
||||
)
|
||||
|
||||
|
||||
def generate_folder_name():
|
||||
now = datetime.datetime.now()
|
||||
folder_name = now.strftime("%Y%m%d_%H%M%S")
|
||||
return folder_name
|
||||
|
||||
|
||||
@click.command()
|
||||
@click.option(
|
||||
"--checkpoint-path",
|
||||
type=click.Path(path_type=Path, exists=True),
|
||||
default="checkpoints/fish-speech-1.4",
|
||||
)
|
||||
@click.option(
|
||||
"--mode", type=str, default="int8", help="type of quantization to perform"
|
||||
)
|
||||
@click.option(
|
||||
"--groupsize", type=int, default=128, help="Group size for int4 quantization."
|
||||
)
|
||||
@click.option("--timestamp", type=str, default="None", help="When to do quantization")
|
||||
def quantize(checkpoint_path: Path, mode: str, groupsize: int, timestamp: str) -> None:
|
||||
|
||||
device = "cpu"
|
||||
precision = torch.bfloat16
|
||||
|
||||
print("Loading model ...")
|
||||
t0 = time.time()
|
||||
|
||||
model, _ = load_model(
|
||||
checkpoint_path=checkpoint_path,
|
||||
device=device,
|
||||
precision=precision,
|
||||
compile=False,
|
||||
)
|
||||
vq_model = "firefly-gan-vq-fsq-8x1024-21hz-generator.pth"
|
||||
now = timestamp if timestamp != "None" else generate_folder_name()
|
||||
|
||||
if mode == "int8":
|
||||
print(
|
||||
"Quantizing model weights for int8 weight-only symmetric per-channel quantization"
|
||||
)
|
||||
quant_handler = WeightOnlyInt8QuantHandler(model)
|
||||
quantized_state_dict = quant_handler.create_quantized_state_dict()
|
||||
|
||||
dir_name = checkpoint_path
|
||||
dst_name = Path(f"checkpoints/fs-1.2-int8-{now}")
|
||||
shutil.copytree(str(dir_name.resolve()), str(dst_name.resolve()))
|
||||
if (dst_name / vq_model).exists():
|
||||
(dst_name / vq_model).unlink()
|
||||
quantize_path = dst_name / "model.pth"
|
||||
|
||||
elif mode == "int4":
|
||||
print(
|
||||
"Quantizing model weights for int4 weight-only affine per-channel groupwise quantization"
|
||||
)
|
||||
quant_handler = WeightOnlyInt4QuantHandler(model, groupsize)
|
||||
quantized_state_dict = quant_handler.create_quantized_state_dict()
|
||||
|
||||
dir_name = checkpoint_path
|
||||
dst_name = Path(f"checkpoints/fs-1.2-int4-g{groupsize}-{now}")
|
||||
shutil.copytree(str(dir_name.resolve()), str(dst_name.resolve()))
|
||||
if (dst_name / vq_model).exists():
|
||||
(dst_name / vq_model).unlink()
|
||||
quantize_path = dst_name / "model.pth"
|
||||
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid quantization mode {mode} needs to be one of [int8, int4, int4-gpptq]"
|
||||
)
|
||||
|
||||
print(f"Writing quantized weights to {quantize_path}")
|
||||
quantize_path.unlink(missing_ok=True) # remove existing file if one already there
|
||||
torch.save(quantized_state_dict, quantize_path)
|
||||
print(f"Quantization complete took {time.time() - t0:.02f} seconds")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
quantize()
|
||||
@@ -0,0 +1,57 @@
|
||||
from tokenizers import Tokenizer, decoders, models, pre_tokenizers, processors, trainers
|
||||
from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast
|
||||
|
||||
# Initialize a tokenizer
|
||||
tokenizer = Tokenizer(models.BPE())
|
||||
|
||||
# Customize pre-tokenization and decoding
|
||||
tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
|
||||
tokenizer.decoder = decoders.ByteLevel()
|
||||
tokenizer.post_processor = processors.ByteLevel(trim_offsets=False)
|
||||
|
||||
# Don't train the tokenizer
|
||||
trainer = trainers.BpeTrainer(
|
||||
vocab_size=0,
|
||||
min_frequency=2,
|
||||
initial_alphabet=pre_tokenizers.ByteLevel.alphabet(),
|
||||
special_tokens=[
|
||||
"<|begin_of_sequence|>",
|
||||
"<|end_of_sequence|>",
|
||||
"<|im_start|>",
|
||||
"<|im_sep|>", # system, user, assistant, etc.
|
||||
"<|im_end|>",
|
||||
"<|semantic|>", # audio features
|
||||
"<|pad|>",
|
||||
],
|
||||
)
|
||||
|
||||
# <|im_start|>user<|im_sep|>...<|im_end|>
|
||||
# <|im_start|>assistant<|im_sep|><|semantic|><|semantic|><|semantic|><|semantic|><|semantic|><|im_end|>
|
||||
tokenizer.train_from_iterator([], trainer=trainer)
|
||||
|
||||
print(len(tokenizer.get_vocab()))
|
||||
x = tokenizer.encode(
|
||||
"Hello, how are you? dfgnviadfjoiviouajeiodfjv 你好世界 🈶<|semantic|>"
|
||||
).ids
|
||||
print(x, len(x))
|
||||
print(tokenizer.decode(x, skip_special_tokens=True))
|
||||
|
||||
|
||||
tokenizer = PreTrainedTokenizerFast(
|
||||
tokenizer_object=tokenizer,
|
||||
pad_token="<|pad|>",
|
||||
bos_token="<|begin_of_sequence|>",
|
||||
eos_token="<|end_of_sequence|>",
|
||||
)
|
||||
|
||||
# Try tokenizing a new sequence
|
||||
sequence = "All around, too, lay vast quantities of the costliest merchandise, and treasures were heaped in every cranny of the rocks, but all these things only added to the desolation of the scene. 测试中文, 你好世界 🈶<|semantic|>"
|
||||
encoded = tokenizer(sequence).input_ids
|
||||
|
||||
print("Test encoding....")
|
||||
print(f"\tSentence: {sequence}")
|
||||
print(f"\tEncoded: {encoded}")
|
||||
print(f"\tDecoded: {tokenizer.batch_decode(encoded)}")
|
||||
print(f"\tDecoded: {tokenizer.decode(encoded)}")
|
||||
|
||||
tokenizer.push_to_hub("fishaudio/fish-speech-1", private=True)
|
||||
@@ -0,0 +1,23 @@
|
||||
from .braceexpand import braceexpand
|
||||
from .context import autocast_exclude_mps
|
||||
from .file import get_latest_checkpoint
|
||||
from .instantiators import instantiate_callbacks, instantiate_loggers
|
||||
from .logger import RankedLogger
|
||||
# from .logging_utils import log_hyperparameters
|
||||
from .rich_utils import enforce_tags, print_config_tree
|
||||
from .utils import extras, get_metric_value, task_wrapper
|
||||
|
||||
__all__ = [
|
||||
"enforce_tags",
|
||||
"extras",
|
||||
"get_metric_value",
|
||||
"RankedLogger",
|
||||
"instantiate_callbacks",
|
||||
"instantiate_loggers",
|
||||
# "log_hyperparameters",
|
||||
"print_config_tree",
|
||||
"task_wrapper",
|
||||
"braceexpand",
|
||||
"get_latest_checkpoint",
|
||||
"autocast_exclude_mps",
|
||||
]
|
||||
@@ -0,0 +1,217 @@
|
||||
"""
|
||||
Bash-style brace expansion
|
||||
Copied from: https://github.com/trendels/braceexpand/blob/main/src/braceexpand/__init__.py
|
||||
License: MIT
|
||||
"""
|
||||
|
||||
import re
|
||||
import string
|
||||
from itertools import chain, product
|
||||
from typing import Iterable, Iterator, Optional
|
||||
|
||||
__all__ = ["braceexpand", "alphabet", "UnbalancedBracesError"]
|
||||
|
||||
|
||||
class UnbalancedBracesError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
alphabet = string.ascii_uppercase + string.ascii_lowercase
|
||||
|
||||
int_range_re = re.compile(r"^(-?\d+)\.\.(-?\d+)(?:\.\.-?(\d+))?$")
|
||||
char_range_re = re.compile(r"^([A-Za-z])\.\.([A-Za-z])(?:\.\.-?(\d+))?$")
|
||||
escape_re = re.compile(r"\\(.)")
|
||||
|
||||
|
||||
def braceexpand(pattern: str, escape: bool = True) -> Iterator[str]:
|
||||
"""braceexpand(pattern) -> iterator over generated strings
|
||||
|
||||
Returns an iterator over the strings resulting from brace expansion
|
||||
of pattern. This function implements Brace Expansion as described in
|
||||
bash(1), with the following limitations:
|
||||
|
||||
* A pattern containing unbalanced braces will raise an
|
||||
UnbalancedBracesError exception. In bash, unbalanced braces will either
|
||||
be partly expanded or ignored.
|
||||
|
||||
* A mixed-case character range like '{Z..a}' or '{a..Z}' will not
|
||||
include the characters '[]^_`' between 'Z' and 'a'.
|
||||
|
||||
When escape is True (the default), characters in pattern can be
|
||||
prefixed with a backslash to cause them not to be interpreted as
|
||||
special characters for brace expansion (such as '{', '}', ',').
|
||||
To pass through a a literal backslash, double it ('\\\\').
|
||||
|
||||
When escape is False, backslashes in pattern have no special
|
||||
meaning and will be preserved in the output.
|
||||
|
||||
Examples:
|
||||
|
||||
>>> from braceexpand import braceexpand
|
||||
|
||||
# Integer range
|
||||
>>> list(braceexpand('item{1..3}'))
|
||||
['item1', 'item2', 'item3']
|
||||
|
||||
# Character range
|
||||
>>> list(braceexpand('{a..c}'))
|
||||
['a', 'b', 'c']
|
||||
|
||||
# Sequence
|
||||
>>> list(braceexpand('index.html{,.backup}'))
|
||||
['index.html', 'index.html.backup']
|
||||
|
||||
# Nested patterns
|
||||
>>> list(braceexpand('python{2.{5..7},3.{2,3}}'))
|
||||
['python2.5', 'python2.6', 'python2.7', 'python3.2', 'python3.3']
|
||||
|
||||
# Prefixing an integer with zero causes all numbers to be padded to
|
||||
# the same width.
|
||||
>>> list(braceexpand('{07..10}'))
|
||||
['07', '08', '09', '10']
|
||||
|
||||
# An optional increment can be specified for ranges.
|
||||
>>> list(braceexpand('{a..g..2}'))
|
||||
['a', 'c', 'e', 'g']
|
||||
|
||||
# Ranges can go in both directions.
|
||||
>>> list(braceexpand('{4..1}'))
|
||||
['4', '3', '2', '1']
|
||||
|
||||
# Numbers can be negative
|
||||
>>> list(braceexpand('{2..-1}'))
|
||||
['2', '1', '0', '-1']
|
||||
|
||||
# Unbalanced braces raise an exception.
|
||||
>>> list(braceexpand('{1{2,3}'))
|
||||
Traceback (most recent call last):
|
||||
...
|
||||
UnbalancedBracesError: Unbalanced braces: '{1{2,3}'
|
||||
|
||||
# By default, the backslash is the escape character.
|
||||
>>> list(braceexpand(r'{1\\{2,3}'))
|
||||
['1{2', '3']
|
||||
|
||||
# Setting 'escape' to False disables backslash escaping.
|
||||
>>> list(braceexpand(r'\\{1,2}', escape=False))
|
||||
['\\\\1', '\\\\2']
|
||||
|
||||
"""
|
||||
return (
|
||||
escape_re.sub(r"\1", s) if escape else s for s in parse_pattern(pattern, escape)
|
||||
)
|
||||
|
||||
|
||||
def parse_pattern(pattern: str, escape: bool) -> Iterator[str]:
|
||||
start = 0
|
||||
pos = 0
|
||||
bracketdepth = 0
|
||||
items: list[Iterable[str]] = []
|
||||
|
||||
# print 'pattern:', pattern
|
||||
while pos < len(pattern):
|
||||
if escape and pattern[pos] == "\\":
|
||||
pos += 2
|
||||
continue
|
||||
elif pattern[pos] == "{":
|
||||
if bracketdepth == 0 and pos > start:
|
||||
# print 'literal:', pattern[start:pos]
|
||||
items.append([pattern[start:pos]])
|
||||
start = pos
|
||||
bracketdepth += 1
|
||||
elif pattern[pos] == "}":
|
||||
bracketdepth -= 1
|
||||
if bracketdepth == 0:
|
||||
# print 'expression:', pattern[start+1:pos]
|
||||
expr = pattern[start + 1 : pos]
|
||||
item = parse_expression(expr, escape)
|
||||
if item is None: # not a range or sequence
|
||||
items.extend([["{"], parse_pattern(expr, escape), ["}"]])
|
||||
else:
|
||||
items.append(item)
|
||||
start = pos + 1 # skip the closing brace
|
||||
pos += 1
|
||||
|
||||
if bracketdepth != 0: # unbalanced braces
|
||||
raise UnbalancedBracesError("Unbalanced braces: '%s'" % pattern)
|
||||
|
||||
if start < pos:
|
||||
items.append([pattern[start:]])
|
||||
|
||||
return ("".join(item) for item in product(*items))
|
||||
|
||||
|
||||
def parse_expression(expr: str, escape: bool) -> Optional[Iterable[str]]:
|
||||
int_range_match = int_range_re.match(expr)
|
||||
if int_range_match:
|
||||
return make_int_range(*int_range_match.groups())
|
||||
|
||||
char_range_match = char_range_re.match(expr)
|
||||
if char_range_match:
|
||||
return make_char_range(*char_range_match.groups())
|
||||
|
||||
return parse_sequence(expr, escape)
|
||||
|
||||
|
||||
def parse_sequence(seq: str, escape: bool) -> Optional[Iterator[str]]:
|
||||
# sequence -> chain(*sequence_items)
|
||||
start = 0
|
||||
pos = 0
|
||||
bracketdepth = 0
|
||||
items: list[Iterable[str]] = []
|
||||
|
||||
# print 'sequence:', seq
|
||||
while pos < len(seq):
|
||||
if escape and seq[pos] == "\\":
|
||||
pos += 2
|
||||
continue
|
||||
elif seq[pos] == "{":
|
||||
bracketdepth += 1
|
||||
elif seq[pos] == "}":
|
||||
bracketdepth -= 1
|
||||
elif seq[pos] == "," and bracketdepth == 0:
|
||||
items.append(parse_pattern(seq[start:pos], escape))
|
||||
start = pos + 1 # skip the comma
|
||||
pos += 1
|
||||
|
||||
if bracketdepth != 0:
|
||||
raise UnbalancedBracesError
|
||||
if not items:
|
||||
return None
|
||||
|
||||
# part after the last comma (may be the empty string)
|
||||
items.append(parse_pattern(seq[start:], escape))
|
||||
return chain(*items)
|
||||
|
||||
|
||||
def make_int_range(left: str, right: str, incr: Optional[str] = None) -> Iterator[str]:
|
||||
if any([s.startswith(("0", "-0")) for s in (left, right) if s not in ("0", "-0")]):
|
||||
padding = max(len(left), len(right))
|
||||
else:
|
||||
padding = 0
|
||||
step = (int(incr) or 1) if incr else 1
|
||||
start = int(left)
|
||||
end = int(right)
|
||||
r = range(start, end + 1, step) if start < end else range(start, end - 1, -step)
|
||||
fmt = "%0{}d".format(padding)
|
||||
return (fmt % i for i in r)
|
||||
|
||||
|
||||
def make_char_range(left: str, right: str, incr: Optional[str] = None) -> str:
|
||||
step = (int(incr) or 1) if incr else 1
|
||||
start = alphabet.index(left)
|
||||
end = alphabet.index(right)
|
||||
if start < end:
|
||||
return alphabet[start : end + 1 : step]
|
||||
else:
|
||||
end = end or -len(alphabet)
|
||||
return alphabet[start : end - 1 : -step]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import doctest
|
||||
import sys
|
||||
|
||||
failed, _ = doctest.testmod(optionflags=doctest.IGNORE_EXCEPTION_DETAIL)
|
||||
if failed:
|
||||
sys.exit(1)
|
||||
@@ -0,0 +1,13 @@
|
||||
from contextlib import nullcontext
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def autocast_exclude_mps(
|
||||
device_type: str, dtype: torch.dtype
|
||||
) -> nullcontext | torch.autocast:
|
||||
return (
|
||||
nullcontext()
|
||||
if torch.backends.mps.is_available()
|
||||
else torch.autocast(device_type, dtype)
|
||||
)
|
||||
@@ -0,0 +1,16 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def get_latest_checkpoint(path: Path | str) -> Path | None:
|
||||
# Find the latest checkpoint
|
||||
ckpt_dir = Path(path)
|
||||
|
||||
if ckpt_dir.exists() is False:
|
||||
return None
|
||||
|
||||
ckpts = sorted(ckpt_dir.glob("*.ckpt"), key=os.path.getmtime)
|
||||
if len(ckpts) == 0:
|
||||
return None
|
||||
|
||||
return ckpts[-1]
|
||||
@@ -0,0 +1,50 @@
|
||||
from typing import List
|
||||
|
||||
import hydra
|
||||
from omegaconf import DictConfig
|
||||
# from pytorch_lightning import Callback
|
||||
# from pytorch_lightning.loggers import Logger
|
||||
|
||||
from .logger import RankedLogger
|
||||
|
||||
log = RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
def instantiate_callbacks(callbacks_cfg ) :
|
||||
"""Instantiates callbacks from config."""
|
||||
|
||||
callbacks = []
|
||||
|
||||
if not callbacks_cfg:
|
||||
log.warning("No callback configs found! Skipping..")
|
||||
return callbacks
|
||||
|
||||
if not isinstance(callbacks_cfg, DictConfig):
|
||||
raise TypeError("Callbacks config must be a DictConfig!")
|
||||
|
||||
for _, cb_conf in callbacks_cfg.items():
|
||||
if isinstance(cb_conf, DictConfig) and "_target_" in cb_conf:
|
||||
log.info(f"Instantiating callback <{cb_conf._target_}>")
|
||||
callbacks.append(hydra.utils.instantiate(cb_conf))
|
||||
|
||||
return callbacks
|
||||
|
||||
|
||||
def instantiate_loggers(logger_cfg ) :
|
||||
"""Instantiates loggers from config."""
|
||||
|
||||
logger = []
|
||||
|
||||
if not logger_cfg:
|
||||
log.warning("No logger configs found! Skipping...")
|
||||
return logger
|
||||
|
||||
if not isinstance(logger_cfg, DictConfig):
|
||||
raise TypeError("Logger config must be a DictConfig!")
|
||||
|
||||
for _, lg_conf in logger_cfg.items():
|
||||
if isinstance(lg_conf, DictConfig) and "_target_" in lg_conf:
|
||||
log.info(f"Instantiating logger <{lg_conf._target_}>")
|
||||
logger.append(hydra.utils.instantiate(lg_conf))
|
||||
|
||||
return logger
|
||||
@@ -0,0 +1,56 @@
|
||||
import logging
|
||||
from typing import Mapping, Optional
|
||||
|
||||
# from lightning_utilities.core.rank_zero import rank_prefixed_message, rank_zero_only
|
||||
|
||||
|
||||
class RankedLogger(logging.LoggerAdapter):
|
||||
"""A multi-GPU-friendly python command line logger."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
name: str = __name__,
|
||||
rank_zero_only: bool = True,
|
||||
extra: Optional[Mapping[str, object]] = None,
|
||||
) -> None:
|
||||
"""Initializes a multi-GPU-friendly python command line logger that logs on all processes
|
||||
with their rank prefixed in the log message.
|
||||
|
||||
:param name: The name of the logger. Default is ``__name__``.
|
||||
:param rank_zero_only: Whether to force all logs to only occur on the rank zero process. Default is `False`.
|
||||
:param extra: (Optional) A dict-like object which provides contextual information. See `logging.LoggerAdapter`.
|
||||
"""
|
||||
logger = logging.getLogger(name)
|
||||
super().__init__(logger=logger, extra=extra)
|
||||
self.rank_zero_only = rank_zero_only
|
||||
|
||||
def log(
|
||||
self, level: int, msg: str, rank: Optional[int] = None, *args, **kwargs
|
||||
) -> None:
|
||||
"""Delegate a log call to the underlying logger, after prefixing its message with the rank
|
||||
of the process it's being logged from. If `'rank'` is provided, then the log will only
|
||||
occur on that rank/process.
|
||||
|
||||
:param level: The level to log at. Look at `logging.__init__.py` for more information.
|
||||
:param msg: The message to log.
|
||||
:param rank: The rank to log at.
|
||||
:param args: Additional args to pass to the underlying logging function.
|
||||
:param kwargs: Any additional keyword args to pass to the underlying logging function.
|
||||
"""
|
||||
self.logger.log(level, msg, *args, **kwargs)
|
||||
# if self.isEnabledFor(level):
|
||||
# msg, kwargs = self.process(msg, kwargs)
|
||||
# current_rank = getattr(rank_zero_only, "rank", None)
|
||||
# if current_rank is None:
|
||||
# raise RuntimeError(
|
||||
# "The `rank_zero_only.rank` needs to be set before use"
|
||||
# )
|
||||
# msg = rank_prefixed_message(msg, current_rank)
|
||||
# if self.rank_zero_only:
|
||||
# if current_rank == 0:
|
||||
# self.logger.log(level, msg, *args, **kwargs)
|
||||
# else:
|
||||
# if rank is None:
|
||||
# self.logger.log(level, msg, *args, **kwargs)
|
||||
# elif current_rank == rank:
|
||||
# self.logger.log(level, msg, *args, **kwargs)
|
||||
@@ -0,0 +1,48 @@
|
||||
from lightning.pytorch.utilities import rank_zero_only
|
||||
|
||||
from fish_speech.utils import logger as log
|
||||
|
||||
|
||||
@rank_zero_only
|
||||
def log_hyperparameters(object_dict: dict) -> None:
|
||||
"""Controls which config parts are saved by lightning loggers.
|
||||
|
||||
Additionally saves:
|
||||
- Number of model parameters
|
||||
"""
|
||||
|
||||
hparams = {}
|
||||
|
||||
cfg = object_dict["cfg"]
|
||||
model = object_dict["model"]
|
||||
trainer = object_dict["trainer"]
|
||||
|
||||
if not trainer.logger:
|
||||
log.warning("Logger not found! Skipping hyperparameter logging...")
|
||||
return
|
||||
|
||||
hparams["model"] = cfg["model"]
|
||||
|
||||
# save number of model parameters
|
||||
hparams["model/params/total"] = sum(p.numel() for p in model.parameters())
|
||||
hparams["model/params/trainable"] = sum(
|
||||
p.numel() for p in model.parameters() if p.requires_grad
|
||||
)
|
||||
hparams["model/params/non_trainable"] = sum(
|
||||
p.numel() for p in model.parameters() if not p.requires_grad
|
||||
)
|
||||
|
||||
hparams["data"] = cfg["data"]
|
||||
hparams["trainer"] = cfg["trainer"]
|
||||
|
||||
hparams["callbacks"] = cfg.get("callbacks")
|
||||
hparams["extras"] = cfg.get("extras")
|
||||
|
||||
hparams["task_name"] = cfg.get("task_name")
|
||||
hparams["tags"] = cfg.get("tags")
|
||||
hparams["ckpt_path"] = cfg.get("ckpt_path")
|
||||
hparams["seed"] = cfg.get("seed")
|
||||
|
||||
# send hparams to all loggers
|
||||
for logger in trainer.loggers:
|
||||
logger.log_hyperparams(hparams)
|
||||
@@ -0,0 +1,100 @@
|
||||
from pathlib import Path
|
||||
from typing import Sequence
|
||||
|
||||
import rich
|
||||
import rich.syntax
|
||||
import rich.tree
|
||||
from hydra.core.hydra_config import HydraConfig
|
||||
# from lightning.pytorch.utilities import rank_zero_only
|
||||
from omegaconf import DictConfig, OmegaConf, open_dict
|
||||
from rich.prompt import Prompt
|
||||
|
||||
from fish_speech.utils import logger as log
|
||||
|
||||
|
||||
|
||||
def print_config_tree(
|
||||
cfg: DictConfig,
|
||||
print_order: Sequence[str] = (
|
||||
"data",
|
||||
"model",
|
||||
"callbacks",
|
||||
"logger",
|
||||
"trainer",
|
||||
"paths",
|
||||
"extras",
|
||||
),
|
||||
resolve: bool = False,
|
||||
save_to_file: bool = False,
|
||||
) -> None:
|
||||
"""Prints content of DictConfig using Rich library and its tree structure.
|
||||
|
||||
Args:
|
||||
cfg (DictConfig): Configuration composed by Hydra.
|
||||
print_order (Sequence[str], optional): Determines in what order config components are printed.
|
||||
resolve (bool, optional): Whether to resolve reference fields of DictConfig.
|
||||
save_to_file (bool, optional): Whether to export config to the hydra output folder.
|
||||
""" # noqa: E501
|
||||
|
||||
style = "dim"
|
||||
tree = rich.tree.Tree("CONFIG", style=style, guide_style=style)
|
||||
|
||||
queue = []
|
||||
|
||||
# add fields from `print_order` to queue
|
||||
for field in print_order:
|
||||
(
|
||||
queue.append(field)
|
||||
if field in cfg
|
||||
else log.warning(
|
||||
f"Field '{field}' not found in config. "
|
||||
+ f"Skipping '{field}' config printing..."
|
||||
)
|
||||
)
|
||||
|
||||
# add all the other fields to queue (not specified in `print_order`)
|
||||
for field in cfg:
|
||||
if field not in queue:
|
||||
queue.append(field)
|
||||
|
||||
# generate config tree from queue
|
||||
for field in queue:
|
||||
branch = tree.add(field, style=style, guide_style=style)
|
||||
|
||||
config_group = cfg[field]
|
||||
if isinstance(config_group, DictConfig):
|
||||
branch_content = OmegaConf.to_yaml(config_group, resolve=resolve)
|
||||
else:
|
||||
branch_content = str(config_group)
|
||||
|
||||
branch.add(rich.syntax.Syntax(branch_content, "yaml"))
|
||||
|
||||
# print config tree
|
||||
rich.print(tree)
|
||||
|
||||
# save config tree to file
|
||||
if save_to_file:
|
||||
with open(Path(cfg.paths.output_dir, "config_tree.log"), "w") as file:
|
||||
rich.print(tree, file=file)
|
||||
|
||||
|
||||
|
||||
def enforce_tags(cfg: DictConfig, save_to_file: bool = False) -> None:
|
||||
"""Prompts user to input tags from command line if no tags are provided in config.""" # noqa: E501
|
||||
|
||||
if not cfg.get("tags"):
|
||||
if "id" in HydraConfig().cfg.hydra.job:
|
||||
raise ValueError("Specify tags before launching a multirun!")
|
||||
|
||||
log.warning("No tags provided in config. Prompting user to input tags...")
|
||||
tags = Prompt.ask("Enter a list of comma separated tags", default="dev")
|
||||
tags = [t.strip() for t in tags.split(",") if t != ""]
|
||||
|
||||
with open_dict(cfg):
|
||||
cfg.tags = tags
|
||||
|
||||
log.info(f"Tags: {cfg.tags}")
|
||||
|
||||
if save_to_file:
|
||||
with open(Path(cfg.paths.output_dir, "tags.log"), "w") as file:
|
||||
rich.print(cfg.tags, file=file)
|
||||
@@ -0,0 +1,122 @@
|
||||
import torch
|
||||
import torchaudio.functional as F
|
||||
from torch import Tensor, nn
|
||||
from torchaudio.transforms import MelScale
|
||||
|
||||
|
||||
class LinearSpectrogram(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_fft=2048,
|
||||
win_length=2048,
|
||||
hop_length=512,
|
||||
center=False,
|
||||
mode="pow2_sqrt",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.n_fft = n_fft
|
||||
self.win_length = win_length
|
||||
self.hop_length = hop_length
|
||||
self.center = center
|
||||
self.mode = mode
|
||||
|
||||
self.register_buffer("window", torch.hann_window(win_length), persistent=False)
|
||||
|
||||
def forward(self, y: Tensor) -> Tensor:
|
||||
if y.ndim == 3:
|
||||
y = y.squeeze(1)
|
||||
|
||||
y = torch.nn.functional.pad(
|
||||
y.unsqueeze(1),
|
||||
(
|
||||
(self.win_length - self.hop_length) // 2,
|
||||
(self.win_length - self.hop_length + 1) // 2,
|
||||
),
|
||||
mode="reflect",
|
||||
).squeeze(1)
|
||||
|
||||
spec = torch.stft(
|
||||
y,
|
||||
self.n_fft,
|
||||
hop_length=self.hop_length,
|
||||
win_length=self.win_length,
|
||||
window=self.window,
|
||||
center=self.center,
|
||||
pad_mode="reflect",
|
||||
normalized=False,
|
||||
onesided=True,
|
||||
return_complex=True,
|
||||
)
|
||||
|
||||
spec = torch.view_as_real(spec)
|
||||
|
||||
if self.mode == "pow2_sqrt":
|
||||
spec = torch.sqrt(spec.pow(2).sum(-1) + 1e-6)
|
||||
|
||||
return spec
|
||||
|
||||
|
||||
class LogMelSpectrogram(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
sample_rate=44100,
|
||||
n_fft=2048,
|
||||
win_length=2048,
|
||||
hop_length=512,
|
||||
n_mels=128,
|
||||
center=False,
|
||||
f_min=0.0,
|
||||
f_max=None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.sample_rate = sample_rate
|
||||
self.n_fft = n_fft
|
||||
self.win_length = win_length
|
||||
self.hop_length = hop_length
|
||||
self.center = center
|
||||
self.n_mels = n_mels
|
||||
self.f_min = f_min
|
||||
self.f_max = f_max or float(sample_rate // 2)
|
||||
|
||||
self.spectrogram = LinearSpectrogram(n_fft, win_length, hop_length, center)
|
||||
|
||||
fb = F.melscale_fbanks(
|
||||
n_freqs=self.n_fft // 2 + 1,
|
||||
f_min=self.f_min,
|
||||
f_max=self.f_max,
|
||||
n_mels=self.n_mels,
|
||||
sample_rate=self.sample_rate,
|
||||
norm="slaney",
|
||||
mel_scale="slaney",
|
||||
)
|
||||
self.register_buffer(
|
||||
"fb",
|
||||
fb,
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
def compress(self, x: Tensor) -> Tensor:
|
||||
return torch.log(torch.clamp(x, min=1e-5))
|
||||
|
||||
def decompress(self, x: Tensor) -> Tensor:
|
||||
return torch.exp(x)
|
||||
|
||||
def apply_mel_scale(self, x: Tensor) -> Tensor:
|
||||
return torch.matmul(x.transpose(-1, -2), self.fb).transpose(-1, -2)
|
||||
|
||||
def forward(
|
||||
self, x: Tensor, return_linear: bool = False, sample_rate: int = None
|
||||
) -> Tensor:
|
||||
if sample_rate is not None and sample_rate != self.sample_rate:
|
||||
x = F.resample(x, orig_freq=sample_rate, new_freq=self.sample_rate)
|
||||
|
||||
linear = self.spectrogram(x)
|
||||
x = self.apply_mel_scale(linear)
|
||||
x = self.compress(x)
|
||||
|
||||
if return_linear:
|
||||
return x, self.compress(linear)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,114 @@
|
||||
import warnings
|
||||
from importlib.util import find_spec
|
||||
from typing import Callable
|
||||
|
||||
from omegaconf import DictConfig
|
||||
|
||||
from .logger import RankedLogger
|
||||
from .rich_utils import enforce_tags, print_config_tree
|
||||
|
||||
log = RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
def extras(cfg: DictConfig) -> None:
|
||||
"""Applies optional utilities before the task is started.
|
||||
|
||||
Utilities:
|
||||
- Ignoring python warnings
|
||||
- Setting tags from command line
|
||||
- Rich config printing
|
||||
"""
|
||||
|
||||
# return if no `extras` config
|
||||
if not cfg.get("extras"):
|
||||
log.warning("Extras config not found! <cfg.extras=null>")
|
||||
return
|
||||
|
||||
# disable python warnings
|
||||
if cfg.extras.get("ignore_warnings"):
|
||||
log.info("Disabling python warnings! <cfg.extras.ignore_warnings=True>")
|
||||
warnings.filterwarnings("ignore")
|
||||
|
||||
# prompt user to input tags from command line if none are provided in the config
|
||||
if cfg.extras.get("enforce_tags"):
|
||||
log.info("Enforcing tags! <cfg.extras.enforce_tags=True>")
|
||||
enforce_tags(cfg, save_to_file=True)
|
||||
|
||||
# pretty print config tree using Rich library
|
||||
if cfg.extras.get("print_config"):
|
||||
log.info("Printing config tree with Rich! <cfg.extras.print_config=True>")
|
||||
print_config_tree(cfg, resolve=True, save_to_file=True)
|
||||
|
||||
|
||||
def task_wrapper(task_func: Callable) -> Callable:
|
||||
"""Optional decorator that controls the failure behavior when executing the task function.
|
||||
|
||||
This wrapper can be used to:
|
||||
- make sure loggers are closed even if the task function raises an exception (prevents multirun failure)
|
||||
- save the exception to a `.log` file
|
||||
- mark the run as failed with a dedicated file in the `logs/` folder (so we can find and rerun it later)
|
||||
- etc. (adjust depending on your needs)
|
||||
|
||||
Example:
|
||||
```
|
||||
@utils.task_wrapper
|
||||
def train(cfg: DictConfig) -> Tuple[dict, dict]:
|
||||
|
||||
...
|
||||
|
||||
return metric_dict, object_dict
|
||||
```
|
||||
""" # noqa: E501
|
||||
|
||||
def wrap(cfg: DictConfig):
|
||||
# execute the task
|
||||
try:
|
||||
metric_dict, object_dict = task_func(cfg=cfg)
|
||||
|
||||
# things to do if exception occurs
|
||||
except Exception as ex:
|
||||
# save exception to `.log` file
|
||||
log.exception("")
|
||||
|
||||
# some hyperparameter combinations might be invalid or
|
||||
# cause out-of-memory errors so when using hparam search
|
||||
# plugins like Optuna, you might want to disable
|
||||
# raising the below exception to avoid multirun failure
|
||||
raise ex
|
||||
|
||||
# things to always do after either success or exception
|
||||
finally:
|
||||
# display output dir path in terminal
|
||||
log.info(f"Output dir: {cfg.paths.run_dir}")
|
||||
|
||||
# always close wandb run (even if exception occurs so multirun won't fail)
|
||||
if find_spec("wandb"): # check if wandb is installed
|
||||
import wandb
|
||||
|
||||
if wandb.run:
|
||||
log.info("Closing wandb!")
|
||||
wandb.finish()
|
||||
|
||||
return metric_dict, object_dict
|
||||
|
||||
return wrap
|
||||
|
||||
|
||||
def get_metric_value(metric_dict: dict, metric_name: str) -> float:
|
||||
"""Safely retrieves value of the metric logged in LightningModule."""
|
||||
|
||||
if not metric_name:
|
||||
log.info("Metric name is None! Skipping metric value retrieval...")
|
||||
return None
|
||||
|
||||
if metric_name not in metric_dict:
|
||||
raise Exception(
|
||||
f"Metric value not found! <metric_name={metric_name}>\n"
|
||||
"Make sure metric name logged in LightningModule is correct!\n"
|
||||
"Make sure `optimized_metric` name in `hparams_search` config is correct!"
|
||||
)
|
||||
|
||||
metric_value = metric_dict[metric_name].item()
|
||||
log.info(f"Retrieved metric value! <{metric_name}={metric_value}>")
|
||||
|
||||
return metric_value
|
||||
@@ -0,0 +1,98 @@
|
||||
|
||||
import hydra
|
||||
from hydra import compose, initialize
|
||||
from hydra.utils import instantiate
|
||||
import torch
|
||||
from loguru import logger
|
||||
import torchaudio
|
||||
|
||||
|
||||
def load_model(config_name, checkpoint_path, device="cuda"):
|
||||
hydra.core.global_hydra.GlobalHydra.instance().clear()
|
||||
with initialize(version_base="1.3", config_path="./configs"):
|
||||
cfg = compose(config_name=config_name)
|
||||
|
||||
model = instantiate(cfg)
|
||||
state_dict = torch.load(
|
||||
checkpoint_path,
|
||||
map_location=device,
|
||||
)
|
||||
if "state_dict" in state_dict:
|
||||
state_dict = state_dict["state_dict"]
|
||||
|
||||
if any("generator" in k for k in state_dict):
|
||||
state_dict = {
|
||||
k.replace("generator.", ""): v
|
||||
for k, v in state_dict.items()
|
||||
if "generator." in k
|
||||
}
|
||||
|
||||
result = model.load_state_dict(state_dict, strict=False)
|
||||
model.eval()
|
||||
model.to(device)
|
||||
|
||||
logger.info(f"Loaded model: {result}")
|
||||
return model
|
||||
|
||||
|
||||
def codes2audio(model, indices, device):
|
||||
# Restore
|
||||
feature_lengths = torch.tensor([indices.shape[1]], device=device)
|
||||
|
||||
fake_audios, _ = model.decode(
|
||||
indices=indices[None], feature_lengths=feature_lengths
|
||||
)
|
||||
|
||||
audio_time = fake_audios.shape[-1] / model.spec_transform.sample_rate
|
||||
|
||||
logger.info(
|
||||
f"Generated audio of shape {fake_audios.shape}, equivalent to {audio_time:.2f} seconds from {indices.shape[1]} features, features/second: {indices.shape[1] / audio_time:.2f}"
|
||||
)
|
||||
|
||||
# Save audio
|
||||
fake_audio = fake_audios[0, 0]
|
||||
|
||||
# to tensor
|
||||
waveform = fake_audio.unsqueeze(0)
|
||||
|
||||
sample_rate = model.spec_transform.sample_rate
|
||||
|
||||
audio_content = {"waveform": waveform.unsqueeze(0), "sample_rate": sample_rate}
|
||||
|
||||
return audio_content
|
||||
|
||||
|
||||
def audio2prompt(model, audio_content, device):
|
||||
logger.info(f"Processing in-place reconstruction of {audio_content}")
|
||||
|
||||
audio = audio_content['waveform'].squeeze(0)
|
||||
sr = audio_content['sample_rate']
|
||||
|
||||
if audio.shape[0] > 1:
|
||||
audio = audio.mean(0, keepdim=True)
|
||||
|
||||
audio = torchaudio.functional.resample(
|
||||
audio, sr, model.spec_transform.sample_rate
|
||||
)
|
||||
|
||||
audios = audio[None].to(device)
|
||||
logger.info(
|
||||
f"Loaded audio with {audios.shape[2] / model.spec_transform.sample_rate:.2f} seconds"
|
||||
)
|
||||
|
||||
# VQ Encoder
|
||||
audio_lengths = torch.tensor([audios.shape[2]], device=device, dtype=torch.long)
|
||||
indices = model.encode(audios, audio_lengths)[0][0]
|
||||
|
||||
logger.info(f"Generated indices of shape {indices.shape}")
|
||||
|
||||
audio_content = codes2audio(model, indices, device)
|
||||
|
||||
return (audio_content, indices.cpu().numpy(), )
|
||||
|
||||
|
||||
def semantic2audio(model, codes, device):
|
||||
logger.info(f"Processing precomputed indices from {codes.shape}")
|
||||
indices = torch.from_numpy(codes).to(device).long()
|
||||
audio_content = codes2audio(model, indices, device)
|
||||
return (audio_content, )
|
||||
@@ -0,0 +1,321 @@
|
||||
import os
|
||||
import folder_paths
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
# import comfy.utils
|
||||
from PIL import Image
|
||||
# from PIL.PngImagePlugin import PngInfo
|
||||
|
||||
import cv2
|
||||
from scenedetect.video_manager import VideoManager
|
||||
from scenedetect.scene_manager import SceneManager
|
||||
from scenedetect.detectors import AdaptiveDetector
|
||||
|
||||
import os
|
||||
import random
|
||||
import string
|
||||
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons."""
|
||||
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
|
||||
any_type = AnyType("*")
|
||||
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def generate_folder_name(directory,video_path):
|
||||
# Get the directory and filename from the video path
|
||||
_, filename = os.path.split(video_path)
|
||||
# Generate a random string of lowercase letters and digits
|
||||
random_string = ''.join(random.choices(string.ascii_lowercase + string.digits, k=8))
|
||||
# Create the folder name by combining the random string and the filename
|
||||
folder_name = random_string + '_' + filename
|
||||
# Create the full folder path by joining the directory and the folder name
|
||||
folder_path = os.path.join(directory, folder_name)
|
||||
return folder_path
|
||||
|
||||
def create_folder(directory,video_path):
|
||||
folder_path = generate_folder_name(directory,video_path)
|
||||
os.makedirs(folder_path)
|
||||
return folder_path
|
||||
|
||||
|
||||
|
||||
def detect_scenes(video_path, min_scene_len=15, adaptive_threshold=3.0,callback=None):
|
||||
# Create a VideoManager object to load the video file.
|
||||
video_manager = VideoManager([video_path])
|
||||
video_manager.set_downscale_factor()
|
||||
|
||||
# Create a SceneManager object to manage the scene detection process.
|
||||
scene_manager = SceneManager()
|
||||
# scene_manager.add_detector(AdaptiveDetector())
|
||||
|
||||
adaptive_detector = AdaptiveDetector(adaptive_threshold=adaptive_threshold,min_scene_len=min_scene_len)
|
||||
|
||||
scene_manager.add_detector(adaptive_detector)
|
||||
|
||||
# Initialize the video processing loop.
|
||||
video_manager.start()
|
||||
if callback:
|
||||
scene_manager.detect_scenes(frame_source=video_manager,callback=callback)
|
||||
else:
|
||||
scene_manager.detect_scenes(frame_source=video_manager)
|
||||
|
||||
# Iterate over the detected scenes and print their start and end timecodes.
|
||||
scenes = []
|
||||
for scene in scene_manager.get_scene_list():
|
||||
# start_time = scene[0].get_timecode()
|
||||
# end_time = scene[1].get_timecode()
|
||||
# scenes.append((start_time, end_time))
|
||||
scenes.append(scene)
|
||||
|
||||
# Release the video manager and scene manager resources.
|
||||
video_manager.release()
|
||||
# scene_manager.release()
|
||||
|
||||
return scenes
|
||||
|
||||
|
||||
# 采样逻辑
|
||||
def calculate_sample_range(start_frame, middle_frame, end_frame, number_of_sample_frames):
|
||||
half_samples = number_of_sample_frames // 2
|
||||
|
||||
# 初始化采样帧列表
|
||||
samples = [middle_frame]
|
||||
|
||||
if number_of_sample_frames==1:
|
||||
return samples
|
||||
|
||||
# 计算间隔
|
||||
interval_before = (middle_frame - start_frame) // half_samples
|
||||
interval_after = (end_frame - middle_frame) // half_samples
|
||||
|
||||
# 添加中间帧前的采样帧
|
||||
for i in range(1, half_samples + 1):
|
||||
sample_before = middle_frame - i * interval_before
|
||||
if sample_before >= start_frame:
|
||||
samples.insert(0, sample_before)
|
||||
|
||||
# 添加中间帧后的采样帧
|
||||
for i in range(1, half_samples + 1):
|
||||
sample_after = middle_frame + i * interval_after
|
||||
if sample_after <= end_frame:
|
||||
samples.append(sample_after)
|
||||
|
||||
# 如果采样帧数是偶数,则需要移除最靠近边界的一个帧
|
||||
if number_of_sample_frames % 2 == 0:
|
||||
if len(samples) > number_of_sample_frames:
|
||||
if abs(samples[0] - start_frame) < abs(samples[-1] - end_frame):
|
||||
samples.pop(0)
|
||||
else:
|
||||
samples.pop()
|
||||
|
||||
return samples
|
||||
|
||||
|
||||
def split_video_by_scenes(video_path, scenes, output_path, number_of_sample_frames=1):
|
||||
# Load the video file
|
||||
video = cv2.VideoCapture(video_path)
|
||||
|
||||
# Get the video properties
|
||||
fps = video.get(cv2.CAP_PROP_FPS)
|
||||
width = int(video.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(video.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
# 视频的总帧数
|
||||
total_frames = int(video.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
|
||||
# Create a list to hold the paths of the scene videos
|
||||
scenes_video = []
|
||||
keyframes = []
|
||||
|
||||
# Iterate over the scenes
|
||||
for scene_num, scene in enumerate(scenes, start=1):
|
||||
start_time = scene[0]
|
||||
end_time = scene[1]
|
||||
|
||||
# Calculate the start and end frames based on the timestamps
|
||||
start_frame = int(start_time.get_seconds() * fps)
|
||||
end_frame = int(end_time.get_seconds() * fps)
|
||||
|
||||
# Calculate the middle frame
|
||||
middle_frame = (start_frame + end_frame) // 2
|
||||
|
||||
sample_frames=[]
|
||||
|
||||
# Calculate the range of frames to sample
|
||||
# sample_range = range(max(start_frame, middle_frame - number_of_sample_frames // 2),
|
||||
# min(end_frame, middle_frame + number_of_sample_frames // 2 + 1))
|
||||
sample_range=calculate_sample_range(start_frame, middle_frame, end_frame, number_of_sample_frames)
|
||||
|
||||
# Set the video file's current frame to the start frame
|
||||
video.set(cv2.CAP_PROP_POS_FRAMES, start_frame)
|
||||
|
||||
# Create a VideoWriter object for the current scene
|
||||
output_path1 = os.path.join(output_path, f"scene{scene_num}.avi")
|
||||
scenes_video.append(output_path1)
|
||||
writer = cv2.VideoWriter(output_path1, cv2.VideoWriter_fourcc(*'XVID'), fps, (width, height))
|
||||
|
||||
# Write the frames of the current scene to the video file
|
||||
for frame_num in range(start_frame, end_frame + 1):
|
||||
ret, frame = video.read()
|
||||
if not ret:
|
||||
break
|
||||
writer.write(frame)
|
||||
|
||||
# If this frame is in the sample range, save it to keyframes
|
||||
if frame_num in sample_range:
|
||||
# Convert the frame to RGB (OpenCV uses BGR by default)
|
||||
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
# Convert the frame to a PIL image
|
||||
pil_image = Image.fromarray(frame_rgb)
|
||||
|
||||
sample_frames.append(pil2tensor(pil_image))
|
||||
|
||||
keyframe_info = {
|
||||
'start_frame': start_frame,
|
||||
'end_frame': end_frame,
|
||||
'sample_frames': sample_frames,
|
||||
'video_path': output_path1
|
||||
}
|
||||
|
||||
keyframes.append(keyframe_info)
|
||||
|
||||
# Release the VideoWriter object
|
||||
writer.release()
|
||||
|
||||
# Release the video file
|
||||
video.release()
|
||||
|
||||
return scenes_video, keyframes,total_frames
|
||||
|
||||
|
||||
def get_files_with_extension(directory, extension):
|
||||
file_list = []
|
||||
for root, dirs, files in os.walk(directory):
|
||||
for file in files:
|
||||
if file.endswith(extension):
|
||||
file = os.path.splitext(file)[0]
|
||||
file_path = os.path.join(root, file)
|
||||
file_name = os.path.relpath(file_path, directory)
|
||||
file_list.append(file_name)
|
||||
return file_list
|
||||
|
||||
# 从list里取中间的元素
|
||||
def get_middle_element(lst):
|
||||
if not lst:
|
||||
return None # 如果列表为空,返回None
|
||||
mid_index = len(lst) // 2
|
||||
index=0
|
||||
if len(lst) % 2 == 0:
|
||||
index=mid_index - 1
|
||||
else:
|
||||
index=mid_index
|
||||
|
||||
if index<0:
|
||||
index=0
|
||||
|
||||
return lst[index] # 返回中间的一个元素
|
||||
|
||||
|
||||
class SceneInfoNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
|
||||
return {"required": {
|
||||
"scenes": ('SCENE_',),
|
||||
"index": ("INT", {"default": 0, "min": -1, "step": 1}),
|
||||
},}
|
||||
|
||||
RETURN_TYPES = ('IMAGE','IMAGE','INT','INT','SCENE_VIDEO',)
|
||||
RETURN_NAMES = ("sample_frames","middle_frames","start_frame","end_frame","scene_video",)
|
||||
# OUTPUT_IS_LIST = (False,)
|
||||
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
INPUT_IS_LIST = False
|
||||
def run(self,scenes,index):
|
||||
|
||||
if index==-1:
|
||||
images_list=[]
|
||||
m_images=[]
|
||||
start_frames=[]
|
||||
end_frames=[]
|
||||
video_paths=[]
|
||||
for i in range(len(scenes)):
|
||||
s=scenes[i]
|
||||
m_images.append(get_middle_element(s['sample_frames']))
|
||||
sample_frames=torch.cat(s['sample_frames'], dim=0)
|
||||
images_list.append(sample_frames)
|
||||
start_frames.append(s['start_frame'])
|
||||
end_frames.append(s['end_frame'])
|
||||
video_paths.append(s['video_path'])
|
||||
# images = torch.cat(images, dim=0)
|
||||
m_images=torch.cat(m_images, dim=0)
|
||||
return (images_list,m_images,start_frames,end_frames,video_paths,)
|
||||
else:
|
||||
s=scenes[index]
|
||||
images=s['sample_frames']
|
||||
images = torch.cat(images, dim=0)
|
||||
m_images=get_middle_element(s['sample_frames'])
|
||||
return ([images],m_images,s['start_frame'],s['end_frame'],s['video_path'],)
|
||||
|
||||
|
||||
# 分割视频
|
||||
class ScenedetectNode_:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
video_extensions = ['webm', 'mp4', 'mkv', 'gif']
|
||||
input_dir = folder_paths.get_input_directory()
|
||||
files = []
|
||||
for f in os.listdir(input_dir):
|
||||
if os.path.isfile(os.path.join(input_dir, f)):
|
||||
file_parts = f.split('.')
|
||||
if len(file_parts) > 1 and (file_parts[-1] in video_extensions):
|
||||
files.append(f)
|
||||
return {"required": {
|
||||
"video": (sorted(files), {"video_upload": True}),
|
||||
"min_scene_len": ("INT", {"default": 10, "min": 1, "step": 1}),
|
||||
"adaptive_threshold": ("FLOAT", {"default": 2.5, "min": 0, "step": 0.1}),
|
||||
"number_of_sample_frames": ("INT", {"default": 1, "min": 1, "step": 1}), # 抽取的帧数,默认是1帧,中间帧
|
||||
},}
|
||||
|
||||
RETURN_TYPES = ("SCENE_VIDEO","SCENE_", "INT","INT",)
|
||||
RETURN_NAMES = ("scenes_video","scenes","scene_len","total_frames",)
|
||||
OUTPUT_IS_LIST = (False,False,False,)
|
||||
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def run(self, video, min_scene_len,adaptive_threshold,number_of_sample_frames):
|
||||
video_path = folder_paths.get_annotated_filepath(video)
|
||||
# Example usage:
|
||||
scenes = detect_scenes(video_path, min_scene_len=min_scene_len, adaptive_threshold=adaptive_threshold)
|
||||
# print("##scenes", scenes)
|
||||
# for start_time, end_time in scenes:
|
||||
# print(f"Scene detected from {start_time} to {end_time}")
|
||||
|
||||
|
||||
tp=folder_paths.get_temp_directory()
|
||||
basename = os.path.basename(video_path) # 获取文件名
|
||||
name_without_extension = os.path.splitext(basename)[0] # 去掉文件后缀
|
||||
|
||||
folder_path = create_folder(tp,name_without_extension)
|
||||
# print("New folder created:", folder_path)
|
||||
|
||||
vs_files,keyframes,total=split_video_by_scenes(video_path,scenes,folder_path,number_of_sample_frames)
|
||||
# print("New folder created:", vs_files)
|
||||
|
||||
return (vs_files,keyframes,len(scenes),total,)
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-mixlab-nodes"
|
||||
description = "3D, ScreenShareNode & FloatingVideoNode, SpeechRecognition & SpeechSynthesis, GPT, LoadImagesFromLocal, Layers, Other Nodes, ..."
|
||||
version = "0.32.0"
|
||||
version = "0.43.0"
|
||||
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"]
|
||||
|
||||
|
||||
@@ -4,17 +4,29 @@ watchdog
|
||||
opencv-python-headless
|
||||
matplotlib
|
||||
openai
|
||||
simple-lama-inpainting
|
||||
torchaudio
|
||||
|
||||
clip-interrogator==0.6.0
|
||||
transformers>=4.36.0
|
||||
lark-parser
|
||||
imageio-ffmpeg
|
||||
rembg[gpu]
|
||||
omegaconf==2.3.0
|
||||
omegaconf>=2.3.0
|
||||
Pillow>=9.5.0
|
||||
einops==0.7.0
|
||||
einops>=0.7.0
|
||||
trimesh>=4.0.5
|
||||
huggingface-hub
|
||||
scikit-image
|
||||
torchaudio
|
||||
soundfile>=0.12.1
|
||||
soundfile>=0.12.1
|
||||
json-repair
|
||||
|
||||
bitsandbytes
|
||||
accelerate
|
||||
scenedetect[opencv-headless]
|
||||
|
||||
hydra-core>=1.3.2
|
||||
loralib>=0.1.2
|
||||
natsort>=8.4.0
|
||||
# simple-lama-inpainting
|
||||
|
||||
git+https://github.com/shadowcz007/SenseVoice-python.git
|
||||
@@ -2,6 +2,62 @@ import { app } from '../../../scripts/app.js'
|
||||
import { api } from '../../../scripts/api.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
import { loadExternalScript, get_position_style } from './common.js'
|
||||
|
||||
function setCameraOrbit (modelview, distant, angles, screenNumber) {
|
||||
//2.1 20
|
||||
// const angles = {
|
||||
// 1: -20.0,
|
||||
// 2: -17.9,
|
||||
// 3: -15.8,
|
||||
// 4: -13.7,
|
||||
// 5: -11.6,
|
||||
// 6: -9.5,
|
||||
// 7: -7.4,
|
||||
// 8: -5.3,
|
||||
// 9: -3.2,
|
||||
// 10: -1.1,
|
||||
// 11: 1.1,
|
||||
// 12: 3.2,
|
||||
// 13: 5.3,
|
||||
// 14: 7.4,
|
||||
// 15: 9.5,
|
||||
// 16: 11.6,
|
||||
// 17: 13.7,
|
||||
// 18: 15.8,
|
||||
// 19: 17.9,
|
||||
// 20: 20.0
|
||||
// };
|
||||
|
||||
// 12 3.6
|
||||
// const angles = {
|
||||
// 1: -20.0,
|
||||
// 2: -16.4,
|
||||
// 3: -12.7,
|
||||
// 4: -9.1,
|
||||
// 5: -5.5,
|
||||
// 6: -1.8,
|
||||
// 7: 1.8,
|
||||
// 8: 5.5,
|
||||
// 9: 9.1,
|
||||
// 10: 12.7,
|
||||
// 11: 16.4,
|
||||
// 12: 20.0
|
||||
// }
|
||||
|
||||
const angle = angles[screenNumber]
|
||||
|
||||
let co=modelview.cameraOrbit.split(" ")
|
||||
|
||||
if (angle !== undefined) {
|
||||
|
||||
modelview.cameraOrbit = `${angle}deg ${co[1]} ${distant}m`
|
||||
console.log(screenNumber, angle)
|
||||
} else {
|
||||
console.error('Invalid screen number')
|
||||
}
|
||||
}
|
||||
|
||||
const getLocalData = key => {
|
||||
let data = {}
|
||||
try {
|
||||
@@ -26,7 +82,8 @@ const setLocalDataOfWin = (key, value) => {
|
||||
localStorage.setItem(key, JSON.stringify(value))
|
||||
// window[key] = value
|
||||
}
|
||||
async function uploadImage (blob, fileType = '.svg', filename) {
|
||||
|
||||
async function uploadImage_ (blob, fileType = '.svg', filename) {
|
||||
// const blob = await (await fetch(src)).blob();
|
||||
const body = new FormData()
|
||||
body.append(
|
||||
@@ -41,13 +98,17 @@ async function uploadImage (blob, fileType = '.svg', filename) {
|
||||
|
||||
// console.log(resp)
|
||||
let data = await resp.json()
|
||||
return data
|
||||
}
|
||||
|
||||
async function uploadImage (blob, fileType = '.svg', filename) {
|
||||
let data = await uploadImage_(blob, fileType, filename)
|
||||
let { name, subfolder } = data
|
||||
let src = api.apiURL(
|
||||
`/view?filename=${encodeURIComponent(
|
||||
name
|
||||
)}&type=input&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
|
||||
)
|
||||
|
||||
return src
|
||||
}
|
||||
|
||||
@@ -78,38 +139,6 @@ const parseImage = url => {
|
||||
})
|
||||
}
|
||||
|
||||
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'
|
||||
}
|
||||
}
|
||||
|
||||
async function extractMaterial (
|
||||
modelViewerVariants,
|
||||
selectMaterial,
|
||||
@@ -171,6 +200,42 @@ async function changeMaterial (
|
||||
targetMaterial.pbrMetallicRoughness.baseColorTexture.setTexture(targetTexture)
|
||||
}
|
||||
|
||||
function inputFileClick (isFileURL = false, isGlb = false) {
|
||||
return new Promise((res, rej) => {
|
||||
// 创建一个input元素
|
||||
var input = document.createElement('input')
|
||||
input.type = 'file'
|
||||
input.accept = isGlb ? '.glb' : 'image/*'
|
||||
|
||||
// 监听input的change事件
|
||||
input.addEventListener('change', function () {
|
||||
// 获取上传的文件
|
||||
var file = input.files[0]
|
||||
|
||||
if (isFileURL) {
|
||||
res(URL.createObjectURL(file))
|
||||
return
|
||||
}
|
||||
|
||||
// 创建一个FileReader对象来读取文件
|
||||
var reader = new FileReader()
|
||||
|
||||
// 监听FileReader的load事件
|
||||
reader.addEventListener('load', async () => {
|
||||
let base64 = reader.result
|
||||
input.remove()
|
||||
res(base64)
|
||||
})
|
||||
|
||||
// 读取文件
|
||||
reader.readAsDataURL(file)
|
||||
})
|
||||
|
||||
// 触发input的点击事件
|
||||
input.click()
|
||||
})
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.3D.3DImage',
|
||||
async getCustomWidgets (app) {
|
||||
@@ -189,7 +254,7 @@ app.registerExtension({
|
||||
let d = getLocalData('_mixlab_3d_image')
|
||||
// console.log('serializeValue', node)
|
||||
if (d && d[node.id]) {
|
||||
let { url, bg, material } = d[node.id]
|
||||
let { url, bg, material, images } = d[node.id]
|
||||
let data = {}
|
||||
if (url) {
|
||||
data.image = await parseImage(url)
|
||||
@@ -205,6 +270,10 @@ app.registerExtension({
|
||||
data.material = await parseImage(material)
|
||||
}
|
||||
|
||||
if (images) {
|
||||
data.images = images
|
||||
}
|
||||
|
||||
return JSON.parse(JSON.stringify(data))
|
||||
} else {
|
||||
return {}
|
||||
@@ -217,66 +286,62 @@ app.registerExtension({
|
||||
}
|
||||
},
|
||||
|
||||
async init () {
|
||||
await loadExternalScript('/mixlab/app/lib/model-viewer.min.js', 'module')
|
||||
},
|
||||
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == '3DImage') {
|
||||
console.log('nodeType.comfyClass', nodeType.comfyClass)
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const uploadWidget = this.widgets.filter(w => w.name == 'upload')[0]
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'upload-preview',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, 88, node.size[1])
|
||||
get_position_style(ctx, widget_width - 122, 88, node.size[1], 44)
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
widget.div.style.width = `120px`
|
||||
|
||||
document.body.appendChild(widget.div)
|
||||
|
||||
const inputDiv = (key, placeholder, preview) => {
|
||||
let div = document.createElement('div')
|
||||
const ip = document.createElement('input')
|
||||
ip.type = 'file'
|
||||
const ip = document.createElement('button')
|
||||
ip.className = `${'comfy-multiline-input'} ${placeholder}`
|
||||
div.style = `display: flex;
|
||||
align-items: center;
|
||||
margin: 6px 8px;
|
||||
margin-top: 0;`
|
||||
ip.placeholder = placeholder
|
||||
// ip.value = value
|
||||
|
||||
ip.style = `outline: none;
|
||||
border: none;
|
||||
padding: 4px;
|
||||
width: 60%;cursor: pointer;
|
||||
width: 100px;cursor: pointer;
|
||||
height: 32px;`
|
||||
const label = document.createElement('label')
|
||||
label.style = 'font-size: 10px;min-width:32px'
|
||||
label.innerText = placeholder
|
||||
div.appendChild(label)
|
||||
ip.innerText = placeholder
|
||||
div.appendChild(ip)
|
||||
|
||||
let that = this,
|
||||
filename = new Date().getTime()
|
||||
let that = this
|
||||
|
||||
ip.addEventListener('change', async event => {
|
||||
const file = event.target.files[0]
|
||||
const reader = new FileReader()
|
||||
filename = new Date().getTime()
|
||||
// 读取文件内容
|
||||
reader.onload = async e => {
|
||||
const fileURL = URL.createObjectURL(file)
|
||||
// console.log('文件URL: ', fileURL)
|
||||
let html = `<model-viewer src="${fileURL}"
|
||||
min-field-of-view="0deg" max-field-of-view="180deg"
|
||||
ip.addEventListener('click', async event => {
|
||||
let fileURL = await inputFileClick(true, true)
|
||||
|
||||
// console.log('文件URL: ', fileURL)
|
||||
let html = `<model-viewer src="${fileURL}"
|
||||
oncontextmenu="return false;"
|
||||
style="outline:1px solid white"
|
||||
min-field-of-view="0deg"
|
||||
max-field-of-view="180deg"
|
||||
min-camera-orbit="auto auto 0m"
|
||||
max-camera-orbit="auto auto 1000m"
|
||||
shadow-intensity="1"
|
||||
camera-controls
|
||||
touch-action="pan-y">
|
||||
@@ -285,230 +350,314 @@ app.registerExtension({
|
||||
<div>Variant: <select class="variant"></select></div>
|
||||
<div>Material: <select class="material"></select></div>
|
||||
<div>Material: <div class="material_img"> </div></div>
|
||||
<div><button class="bg">BG</button></div>
|
||||
<div>
|
||||
<button class="bg">BG</button>
|
||||
|
||||
</div>
|
||||
<div>
|
||||
<input class="ddcap_distant" type="number" min="1" step="1" value="55">
|
||||
<input class="total_images" type="number" min="1" max="180" step="1" value="20">
|
||||
<input class="ddcap_range" type="number" min="0" max="20" step="0.1" value="2.1">
|
||||
<button class="ddcap">Capture Rotational Screenshots</button></div>
|
||||
|
||||
<div><button class="export">Export GLB</button></div>
|
||||
|
||||
</div></model-viewer>`
|
||||
|
||||
preview.innerHTML = html
|
||||
if (that.size[1] < 400) {
|
||||
that.setSize([that.size[0], that.size[1] + 300])
|
||||
app.canvas.draw(true, true)
|
||||
}
|
||||
|
||||
const modelViewerVariants = preview.querySelector('model-viewer')
|
||||
const select = preview.querySelector('.variant')
|
||||
const selectMaterial = preview.querySelector('.material')
|
||||
const material_img = preview.querySelector('.material_img')
|
||||
const bg = preview.querySelector('.bg')
|
||||
const exportGLB = preview.querySelector('.export')
|
||||
|
||||
if (modelViewerVariants) {
|
||||
modelViewerVariants.style.width = `${that.size[0] - 24}px`
|
||||
modelViewerVariants.style.height = `${that.size[1] - 48}px`
|
||||
}
|
||||
|
||||
modelViewerVariants.addEventListener('load', async () => {
|
||||
const names = modelViewerVariants.availableVariants
|
||||
|
||||
// 变量
|
||||
for (const name of names) {
|
||||
const option = document.createElement('option')
|
||||
option.value = name
|
||||
option.textContent = name
|
||||
select.appendChild(option)
|
||||
}
|
||||
// Adds a default option.
|
||||
if (names.length === 0) {
|
||||
const option = document.createElement('option')
|
||||
option.value = 'default'
|
||||
option.textContent = 'Default'
|
||||
select.appendChild(option)
|
||||
}
|
||||
|
||||
// 材质
|
||||
extractMaterial(
|
||||
modelViewerVariants,
|
||||
selectMaterial,
|
||||
material_img
|
||||
)
|
||||
})
|
||||
|
||||
let timer = null
|
||||
const delay = 500 // 延迟时间,单位为毫秒
|
||||
|
||||
async function checkCameraChange () {
|
||||
let dd = getLocalData(key)
|
||||
let base64Data = modelViewerVariants.toDataURL()
|
||||
|
||||
const contentType = getContentTypeFromBase64(base64Data)
|
||||
|
||||
const blob = await base64ToBlobFromURL(base64Data, contentType)
|
||||
|
||||
// const fileBlob = new Blob([e.target.result], { type: file.type });
|
||||
let url = await uploadImage(blob, '.png')
|
||||
// console.log(url)
|
||||
|
||||
let bg_blob = await base64ToBlobFromURL(
|
||||
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mN88uXrPQAFwwK/6xJ6CQAAAABJRU5ErkJggg=='
|
||||
)
|
||||
let url_bg = await uploadImage(bg_blob, '.png')
|
||||
// console.log('url_bg',url_bg)
|
||||
|
||||
if (!dd[that.id]) {
|
||||
dd[that.id] = { url, bg: url_bg }
|
||||
} else {
|
||||
dd[that.id] = { ...dd[that.id], url }
|
||||
}
|
||||
|
||||
// 材质贴图
|
||||
let thumbUrl = material_img.getAttribute('src')
|
||||
if (thumbUrl) {
|
||||
let tb = await base64ToBlobFromURL(thumbUrl)
|
||||
let tUrl = await uploadImage(tb, '.png')
|
||||
// console.log('材质贴图', tUrl, thumbUrl)
|
||||
dd[that.id].material = tUrl
|
||||
}
|
||||
|
||||
setLocalDataOfWin(key, dd)
|
||||
}
|
||||
|
||||
function startTimer () {
|
||||
if (timer) clearTimeout(timer)
|
||||
timer = setTimeout(checkCameraChange, delay)
|
||||
}
|
||||
|
||||
modelViewerVariants.addEventListener('camera-change', startTimer)
|
||||
|
||||
select.addEventListener('input', async event => {
|
||||
modelViewerVariants.variantName =
|
||||
event.target.value === 'default' ? null : event.target.value
|
||||
// 材质
|
||||
await extractMaterial(
|
||||
modelViewerVariants,
|
||||
selectMaterial,
|
||||
material_img
|
||||
)
|
||||
checkCameraChange()
|
||||
})
|
||||
|
||||
selectMaterial.addEventListener('input', event => {
|
||||
// console.log(selectMaterial.value)
|
||||
material_img.setAttribute('src', selectMaterial.value)
|
||||
|
||||
if (selectMaterial.getAttribute('data-new-material')) {
|
||||
let index =
|
||||
~~selectMaterial.selectedOptions[0].getAttribute(
|
||||
'data-index'
|
||||
)
|
||||
changeMaterial(
|
||||
modelViewerVariants,
|
||||
modelViewerVariants.model.materials[index],
|
||||
selectMaterial.getAttribute('data-new-material')
|
||||
)
|
||||
}
|
||||
|
||||
checkCameraChange()
|
||||
})
|
||||
|
||||
bg.addEventListener('click', () => {
|
||||
// 创建一个input元素
|
||||
var input = document.createElement('input')
|
||||
input.type = 'file'
|
||||
|
||||
// 监听input的change事件
|
||||
input.addEventListener('change', function () {
|
||||
// 获取上传的文件
|
||||
var file = input.files[0]
|
||||
|
||||
// 创建一个FileReader对象来读取文件
|
||||
var reader = new FileReader()
|
||||
|
||||
// 监听FileReader的load事件
|
||||
reader.addEventListener('load', async () => {
|
||||
let base64 = reader.result
|
||||
// 将读取的文件内容设置为div的背景
|
||||
preview.style.backgroundImage = 'url(' + base64 + ')'
|
||||
|
||||
const contentType = getContentTypeFromBase64(base64)
|
||||
|
||||
const blob = await base64ToBlobFromURL(base64, contentType)
|
||||
|
||||
// const fileBlob = new Blob([e.target.result], { type: file.type });
|
||||
let bg_url = await uploadImage(blob, '.png')
|
||||
let bg_img = await createImage(base64)
|
||||
|
||||
let dd = getLocalData(key)
|
||||
// console.log(dd[that.id],bg_url)
|
||||
if (!dd[that.id]) dd[that.id] = { url: '', bg: bg_url }
|
||||
dd[that.id] = {
|
||||
...dd[that.id],
|
||||
bg: bg_url,
|
||||
bg_w: bg_img.naturalWidth,
|
||||
bg_h: bg_img.naturalHeight
|
||||
}
|
||||
|
||||
setLocalDataOfWin(key, dd)
|
||||
|
||||
// 更新尺寸
|
||||
let w = that.size[0] - 24,
|
||||
h = (w * bg_img.naturalHeight) / bg_img.naturalWidth
|
||||
|
||||
if (modelViewerVariants) {
|
||||
modelViewerVariants.style.width = `${w}px`
|
||||
modelViewerVariants.style.height = `${h}px`
|
||||
}
|
||||
preview.style.width = `${w}px`
|
||||
})
|
||||
|
||||
// 读取文件
|
||||
reader.readAsDataURL(file)
|
||||
})
|
||||
|
||||
// 触发input的点击事件
|
||||
input.click()
|
||||
})
|
||||
|
||||
exportGLB.addEventListener('click', async () => {
|
||||
const glTF = await modelViewerVariants.exportScene()
|
||||
const file = new File([glTF], 'export.glb')
|
||||
const link = document.createElement('a')
|
||||
link.download = file.name
|
||||
link.href = URL.createObjectURL(file)
|
||||
link.click()
|
||||
})
|
||||
|
||||
uploadWidget.value = await uploadWidget.serializeValue()
|
||||
|
||||
// 更新尺寸
|
||||
let dd = getLocalData(key)
|
||||
// console.log(dd[that.id],bg_url)
|
||||
if (dd[that.id]) {
|
||||
const { bg_w, bg_h } = dd[that.id]
|
||||
if (bg_h && bg_w) {
|
||||
let w = that.size[0] - 24,
|
||||
h = (w * bg_h) / bg_w
|
||||
|
||||
if (modelViewerVariants) {
|
||||
modelViewerVariants.style.width = `${w}px`
|
||||
modelViewerVariants.style.height = `${h}px`
|
||||
}
|
||||
preview.style.width = `${w}px`
|
||||
}
|
||||
}
|
||||
preview.innerHTML = html
|
||||
if (that.size[1] < 400) {
|
||||
that.setSize([that.size[0], that.size[1] + 300])
|
||||
app.canvas.draw(true, true)
|
||||
}
|
||||
|
||||
// 以文本形式读取文件
|
||||
reader.readAsDataURL(file)
|
||||
const modelViewerVariants = preview.querySelector('model-viewer')
|
||||
const select = preview.querySelector('.variant')
|
||||
const selectMaterial = preview.querySelector('.material')
|
||||
const material_img = preview.querySelector('.material_img')
|
||||
const bg = preview.querySelector('.bg')
|
||||
|
||||
const exportGLB = preview.querySelector('.export')
|
||||
|
||||
const ddcap_distant = preview.querySelector('.ddcap_distant')
|
||||
const total_images = preview.querySelector('.total_images')
|
||||
const ddcap_range = preview.querySelector('.ddcap_range')
|
||||
const ddCap = preview.querySelector('.ddcap')
|
||||
const sleep = (t = 1000) => {
|
||||
return new Promise((res, rej) => {
|
||||
return setTimeout(() => {
|
||||
res(t)
|
||||
}, t)
|
||||
})
|
||||
}
|
||||
|
||||
async function captureImage (isUrl = true) {
|
||||
let base64Data = modelViewerVariants.toDataURL()
|
||||
|
||||
const contentType = getContentTypeFromBase64(base64Data)
|
||||
|
||||
const blob = await base64ToBlobFromURL(base64Data, contentType)
|
||||
|
||||
if (isUrl) return await uploadImage(blob, '.png')
|
||||
return await uploadImage_(blob, '.png')
|
||||
}
|
||||
|
||||
async function captureImages (
|
||||
ddcap_range = 1,
|
||||
total_images = 12,
|
||||
distant = 0.23
|
||||
) {
|
||||
// 初始 角度
|
||||
var center = modelViewerVariants.getBoundingBoxCenter().toString()
|
||||
modelViewerVariants.cameraTarget = center
|
||||
|
||||
const startAngle = -((total_images - 1) / 2) * ddcap_range
|
||||
const angles = {}
|
||||
|
||||
for (let i = 0; i < total_images; i++) {
|
||||
angles[i + 1] = startAngle + i * ddcap_range
|
||||
}
|
||||
console.log(angles)
|
||||
|
||||
let frames = []
|
||||
|
||||
modelViewerVariants.removeAttribute('camera-controls')
|
||||
|
||||
for (let i = 0; i < total_images; i++) {
|
||||
setCameraOrbit(modelViewerVariants, distant, angles, i + 1)
|
||||
|
||||
// modelViewerVariants.cameraOrbit = `${currentAngle}deg ${initialCameraOrbit[1]} ${initialCameraOrbit[2]}`
|
||||
await sleep(1000)
|
||||
// console.log(`Capturing image at angle: ${currentAngle}deg`)
|
||||
let file = await captureImage(false)
|
||||
frames.push(file)
|
||||
// currentAngle += angleIncrement
|
||||
}
|
||||
await sleep(1000)
|
||||
// 恢复到初始旋转角度
|
||||
// modelViewerVariants.cameraOrbit = initialCameraOrbit.join(' ')
|
||||
modelViewerVariants.setAttribute('camera-controls', '')
|
||||
return frames
|
||||
}
|
||||
|
||||
ddCap.addEventListener('click', async e => {
|
||||
const distant = Number(ddcap_distant.value), // 23m
|
||||
totalImages = Number(total_images.value),
|
||||
angleIncrement = Number(ddcap_range.value)
|
||||
console.log(angleIncrement, totalImages)
|
||||
let images = await captureImages(
|
||||
angleIncrement,
|
||||
totalImages,
|
||||
distant
|
||||
)
|
||||
|
||||
let dd = getLocalData(key)
|
||||
dd[that.id].images = images
|
||||
setLocalDataOfWin(key, dd)
|
||||
})
|
||||
|
||||
ddcap_distant.addEventListener('input', async e => {
|
||||
// console.log(ddcap_distant.value)
|
||||
const center = modelViewerVariants.getBoundingBoxCenter().toString()
|
||||
modelViewerVariants.cameraTarget = center;
|
||||
const initialCameraOrbit =
|
||||
modelViewerVariants.cameraOrbit.split(' ')
|
||||
modelViewerVariants.cameraOrbit = `${initialCameraOrbit[2]} ${initialCameraOrbit[1]} ${ddcap_distant.value}m`
|
||||
modelViewerVariants.setAttribute('camera-controls', '')
|
||||
})
|
||||
|
||||
// ddcap_range_top.addEventListener('input', async e => {
|
||||
// // console.log(ddcap_range.value)
|
||||
// const initialCameraOrbit =
|
||||
// modelViewerVariants.cameraOrbit.split(' ')
|
||||
// modelViewerVariants.cameraOrbit = `${initialCameraOrbit[0]} ${ddcap_range_top.value}deg ${initialCameraOrbit[2]}`
|
||||
// modelViewerVariants.setAttribute('camera-controls', '')
|
||||
// })
|
||||
|
||||
if (modelViewerVariants) {
|
||||
modelViewerVariants.style.width = `${that.size[0] - 48}px`
|
||||
modelViewerVariants.style.height = `${that.size[1] - 48}px`
|
||||
}
|
||||
|
||||
modelViewerVariants.addEventListener('load', async () => {
|
||||
const names = modelViewerVariants.availableVariants
|
||||
|
||||
// 变量
|
||||
for (const name of names) {
|
||||
const option = document.createElement('option')
|
||||
option.value = name
|
||||
option.textContent = name
|
||||
select.appendChild(option)
|
||||
}
|
||||
// Adds a default option.
|
||||
if (names.length === 0) {
|
||||
const option = document.createElement('option')
|
||||
option.value = 'default'
|
||||
option.textContent = 'Default'
|
||||
select.appendChild(option)
|
||||
}
|
||||
|
||||
// 材质
|
||||
extractMaterial(modelViewerVariants, selectMaterial, material_img)
|
||||
})
|
||||
|
||||
let timer = null
|
||||
const delay = 500 // 延迟时间,单位为毫秒
|
||||
|
||||
async function checkCameraChange () {
|
||||
let dd = getLocalData(key)
|
||||
|
||||
// const fileBlob = new Blob([e.target.result], { type: file.type });
|
||||
let url = await captureImage()
|
||||
|
||||
let bg_blob = await base64ToBlobFromURL(
|
||||
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mN88uXrPQAFwwK/6xJ6CQAAAABJRU5ErkJggg=='
|
||||
)
|
||||
let url_bg = await uploadImage(bg_blob, '.png')
|
||||
// console.log('url_bg',url_bg)
|
||||
|
||||
if (!dd[that.id]) {
|
||||
dd[that.id] = { url, bg: url_bg }
|
||||
} else {
|
||||
dd[that.id] = { ...dd[that.id], url }
|
||||
}
|
||||
|
||||
// 材质贴图
|
||||
let thumbUrl = material_img.getAttribute('src')
|
||||
if (thumbUrl) {
|
||||
let tb = await base64ToBlobFromURL(thumbUrl)
|
||||
let tUrl = await uploadImage(tb, '.png')
|
||||
// console.log('材质贴图', tUrl, thumbUrl)
|
||||
dd[that.id].material = tUrl
|
||||
}
|
||||
|
||||
setLocalDataOfWin(key, dd)
|
||||
}
|
||||
|
||||
function startTimer () {
|
||||
if (timer) clearTimeout(timer)
|
||||
timer = setTimeout(checkCameraChange, delay)
|
||||
}
|
||||
|
||||
modelViewerVariants.addEventListener('camera-change', startTimer)
|
||||
|
||||
select.addEventListener('input', async event => {
|
||||
modelViewerVariants.variantName =
|
||||
event.target.value === 'default' ? null : event.target.value
|
||||
// 材质
|
||||
await extractMaterial(
|
||||
modelViewerVariants,
|
||||
selectMaterial,
|
||||
material_img
|
||||
)
|
||||
checkCameraChange()
|
||||
})
|
||||
|
||||
selectMaterial.addEventListener('input', event => {
|
||||
// console.log(selectMaterial.value)
|
||||
material_img.setAttribute('src', selectMaterial.value)
|
||||
|
||||
if (selectMaterial.getAttribute('data-new-material')) {
|
||||
let index =
|
||||
~~selectMaterial.selectedOptions[0].getAttribute('data-index')
|
||||
changeMaterial(
|
||||
modelViewerVariants,
|
||||
modelViewerVariants.model.materials[index],
|
||||
selectMaterial.getAttribute('data-new-material')
|
||||
)
|
||||
}
|
||||
|
||||
checkCameraChange()
|
||||
})
|
||||
|
||||
//更新bg
|
||||
const updateBgData = (id, key, url, w, h) => {
|
||||
let dd = getLocalData(key)
|
||||
// console.log(dd[that.id],url)
|
||||
if (!dd[id]) dd[id] = { url: '', bg: url }
|
||||
dd[id] = {
|
||||
...dd[id],
|
||||
bg: url,
|
||||
bg_w: w,
|
||||
bg_h: h
|
||||
}
|
||||
setLocalDataOfWin(key, dd)
|
||||
}
|
||||
|
||||
bg.addEventListener('click', async () => {
|
||||
//更新bg
|
||||
updateBgData(that.id, key, '', 0, 0)
|
||||
preview.style.backgroundImage = 'none'
|
||||
|
||||
let base64 = await inputFileClick(false, false)
|
||||
// 将读取的文件内容设置为div的背景
|
||||
preview.style.backgroundImage = 'url(' + base64 + ')'
|
||||
|
||||
const contentType = getContentTypeFromBase64(base64)
|
||||
|
||||
const blob = await base64ToBlobFromURL(base64, contentType)
|
||||
|
||||
// const fileBlob = new Blob([e.target.result], { type: file.type });
|
||||
let bg_url = await uploadImage(blob, '.png')
|
||||
let bg_img = await createImage(base64)
|
||||
|
||||
//更新bg
|
||||
updateBgData(
|
||||
that.id,
|
||||
key,
|
||||
bg_url,
|
||||
bg_img.naturalWidth,
|
||||
bg_img.naturalHeight
|
||||
)
|
||||
|
||||
// 更新尺寸
|
||||
let w = that.size[0] - 128,
|
||||
h = (w * bg_img.naturalHeight) / bg_img.naturalWidth
|
||||
|
||||
if (modelViewerVariants) {
|
||||
modelViewerVariants.style.width = `${w}px`
|
||||
modelViewerVariants.style.height = `${h}px`
|
||||
}
|
||||
preview.style.width = `${w}px`
|
||||
})
|
||||
|
||||
exportGLB.addEventListener('click', async () => {
|
||||
const glTF = await modelViewerVariants.exportScene()
|
||||
const file = new File([glTF], 'export.glb')
|
||||
const link = document.createElement('a')
|
||||
link.download = file.name
|
||||
link.href = URL.createObjectURL(file)
|
||||
link.click()
|
||||
})
|
||||
|
||||
uploadWidget.value = await uploadWidget.serializeValue()
|
||||
|
||||
// 更新尺寸
|
||||
let dd = getLocalData(key)
|
||||
// console.log(dd[that.id],bg_url)
|
||||
if (dd[that.id]) {
|
||||
const { bg_w, bg_h } = dd[that.id]
|
||||
if (bg_h && bg_w) {
|
||||
let w = that.size[0] - 48,
|
||||
h = (w * bg_h) / bg_w
|
||||
|
||||
if (modelViewerVariants) {
|
||||
modelViewerVariants.style.width = `${w}px`
|
||||
modelViewerVariants.style.height = `${h}px`
|
||||
}
|
||||
preview.style.width = `${w}px`
|
||||
}
|
||||
}
|
||||
})
|
||||
return div
|
||||
}
|
||||
|
||||
let preview = document.createElement('div')
|
||||
preview.className = 'preview'
|
||||
preview.style = `margin-top: 12px;display: flex;
|
||||
preview.style = `margin-top: 12px;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;background-repeat: no-repeat;background-size: contain;`
|
||||
align-items: center;background-repeat: no-repeat;
|
||||
background-size: contain;`
|
||||
|
||||
let upload = inputDiv('_mixlab_3d_image', '3D Model', preview)
|
||||
|
||||
@@ -523,18 +672,25 @@ app.registerExtension({
|
||||
|
||||
// 更新尺寸
|
||||
let dd = getLocalData('_mixlab_3d_image')
|
||||
// console.log(dd[that.id],bg_url)
|
||||
|
||||
if (dd[that.id]) {
|
||||
const { bg_w, bg_h } = dd[that.id]
|
||||
if (bg_h && bg_w) {
|
||||
let w = that.size[0] - 24,
|
||||
h = (w * bg_h) / bg_w
|
||||
let w = that.size[0] - 128
|
||||
preview.style.width = `${w}px`
|
||||
console.log('更新尺寸', w)
|
||||
|
||||
if (modelViewerVariants) {
|
||||
modelViewerVariants.style.width = `${w}px`
|
||||
modelViewerVariants.style.height = `${Math.round(
|
||||
that.size[1] * 0.8
|
||||
)}px`
|
||||
}
|
||||
|
||||
if (bg_h && bg_w) {
|
||||
let h = (w * bg_h) / bg_w
|
||||
if (modelViewerVariants) {
|
||||
modelViewerVariants.style.width = `${w}px`
|
||||
modelViewerVariants.style.height = `${h}px`
|
||||
}
|
||||
preview.style.width = `${w}px`
|
||||
}
|
||||
}
|
||||
|
||||
@@ -561,7 +717,7 @@ app.registerExtension({
|
||||
const r = onExecuted?.apply?.(this, arguments)
|
||||
|
||||
let div = this.widgets.filter(d => d.div)[0]?.div
|
||||
console.log('Test', this.widgets)
|
||||
// console.log('Test', this.widgets)
|
||||
|
||||
let material = message.material[0]
|
||||
if (material) {
|
||||
@@ -617,7 +773,7 @@ app.registerExtension({
|
||||
// let base64 = await parseImage(url)
|
||||
|
||||
let pre = widget.div.querySelector('.preview')
|
||||
pre.style.width = `${node.size[0]}px`
|
||||
pre.style.width = `${node.size[0] - 24}px`
|
||||
pre.innerHTML = `
|
||||
${url ? `<img src="${url}" style="width:100%"/>` : ''}
|
||||
`
|
||||
|
||||
@@ -3,28 +3,17 @@ 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)
|
||||
// console.log('td_bg', td_bg)
|
||||
import {
|
||||
getUrl,
|
||||
base64Df,
|
||||
get_position_style,
|
||||
getObjectInfo
|
||||
} from './common.js'
|
||||
|
||||
//本机安装的插件节点全集
|
||||
window._nodesAll = null
|
||||
|
||||
//获取当前系统的插件,节点清单
|
||||
function getObjectInfo () {
|
||||
return new Promise(async (resolve, reject) => {
|
||||
let url = getUrl()
|
||||
|
||||
try {
|
||||
const response = await fetch(`${url}/object_info`)
|
||||
const data = await response.json()
|
||||
resolve(data)
|
||||
} catch (error) {
|
||||
reject(error)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
const base64Df =
|
||||
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
|
||||
|
||||
const parseImageToBase64 = url => {
|
||||
return new Promise((res, rej) => {
|
||||
fetch(url)
|
||||
@@ -44,39 +33,6 @@ const parseImageToBase64 = url => {
|
||||
})
|
||||
}
|
||||
|
||||
function get_position_style (ctx, widget_width, y, node_height) {
|
||||
const MARGIN = 12 // 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: 'flex-start',
|
||||
zIndex: 9999999
|
||||
}
|
||||
}
|
||||
|
||||
async function drawImageToCanvas (imageUrl, sFactor = 320) {
|
||||
var canvas = document.createElement('canvas')
|
||||
var ctx = canvas.getContext('2d')
|
||||
@@ -256,7 +212,10 @@ async function extractInputAndOutputData (
|
||||
node.type === 'KSampler' ||
|
||||
node.type == 'SamplerCustom' ||
|
||||
node.type === 'ChinesePrompt_Mix' ||
|
||||
node.type === 'Seed_'
|
||||
node.type === 'Seed_' ||
|
||||
node.type === 'SiliconflowLLM' ||
|
||||
node.type === 'ChatGPTOpenAI' ||
|
||||
node.type === 'SiliconflowTextToImageNode'
|
||||
) {
|
||||
// seed 的类型收集
|
||||
try {
|
||||
@@ -276,23 +235,6 @@ async function extractInputAndOutputData (
|
||||
return { input, output, seed, seedTitle }
|
||||
}
|
||||
|
||||
function getUrl () {
|
||||
let api_host = `${window.location.hostname}:${window.location.port}`
|
||||
let api_base = ''
|
||||
let url = `${window.location.protocol}//${api_host}${api_base}`
|
||||
return url
|
||||
}
|
||||
|
||||
const getLocalData = key => {
|
||||
let data = {}
|
||||
try {
|
||||
data = JSON.parse(localStorage.getItem(key)) || {}
|
||||
} catch (error) {
|
||||
return {}
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
async function save_app (json) {
|
||||
let url = getUrl()
|
||||
|
||||
@@ -327,18 +269,20 @@ function downloadJsonFile (jsonData, fileName = 'mix_app.json') {
|
||||
async function save (json, download = false, showInfo = true) {
|
||||
let nodesAll = window._nodesAll || (await getObjectInfo())
|
||||
|
||||
console.log('####SAVE', nodesAll, json[0])
|
||||
console.log('####SAVE', nodesAll, json)
|
||||
|
||||
const name = json[0],
|
||||
version = json[5],
|
||||
share_prefix = json[6], //用于分享的功能扩展
|
||||
link = json[7], //用于创建界面上的跳转链接
|
||||
category = json[8] || '', //用于分类
|
||||
idle_animation = json[9], //用于动画,比如数字人her
|
||||
description = json[4],
|
||||
inputIds = json[2].split('\n').filter(f => f),
|
||||
outputIds = json[3].split('\n').filter(f => f)
|
||||
|
||||
const iconData = json[1][0]
|
||||
|
||||
let { filename, subfolder, type } = iconData
|
||||
let iconUrl = api.apiURL(
|
||||
`/view?filename=${encodeURIComponent(
|
||||
@@ -391,6 +335,27 @@ async function save (json, download = false, showInfo = true) {
|
||||
try {
|
||||
data.app.icon = await drawImageToCanvas(iconUrl)
|
||||
} catch (error) {}
|
||||
|
||||
let images = []
|
||||
if (json[1].length > 1 && idle_animation) {
|
||||
images = Array.from(json[1], j => {
|
||||
let { filename, subfolder, type } = j
|
||||
return api.apiURL(
|
||||
`/view?filename=${encodeURIComponent(
|
||||
filename
|
||||
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
try {
|
||||
for (let index = 0; index < images.length; index++) {
|
||||
const imgurl = images[index]
|
||||
images[index] = await drawImageToCanvas(imgurl)
|
||||
}
|
||||
if (idle_animation) data.app.idle_animation = images
|
||||
} catch (error) {}
|
||||
|
||||
// console.log(data.app)
|
||||
// let http_workflow = app.graph.serialize()
|
||||
await save_app(data)
|
||||
@@ -400,13 +365,17 @@ async function save (json, download = false, showInfo = true) {
|
||||
|
||||
if (showInfo) {
|
||||
let open = window.confirm(
|
||||
`You can now access the standalone application on a new page!\n${getUrl()}/mixlab/app?filename=${encodeURIComponent(
|
||||
`You can now access the standalone application on a new page!\n${getUrl()}/mixlab/app${
|
||||
data.app.idle_animation ? '/her.html' : ''
|
||||
}?filename=${encodeURIComponent(
|
||||
data.app.filename
|
||||
)}&category=${encodeURIComponent(data.app.category)}`
|
||||
)
|
||||
if (open)
|
||||
window.open(
|
||||
`${getUrl()}/mixlab/app?filename=${encodeURIComponent(
|
||||
`${getUrl()}/mixlab/app${
|
||||
data.app.idle_animation ? '/her.html' : ''
|
||||
}?filename=${encodeURIComponent(
|
||||
data.app.filename
|
||||
)}&category=${encodeURIComponent(data.app.category)}`
|
||||
)
|
||||
@@ -470,15 +439,15 @@ app.registerExtension({
|
||||
type: 'div',
|
||||
name: 'AppInfoRun',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
{...get_position_style(
|
||||
Object.assign(this.div.style, {
|
||||
...get_position_style(
|
||||
ctx,
|
||||
widget_width,
|
||||
node.size[1] - widget_height,
|
||||
node.size[1]
|
||||
),zIndex:1}
|
||||
)
|
||||
),
|
||||
zIndex: 1
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -708,7 +677,6 @@ app.registerExtension({
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
|
||||
window._mixlab_app_json = null
|
||||
|
||||
}
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
|
||||
@@ -19,7 +19,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
|
||||
return {
|
||||
transformOrigin: '0 0',
|
||||
transform: transform,
|
||||
left: `0`,
|
||||
left:
|
||||
document.querySelector('.comfy-menu').style.display === 'none'
|
||||
? `60px`
|
||||
: `0`,
|
||||
top: `0`,
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import { getUrl } from './common.js'
|
||||
|
||||
async function* completion (url, messages, controller) {
|
||||
let data = {
|
||||
model: 'gpt-3.5-turbo-16k',
|
||||
@@ -91,16 +93,116 @@ async function* completion (url, messages, controller) {
|
||||
return content
|
||||
// return (await response.json()).content
|
||||
}
|
||||
|
||||
export async function completion_ (url, messages, controller, callback) {
|
||||
let request = await completion(url, messages, controller)
|
||||
export async function completion_ (
|
||||
apiKey,
|
||||
url,
|
||||
model_name,
|
||||
messages,
|
||||
controller,
|
||||
callback
|
||||
) {
|
||||
let request = await chatCompletion(
|
||||
apiKey,
|
||||
url,
|
||||
model_name,
|
||||
messages,
|
||||
controller
|
||||
)
|
||||
for await (const chunk of request) {
|
||||
let content = chunk.data.choices[0].delta.content || ''
|
||||
if (chunk.data.choices[0].role == 'assistant') {
|
||||
//开始
|
||||
content = ''
|
||||
}
|
||||
|
||||
if (callback) callback(content)
|
||||
if (callback) callback(chunk)
|
||||
}
|
||||
}
|
||||
|
||||
export async function* chatCompletion (
|
||||
apiKey,
|
||||
api_url,
|
||||
model_name,
|
||||
messages,
|
||||
controller
|
||||
) {
|
||||
const mixlabAPI = `${getUrl()}/chat/completions`
|
||||
|
||||
const requestBody = {
|
||||
messages: messages,
|
||||
stream: true,
|
||||
key: apiKey,
|
||||
model_name: model_name,
|
||||
api_url
|
||||
}
|
||||
|
||||
let response = await fetch(mixlabAPI, {
|
||||
method: 'POST',
|
||||
headers: {
|
||||
'Content-Type': 'application/json'
|
||||
// Authorization: `Bearer ${apiKey}`
|
||||
},
|
||||
body: JSON.stringify(requestBody),
|
||||
mode: 'cors', // This is to ensure the request is made with CORS
|
||||
signal: controller.signal
|
||||
})
|
||||
|
||||
const reader = response.body.getReader()
|
||||
const decoder = new TextDecoder()
|
||||
|
||||
let content = ''
|
||||
let leftover = '' // Buffer for partially read lines
|
||||
|
||||
try {
|
||||
let cont = true
|
||||
while (cont) {
|
||||
let result = await reader.read()
|
||||
if (result.done) {
|
||||
break
|
||||
}
|
||||
const text = leftover + decoder.decode(result.value)
|
||||
// Check if the last character is a line break
|
||||
const endsWithLineBreak = text.endsWith('\n')
|
||||
|
||||
// Split the text into lines
|
||||
let lines = text.split('\n')
|
||||
|
||||
// If the text doesn't end with a line break, then the last line is incomplete
|
||||
// Store it in leftover to be added to the next chunk of data
|
||||
if (!endsWithLineBreak) {
|
||||
leftover = lines.pop()
|
||||
} else {
|
||||
leftover = '' // Reset leftover if we have a line break at the end
|
||||
}
|
||||
|
||||
// Parse all sse events and add them to result
|
||||
const regex = /^(\S+):\s(.*)$/gm
|
||||
for (const line of lines) {
|
||||
const match = regex.exec(line)
|
||||
if (match) {
|
||||
result[match[1]] = match[2]
|
||||
// since we know this is llama.cpp, let's just decode the json in data
|
||||
if (result.data) {
|
||||
result.data = JSON.parse(result.data)
|
||||
|
||||
|
||||
content += result.data.choices[0].delta?.content || ''
|
||||
// console.log('#result.content',content)
|
||||
// yield
|
||||
yield result
|
||||
|
||||
// if we got a stop token from server, we will break here
|
||||
if (result.data.choices[0].finish_reason == 'stop') {
|
||||
if (result.data.generation_settings) {
|
||||
// generation_settings = result.data.generation_settings;
|
||||
}
|
||||
cont = false
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} catch (e) {
|
||||
console.error('chat error: ', e)
|
||||
throw e
|
||||
} finally {
|
||||
controller.abort()
|
||||
}
|
||||
|
||||
return content
|
||||
}
|
||||
|
||||
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
|
||||
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
|
||||
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
|
||||
|
||||
const version = 'v0.32.0'
|
||||
const version = 'v0.43.0'
|
||||
|
||||
fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
|
||||
.then(response => response.json())
|
||||
|
||||
@@ -0,0 +1,227 @@
|
||||
export const base64Df =
|
||||
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
|
||||
|
||||
export function getUrl () {
|
||||
let api_host = `${window.location.hostname}:${window.location.port}`
|
||||
let api_base = ''
|
||||
let url = `${window.location.protocol}//${api_host}${api_base}`
|
||||
return url
|
||||
}
|
||||
|
||||
// 获得插件/节点的索引数据
|
||||
export async function get_nodes_map () {
|
||||
let url = getUrl()
|
||||
|
||||
const res = await fetch(`${url}/mixlab/nodes_map`, {
|
||||
method: 'POST',
|
||||
body: JSON.stringify({
|
||||
data: 'json'
|
||||
})
|
||||
})
|
||||
return await res.json()
|
||||
}
|
||||
|
||||
// 更新或者获取key
|
||||
export const updateLLMAPIKey = async key => {
|
||||
try {
|
||||
const res = await fetch(`${getUrl()}/mixlab/llm_api_key`, {
|
||||
method: 'POST',
|
||||
headers: {
|
||||
'Content-Type': 'application/json'
|
||||
},
|
||||
body: JSON.stringify({
|
||||
key: key || null
|
||||
})
|
||||
})
|
||||
|
||||
const data = await res.json()
|
||||
|
||||
if (!res.ok) {
|
||||
console.error('Error:', data.error)
|
||||
return
|
||||
}
|
||||
|
||||
if (key) {
|
||||
console.log('API key saved successfully:', data.message)
|
||||
return key
|
||||
} else {
|
||||
console.log('Retrieved API key:', data.key)
|
||||
return data.key
|
||||
}
|
||||
} catch (error) {
|
||||
console.error('Request failed:', error)
|
||||
}
|
||||
}
|
||||
|
||||
//获取当前系统的插件,节点清单
|
||||
export function getObjectInfo () {
|
||||
return new Promise(async (resolve, reject) => {
|
||||
let url = getUrl()
|
||||
|
||||
try {
|
||||
const response = await fetch(`${url}/object_info`)
|
||||
const data = await response.json()
|
||||
resolve(data)
|
||||
} catch (error) {
|
||||
reject(error)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
export function get_position_style (
|
||||
ctx,
|
||||
widget_width,
|
||||
y,
|
||||
node_height,
|
||||
left = 44
|
||||
) {
|
||||
const MARGIN = 0 // the margin around the html element
|
||||
|
||||
/* Create a transform that deals with all the scrolling and zooming */
|
||||
const elRect = ctx.canvas.getBoundingClientRect()
|
||||
|
||||
const scaleX = elRect.width / ctx.canvas.width
|
||||
const scaleY = elRect.height / ctx.canvas.height
|
||||
|
||||
const transform = new DOMMatrix()
|
||||
.scaleSelf(scaleX, scaleY)
|
||||
.multiplySelf(ctx.getTransform())
|
||||
.translateSelf(MARGIN, MARGIN + y)
|
||||
|
||||
return {
|
||||
transformOrigin: '0 0',
|
||||
transform: transform,
|
||||
left:
|
||||
document.querySelector('.comfy-menu').style.display === 'none'
|
||||
? `${left}px`
|
||||
: `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: 'flex-start',
|
||||
zIndex: 99
|
||||
}
|
||||
}
|
||||
|
||||
export function loadCSS (url) {
|
||||
var link = document.createElement('link')
|
||||
link.rel = 'stylesheet'
|
||||
link.type = 'text/css'
|
||||
link.href = url
|
||||
document.getElementsByTagName('head')[0].appendChild(link)
|
||||
}
|
||||
|
||||
|
||||
export function injectCSS (css) {
|
||||
// 检查页面中是否已经存在具有相同内容的style标签
|
||||
const existingStyle = document.querySelector('style')
|
||||
if (existingStyle && existingStyle.textContent === css) {
|
||||
return // 如果已经存在相同的样式,则不进行注入
|
||||
}
|
||||
|
||||
// 创建一个新的style标签,并将CSS内容注入其中
|
||||
const style = document.createElement('style')
|
||||
style.textContent = css
|
||||
|
||||
// 将style标签插入到页面的head元素中
|
||||
const head = document.querySelector('head')
|
||||
head.appendChild(style)
|
||||
}
|
||||
|
||||
|
||||
export function loadExternalScript (url, type) {
|
||||
return new Promise((resolve, reject) => {
|
||||
const existingScript = document.querySelector(`script[src="${url}"]`)
|
||||
if (existingScript) {
|
||||
existingScript.onload = () => {
|
||||
resolve()
|
||||
}
|
||||
existingScript.onerror = reject
|
||||
return
|
||||
}
|
||||
|
||||
const script = document.createElement('script')
|
||||
script.src = url
|
||||
if (type) script.type = type // Add this line to load the script as an ES module
|
||||
script.onload = () => {
|
||||
resolve()
|
||||
}
|
||||
script.onerror = reject
|
||||
document.head.appendChild(script)
|
||||
})
|
||||
}
|
||||
|
||||
export async function getQueue () {
|
||||
try {
|
||||
const res = await fetch(`${getUrl()}/queue`)
|
||||
const data = await res.json()
|
||||
// console.log(data.queue_running,data.queue_pending)
|
||||
return {
|
||||
// Running action uses a different endpoint for cancelling
|
||||
Running: data.queue_running.length,
|
||||
Pending: data.queue_pending.length
|
||||
}
|
||||
} catch (error) {
|
||||
console.error(error)
|
||||
return { Running: 0, Pending: 0 }
|
||||
}
|
||||
}
|
||||
|
||||
export async function interrupt () {
|
||||
const resp = await fetch(`${getUrl()}/interrupt`, {
|
||||
method: 'POST'
|
||||
})
|
||||
}
|
||||
|
||||
export async function sleep (t = 200) {
|
||||
return new Promise((res, rej) => {
|
||||
setTimeout(() => {
|
||||
res(true)
|
||||
}, t)
|
||||
})
|
||||
}
|
||||
|
||||
export function createImage (url) {
|
||||
let im = new Image()
|
||||
return new Promise((res, rej) => {
|
||||
im.onload = () => res(im)
|
||||
im.src = url
|
||||
})
|
||||
}
|
||||
|
||||
export function convertImageUrlToBase64 (imageUrl) {
|
||||
return fetch(imageUrl)
|
||||
.then(response => response.blob())
|
||||
.then(blob => {
|
||||
return new Promise((resolve, reject) => {
|
||||
const reader = new FileReader()
|
||||
reader.onloadend = () => resolve(reader.result)
|
||||
reader.onerror = reject
|
||||
reader.readAsDataURL(blob)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
export const getLocalData = key => {
|
||||
let data = {}
|
||||
try {
|
||||
data = JSON.parse(localStorage.getItem(key)) || {}
|
||||
} catch (error) {
|
||||
return {}
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
export const saveLocalData = (key, id, val) => {
|
||||
let data = getLocalData(key)
|
||||
data[id] = val
|
||||
localStorage.setItem(key, JSON.stringify(data))
|
||||
}
|
||||
@@ -1,319 +1,5 @@
|
||||
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.SiliconflowLLM',
|
||||
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.
|
||||
},
|
||||
}
|
||||
},
|
||||
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'SiliconflowLLM') {
|
||||
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 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')
|
||||
|
||||
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) {
|
||||
// Fires every time a node is constructed
|
||||
// You can modify widgets/add handlers/etc here
|
||||
|
||||
if (node.type === 'SiliconflowLLM') {
|
||||
let widget = node.widgets.filter(w => w.div)[0]
|
||||
|
||||
let apiKey = getLocalData('_mixlab_api_key');
|
||||
|
||||
let id = node.id
|
||||
|
||||
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.GPT.ShowTextForGPT',
|
||||
@@ -328,7 +14,7 @@ app.registerExtension({
|
||||
for (let i = 0; i < this.widgets.length; i++) {
|
||||
if (this.widgets[i].name == 'show_text')
|
||||
this.widgets[i].onRemove?.()
|
||||
console.log('#ShowTextForGPT', this.widgets[i])
|
||||
|
||||
}
|
||||
this.widgets.length = 2
|
||||
}
|
||||
@@ -399,24 +85,5 @@ app.registerExtension({
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'ShowTextForGPT') {
|
||||
let widget = node.widgets.filter(w => w.name == 'show_text')[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)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
@@ -4,6 +4,8 @@ import { api } from '../../../scripts/api.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
import { applyTextReplacements } from '../../../scripts/utils.js'
|
||||
|
||||
import { loadExternalScript, get_position_style } from './common.js'
|
||||
|
||||
function loadImageToCanvas (base64Image) {
|
||||
var img = new Image()
|
||||
var canvas = document.createElement('canvas')
|
||||
@@ -88,37 +90,40 @@ function getContentTypeFromBase64 (base64Data) {
|
||||
// const blob = base64ToBlob(base64Data, contentType);
|
||||
// console.log(blob);
|
||||
|
||||
function get_position_style (ctx, widget_width, y, node_height) {
|
||||
const MARGIN = 4 // the margin around the html element
|
||||
// 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)
|
||||
// /* 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'
|
||||
}
|
||||
}
|
||||
// return {
|
||||
// transformOrigin: '0 0',
|
||||
// transform: transform,
|
||||
// left:
|
||||
// document.querySelector('.comfy-menu').style.display === 'none'
|
||||
// ? `60px`
|
||||
// : `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 = {}
|
||||
@@ -344,7 +349,7 @@ app.registerExtension({
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, 44, node.size[1])
|
||||
get_position_style(ctx, widget_width, 44, node.size[1], 60)
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -530,7 +535,7 @@ app.registerExtension({
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1])
|
||||
get_position_style(ctx, widget_width, y, node.size[1], 36)
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -675,6 +680,18 @@ const createInputImageForBatch = (base64, widget) => {
|
||||
return im
|
||||
}
|
||||
|
||||
// 添加新图片
|
||||
const addBase64ToWidgetForLoadImagesToBatch = (
|
||||
base64,
|
||||
imagesWidget,
|
||||
imagesDiv
|
||||
) => {
|
||||
if (!imagesWidget.value.base64) imagesWidget.value.base64 = []
|
||||
imagesWidget.value.base64.push(base64)
|
||||
let im = createInputImageForBatch(base64, imagesWidget)
|
||||
imagesDiv.appendChild(im)
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.Comfy.LoadImagesToBatch',
|
||||
async getCustomWidgets (app) {
|
||||
@@ -705,7 +722,6 @@ app.registerExtension({
|
||||
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'LoadImagesToBatch') {
|
||||
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
@@ -717,10 +733,10 @@ app.registerExtension({
|
||||
type: 'div',
|
||||
name: 'image_base64',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, 44, node.size[1])
|
||||
)
|
||||
Object.assign(this.div.style, {
|
||||
...get_position_style(ctx, widget_width, y, node.size[1], 72),
|
||||
top: `${widget_height}px`
|
||||
})
|
||||
},
|
||||
serialize: false
|
||||
}
|
||||
@@ -751,13 +767,18 @@ app.registerExtension({
|
||||
base64 = await loadImageToCanvas(base64)
|
||||
// console.log(base64)
|
||||
if (!imagesWidget.value) imagesWidget.value = { base64: [] }
|
||||
imagesWidget.value.base64.push(base64)
|
||||
let im = createInputImageForBatch(base64, imagesWidget)
|
||||
imagesDiv.appendChild(im)
|
||||
addBase64ToWidgetForLoadImagesToBatch(
|
||||
base64,
|
||||
imagesWidget,
|
||||
imagesDiv
|
||||
)
|
||||
}
|
||||
reader.readAsDataURL(file)
|
||||
})
|
||||
|
||||
// 如果是复制的,有数据 , 这个不生效,取不到数据, 需要在nodeCreated里获取
|
||||
// console.log('#LoadImagesToBatch', imagesWidget.value?.base64)
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Upload Image'
|
||||
|
||||
@@ -829,18 +850,36 @@ 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]
|
||||
|
||||
let pre = imagePreview.div.querySelector('.images_preview')
|
||||
// console.log('#LoadImagesToBatch', imagesWidget.value?.base64)
|
||||
let imagesDiv = imagePreview.div.querySelector('.images_preview')
|
||||
imagesDiv.innerHTML = ''
|
||||
for (const d of imagesWidget.value?.base64 || []) {
|
||||
let im = createInputImageForBatch(d, imagesWidget)
|
||||
pre.appendChild(im)
|
||||
imagesDiv.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')
|
||||
imagesDiv.innerHTML = ''
|
||||
for (const d of imagesWidget.value?.base64 || []) {
|
||||
let im = createInputImageForBatch(d, imagesWidget)
|
||||
imagesDiv.appendChild(im)
|
||||
}
|
||||
}
|
||||
}, 1000)
|
||||
}
|
||||
})
|
||||
|
||||
@@ -848,9 +887,11 @@ app.registerExtension({
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.output.ComparingTwoFrames_',
|
||||
init () {
|
||||
loadExternalScript('/mixlab/app/lib/juxtapose.min.js')
|
||||
|
||||
$el('link', {
|
||||
rel: 'stylesheet',
|
||||
href: '/extensions/comfyui-mixlab-nodes/lib/juxtapose.css',
|
||||
href: '/mixlab/app/lib/juxtapose.css',
|
||||
parent: document.head
|
||||
})
|
||||
|
||||
@@ -868,8 +909,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)
|
||||
@@ -877,10 +918,10 @@ app.registerExtension({
|
||||
type: 'div',
|
||||
name: 'preview',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, 400, 44, node.size[1])
|
||||
)
|
||||
let s = get_position_style(ctx, widget_width, 44, node.size[1], 36)
|
||||
delete s.height
|
||||
|
||||
Object.assign(this.div.style, s)
|
||||
},
|
||||
serialize: false
|
||||
}
|
||||
@@ -891,20 +932,15 @@ 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)
|
||||
@@ -964,7 +1000,7 @@ app.registerExtension({
|
||||
label: 'After'
|
||||
}
|
||||
]
|
||||
this.size=[this.size[0],300]
|
||||
this.size = [this.size[0], 300]
|
||||
}
|
||||
}
|
||||
},
|
||||
@@ -974,7 +1010,6 @@ 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,
|
||||
|
||||
@@ -81,7 +81,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
|
||||
return {
|
||||
transformOrigin: '0 0',
|
||||
transform: transform,
|
||||
left: `0`,
|
||||
left:
|
||||
document.querySelector('.comfy-menu').style.display === 'none'
|
||||
? `60px`
|
||||
: `0`,
|
||||
top: `0`,
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
|
||||
@@ -3,31 +3,17 @@ import { app } from '../../../scripts/app.js'
|
||||
import { ComfyWidgets } from '../../../scripts/widgets.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
let api_host = `${window.location.hostname}:${window.location.port}`
|
||||
let api_base = ''
|
||||
let url = `${window.location.protocol}//${api_host}${api_base}`
|
||||
import {
|
||||
getQueue,
|
||||
interrupt,
|
||||
get_position_style,
|
||||
base64Df,
|
||||
getUrl,
|
||||
createImage,
|
||||
sleep
|
||||
} from './common.js'
|
||||
|
||||
async function getQueue () {
|
||||
try {
|
||||
const res = await fetch(`${url}/queue`)
|
||||
const data = await res.json()
|
||||
// console.log(data.queue_running,data.queue_pending)
|
||||
return {
|
||||
// Running action uses a different endpoint for cancelling
|
||||
Running: data.queue_running.length,
|
||||
Pending: data.queue_pending.length
|
||||
}
|
||||
} catch (error) {
|
||||
console.error(error)
|
||||
return { Running: 0, Pending: 0 }
|
||||
}
|
||||
}
|
||||
|
||||
async function interrupt () {
|
||||
const resp = await fetch(`${url}/interrupt`, {
|
||||
method: 'POST'
|
||||
})
|
||||
}
|
||||
// let url = getUrl()
|
||||
|
||||
async function clipboardWriteImage (win, url) {
|
||||
const canvas = document.createElement('canvas')
|
||||
@@ -208,22 +194,6 @@ async function shareScreen (
|
||||
}
|
||||
}
|
||||
|
||||
async function sleep (t = 200) {
|
||||
return new Promise((res, rej) => {
|
||||
setTimeout(() => {
|
||||
res(true)
|
||||
}, t)
|
||||
})
|
||||
}
|
||||
|
||||
function createImage (url) {
|
||||
let im = new Image()
|
||||
return new Promise((res, rej) => {
|
||||
im.onload = () => res(im)
|
||||
im.src = url
|
||||
})
|
||||
}
|
||||
|
||||
async function compareImages (threshold, previousImage, currentImage) {
|
||||
// 将 base64 转换为 Image 对象
|
||||
var previousImg = await createImage(previousImage)
|
||||
@@ -458,44 +428,6 @@ async function requestCamera () {
|
||||
return false
|
||||
}
|
||||
|
||||
/*
|
||||
A method that returns the required style for the html
|
||||
*/
|
||||
function get_position_style (ctx, widget_width, y, node_height, top) {
|
||||
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: `${top}px`,
|
||||
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 - MARGIN * 2}px`,
|
||||
// background: '#EEEEEE',
|
||||
display: 'flex',
|
||||
flexDirection: 'column',
|
||||
// alignItems: 'center',
|
||||
justifyContent: 'space-around'
|
||||
}
|
||||
}
|
||||
|
||||
const base64Df =
|
||||
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.image.ScreenShareNode',
|
||||
async getCustomWidgets (app) {
|
||||
@@ -593,17 +525,12 @@ app.registerExtension({
|
||||
type: 'HTML', // whatever
|
||||
name: 'sreen_share', // whatever
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
// console.log('ScreenSHare', y, widget_height)
|
||||
// console.log('ScreenSHare', node)
|
||||
Object.assign(
|
||||
this.card.style,
|
||||
get_position_style(
|
||||
ctx,
|
||||
widget_width,
|
||||
widget_height * 5,
|
||||
node.size[1],
|
||||
40
|
||||
)
|
||||
get_position_style(ctx, widget_width, y, node.size[1], 40)
|
||||
)
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1043,12 +970,13 @@ async function setArea (src) {
|
||||
div.innerHTML = `
|
||||
<div id='ml_overlay' style='position: absolute;top:0;background: #251f1fc4;
|
||||
height: 100vh;
|
||||
z-index:999999;
|
||||
z-index:99999999999999;
|
||||
width: 100%;'>
|
||||
<img id='ml_video' style='position: absolute;
|
||||
height: ${displayHeight}px;user-select: none;
|
||||
-webkit-user-drag: none;
|
||||
outline: 2px solid #eaeaea;
|
||||
left: 0;
|
||||
box-shadow: 8px 9px 17px #575757;' />
|
||||
<div id='ml_selection' style='position: absolute;
|
||||
border: 2px dashed red;
|
||||
@@ -1219,10 +1147,10 @@ app.registerExtension({
|
||||
type: 'video',
|
||||
name: 'FloatingVideo',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.card.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1], 0)
|
||||
)
|
||||
Object.assign(this.card.style, {
|
||||
...get_position_style(ctx, widget_width, y, node.size[1], 40),
|
||||
top: `${widget_height}px`
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
import { api } from '../../../scripts/api.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
import { get_position_style } from './common.js'
|
||||
|
||||
function base64ToBlobFromURL (base64URL, contentType) {
|
||||
return fetch(base64URL).then(response => response.blob())
|
||||
}
|
||||
|
||||
async function uploadImage (blob, fileType = '.svg', filename) {
|
||||
// const blob = await (await fetch(src)).blob();
|
||||
const body = new FormData()
|
||||
body.append(
|
||||
'image',
|
||||
new File([blob], (filename || new Date().getTime()) + fileType)
|
||||
)
|
||||
|
||||
const resp = await api.fetchApi('/upload/image', {
|
||||
method: 'POST',
|
||||
body
|
||||
})
|
||||
|
||||
// console.log(resp)
|
||||
let data = await resp.json()
|
||||
let { name, subfolder } = data
|
||||
// let src = api.apiURL(
|
||||
// `/view?filename=${encodeURIComponent(
|
||||
// name
|
||||
// )}&type=input&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
|
||||
// )
|
||||
|
||||
return data
|
||||
}
|
||||
// 上传得到url
|
||||
async function uploadBase64ToFile (base64) {
|
||||
let bg_blob = await base64ToBlobFromURL(base64)
|
||||
let url = await uploadImage(bg_blob, '.png')
|
||||
return url
|
||||
}
|
||||
|
||||
const p5InputNode = {
|
||||
name: 'Mixlab.Comfy.P5Input',
|
||||
async getCustomWidgets (app) {
|
||||
return {
|
||||
IMAGEBASE64 (node, inputName, inputData, app) {
|
||||
const widget = {
|
||||
value: {
|
||||
images: []
|
||||
}, // 不能[x,x,x]
|
||||
type: inputData[0], // the type
|
||||
name: inputName, // the name, slice
|
||||
size: [320, 120], // a default size
|
||||
draw (ctx, node, width, y) {},
|
||||
computeSize (...args) {
|
||||
return [128, 32] // a method to compute the current size of the widget
|
||||
}
|
||||
}
|
||||
node.addCustomWidget(widget)
|
||||
return widget
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'P5Input') {
|
||||
// console.log('P5Input')
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'image_base64',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(
|
||||
ctx,
|
||||
widget_width - 24,
|
||||
44,
|
||||
node.size[1] * 2.8,
|
||||
44
|
||||
)
|
||||
)
|
||||
},
|
||||
serialize: false
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
widget.div.style = `margin:12px;width:400px;height:480px;background:white`
|
||||
|
||||
document.body.appendChild(widget.div)
|
||||
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
// document.addEventListener('wheel', handleMouseWheel)
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
// window.removeEventListener('message', ms)
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
// 节点的大小控制
|
||||
this.setSize([480, 560])
|
||||
app.canvas.draw(true, true)
|
||||
|
||||
const onResize = this.onResize
|
||||
this.onResize = () => {
|
||||
// 设置最小尺寸
|
||||
if (
|
||||
Math.max(this.size[0], 480) != this.size[0] &&
|
||||
Math.max(this.size[1], 560) != this.size[1]
|
||||
) {
|
||||
this.setSize([
|
||||
Math.max(this.size[0], 480),
|
||||
Math.max(this.size[1], 560)
|
||||
])
|
||||
}
|
||||
|
||||
return onResize?.apply(this, arguments)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
onExecuted?.apply(this, arguments)
|
||||
// console.log('##onExecuted', this, message._info)
|
||||
// app.graph.getNodeById(8).widgets[1].div.querySelector('iframe').contentWindow.postMessage('Hello from parent', '*');
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'P5Input') {
|
||||
}
|
||||
},
|
||||
nodeCreated (node, app) {
|
||||
//数据延迟??
|
||||
setTimeout(() => {
|
||||
let widget = node.widgets?.filter(w => w.name == 'image_base64')[0]
|
||||
let framesWidget = node.widgets?.filter(w => w.name == 'frames')[0]
|
||||
if (node.type === 'P5Input' && widget) {
|
||||
console.log('#nodeCreated P5Input')
|
||||
if (framesWidget && !framesWidget.value)
|
||||
framesWidget.value = { images: [] }
|
||||
|
||||
framesWidget.value._seed = Math.random()
|
||||
|
||||
let nodeId = node.id
|
||||
//延迟才能获得this.id
|
||||
widget.div.innerHTML = `<iframe src="mixlab/app/p5_export/p5.html?id=${nodeId}"
|
||||
style="border:0;width:100%;height:100%;"
|
||||
></iframe>`
|
||||
|
||||
// 监听来自iframe的消息
|
||||
const ms = async event => {
|
||||
const data = event.data
|
||||
console.log('#P5 Input #', data)
|
||||
if (
|
||||
data.from === 'p5.widget' &&
|
||||
data.status === 'save' &&
|
||||
data.frames &&
|
||||
data.frames.length >= 0 &&
|
||||
data.nodeId == nodeId &&
|
||||
data.id != framesWidget.value.id
|
||||
) {
|
||||
const frames = data.frames
|
||||
|
||||
//workflow会存储到local,会卡死
|
||||
framesWidget.value.images = []
|
||||
for (const f of frames) {
|
||||
let file = await uploadBase64ToFile(f)
|
||||
framesWidget.value.images.push(file)
|
||||
}
|
||||
// framesWidget.value.base64 = frames
|
||||
// framesWidget.value._seed = Math.random()
|
||||
node.title = 'P5 Input #' + frames.length
|
||||
framesWidget.value.id = data.id
|
||||
}
|
||||
}
|
||||
|
||||
window.addEventListener('message', ms)
|
||||
}
|
||||
}, 1000)
|
||||
}
|
||||
}
|
||||
|
||||
app.registerExtension(p5InputNode)
|
||||
@@ -3,51 +3,37 @@ import { api } from '../../../scripts/api.js'
|
||||
import { ComfyWidgets } from '../../../scripts/widgets.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
import PhotoSwipeLightbox from '/extensions/comfyui-mixlab-nodes/lib/photoswipe-lightbox.esm.min.js'
|
||||
function loadCSS (url) {
|
||||
var link = document.createElement('link')
|
||||
link.rel = 'stylesheet'
|
||||
link.type = 'text/css'
|
||||
link.href = url
|
||||
document.getElementsByTagName('head')[0].appendChild(link)
|
||||
import { loadCSS, injectCSS } from './common.js'
|
||||
|
||||
// Create a style element
|
||||
const style = document.createElement('style')
|
||||
// Define the CSS rule for scrollbar width
|
||||
const cssRule = `.pswp__custom-caption {
|
||||
background: rgb(20 27 70);
|
||||
font-size: 16px;
|
||||
color: #fff;
|
||||
width: calc(100% - 32px);
|
||||
max-width: 980px;
|
||||
padding: 2px 8px;
|
||||
border-radius: 4px;
|
||||
position: absolute;
|
||||
left: 50%;
|
||||
bottom: 16px;
|
||||
transform: translateX(-50%);
|
||||
}
|
||||
.pswp__custom-caption a {
|
||||
color: #fff;
|
||||
text-decoration: underline;
|
||||
}
|
||||
.hidden-caption-content {
|
||||
display: none;
|
||||
}`
|
||||
// Add the CSS rule to the style element
|
||||
style.appendChild(document.createTextNode(cssRule))
|
||||
import PhotoSwipeLightbox from '/mixlab/app/lib/photoswipe-lightbox.esm.min.js'
|
||||
|
||||
// Append the style element to the document head
|
||||
document.head.appendChild(style)
|
||||
}
|
||||
loadCSS('/extensions/comfyui-mixlab-nodes/lib/photoswipe.min.css')
|
||||
loadCSS('/mixlab/app/lib/photoswipe.min.css')
|
||||
injectCSS(`.pswp__custom-caption {
|
||||
background: rgb(20 27 70);
|
||||
font-size: 16px;
|
||||
color: #fff;
|
||||
width: calc(100% - 32px);
|
||||
max-width: 980px;
|
||||
padding: 2px 8px;
|
||||
border-radius: 4px;
|
||||
position: absolute;
|
||||
left: 50%;
|
||||
bottom: 16px;
|
||||
transform: translateX(-50%);
|
||||
}
|
||||
.pswp__custom-caption a {
|
||||
color: #fff;
|
||||
text-decoration: underline;
|
||||
}
|
||||
.hidden-caption-content {
|
||||
display: none;
|
||||
}`)
|
||||
|
||||
function initLightBox () {
|
||||
const lightbox = new PhotoSwipeLightbox({
|
||||
gallery: '.prompt_image_output',
|
||||
children: 'a',
|
||||
pswpModule: () =>
|
||||
import('/extensions/comfyui-mixlab-nodes/lib/photoswipe.esm.min.js')
|
||||
pswpModule: () => import('/mixlab/app/lib/photoswipe.esm.min.js')
|
||||
})
|
||||
|
||||
lightbox.on('uiRegister', function () {
|
||||
@@ -100,7 +86,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
|
||||
return {
|
||||
transformOrigin: '0 0',
|
||||
transform: transform,
|
||||
left: `0`,
|
||||
left:
|
||||
document.querySelector('.comfy-menu').style.display === 'none'
|
||||
? `60px`
|
||||
: `0`,
|
||||
top: `0`,
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
@@ -178,7 +167,7 @@ app.registerExtension({
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
|
||||
const mutable_prompt = this.widgets.filter(
|
||||
w => w.name == 'mutable_prompt'
|
||||
)[0]
|
||||
@@ -190,7 +179,12 @@ app.registerExtension({
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1])
|
||||
get_position_style(
|
||||
ctx,
|
||||
widget_width,
|
||||
y + widget_height + 24,
|
||||
node.size[1]
|
||||
)
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -207,7 +201,7 @@ app.registerExtension({
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid; height: 30px;min-width: 122px;
|
||||
border-style: solid;height: 30px;min-width: 122px;
|
||||
`
|
||||
|
||||
// const btn=document.createElement('button');
|
||||
@@ -266,7 +260,6 @@ app.registerExtension({
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'RandomPrompt') {
|
||||
|
||||
}
|
||||
}
|
||||
})
|
||||
@@ -408,7 +401,7 @@ const _createResult = async (node, widget, message) => {
|
||||
const width = node.size[0] * 0.5 - 12
|
||||
|
||||
let height_add = 0
|
||||
|
||||
|
||||
for (let index = 0; index < message._images.length; index++) {
|
||||
const imgs = message._images[index]
|
||||
|
||||
@@ -559,7 +552,7 @@ app.registerExtension({
|
||||
|
||||
let cards = widget.div.querySelectorAll('.card')
|
||||
if (cards.length == 0) node.size = [280, 120]
|
||||
if(widget.value) _createResult(node, widget, widget.value)
|
||||
if (widget.value) _createResult(node, widget, widget.value)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
@@ -19,7 +19,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
|
||||
return {
|
||||
transformOrigin: '0 0',
|
||||
transform: transform,
|
||||
left: `0`,
|
||||
left:
|
||||
document.querySelector('.comfy-menu').style.display === 'none'
|
||||
? `60px`
|
||||
: `0`,
|
||||
top: `0`,
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
|
||||
@@ -21,7 +21,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
|
||||
return {
|
||||
transformOrigin: '0 0',
|
||||
transform: transform,
|
||||
left: `0`,
|
||||
left:
|
||||
document.querySelector('.comfy-menu').style.display === 'none'
|
||||
? `60px`
|
||||
: `0`,
|
||||
top: '0',
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
|
||||
@@ -6,6 +6,10 @@ window._bg_img = null
|
||||
* draws the back canvas (the one containing the background and the connections)
|
||||
* @method drawBackCanvas
|
||||
**/
|
||||
|
||||
// 判断是否是新版的,LGraphCanvas.prototype.drawBackCanvas.toString().match('window.devicePixelRatio')
|
||||
let scale=LGraphCanvas.prototype.drawBackCanvas.toString().match('window.devicePixelRatio')?window.devicePixelRatio:1;
|
||||
|
||||
LGraphCanvas.prototype.drawBackCanvas = function () {
|
||||
var canvas = this.bgcanvas
|
||||
if (
|
||||
@@ -59,7 +63,8 @@ LGraphCanvas.prototype.drawBackCanvas = function () {
|
||||
//reset in case of error
|
||||
if (!this.viewport) {
|
||||
ctx.restore()
|
||||
ctx.setTransform(1, 0, 0, 1, 0, 0)
|
||||
// ctx.setTransform(1, 0, 0, 1, 0, 0)
|
||||
ctx.setTransform(scale, 0, 0, scale, 0, 0)
|
||||
}
|
||||
this.visible_links.length = 0
|
||||
|
||||
|
||||
@@ -1,46 +1,13 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
import {
|
||||
loadExternalScript,
|
||||
updateLLMAPIKey,
|
||||
get_position_style,
|
||||
getLocalData
|
||||
} from './common.js'
|
||||
|
||||
const getLocalData = key => {
|
||||
let data = {}
|
||||
try {
|
||||
data = JSON.parse(localStorage.getItem(key)) || {}
|
||||
} catch (error) {
|
||||
return {}
|
||||
}
|
||||
return data
|
||||
}
|
||||
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'
|
||||
}
|
||||
}
|
||||
loadExternalScript('/mixlab/app/lib/pickr.min.js')
|
||||
|
||||
function hexToRGBA (hexColor) {
|
||||
var hex = hexColor.replace('#', '')
|
||||
@@ -62,7 +29,7 @@ app.registerExtension({
|
||||
init () {
|
||||
$el('link', {
|
||||
rel: 'stylesheet',
|
||||
href: '/extensions/comfyui-mixlab-nodes/lib/classic.min.css',
|
||||
href: '/mixlab/app/lib/classic.min.css',
|
||||
parent: document.head
|
||||
})
|
||||
|
||||
@@ -122,7 +89,7 @@ app.registerExtension({
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
// console.log('Color nodeData', this.widgets)
|
||||
// console.log('Color nodeData', this.div)
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
@@ -273,19 +240,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,22 +264,18 @@ 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') {
|
||||
@@ -323,7 +286,6 @@ 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 () {
|
||||
@@ -331,7 +293,6 @@ app.registerExtension({
|
||||
min_max(this)
|
||||
}
|
||||
}
|
||||
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'IntNumber') {
|
||||
@@ -340,22 +301,157 @@ 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_llm_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 rowHeight = this.rowHeight
|
||||
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))
|
||||
updateLLMAPIKey(data[this.id])
|
||||
})
|
||||
|
||||
return div
|
||||
}
|
||||
|
||||
let inputKey = inputDiv('_mixlab_llm_api_key', 'Key')
|
||||
|
||||
widget.div.appendChild(inputKey)
|
||||
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
inputKey.remove()
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
// const processMouseWheel=app.canvas.processMouseWheel
|
||||
// app.canvas.processMouseWheel=()=>{
|
||||
// console.log(app.canvas.ds.scale)
|
||||
// return processMouseWheel?.()
|
||||
// }
|
||||
|
||||
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_llm_api_key')
|
||||
|
||||
let id = node.id
|
||||
if (widget.div.querySelector('.Key'))
|
||||
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
|
||||
|
||||
if (apiKey[id]) updateLLMAPIKey(apiKey[id])
|
||||
}
|
||||
},
|
||||
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_llm_api_key')
|
||||
|
||||
let id = node.id
|
||||
|
||||
if (widget.div.querySelector('.Key'))
|
||||
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
|
||||
|
||||
if (apiKey[id]) updateLLMAPIKey(apiKey[id])
|
||||
}
|
||||
}, 1000)
|
||||
}
|
||||
})
|
||||
|
||||