Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c5392aa237 | ||
|
|
16a2b55fa1 | ||
|
|
b423b09ff3 | ||
|
|
e66add88cb | ||
|
|
32b22c39cb | ||
|
|
259baac177 | ||
|
|
67ef8c13a8 | ||
|
|
b2bb1876de | ||
|
|
cda4e626e7 | ||
|
|
d835aff0cb | ||
|
|
21e1967c5e | ||
|
|
c9b5baf4d9 | ||
|
|
67c974c96e | ||
|
|
b46ccb03c9 | ||
|
|
0ecf98e08b | ||
|
|
f024034724 | ||
|
|
327a21f009 | ||
|
|
cfc51532b8 | ||
|
|
00988f92e4 | ||
|
|
868c6085a8 | ||
|
|
a47a56bda0 | ||
|
|
3667b42b2f | ||
|
|
7d142d7d62 | ||
|
|
24863e2ed3 | ||
|
|
fe8b526bbb | ||
|
|
6298be393a | ||
|
|
3a7853f9cc | ||
|
|
4a9413c83d | ||
|
|
21b04d62ae | ||
|
|
96929b6d7c | ||
|
|
07712d80a5 | ||
|
|
10c9eff16f | ||
|
|
edd7af986d | ||
|
|
1dc31927e3 | ||
|
|
36ef7d25ef | ||
|
|
b766b8b65d | ||
|
|
6579ff20b4 | ||
|
|
2fbee59c3e | ||
|
|
d3aaa19148 | ||
|
|
e32a3675fc | ||
|
|
b72e7dda08 | ||
|
|
0f77f28a95 | ||
|
|
289f83675b | ||
|
|
36633b4c72 | ||
|
|
4f45457811 | ||
|
|
c39890cd64 | ||
|
|
90f1e49263 | ||
|
|
9a1cf205db | ||
|
|
45eacb6a50 | ||
|
|
6cb2b57463 | ||
|
|
59f654fa39 | ||
|
|
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 | ||
|
|
7ada28258c | ||
|
|
bd312afd00 |
@@ -7,15 +7,19 @@ on:
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'shadowcz007' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
|
||||
@@ -6,8 +6,20 @@
|
||||
商务合作请联系 389570357@qq.com
|
||||
For business cooperation, please contact email 389570357@qq.com
|
||||
|
||||

|
||||
|
||||
##### `最新`:
|
||||
|
||||
- 新增[fal.ai](https://fal.ai/dashboard)的视频生成:Kling、RunwayGen3、LumaDreamMachine,[工作流下载](./workflow/video-all-in-one-test-workflow.json)
|
||||
|
||||
- 新增 SimulateDevDesignDiscussions,需要安装[swarm](https://github.com/openai/swarm)和[Comfyui-ChatTTS](https://github.com/shadowcz007/Comfyui-ChatTTS),[工作流下载](./workflow/swarm制作的播客节点workflow.json)
|
||||
|
||||
- 新增 SenseVoice
|
||||
|
||||
- [新增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。
|
||||
|
||||
+306
-141
@@ -32,7 +32,7 @@ _URL_=None
|
||||
# except:
|
||||
# print("##nodes.ChatGPT ImportError")
|
||||
|
||||
from .nodes.ChatGPT import openai_client
|
||||
# from .nodes.ChatGPT import openai_client
|
||||
|
||||
from .nodes.RembgNode import get_rembg_models,U2NET_HOME,run_briarmbg,run_rembg
|
||||
|
||||
@@ -64,7 +64,7 @@ def is_installed(package, package_overwrite=None,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
|
||||
@@ -76,7 +76,7 @@ def is_installed(package, package_overwrite=None,auto_install=True):
|
||||
print(package+'## OK')
|
||||
|
||||
return is_has
|
||||
|
||||
|
||||
try:
|
||||
import OpenSSL
|
||||
except ImportError:
|
||||
@@ -211,11 +211,11 @@ def read_workflow_json_files_all(folder_path):
|
||||
data.append(file_info)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
|
||||
sorted_data = sorted(data, key=lambda x: x['date'], reverse=True)
|
||||
return sorted_data
|
||||
|
||||
# workflow
|
||||
# workflow
|
||||
def read_workflow_json_files(folder_path ):
|
||||
json_files = []
|
||||
for filename in os.listdir(folder_path):
|
||||
@@ -263,12 +263,12 @@ def get_my_workflow_for_app(filename="my_workflow_app.json",category="",is_all=F
|
||||
apps=[]
|
||||
if filename==None:
|
||||
|
||||
#TODO 支持目录内遍历
|
||||
#TODO 支持目录内遍历
|
||||
if is_all:
|
||||
data=read_workflow_json_files_all(category_path)
|
||||
else:
|
||||
data=read_workflow_json_files(category_path)
|
||||
|
||||
|
||||
i=0
|
||||
for item in data:
|
||||
# print(item)
|
||||
@@ -325,11 +325,11 @@ def get_my_workflow_for_app(filename="my_workflow_app.json",category="",is_all=F
|
||||
}]
|
||||
except Exception as e:
|
||||
print("发生异常:", str(e))
|
||||
|
||||
|
||||
# 这个代码不需要
|
||||
# if len(apps)==1 and category!='' and category!=None:
|
||||
data=read_workflow_json_files(category_path)
|
||||
|
||||
|
||||
for item in data:
|
||||
x=item["data"]
|
||||
# print(apps[0]['filename'] ,item["filename"])
|
||||
@@ -371,9 +371,9 @@ def save_prompt_result(id,data):
|
||||
if os.path.exists(prompt_result_path):
|
||||
with open(prompt_result_path) as json_file:
|
||||
prompt_result = json.load(json_file)
|
||||
|
||||
|
||||
prompt_result[id]=data
|
||||
|
||||
|
||||
with open(prompt_result_path, 'w') as file:
|
||||
json.dump(prompt_result, file)
|
||||
return prompt_result_path
|
||||
@@ -403,16 +403,16 @@ def save_workflow_for_app(data,filename="my_workflow_app.json",category=""):
|
||||
category_path=os.path.join(app_path,category)
|
||||
if not os.path.exists(category_path):
|
||||
os.mkdir(category_path)
|
||||
|
||||
|
||||
app_workflow_path=os.path.join(category_path, filename)
|
||||
|
||||
|
||||
try:
|
||||
output_str = json.dumps(data['output'])
|
||||
data['app']['id']=calculate_md5(output_str)
|
||||
# id=data['app']['id']
|
||||
except Exception as e:
|
||||
print("发生异常:", str(e))
|
||||
|
||||
|
||||
with open(app_workflow_path, 'w') as file:
|
||||
json.dump(data, file)
|
||||
return filename
|
||||
@@ -439,7 +439,7 @@ _original_request = aiohttp.ClientSession._request
|
||||
# 定义新的 get 方法
|
||||
async def new_request(self, method, url, *args, **kwargs):
|
||||
# 检查环境变量以确定是否使用代理
|
||||
proxy = os.environ.get('HTTP_PROXY') or os.environ.get('HTTPS_PROXY') or os.environ.get('http_proxy') or os.environ.get('https_proxy')
|
||||
proxy = os.environ.get('HTTP_PROXY') or os.environ.get('HTTPS_PROXY') or os.environ.get('http_proxy') or os.environ.get('https_proxy')
|
||||
# print('Proxy Config:',proxy)
|
||||
if proxy and 'proxy' not in kwargs:
|
||||
kwargs['proxy'] = proxy
|
||||
@@ -470,7 +470,7 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
|
||||
|
||||
# if not await check_port_available(address, port):
|
||||
# raise RuntimeError(f"Port {port} is already in use.")
|
||||
|
||||
|
||||
http_success = False
|
||||
http_port=port
|
||||
for i in range(11): # 尝试最多11次
|
||||
@@ -483,7 +483,7 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
|
||||
|
||||
if not http_success:
|
||||
raise RuntimeError(f"Ports {port} to {port + 10} are all in use.")
|
||||
|
||||
|
||||
|
||||
# site = web.TCPSite(runner, address, port)
|
||||
# await site.start()
|
||||
@@ -526,7 +526,7 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
|
||||
address = '127.0.0.1'
|
||||
if address=='0.0.0.0':
|
||||
address = '127.0.0.1'
|
||||
|
||||
|
||||
if verbose:
|
||||
|
||||
logging.info("\n")
|
||||
@@ -535,10 +535,14 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
|
||||
import socket
|
||||
|
||||
hostname = socket.gethostname()
|
||||
ip_address = socket.gethostbyname(hostname)
|
||||
|
||||
# print(f"本机的IP地址是: {ip_address}")
|
||||
# logging.debug("hostname:", hostname)
|
||||
try:
|
||||
ip_address = socket.gethostbyname(hostname)
|
||||
except Exception as e:
|
||||
logging.debug("[mixlab]gethostbyname() downgraded due to exception:", e)
|
||||
ip_address = socket.gethostbyname("")
|
||||
|
||||
# print(f"本机的 IP 地址是:{ip_address}")
|
||||
|
||||
# print("\033[93mStarting server\n")
|
||||
logging.info("\033[93mTo see the GUI go to: http://{}:{} or http://{}:{}".format(ip_address, http_port,address,http_port))
|
||||
@@ -556,7 +560,7 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
|
||||
call_on_start(scheme,address, http_port)
|
||||
except:
|
||||
call_on_start(address,http_port)
|
||||
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error starting the server: {e}")
|
||||
@@ -586,9 +590,9 @@ async def mixlab_hander(request):
|
||||
return web.json_response(data)
|
||||
|
||||
# llm的api key,使用硅基流动
|
||||
@routes.post('/mixlab/llm_api_key')
|
||||
async def mixlab_llm_api_key_handler(request):
|
||||
data = await request.json()
|
||||
@routes.post('/mixlab/llm_api_key')
|
||||
async def mixlab_llm_api_key_handler(request):
|
||||
data = await request.json()
|
||||
api_key = data.get('key')
|
||||
|
||||
app_folder = os.path.join(current_path, "app")
|
||||
@@ -670,7 +674,7 @@ async def static_file_handler(request):
|
||||
filename = request.match_info['filename']
|
||||
file_path = os.path.join(current_path, "webApp", filename)
|
||||
print(file_path)
|
||||
|
||||
|
||||
if os.path.exists(file_path) and os.path.isfile(file_path):
|
||||
if filename.endswith('.js'):
|
||||
content_type = 'application/javascript'
|
||||
@@ -682,7 +686,7 @@ async def static_file_handler(request):
|
||||
content_type = 'image/svg+xml'
|
||||
else:
|
||||
content_type = 'application/octet-stream'
|
||||
|
||||
|
||||
with open(file_path, 'r', encoding='utf-8', errors='ignore') as f:
|
||||
file_data = f.read()
|
||||
return web.Response(text=file_data, content_type=content_type)
|
||||
@@ -723,7 +727,7 @@ async def mixlab_workflow_hander(request):
|
||||
admin=data['admin']
|
||||
|
||||
ds=get_my_workflow_for_app(filename,category,admin)
|
||||
data=[]
|
||||
data=[]
|
||||
for json_data in ds:
|
||||
# 不传给前端
|
||||
if 'output' in json_data['data']:
|
||||
@@ -760,7 +764,7 @@ async def mixlab_workflow_hander(request):
|
||||
async def nodes_map_hander(request):
|
||||
data = await request.json()
|
||||
result={}
|
||||
try:
|
||||
try:
|
||||
result={
|
||||
'data':get_nodes_map(),
|
||||
'status':'success',
|
||||
@@ -793,7 +797,7 @@ async def get_checkpoints(request):
|
||||
names=get_rembg_models(U2NET_HOME)
|
||||
except:
|
||||
print("rembg none")
|
||||
|
||||
|
||||
return web.json_response({"names":names,"types":list(folder_paths.folder_names_and_paths.keys())})
|
||||
|
||||
|
||||
@@ -805,7 +809,7 @@ async def rembg_hander(request):
|
||||
|
||||
data_base64=remove_base64_prefix(data['base64'])
|
||||
image_data = base64.b64decode(data_base64)
|
||||
|
||||
|
||||
# 创建一个BytesIO对象
|
||||
image_stream = io.BytesIO(image_data)
|
||||
|
||||
@@ -821,8 +825,8 @@ async def rembg_hander(request):
|
||||
rgba_images[0].save(buf, format='PNG')
|
||||
img_bytes = buf.getvalue()
|
||||
img_base64 = base64.b64encode(img_bytes).decode('utf-8')
|
||||
|
||||
try:
|
||||
|
||||
try:
|
||||
result={
|
||||
'data':img_base64,
|
||||
'model':model,
|
||||
@@ -848,7 +852,7 @@ async def rembg_hander(request):
|
||||
# res=get_prompt_result()
|
||||
# except Exception as e:
|
||||
# print('/mixlab/prompt_result',False,e)
|
||||
|
||||
|
||||
# return web.json_response({"result":res})
|
||||
|
||||
# 种子设置
|
||||
@@ -860,15 +864,15 @@ def random_seed(seed, data):
|
||||
if id in seed:
|
||||
if 'seed' in value['inputs'] and not isinstance(value['inputs']['seed'], list) and seed[id] in ['increment', 'decrement', 'randomize']:
|
||||
value['inputs']['seed'] = round(random.random() * max_seed)
|
||||
|
||||
|
||||
if 'noise_seed' in value['inputs'] and not isinstance(value['inputs']['noise_seed'], list) and seed[id] in ['increment', 'decrement', 'randomize']:
|
||||
value['inputs']['noise_seed'] = round(random.random() * max_seed)
|
||||
|
||||
|
||||
if value.get('class_type') == "Seed_" and seed[id] in ['increment', 'decrement', 'randomize']:
|
||||
value['inputs']['seed'] = round(random.random() * max_seed)
|
||||
|
||||
|
||||
# print('new Seed', value)
|
||||
|
||||
|
||||
return data
|
||||
|
||||
|
||||
@@ -893,9 +897,9 @@ async def mixlab_post_prompt(request):
|
||||
apps=json_data['apps']
|
||||
except:
|
||||
apps=get_my_workflow_for_app(json_data['filename'],json_data['category'],False)
|
||||
|
||||
|
||||
prompt=json_data['prompt'] if 'prompt' in json_data else None
|
||||
|
||||
|
||||
if len(apps)>0:
|
||||
# 取到prompt
|
||||
prompt=apps[0]['data']['output']
|
||||
@@ -917,13 +921,13 @@ async def mixlab_post_prompt(request):
|
||||
for inp in input_data:
|
||||
id=inp['id']
|
||||
if prompt[id]['class_type']==inp['class_type']:
|
||||
prompt[id]['inputs'].update(inp['inputs'])
|
||||
prompt[id]['inputs'].update(inp['inputs'])
|
||||
|
||||
|
||||
if prompt==None:
|
||||
return web.json_response({"error": "no prompt", "node_errors": []}, status=400)
|
||||
else:
|
||||
# 种子更新
|
||||
# 种子更新
|
||||
'''
|
||||
"seed": {
|
||||
"45": "randomize",
|
||||
@@ -935,7 +939,7 @@ async def mixlab_post_prompt(request):
|
||||
# print("#json_data",prompt)
|
||||
# 需要把apps处理成 prompt
|
||||
# 注意seed的处理
|
||||
|
||||
|
||||
if "number" in json_data:
|
||||
number = float(json_data['number'])
|
||||
else:
|
||||
@@ -1002,8 +1006,8 @@ from .nodes.ImageNode import DepthViewer_,ImageBatchToList_,ImageListToBatch_,Co
|
||||
# from .nodes.Vae import VAELoader,VAEDecode
|
||||
from .nodes.ScreenShareNode import ScreenShareNode,FloatingVideo
|
||||
|
||||
from .nodes.Audio import AudioPlayNode,SpeechRecognition,SpeechSynthesis
|
||||
from .nodes.Utils import KeyInput,IncrementingListNode,ListSplit,CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,TESTNODE_TOKEN,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
|
||||
from .nodes.Audio import AudioPlayNode,SpeechRecognition,SpeechSynthesis,AnalyzeAudioNone
|
||||
from .nodes.Utils import CreateJsonNode,KeyInput,IncrementingListNode,ListSplit,CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,TESTNODE_TOKEN,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
|
||||
from .nodes.Mask import PreviewMask_,MaskListReplace,MaskListMerge,OutlineMask,FeatheredMask
|
||||
|
||||
from .nodes.Style import ApplyVisualStylePrompting,StyleAlignedReferenceSampler,StyleAlignedBatchAlign,StyleAlignedSampleReferenceLatents
|
||||
@@ -1017,13 +1021,33 @@ NODE_CLASS_MAPPINGS = {
|
||||
"AppInfo":AppInfo,
|
||||
"TESTNODE_":TESTNODE_,
|
||||
"TESTNODE_TOKEN":TESTNODE_TOKEN,
|
||||
|
||||
# Prompt
|
||||
"RandomPrompt":RandomPrompt,
|
||||
# "LoraPrompt":LoraPrompt,
|
||||
"EmbeddingPrompt":EmbeddingPrompt,
|
||||
"PromptSlide":PromptSlide,
|
||||
"GLIGENTextBoxApply_Advanced":GLIGENTextBoxApply_Advanced,
|
||||
"PromptSimplification":PromptSimplification,
|
||||
|
||||
# Input
|
||||
"GridInput":GridInput,
|
||||
"ImagesPrompt_":ImagesPrompt,
|
||||
"KeyInput":KeyInput,
|
||||
"FloatSlider":FloatSlider,
|
||||
"IntNumber":IntNumber,
|
||||
"TextInput_":TextInput,
|
||||
"Font":FontInput,
|
||||
"LimitNumber":LimitNumber,
|
||||
|
||||
# Output
|
||||
"PromptImage":PromptImage,
|
||||
"SaveImageToLocal":SaveImageToLocal,
|
||||
"SaveImageAndMetadata_":SaveImageAndMetadata,
|
||||
"ComparingTwoFrames_":ComparingTwoFrames,
|
||||
"CreateJsonNode":CreateJsonNode,
|
||||
|
||||
# Image
|
||||
"MirroredImage":MirroredImage,
|
||||
"NoiseImage":NoiseImage,
|
||||
"GradientImage":GradientImage,
|
||||
@@ -1035,131 +1059,198 @@ NODE_CLASS_MAPPINGS = {
|
||||
"TextImage":TextImage,
|
||||
"EnhanceImage":EnhanceImage,
|
||||
"SvgImage":SvgImage,
|
||||
"3DImage":Image3D,
|
||||
"ImageColorTransfer":ImageColorTransfer,
|
||||
"ShowLayer":ShowLayer,
|
||||
"NewLayer":NewLayer,
|
||||
"ImageListToBatch_":ImageListToBatch_,
|
||||
"ImageBatchToList_":ImageBatchToList_,
|
||||
"CompositeImages_":CompositeImages,
|
||||
"ImageCropByAlpha":ImageCropByAlpha,
|
||||
"GetImageSize_":GetImageSize_,
|
||||
|
||||
# 3D
|
||||
"3DImage":Image3D,
|
||||
"DepthViewer": DepthViewer_,
|
||||
|
||||
# Color
|
||||
"ImageColorTransfer":ImageColorTransfer,
|
||||
"Color":ColorInput,
|
||||
|
||||
# Layer
|
||||
"ShowLayer":ShowLayer,
|
||||
"NewLayer":NewLayer,
|
||||
"MergeLayers":MergeLayers,
|
||||
"CompositeImages_":CompositeImages,
|
||||
"SplitImage":SplitImage,
|
||||
"CenterImage":CenterImage,
|
||||
"GridOutput":GridOutput,
|
||||
"GridDisplayAndSave":GridDisplayAndSave,
|
||||
"GridInput":GridInput,
|
||||
"MergeLayers":MergeLayers,
|
||||
|
||||
# Mask
|
||||
"SplitLongMask":SplitLongMask,
|
||||
"FeatheredMask":FeatheredMask,
|
||||
"SmoothMask":SmoothMask,
|
||||
"FaceToMask":FaceToMask,
|
||||
"AreaToMask":AreaToMask,
|
||||
"ImageCropByAlpha":ImageCropByAlpha,
|
||||
"ImagesPrompt_":ImagesPrompt,
|
||||
# "VAELoaderConsistencyDecoder":VAELoader,
|
||||
"SaveImageToLocal":SaveImageToLocal,
|
||||
"SaveImageAndMetadata_":SaveImageAndMetadata,
|
||||
"ComparingTwoFrames_":ComparingTwoFrames,
|
||||
# "VAEDecodeConsistencyDecoder":VAEDecode,
|
||||
"ScreenShare":ScreenShareNode,
|
||||
"FloatingVideo":FloatingVideo,
|
||||
|
||||
"SpeechRecognition":SpeechRecognition,
|
||||
"SpeechSynthesis":SpeechSynthesis,
|
||||
"KeyInput":KeyInput,
|
||||
"Color":ColorInput,
|
||||
"FloatSlider":FloatSlider,
|
||||
"IntNumber":IntNumber,
|
||||
"TextInput_":TextInput,
|
||||
"Font":FontInput,
|
||||
"TextToNumber":TextToNumber,
|
||||
"DynamicDelayProcessor":DynamicDelayProcessor,
|
||||
"MultiplicationNode":MultiplicationNode,
|
||||
"GetImageSize_":GetImageSize_,
|
||||
"SwitchByIndex":SwitchByIndex,
|
||||
"LimitNumber":LimitNumber,
|
||||
"OutlineMask":OutlineMask,
|
||||
"MaskListMerge_":MaskListMerge,
|
||||
"PreviewMask_":PreviewMask_,
|
||||
|
||||
# "VAELoaderConsistencyDecoder":VAELoader,
|
||||
# "VAEDecodeConsistencyDecoder":VAEDecode,
|
||||
|
||||
# Screen
|
||||
"ScreenShare":ScreenShareNode,
|
||||
"FloatingVideo":FloatingVideo,
|
||||
|
||||
# Audio
|
||||
"SpeechRecognition":SpeechRecognition,
|
||||
"SpeechSynthesis":SpeechSynthesis,
|
||||
"AudioPlay":AudioPlayNode,
|
||||
"AnalyzeAudio":AnalyzeAudioNone,
|
||||
|
||||
# Text
|
||||
"TextToNumber":TextToNumber,
|
||||
"JoinWithDelimiter":JoinWithDelimiter,
|
||||
|
||||
# Utils
|
||||
"MultiplicationNode":MultiplicationNode,
|
||||
"DynamicDelayProcessor":DynamicDelayProcessor,
|
||||
"SwitchByIndex":SwitchByIndex,
|
||||
"ListSplit_":ListSplit,
|
||||
|
||||
# Experiment
|
||||
"Seed_":CreateSeedNode,
|
||||
"CkptNames_":CreateCkptNames,
|
||||
"SamplerNames_":CreateSampler_names,
|
||||
"LoraNames_":CreateLoraNames,
|
||||
|
||||
# Style
|
||||
"ApplyVisualStylePrompting_":ApplyVisualStylePrompting,
|
||||
"StyleAlignedReferenceSampler_": StyleAlignedReferenceSampler,
|
||||
"StyleAlignedSampleReferenceLatents_": StyleAlignedSampleReferenceLatents,
|
||||
"StyleAlignedBatchAlign_": StyleAlignedBatchAlign,
|
||||
"ListSplit_":ListSplit,
|
||||
"MaskListReplace_":MaskListReplace,
|
||||
"StyleAlignedBatchAlign_": StyleAlignedBatchAlign,
|
||||
|
||||
# Video
|
||||
"MaskListReplace_":MaskListReplace,
|
||||
"IncrementingListNode_":IncrementingListNode,
|
||||
"PreviewMask_":PreviewMask_,
|
||||
"AudioPlay":AudioPlayNode,
|
||||
|
||||
"P5Input":P5Input
|
||||
}
|
||||
|
||||
# 一个包含节点友好/可读的标题的字典
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AppInfo":"App Info ♾️MixlabApp",
|
||||
"ScreenShare":"Screen Share ♾️Mixlab",
|
||||
"FloatingVideo":"Floating Video ♾️Mixlab",
|
||||
"TextImage":"Text Image ♾️Mixlab",
|
||||
|
||||
"Color":"Color Input ♾️MixlabApp",
|
||||
"TextInput_":"Text Input ♾️MixlabApp",
|
||||
"KeyInput":"API Key Input ♾️MixlabApp",
|
||||
"FloatSlider":"Float Slider Input ♾️MixlabApp",
|
||||
"IntNumber":"Int Input ♾️MixlabApp",
|
||||
"ImagesPrompt_":"Images Input ♾️MixlabApp",
|
||||
"SaveImageAndMetadata_":"Save Image Output ♾️MixlabApp",
|
||||
"ComparingTwoFrames_":"Comparing Two Frames ♾️MixlabApp",
|
||||
"ResizeImageMixlab":"Resize Image ♾️Mixlab",
|
||||
"AppInfo":"App Info ♾️Mixlab",
|
||||
"TESTNODE_":"TESTNODE_ ♾️Mixlab",
|
||||
"TESTNODE_TOKEN":"TESTNODE_TOKEN ♾️Mixlab",
|
||||
|
||||
# Prompt
|
||||
"RandomPrompt": "Random Prompt ♾️Mixlab",
|
||||
"PromptImage":"Output Prompt and Image ♾️Mixlab",
|
||||
"SplitLongMask":"Splitting a long image into sections",
|
||||
"VAELoaderConsistencyDecoder":"Consistency Decoder Loader",
|
||||
"VAEDecodeConsistencyDecoder":"Consistency Decoder Decode",
|
||||
|
||||
|
||||
"MergeLayers":"Merge Layers ♾️Mixlab",
|
||||
"SpeechSynthesis":"SpeechSynthesis ♾️Mixlab",
|
||||
"SpeechRecognition":"SpeechRecognition ♾️Mixlab",
|
||||
"3DImage":"3DImage ♾️Mixlab",
|
||||
"ImageListToBatch_":"Image List To Batch",
|
||||
"ImageBatchToList_":"Image Batch To List",
|
||||
"CompositeImages_":"Composite Images ♾️Mixlab",
|
||||
"DynamicDelayProcessor":"DynamicDelayByText ♾️Mixlab",
|
||||
"LaMaInpainting":"LaMaInpainting ♾️Mixlab",
|
||||
"EmbeddingPrompt":"Embedding Prompt ♾️Mixlab",
|
||||
"PromptSlide":"Prompt Slide ♾️Mixlab",
|
||||
"GLIGENTextBoxApply_Advanced":"GLIGEN TextBox Apply ♾️Mixlab",
|
||||
"PromptSimplification":"PromptSimplification ♾️Mixlab",
|
||||
"PromptGenerate_Mix":"Prompt Generate ♾️Mixlab",
|
||||
"ChinesePrompt_Mix":"Chinese Prompt ♾️Mixlab",
|
||||
"GamePal":"GamePal ♾️Mixlab",
|
||||
"RembgNode_Mix":"Remove Background ♾️Mixlab",
|
||||
|
||||
# Input
|
||||
"GridInput":"Grid Input ♾️Mixlab",
|
||||
"ImagesPrompt_":"Images Input ♾️Mixlab",
|
||||
"KeyInput":"API Key Input ♾️Mixlab",
|
||||
"FloatSlider":"Float Slider Input ♾️Mixlab",
|
||||
"IntNumber":"Int Input ♾️Mixlab",
|
||||
"TextInput_":"Text Input ♾️Mixlab",
|
||||
"Font":"Font Input ♾️Mixlab",
|
||||
"LimitNumber":"LimitNumber Input ♾️Mixlab",
|
||||
|
||||
# Output
|
||||
"PromptImage":"Output Prompt and Image ♾️Mixlab",
|
||||
"SaveImageToLocal":"Save Image To Local ♾️Mixlab",
|
||||
"SaveImageAndMetadata_":"Save Image Output ♾️Mixlab",
|
||||
"ComparingTwoFrames_":"Comparing Two Frames ♾️Mixlab",
|
||||
|
||||
# Image
|
||||
"MirroredImage":"MirroredImage ♾️Mixlab",
|
||||
"NoiseImage":"NoiseImage ♾️Mixlab",
|
||||
"GradientImage":"GradientImage ♾️Mixlab",
|
||||
"TransparentImage":"TransparentImage ♾️Mixlab",
|
||||
"ResizeImageMixlab":"Resize Image ♾️Mixlab",
|
||||
"LoadImagesFromPath":"Load Images From Path ♾️Mixlab",
|
||||
"LoadImagesFromURL":"Load Images From URL ♾️Mixlab",
|
||||
"LoadImagesToBatch":"Load Images(base64) ♾️Mixlab",
|
||||
"TextImage":"Text Image ♾️Mixlab",
|
||||
"EnhanceImage":"Enhance Image ♾️Mixlab",
|
||||
"SvgImage":"Svg Image ♾️Mixlab",
|
||||
"ImageListToBatch_":"Image List To Batch ♾️Mixlab",
|
||||
"ImageBatchToList_":"Image Batch To List ♾️Mixlab",
|
||||
"ImageCropByAlpha":"ImageCropByAlpha ♾️Mixlab",
|
||||
"GetImageSize_":"Get Image Size ♾️Mixlab",
|
||||
|
||||
# 3D
|
||||
"3DImage":"3DImage ♾️Mixlab",
|
||||
"DepthViewer": "Depth Viewer ♾️Mixlab",
|
||||
|
||||
# "VAELoaderConsistencyDecoder":"Consistency Decoder Loader",
|
||||
# "VAEDecodeConsistencyDecoder":"Consistency Decoder Decode",
|
||||
|
||||
# Color
|
||||
"ImageColorTransfer":"Image Color Transfer ♾️Mixlab",
|
||||
"Color":"Color Input ♾️MixlabApp",
|
||||
|
||||
# Layer
|
||||
"ShowLayer":"Show Layer ♾️Mixlab",
|
||||
"NewLayer":"New Layer ♾️Mixlab",
|
||||
"MergeLayers":"Merge Layers ♾️Mixlab",
|
||||
"CompositeImages_":"Composite Images ♾️Mixlab",
|
||||
"SplitImage":"Split Image ♾️Mixlab",
|
||||
"CenterImage":"Center Image ♾️Mixlab",
|
||||
"GridDisplayAndSave":"Grid Display And Save ♾️Mixlab",
|
||||
"GridOutput":"Grid Output ♾️Mixlab",
|
||||
|
||||
# Mask
|
||||
"SplitLongMask":"Splitting a long image into sections",
|
||||
"FeatheredMask":"Feathered Mask ♾️Mixlab",
|
||||
"SmoothMask":"Smooth Mask ♾️Mixlab",
|
||||
"FaceToMask":"Face To Mask ♾️Mixlab",
|
||||
"AreaToMask":"Area To Mask ♾️Mixlab",
|
||||
"OutlineMask":"Outline Mask ♾️Mixlab",
|
||||
"MaskListMerge_":"MaskList to Mask ♾️Mixlab",
|
||||
"PreviewMask_":"Preview Mask ♾️Mixlab",
|
||||
|
||||
# Screen
|
||||
"ScreenShare":"Screen Share ♾️Mixlab",
|
||||
"FloatingVideo":"Floating Video ♾️Mixlab",
|
||||
|
||||
# Audio
|
||||
"SpeechSynthesis":"SpeechSynthesis ♾️Mixlab",
|
||||
"SpeechRecognition":"SpeechRecognition ♾️Mixlab",
|
||||
"AudioPlay":"Preview Audio ♾️Mixlab",
|
||||
"AnalyzeAudio":"Analyze Audio ♾️Mixlab",
|
||||
|
||||
# Utils
|
||||
"DynamicDelayProcessor":"DynamicDelayByText ♾️Mixlab",
|
||||
"MultiplicationNode":"Math Operation ♾️Mixlab",
|
||||
"ListSplit_":"Split List ♾️Mixlab",
|
||||
"SwitchByIndex":"List Switch By Index ♾️Mixlab",
|
||||
"CreateJsonNode":"Create Json",
|
||||
|
||||
# "GamePal":"GamePal ♾️Mixlab",
|
||||
# Experiment
|
||||
"Seed_":"CreateSeedNode ♾️Mixlab",
|
||||
"CkptNames_":"CreateCkptNames ♾️Mixlab",
|
||||
"SamplerNames_":"CreateSampler_names ♾️Mixlab",
|
||||
"LoraNames_":"LoraName ♾️Mixlab",
|
||||
|
||||
# Style
|
||||
"ApplyVisualStylePrompting_":"Apply VisualStyle Prompting ♾️Mixlab",
|
||||
"StyleAlignedReferenceSampler_": "StyleAligned Reference Sampler ♾️Mixlab",
|
||||
"StyleAlignedSampleReferenceLatents_": "StyleAligned Sample Reference Latents ♾️Mixlab",
|
||||
"StyleAlignedBatchAlign_": "StyleAligned Batch Align ♾️Mixlab",
|
||||
|
||||
# Video
|
||||
"MaskListReplace_":"MaskList Replace ♾️Mixlab",
|
||||
"IncrementingListNode_":"Create Incrementing Number List ♾️Mixlab",
|
||||
"LoadVideoAndSegment_":"Load Video And Segment ♾️Mixlab",
|
||||
"VideoCombine_Adv":"Video Combine ♾️Mixlab",
|
||||
"MaskListMerge_":"MaskList to Mask ♾️Mixlab",
|
||||
"ListSplit_":"Split List ♾️Mixlab",
|
||||
"MaskListReplace_":"MaskList Replace ♾️Mixlab",
|
||||
"ImageListReplace_":"ImageList Replace ♾️Mixlab",
|
||||
"SwitchByIndex":"List Switch By Index ♾️Mixlab",
|
||||
"GLIGENTextBoxApply_Advanced":"GLIGEN TextBox Apply ♾️Mixlab",
|
||||
"GridDisplayAndSave":"Grid Display And Save ♾️Mixlab",
|
||||
"GridInput":"Grid Input ♾️Mixlab",
|
||||
"GridOutput":"Grid Output ♾️Mixlab",
|
||||
"GetImageSize_":"Get Image Size ♾️Mixlab",
|
||||
"IncrementingListNode_":"Create Incrementing Number List ♾️Mixlab",
|
||||
"LoadImagesToBatch":"Load Images(base64) ♾️Mixlab",
|
||||
"PreviewMask_":"Preview Mask",
|
||||
"AudioPlay":"Preview Audio ♾️Mixlab",
|
||||
|
||||
"MultiplicationNode":"Math Operation ♾️Mixlab",
|
||||
|
||||
"P5Input":"P5 Input ♾️Mixlab for test"
|
||||
"P5Input":"P5 Input ♾️Mixlab for test"
|
||||
}
|
||||
|
||||
# web ui的节点功能
|
||||
@@ -1170,26 +1261,32 @@ logging.info('\033[91m ### Mixlab Nodes: \033[93mLoaded')
|
||||
# print('\033[91m ### Mixlab Nodes: \033[93mLoaded')
|
||||
|
||||
try:
|
||||
from .nodes.ChatGPT import JsonRepair,ChatGPTNode,ShowTextForGPT,CharacterInText,TextSplitByDelimiter,SiliconflowFreeNode
|
||||
from .nodes.ChatGPT import SimulateDevDesignDiscussions,SiliconflowTextToImageNode,JsonRepair,ChatGPTNode,ShowTextForGPT,CharacterInText,TextSplitByDelimiter,SiliconflowFreeNode
|
||||
logging.info('ChatGPT.available True')
|
||||
|
||||
NODE_CLASS_MAPPINGS_V = {
|
||||
NODE_CLASS_MAPPINGS_V = {
|
||||
"ChatGPTOpenAI":ChatGPTNode,
|
||||
"SiliconflowLLM":SiliconflowFreeNode,
|
||||
"SiliconflowTextToImageNode":SiliconflowTextToImageNode,
|
||||
"ShowTextForGPT":ShowTextForGPT,
|
||||
"CharacterInText":CharacterInText,
|
||||
"TextSplitByDelimiter":TextSplitByDelimiter,
|
||||
"JsonRepair":JsonRepair
|
||||
"JsonRepair":JsonRepair,
|
||||
|
||||
"SimulateDevDesignDiscussions":SimulateDevDesignDiscussions
|
||||
}
|
||||
|
||||
# 一个包含节点友好/可读的标题的字典
|
||||
NODE_DISPLAY_NAME_MAPPINGS_V = {
|
||||
NODE_DISPLAY_NAME_MAPPINGS_V = {
|
||||
"ChatGPTOpenAI":"ChatGPT & Local LLM ♾️Mixlab",
|
||||
"SiliconflowLLM":"LLM Siliconflow ♾️Mixlab",
|
||||
"SiliconflowTextToImageNode":"TextToImage Siliconflow ♾️Mixlab",
|
||||
"ShowTextForGPT":"Show Text ♾️MixlabApp",
|
||||
"CharacterInText":"Character In Text",
|
||||
"TextSplitByDelimiter":"Text Split By Delimiter",
|
||||
"JsonRepair":"Json Repair"
|
||||
"JsonRepair":"Json Repair",
|
||||
|
||||
"SimulateDevDesignDiscussions":"SimulateDevDesignDiscussions ♾️Mixlab Podcast"
|
||||
}
|
||||
|
||||
|
||||
@@ -1215,6 +1312,7 @@ try:
|
||||
logging.info('LaMaInpainting.available {}'.format(LaMaInpainting.available))
|
||||
if LaMaInpainting.available:
|
||||
NODE_CLASS_MAPPINGS['LaMaInpainting']=LaMaInpainting
|
||||
NODE_DISPLAY_NAME_MAPPINGS['LaMaInpainting']="LaMaInpainting ♾️Mixlab"
|
||||
except Exception as e:
|
||||
logging.info('LaMaInpainting.available False')
|
||||
|
||||
@@ -1223,6 +1321,7 @@ try:
|
||||
logging.info('ClipInterrogator.available {}'.format(ClipInterrogator.available))
|
||||
if ClipInterrogator.available:
|
||||
NODE_CLASS_MAPPINGS['ClipInterrogator']=ClipInterrogator
|
||||
NODE_DISPLAY_NAME_MAPPINGS['ClipInterrogator']="Clip Interrogator ♾️Mixlab"
|
||||
except Exception as e:
|
||||
logging.info('ClipInterrogator.available False')
|
||||
|
||||
@@ -1242,13 +1341,14 @@ try:
|
||||
logging.info('RembgNode_.available {}'.format(RembgNode_.available))
|
||||
if RembgNode_.available:
|
||||
NODE_CLASS_MAPPINGS['RembgNode_Mix']=RembgNode_
|
||||
NODE_DISPLAY_NAME_MAPPINGS['RembgNode_Mix']="Remove Background ♾️Mixlab"
|
||||
except Exception as e:
|
||||
logging.info('RembgNode_.available False' )
|
||||
|
||||
|
||||
try:
|
||||
from .nodes.Video import GenerateFramesByCount,scenesNode_,CombineAudioVideo,VideoCombine_Adv,LoadVideoAndSegment,ImageListReplace,VAEEncodeForInpaint_Frames,LoadAndCombinedAudio_
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS_V = {
|
||||
"VAEEncodeForInpaint_Frames":VAEEncodeForInpaint_Frames,
|
||||
"ImageListReplace_":ImageListReplace,
|
||||
@@ -1257,7 +1357,7 @@ try:
|
||||
"LoadAndCombinedAudio_":LoadAndCombinedAudio_,
|
||||
"CombineAudioVideo":CombineAudioVideo,
|
||||
"ScenesNode_":scenesNode_,
|
||||
"GenerateFramesByCount":GenerateFramesByCount
|
||||
"GenerateFramesByCount":GenerateFramesByCount
|
||||
}
|
||||
|
||||
# 一个包含节点友好/可读的标题的字典
|
||||
@@ -1268,7 +1368,7 @@ try:
|
||||
"VideoCombine_Adv":"Video Combine",
|
||||
"LoadAndCombinedAudio_":"Load And Combined Audio",
|
||||
"CombineAudioVideo":"Combine Audio Video",
|
||||
"ScenesNode_":"Select Scene",
|
||||
"ScenesNode_":"Select Scene",
|
||||
"GenerateFramesByCount":"Generate Frames By Count"
|
||||
}
|
||||
|
||||
@@ -1286,27 +1386,92 @@ try:
|
||||
# logging.info( folder_paths.get_temp_directory())
|
||||
NODE_CLASS_MAPPINGS['LoadTripoSRModel_']=LoadTripoSRModel
|
||||
NODE_DISPLAY_NAME_MAPPINGS["LoadTripoSRModel_"]= "Load TripoSR Model"
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS['TripoSRSampler_']=TripoSRSampler
|
||||
NODE_DISPLAY_NAME_MAPPINGS["TripoSRSampler_"]= "TripoSR Sampler"
|
||||
|
||||
NODE_CLASS_MAPPINGS['SaveTripoSRMesh']=SaveTripoSRMesh
|
||||
NODE_DISPLAY_NAME_MAPPINGS["SaveTripoSRMesh"]= "Save TripoSR Mesh"
|
||||
|
||||
|
||||
|
||||
except Exception as e:
|
||||
logging.info('TripoSR.available False' )
|
||||
|
||||
from .nodes.MiniCPMNode import MiniCPM_VQA_Simple
|
||||
try:
|
||||
|
||||
from .nodes.MiniCPMNode import MiniCPM_VQA_Simple
|
||||
logging.info('MiniCPMNode.available')
|
||||
# logging.info( folder_paths.get_temp_directory())
|
||||
NODE_CLASS_MAPPINGS['MiniCPM_VQA_Simple']=MiniCPM_VQA_Simple
|
||||
NODE_DISPLAY_NAME_MAPPINGS["MiniCPM_VQA_Simple"]= "MiniCPM VQA Simple"
|
||||
|
||||
|
||||
except Exception as e:
|
||||
logging.info('MiniCPMNode.available False' )
|
||||
|
||||
|
||||
try:
|
||||
from .nodes.scenedetectNode import ScenedetectNode_,SceneInfoNode
|
||||
logging.info('Scenedetect.available')
|
||||
NODE_CLASS_MAPPINGS['ScenedetectNode_']=ScenedetectNode_
|
||||
NODE_CLASS_MAPPINGS['SceneInfoNode']=SceneInfoNode
|
||||
NODE_DISPLAY_NAME_MAPPINGS["ScenedetectNode_"]= "Video Scene Detect"
|
||||
NODE_DISPLAY_NAME_MAPPINGS["SceneInfoNode"]= "Scene Info"
|
||||
|
||||
except Exception as e:
|
||||
logging.info('Scenedetect.available False' )
|
||||
|
||||
|
||||
try:
|
||||
from .nodes.FishSpeech import LoadVQGAN,AudioToPrompt,Prompt2Semantic,Semantic2Audio
|
||||
logging.info('FishSpeech.available')
|
||||
NODE_CLASS_MAPPINGS['LoadVQGAN']=LoadVQGAN
|
||||
NODE_CLASS_MAPPINGS['AudioToPrompt']=AudioToPrompt
|
||||
NODE_CLASS_MAPPINGS['Prompt2Semantic']=Prompt2Semantic
|
||||
NODE_CLASS_MAPPINGS['Semantic2Audio']=Semantic2Audio
|
||||
NODE_DISPLAY_NAME_MAPPINGS["LoadVQGAN"]= "Load VQGAN"
|
||||
NODE_DISPLAY_NAME_MAPPINGS["AudioToPrompt"]= "Audio To Prompt"
|
||||
NODE_DISPLAY_NAME_MAPPINGS["Prompt2Semantic"]= "Prompt To Semantic"
|
||||
NODE_DISPLAY_NAME_MAPPINGS["Semantic2Audio"]= "Semantic To Audio"
|
||||
|
||||
except Exception as e:
|
||||
logging.info('FishSpeech.available False' )
|
||||
|
||||
try:
|
||||
from .nodes.SenseVoice import SenseVoiceNode
|
||||
logging.info('SenseVoice.available')
|
||||
NODE_CLASS_MAPPINGS['SenseVoiceNode']=SenseVoiceNode
|
||||
NODE_DISPLAY_NAME_MAPPINGS["SenseVoiceNode"]= "Sense Voice ♾️Mixlab"
|
||||
|
||||
except Exception as e:
|
||||
logging.info('SenseVoice.available False' )
|
||||
|
||||
try:
|
||||
from .nodes.Whisper import LoadWhisperModel,WhisperTranscribe
|
||||
logging.info('Whisper.available')
|
||||
NODE_CLASS_MAPPINGS['LoadWhisperModel_']=LoadWhisperModel
|
||||
NODE_CLASS_MAPPINGS['WhisperTranscribe_']=WhisperTranscribe
|
||||
NODE_DISPLAY_NAME_MAPPINGS["LoadWhisperModel_"]= "Load Whisper Model ♾️Mixlab"
|
||||
NODE_DISPLAY_NAME_MAPPINGS["WhisperTranscribe_"]= "Whisper Transcribe ♾️Mixlab"
|
||||
|
||||
except Exception as e:
|
||||
logging.info('Whisper.available False' )
|
||||
|
||||
|
||||
try:
|
||||
from .nodes.FalVideo import VideoGenKlingNode,VideoGenLumaDreamMachineNode,VideoGenRunwayGen3Node,LoadVideoFromURL
|
||||
logging.info('FalVideo.available')
|
||||
# Update Node class mappings
|
||||
NODE_CLASS_MAPPINGS['VideoGenKlingNode']=VideoGenKlingNode
|
||||
NODE_CLASS_MAPPINGS['VideoGenRunwayGen3Node']=VideoGenRunwayGen3Node
|
||||
NODE_CLASS_MAPPINGS['VideoGenLumaDreamMachineNode']=VideoGenLumaDreamMachineNode
|
||||
NODE_CLASS_MAPPINGS['LoadVideoFromURL']=LoadVideoFromURL
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS["VideoGenKlingNode"]= "Kling Video Generation @fal"
|
||||
NODE_DISPLAY_NAME_MAPPINGS["VideoGenRunwayGen3Node"]= "Runway Gen3 Image-to-Video @fal"
|
||||
NODE_DISPLAY_NAME_MAPPINGS["VideoGenLumaDreamMachineNode"]= "Luma Dream Machine @fal"
|
||||
NODE_DISPLAY_NAME_MAPPINGS["LoadVideoFromURL"]= "Load Video from URL"
|
||||
|
||||
|
||||
except Exception as e:
|
||||
logging.info('FalVideo.available False' )
|
||||
|
||||
logging.info('\033[93m -------------- \033[0m')
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 366 KiB |
+2702
-399
File diff suppressed because it is too large
Load Diff
@@ -3,6 +3,101 @@ import os
|
||||
import folder_paths
|
||||
import torchaudio
|
||||
|
||||
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 analyze_audio_data(audio_data):
|
||||
total_duration = 0
|
||||
total_gap_duration = 0
|
||||
emotion_counts = {}
|
||||
audio_types = set()
|
||||
languages = set()
|
||||
|
||||
for i, entry in enumerate(audio_data):
|
||||
# Calculate the duration of each audio segment
|
||||
start_time = entry['start_time']
|
||||
end_time = entry['end_time']
|
||||
duration = end_time - start_time
|
||||
total_duration += duration
|
||||
|
||||
# Count the emotions
|
||||
if "emotion" in entry:
|
||||
emotion = entry['emotion']
|
||||
if emotion in emotion_counts:
|
||||
emotion_counts[emotion] += 1
|
||||
else:
|
||||
emotion_counts[emotion] = 1
|
||||
|
||||
# Collect the audio types
|
||||
if "audio_type" in entry:
|
||||
audio_types.add(entry['audio_type'])
|
||||
|
||||
if "language" in entry:
|
||||
languages.add(entry['language'])
|
||||
|
||||
# Calculate gap duration if not the last entry
|
||||
if i < len(audio_data) - 1:
|
||||
next_start_time = audio_data[i + 1]['start_time']
|
||||
gap_duration = next_start_time - end_time
|
||||
if gap_duration > 0:
|
||||
total_gap_duration += gap_duration
|
||||
|
||||
# Get the most frequent emotion
|
||||
if len(emotion_counts.keys())>0:
|
||||
most_frequent_emotion = max(emotion_counts, key=emotion_counts.get)
|
||||
else:
|
||||
most_frequent_emotion=None
|
||||
|
||||
# Convert audio_types set to list for better readability
|
||||
audio_types = list(audio_types)
|
||||
|
||||
languages=list(languages)
|
||||
|
||||
# Print the results
|
||||
print(f"Total Effective Duration: {total_duration:.2f} seconds")
|
||||
print(f"Total Gap Duration: {total_gap_duration:.2f} seconds")
|
||||
print(f"Emotion Changes: {emotion_counts}")
|
||||
print(f"Most Frequent Emotion: {most_frequent_emotion}")
|
||||
print(f"Audio Types: {audio_types}")
|
||||
|
||||
|
||||
return {
|
||||
"total_duration": total_duration,
|
||||
"total_gap_duration": total_gap_duration,
|
||||
"emotion_changes": emotion_counts,
|
||||
"most_frequent_emotion": most_frequent_emotion,
|
||||
"audio_types": audio_types,
|
||||
"languages":languages
|
||||
}
|
||||
|
||||
|
||||
# 分析音频数据
|
||||
class AnalyzeAudioNone:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"json":(any_type,),},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ("result",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
def run(self,json):
|
||||
result=analyze_audio_data(json)
|
||||
return (result,)
|
||||
|
||||
|
||||
|
||||
class SpeechRecognition:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
+375
-14
@@ -1,4 +1,6 @@
|
||||
import openai
|
||||
from swarm import Swarm, Agent
|
||||
|
||||
import time
|
||||
import urllib.error
|
||||
import re,json,os,string,random
|
||||
@@ -7,9 +9,19 @@ 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)
|
||||
|
||||
|
||||
# 从文本中提取json
|
||||
def extract_json_strings(text):
|
||||
json_strings = []
|
||||
@@ -204,7 +216,7 @@ if is_installed('json_repair'):
|
||||
from json_repair import repair_json
|
||||
|
||||
|
||||
def chat(client, model_name,messages ):
|
||||
def chat(client, model_name,messages,max_tokens=4096,temperature=0.6 ):
|
||||
print('#chat',model_name,messages)
|
||||
try_count = 0
|
||||
while True:
|
||||
@@ -213,7 +225,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的
|
||||
@@ -448,7 +462,8 @@ 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"
|
||||
@@ -466,10 +481,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}),
|
||||
},
|
||||
"optional":{
|
||||
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING","STRING","STRING",)
|
||||
@@ -481,11 +497,14 @@ class SiliconflowFreeNode:
|
||||
|
||||
|
||||
def generate_contextual_text(self,
|
||||
api_key,
|
||||
prompt,
|
||||
system_content,
|
||||
api_key,
|
||||
prompt,
|
||||
system_content,
|
||||
model,
|
||||
seed,context_size,custom_model_name=None):
|
||||
seed,
|
||||
context_size,
|
||||
max_tokens,
|
||||
custom_model_name=None):
|
||||
|
||||
if custom_model_name!=None:
|
||||
model=custom_model_name
|
||||
@@ -517,7 +536,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}]
|
||||
|
||||
@@ -525,6 +544,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
|
||||
@@ -693,9 +788,12 @@ class JsonRepair:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"json_string":("STRING", {"forceInput": True,}),
|
||||
"key":("STRING", {"multiline": False,"dynamicPrompts": False,"default": ""}),
|
||||
}
|
||||
"json_string":("STRING", {"forceInput": True,}),
|
||||
"key":("STRING", {"multiline": False,"dynamicPrompts": False,"default": ""}),
|
||||
},
|
||||
"optional":{
|
||||
"json_string2":("STRING", {"forceInput": True,})
|
||||
},
|
||||
}
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
@@ -707,8 +805,11 @@ class JsonRepair:
|
||||
|
||||
CATEGORY = "♾️Mixlab/GPT"
|
||||
|
||||
def run(self, json_string,key=""):
|
||||
def run(self, json_string,key="",json_string2=None):
|
||||
|
||||
if not isinstance(json_string, str):
|
||||
json_string=json.dumps(json_string)
|
||||
|
||||
json_string=extract_json_strings(json_string)
|
||||
# print(json_string)
|
||||
good_json_string = repair_json(json_string)
|
||||
@@ -716,6 +817,20 @@ class JsonRepair:
|
||||
# 将 JSON 字符串解析为 Python 对象
|
||||
data = json.loads(good_json_string)
|
||||
|
||||
if json_string2!=None:
|
||||
if not isinstance(json_string2, str):
|
||||
json_string2=json.dumps(json_string2)
|
||||
|
||||
json_string2=extract_json_strings(json_string2)
|
||||
# print(json_string)
|
||||
good_json_string2 = repair_json(json_string2)
|
||||
|
||||
# 将 JSON 字符串解析为 Python 对象
|
||||
data2 = json.loads(good_json_string2)
|
||||
|
||||
data={**data, **data2}
|
||||
|
||||
|
||||
v=""
|
||||
if key!="" and (key in data):
|
||||
v=data[key]
|
||||
@@ -723,4 +838,250 @@ class JsonRepair:
|
||||
# 将 Python 对象转换回 JSON 字符串,确保中文字符不被转义
|
||||
json_str_with_chinese = json.dumps(data, ensure_ascii=False)
|
||||
|
||||
return (json_str_with_chinese,v,)
|
||||
return (json_str_with_chinese,v,)
|
||||
|
||||
|
||||
# 以下为固定提示词的LLM节点示例
|
||||
class SimulateDevDesignDiscussions:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
|
||||
model_list=[
|
||||
"gpt-4o",
|
||||
"gpt-4o-2024-05-13",
|
||||
"gpt-4",
|
||||
"gpt-4-0314",
|
||||
"gpt-4-0613",
|
||||
"qwen-turbo",
|
||||
"qwen-plus",
|
||||
"qwen-long",
|
||||
"qwen-max",
|
||||
"qwen-max-longcontext",
|
||||
"glm-4",
|
||||
"glm-3-turbo",
|
||||
"moonshot-v1-8k",
|
||||
"moonshot-v1-32k",
|
||||
"moonshot-v1-128k",
|
||||
"deepseek-chat",
|
||||
"Qwen/Qwen2-7B-Instruct",
|
||||
"THUDM/glm-4-9b-chat",
|
||||
"01-ai/Yi-1.5-9B-Chat-16K"
|
||||
]
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"subject": ("STRING", {"multiline": True,"dynamicPrompts": False}),
|
||||
"model": ( model_list,
|
||||
{"default": model_list[0]}),
|
||||
"api_url":(list(llm_apis_dict.keys()),
|
||||
{"default": list(llm_apis_dict.keys())[0]}),
|
||||
},
|
||||
"optional":{
|
||||
"api_key":("STRING", {"forceInput": True,}),
|
||||
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
|
||||
"custom_api_url":("STRING", {"forceInput": True,}), #适合自定义model
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("text",)
|
||||
FUNCTION = "generate_contextual_text"
|
||||
CATEGORY = "♾️Mixlab/GPT"
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def generate_contextual_text(self,
|
||||
subject,
|
||||
model,
|
||||
api_url,
|
||||
api_key=None,
|
||||
custom_model_name=None,
|
||||
custom_api_url=None,
|
||||
):
|
||||
|
||||
# 设置黄色文本的ANSI转义序列
|
||||
YELLOW = "\033[33m"
|
||||
# 重置文本颜色的ANSI转义序列
|
||||
RESET = "\033[0m"
|
||||
|
||||
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",api_key,api_url)
|
||||
#
|
||||
if is_azure_url(api_url):
|
||||
client=azure_client(api_key,api_url)
|
||||
else:
|
||||
# 根据用户选择的模型,设置相应的接口和模型名称
|
||||
if model == "glm-4" :
|
||||
client = ZhipuAI_client(api_key) # 使用 Zhipuai 的接口
|
||||
print('using Zhipuai interface')
|
||||
else :
|
||||
client = openai_client(api_key,api_url) # 使用 ChatGPT 的接口
|
||||
|
||||
|
||||
|
||||
# 以下为多智能体框架
|
||||
client = Swarm(client=client)
|
||||
|
||||
# 定义两个代理:软件系统架构师和设计师
|
||||
software_architect_agent = Agent(
|
||||
name="Software Architect",
|
||||
instructions='''用脱口秀的风格回答编程问题,简短且口语化。
|
||||
|
||||
输出格式
|
||||
====
|
||||
|
||||
* 答案格式:`程序员:xxxxxxxxx`
|
||||
|
||||
示例
|
||||
==
|
||||
|
||||
**输入:**
|
||||
如何优化代码性能?
|
||||
|
||||
**输出:**
|
||||
程序员:兄弟,先把那些循环里的debug信息删掉,CPU都快哭了。'''
|
||||
)
|
||||
|
||||
designer_agent = Agent(
|
||||
name="Designer",
|
||||
instructions='''回答问题时,请扮演一位具有多年空间设计和用户体验设计经验的设计师。你的回答应当天马行空,但又富有深度,带有苏格拉底的思考方式,并且使用脱口秀的风格。回答要简短且非常口语化。格式如下:
|
||||
|
||||
设计师:\[回答内容\]
|
||||
|
||||
Output Format
|
||||
=============
|
||||
|
||||
* 回答应当使用“设计师:\[回答内容\]”的格式。
|
||||
* 回答应当简短、口语化,富有创意和深度。
|
||||
|
||||
Examples
|
||||
========
|
||||
|
||||
**Example 1:**
|
||||
|
||||
主持人:你觉得未来的家会是什么样子?
|
||||
|
||||
设计师:未来的家?想象一下,房子会像变形金刚一样,随时变形满足你的需求。今天是健身房,明天是电影院,后天是游戏场。家不再是四面墙,而是一个随心所欲的魔法空间。
|
||||
|
||||
**Example 2:**
|
||||
|
||||
主持人:你怎么看待极简主义设计?
|
||||
|
||||
设计师:极简主义?就像吃寿司,去掉所有不必要的装饰,只留下最精华的部分。让空间呼吸,让心灵自由。
|
||||
|
||||
**Example 3:**
|
||||
|
||||
主持人:你觉得色彩在设计中有多重要?
|
||||
|
||||
设计师:色彩?哦,那可是设计的灵魂!就像人生中的调味料,一点红色让你激情澎湃,一点蓝色让你心如止水。色彩决定了空间的情绪基调。'''
|
||||
)
|
||||
|
||||
# 定义一个函数,用于转移问题到designer_agent
|
||||
def transfer_to_designer_agent():
|
||||
return designer_agent
|
||||
|
||||
# 将转移函数添加到软件系统架构师和设计师的函数列表中
|
||||
software_architect_agent.functions.append(transfer_to_designer_agent)
|
||||
|
||||
# 问题生成
|
||||
host_agent = Agent(
|
||||
name="Host",
|
||||
instructions='''
|
||||
为播客的主持人生成4到5个问题,这些问题有些是针对设计师问的,有些是针对程序员问的。
|
||||
|
||||
* 主持人:你知道如何开发一款APP产品,从想法到上线吗?
|
||||
* 主持人:站在设计师的角度,你怎么看?
|
||||
* 主持人:不知道程序员又是怎么想的呢?
|
||||
* 主持人:感谢大家的参与,今天收获蛮大的
|
||||
|
||||
Steps
|
||||
=====
|
||||
|
||||
1. 确定问题的对象:设计师或程序员。
|
||||
2. 根据对象设计相关的问题,确保问题的多样性和深度。
|
||||
3. 整理问题,使其符合播客主持人的风格和语气。
|
||||
|
||||
Output Format
|
||||
=============
|
||||
|
||||
问题列表,每个问题以“主持人:”开头,不要出现序号。
|
||||
|
||||
Examples
|
||||
========
|
||||
|
||||
* 主持人:作为一名设计师,你是如何开始一个新项目的?
|
||||
* 主持人:程序员在开发过程中遇到的最大挑战是什么?
|
||||
* 主持人:设计师在团队协作中扮演什么角色?
|
||||
* 主持人:程序员如何确保代码的质量和稳定性?
|
||||
* 主持人:感谢大家的参与,今天的讨论非常有意义。
|
||||
|
||||
Notes
|
||||
=====
|
||||
|
||||
* 确保问题针对不同的角色(设计师和程序员)。
|
||||
* 保持问题的多样性,涵盖从项目开始到完成的各个阶段。
|
||||
* 确保问题能引导出深入的讨论和见解。
|
||||
''')
|
||||
|
||||
|
||||
response = client.run(agent=host_agent, messages=[{
|
||||
"role":"user",
|
||||
"content":f"主题是‘{subject}’"
|
||||
}],model_override=model)
|
||||
|
||||
content=response.messages[-1]["content"]
|
||||
print(f"{YELLOW}{content}{RESET}")
|
||||
|
||||
texts=content.split("\n")
|
||||
|
||||
# texts='''
|
||||
# 主持人:你知道如何开发一款APP产品,从想法到上线吗?
|
||||
# 主持人:站在设计师的角度,你怎么看?
|
||||
# 主持人:不知道程序员又是怎么想的呢?
|
||||
# 主持人:感谢大家的参与,今天收获蛮大的
|
||||
# '''.split("\n")
|
||||
|
||||
messages=[]
|
||||
|
||||
texts = [text.strip() for text in texts if text.strip()]
|
||||
|
||||
result=[]
|
||||
|
||||
for text in texts:
|
||||
messages.append({
|
||||
"role": "user",
|
||||
"content": text
|
||||
})
|
||||
|
||||
# 运行客户端,使用软件系统架构师作为初始代理
|
||||
response = client.run(agent=software_architect_agent, messages=messages,model_override=model)
|
||||
|
||||
print(f"{text}")
|
||||
result.append(text)
|
||||
|
||||
# 输出最后一个响应消息的内容
|
||||
content=response.messages[-1]["content"]
|
||||
print(f"{YELLOW}{content}{RESET}")
|
||||
|
||||
result.append(content)
|
||||
|
||||
messages.append({
|
||||
"role":"assistant",
|
||||
"content":content
|
||||
})
|
||||
|
||||
|
||||
return ("\n".join(result),)
|
||||
|
||||
|
||||
@@ -0,0 +1,332 @@
|
||||
# 修改自 https://github.com/gokayfem/ComfyUI-fal-API/blob/main/nodes/video_node.py
|
||||
# image-to-video all in one
|
||||
|
||||
import os,sys
|
||||
import torch
|
||||
from PIL import Image
|
||||
import tempfile
|
||||
import numpy as np
|
||||
import requests
|
||||
import cv2
|
||||
import subprocess
|
||||
import importlib.util
|
||||
python = sys.executable
|
||||
|
||||
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:
|
||||
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
|
||||
|
||||
|
||||
try:
|
||||
if is_installed('fal_client','fal-client')==True:
|
||||
from fal_client import submit, upload_file
|
||||
except:
|
||||
print("#install fal-client error")
|
||||
|
||||
|
||||
def upload_image(image):
|
||||
try:
|
||||
# Convert the image tensor to a numpy array
|
||||
if isinstance(image, torch.Tensor):
|
||||
image_np = image.cpu().numpy()
|
||||
else:
|
||||
image_np = np.array(image)
|
||||
|
||||
# Ensure the image is in the correct format (H, W, C)
|
||||
if image_np.ndim == 4:
|
||||
image_np = image_np.squeeze(0) # Remove batch dimension if present
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.stack([image_np] * 3, axis=-1) # Convert grayscale to RGB
|
||||
elif image_np.shape[0] == 3:
|
||||
image_np = np.transpose(image_np, (1, 2, 0)) # Change from (C, H, W) to (H, W, C)
|
||||
|
||||
# Normalize the image data to 0-255 range
|
||||
if image_np.dtype == np.float32 or image_np.dtype == np.float64:
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
|
||||
# Convert to PIL Image
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
# Save the image to a temporary file
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as temp_file:
|
||||
pil_image.save(temp_file, format="PNG")
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
# Upload the temporary file
|
||||
image_url = upload_file(temp_file_path)
|
||||
return image_url
|
||||
except Exception as e:
|
||||
print(f"Error uploading image: {str(e)}")
|
||||
return None
|
||||
finally:
|
||||
# Clean up the temporary file
|
||||
if 'temp_file_path' in locals():
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
|
||||
class VideoGenKlingNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"duration": (["5", "10"], {"default": "5"}),
|
||||
"aspect_ratio": (["16:9", "9:16", "1:1"], {"default": "16:9"}),
|
||||
"mode": (["standard", "pro"], {"default": "standard"}),
|
||||
"fal_key":("STRING", {"forceInput": True,}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def generate_video(self, prompt, duration, aspect_ratio,mode,fal_key, image=None):
|
||||
arguments = {
|
||||
"prompt": prompt,
|
||||
"duration": duration,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
}
|
||||
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
|
||||
api_url="fal-ai/kling-video/v1/"+mode
|
||||
|
||||
try:
|
||||
if image is not None:
|
||||
image_url = upload_image(image)
|
||||
if image_url:
|
||||
arguments["image_url"] = image_url
|
||||
handler = submit(api_url+"/image-to-video", arguments=arguments)
|
||||
else:
|
||||
return ("Error: Unable to upload image.",)
|
||||
else:
|
||||
handler = submit(api_url+"/text-to-video", arguments=arguments)
|
||||
|
||||
result = handler.get()
|
||||
video_url = result["video"]["url"]
|
||||
return (video_url,)
|
||||
except Exception as e:
|
||||
print(f"Error generating video: {str(e)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
|
||||
class VideoGenRunwayGen3Node:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"image": ("IMAGE",),
|
||||
"duration": (["5", "10"], {"default": "5"}),
|
||||
"aspect_ratio": (["16:9", "9:16"], {"default": "16:9"}),
|
||||
"fal_key":("STRING", {"forceInput": True,}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def generate_video(self, prompt, image, duration,aspect_ratio,fal_key):
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
try:
|
||||
image_url = upload_image(image)
|
||||
if not image_url:
|
||||
return ("Error: Unable to upload image.",)
|
||||
|
||||
arguments = {
|
||||
"prompt": prompt,
|
||||
"image_url": image_url,
|
||||
"duration": duration,
|
||||
"ratio":aspect_ratio
|
||||
}
|
||||
|
||||
handler = submit("fal-ai/runway-gen3/turbo/image-to-video", arguments=arguments)
|
||||
result = handler.get()
|
||||
video_url = result["video"]["url"]
|
||||
return (video_url,)
|
||||
except Exception as e:
|
||||
print(f"Error generating video: {str(e)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
class VideoGenLumaDreamMachineNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"aspect_ratio": (["16:9", "9:16", "4:3", "3:4", "21:9", "9:21"], {"default": "16:9"}),
|
||||
"fal_key":("STRING", {"forceInput": True,}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
"loop": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def generate_video(self, prompt, aspect_ratio,fal_key, image=None, loop=True):
|
||||
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
|
||||
arguments = {
|
||||
"prompt": prompt,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"loop": loop,
|
||||
}
|
||||
|
||||
try:
|
||||
if image is not None:
|
||||
image_url = upload_image(image)
|
||||
if not image_url:
|
||||
return ("Error: Unable to upload image.",)
|
||||
arguments["image_url"] = image_url
|
||||
endpoint = "fal-ai/luma-dream-machine/image-to-video"
|
||||
else:
|
||||
endpoint = "fal-ai/luma-dream-machine"
|
||||
|
||||
handler = submit(endpoint, arguments=arguments)
|
||||
result = handler.get()
|
||||
video_url = result["video"]["url"]
|
||||
return (video_url,)
|
||||
except Exception as e:
|
||||
print(f"Error generating video: {str(e)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
class LoadVideoFromURL:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"url": ("STRING", {"default": "https://example.com/video.mp4"}),
|
||||
"force_rate": ("INT", {"default": 0, "min": 0, "max": 60, "step": 1}),
|
||||
"force_size": (["Disabled", "Custom Height", "Custom Width", "Custom", "256x?", "?x256", "256x256", "512x?", "?x512", "512x512"],),
|
||||
"custom_width": ("INT", {"default": 512, "min": 0, "max": 8192, "step": 8}),
|
||||
"custom_height": ("INT", {"default": 512, "min": 0, "max": 8192, "step": 8}),
|
||||
"frame_load_cap": ("INT", {"default": 0, "min": 0, "max": 1000000, "step": 1}),
|
||||
"skip_first_frames": ("INT", {"default": 0, "min": 0, "max": 1000000, "step": 1}),
|
||||
"select_every_nth": ("INT", {"default": 1, "min": 1, "max": 1000000, "step": 1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "INT", "VHS_VIDEOINFO")
|
||||
RETURN_NAMES = ("frames", "frame_count", "video_info")
|
||||
FUNCTION = "load_video_from_url"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def load_video_from_url(self, url, force_rate, force_size, custom_width, custom_height, frame_load_cap, skip_first_frames, select_every_nth):
|
||||
# Download the video to a temporary file
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=".mp4") as temp_file:
|
||||
response = requests.get(url, stream=True)
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
temp_file.write(chunk)
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
# Load the video using OpenCV
|
||||
cap = cv2.VideoCapture(temp_file_path)
|
||||
|
||||
# Get video properties
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
duration = total_frames / fps
|
||||
|
||||
# Calculate target size
|
||||
if force_size != "Disabled":
|
||||
if force_size == "Custom Width":
|
||||
new_height = int(height * (custom_width / width))
|
||||
new_width = custom_width
|
||||
elif force_size == "Custom Height":
|
||||
new_width = int(width * (custom_height / height))
|
||||
new_height = custom_height
|
||||
elif force_size == "Custom":
|
||||
new_width, new_height = custom_width, custom_height
|
||||
else:
|
||||
target_width, target_height = map(int, force_size.replace("?", "0").split("x"))
|
||||
if target_width == 0:
|
||||
new_width = int(width * (target_height / height))
|
||||
new_height = target_height
|
||||
else:
|
||||
new_height = int(height * (target_width / width))
|
||||
new_width = target_width
|
||||
else:
|
||||
new_width, new_height = width, height
|
||||
|
||||
frames = []
|
||||
frame_count = 0
|
||||
|
||||
for i in range(total_frames):
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
|
||||
if i < skip_first_frames:
|
||||
continue
|
||||
|
||||
if (i - skip_first_frames) % select_every_nth != 0:
|
||||
continue
|
||||
|
||||
if force_size != "Disabled":
|
||||
frame = cv2.resize(frame, (new_width, new_height))
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frame = torch.from_numpy(frame).float() / 255.0
|
||||
frames.append(frame)
|
||||
|
||||
frame_count += 1
|
||||
|
||||
if frame_load_cap > 0 and frame_count >= frame_load_cap:
|
||||
break
|
||||
|
||||
cap.release()
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
frames = torch.stack(frames)
|
||||
|
||||
video_info = {
|
||||
"source_fps": fps,
|
||||
"source_frame_count": total_frames,
|
||||
"source_duration": duration,
|
||||
"source_width": width,
|
||||
"source_height": height,
|
||||
"loaded_fps": fps if force_rate == 0 else force_rate,
|
||||
"loaded_frame_count": frame_count,
|
||||
"loaded_duration": frame_count / (fps if force_rate == 0 else force_rate),
|
||||
"loaded_width": new_width,
|
||||
"loaded_height": new_height,
|
||||
}
|
||||
|
||||
return (frames, frame_count, video_info)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
+150
-37
@@ -13,12 +13,26 @@ 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)
|
||||
@@ -109,8 +123,13 @@ def composite_images(foreground, background, mask, is_multiply_blend=False, posi
|
||||
}
|
||||
|
||||
# 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)
|
||||
try:
|
||||
resampling_method = Image.Resampling.LANCZOS
|
||||
except AttributeError:
|
||||
resampling_method = Image.ANTIALIAS
|
||||
|
||||
layer_image = layer['image'].resize((layer['width'], layer['height']), resampling_method)
|
||||
layer_mask = layer['mask'].resize((layer['width'], layer['height']), resampling_method)
|
||||
|
||||
bg_image.paste(layer_image, (layer['x'], layer['y']), layer_mask)
|
||||
|
||||
@@ -768,8 +787,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代码
|
||||
@@ -949,9 +967,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)
|
||||
@@ -969,33 +1046,34 @@ def generate_text_image(text, font_path, font_size, text_color, vertical=True, s
|
||||
|
||||
if layout == "vertical":
|
||||
for line in lines:
|
||||
max_char_width = max(font.getsize(char)[0] for char in line)
|
||||
max_char_width = max(font.getbbox(char)[2] - font.getbbox(char)[0] for char in line)
|
||||
for char in line:
|
||||
char_width, char_height = font.getsize(char)
|
||||
left, top, right, bottom = font.getbbox(char)
|
||||
char_width = right - left
|
||||
char_height = bottom - top
|
||||
char_coordinates.append((x, y))
|
||||
y += char_height + spacing
|
||||
max_height = max(max_height, y + padding)
|
||||
x += max_char_width + line_spacing
|
||||
y = padding
|
||||
max_width = x
|
||||
total_line_width = sum(font.getsize(line)[1] for line in lines)
|
||||
total_line_width = sum(font.getbbox(line)[2] - font.getbbox(line)[0] 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:
|
||||
line_width, line_height = font.getsize(line)
|
||||
line_width, line_height = font.getbbox(line)[2] - font.getbbox(line)[0], font.getbbox(line)[3] - font.getbbox(line)[1]
|
||||
for char in line:
|
||||
char_width, char_height = font.getsize(char)
|
||||
left, top, right, bottom = font.getbbox(char)
|
||||
char_width = right - left
|
||||
char_height = bottom - top
|
||||
char_coordinates.append((x, y))
|
||||
x += char_width + spacing
|
||||
max_width = max(max_width, x + padding)
|
||||
y += line_height + line_spacing
|
||||
x = padding
|
||||
# max_height = y
|
||||
total_line_heights = sum(font.getsize(line)[1] for line in lines)
|
||||
total_line_heights = sum(font.getbbox(line)[3] - font.getbbox(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
|
||||
@@ -1008,10 +1086,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
|
||||
@@ -1025,10 +1103,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)
|
||||
@@ -1369,7 +1455,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
|
||||
@@ -1385,7 +1471,7 @@ class LoadImagesFromPath:
|
||||
}
|
||||
|
||||
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"
|
||||
|
||||
@@ -1398,7 +1484,7 @@ class LoadImagesFromPath:
|
||||
watcher_folder=None
|
||||
|
||||
# 运行的函数
|
||||
def run(self,file_path,white_bg,newest_files,index_variable,watcher,result,prompt,seed=1):
|
||||
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)
|
||||
|
||||
@@ -1421,19 +1507,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)
|
||||
|
||||
@@ -1441,13 +1531,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 (imgs,masks,prompt,file_names,)
|
||||
return (imgs,masks,prompt,file_paths,)
|
||||
|
||||
|
||||
# TODO 扩大选区的功能,重新输出mask
|
||||
@@ -1582,6 +1672,20 @@ class TextImage:
|
||||
"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"
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -1595,14 +1699,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)
|
||||
@@ -2789,14 +2899,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'],),
|
||||
@@ -2812,7 +2922,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"
|
||||
|
||||
@@ -2850,11 +2960,14 @@ class ResizeImage:
|
||||
im=tensor2pil(im)
|
||||
|
||||
im=im.convert('RGB')
|
||||
a_im,hex=get_average_color_image(im)
|
||||
|
||||
a_im,hex=get_average_color_image(im)
|
||||
|
||||
if average_color=='on':
|
||||
fill_color=hex
|
||||
|
||||
|
||||
a_im=resize_image(a_im,scale_option,w,h,fill_color)
|
||||
|
||||
im=resize_image(im,scale_option,w,h,fill_color)
|
||||
|
||||
im=pil2tensor(im)
|
||||
|
||||
+34
-4
@@ -1,4 +1,5 @@
|
||||
# Referenced some code:https://github.com/IuvenisSapiens/ComfyUI_MiniCPM-V-2_6-int4
|
||||
# https://github.com/CY-CHENYUE/ComfyUI-MiniCPM-Plus
|
||||
|
||||
import os
|
||||
import torch
|
||||
@@ -6,7 +7,7 @@ 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
|
||||
# from PIL import Image
|
||||
|
||||
def get_model_path(n=""):
|
||||
try:
|
||||
@@ -35,6 +36,7 @@ class MiniCPM_VQA_Simple:
|
||||
"images": ("IMAGE",),
|
||||
"text": ("STRING", {"default": "", "multiline": True}),
|
||||
"seed": ("INT", {"default": -1}), # add seed parameter, default is -1
|
||||
"extract_keywords":("BOOLEAN", {"default": False}),
|
||||
"temperature": (
|
||||
"FLOAT",
|
||||
{
|
||||
@@ -46,7 +48,9 @@ class MiniCPM_VQA_Simple:
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_TYPES = ("STRING","STRING",)
|
||||
RETURN_NAMES = ("result","keywords",)
|
||||
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
|
||||
@@ -55,6 +59,7 @@ class MiniCPM_VQA_Simple:
|
||||
images,
|
||||
text,
|
||||
seed, # add seed parameter, default is -1
|
||||
extract_keywords,
|
||||
temperature,
|
||||
keep_model_loaded,
|
||||
):
|
||||
@@ -90,6 +95,7 @@ class MiniCPM_VQA_Simple:
|
||||
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]
|
||||
@@ -113,6 +119,30 @@ class MiniCPM_VQA_Simple:
|
||||
# max_new_tokens=max_new_tokens,
|
||||
**params,
|
||||
)
|
||||
|
||||
keyword_result=""
|
||||
|
||||
if extract_keywords:#extract_keywords
|
||||
keyword_prompt = f"""Please extract keywords from the following text, including all occurrences of language (e.g. Chinese, English, etc.):
|
||||
[[[{result}]]]
|
||||
Please list the keywords extracted, separated by commas. Make sure to include all important words, no matter what language. For English words, please keep the original case."""
|
||||
|
||||
keyword_msgs = [{'role': 'user', 'content': keyword_prompt}]
|
||||
keyword_result = self.model.chat(
|
||||
image=None,
|
||||
msgs=keyword_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,
|
||||
)
|
||||
print("keyword_result",keyword_result)
|
||||
|
||||
|
||||
# offload model to GPU
|
||||
# self.model = self.model.to(torch.device("cpu"))
|
||||
# self.model.eval()
|
||||
@@ -123,5 +153,5 @@ class MiniCPM_VQA_Simple:
|
||||
self.model = None # set model to None
|
||||
torch.cuda.empty_cache() # release GPU memory
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
return (result,)
|
||||
# print(result)
|
||||
return (result,keyword_result,)
|
||||
|
||||
+26
-6
@@ -187,7 +187,8 @@ class PromptImage:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("json_str",)
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@@ -202,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]
|
||||
@@ -215,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,226 @@
|
||||
# -*- 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)
|
||||
print('#format_to_srt',match,asr_result)
|
||||
if match==None:
|
||||
return None, None, None, None,None,start_time,end_time,None
|
||||
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)
|
||||
|
||||
if lang!=None:
|
||||
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,"STRING","STRING","FLOAT",)
|
||||
RETURN_NAMES = ("result","srt","text","total_seconds",)
|
||||
|
||||
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']
|
||||
sample_rate = audio['sample_rate']
|
||||
# print("Original shape:", waveform.shape) # 打印原始形状
|
||||
if waveform.ndim == 3 and waveform.shape[0] == 1: # 检查是否为三维且 batch_size 为 1
|
||||
waveform = waveform.squeeze(0) # 移除 batch_size 维度
|
||||
else:
|
||||
raise ValueError("Unexpected waveform dimensions")
|
||||
|
||||
print("waveform.shape:", waveform.shape)
|
||||
total_length_seconds = waveform.shape[1] / sample_rate
|
||||
|
||||
waveform_numpy = waveform.numpy().transpose(1, 0) # 转换为 (num_samples, num_channels)
|
||||
|
||||
results=self.processor.process_audio(waveform_numpy, sample_rate, language, use_itn)
|
||||
|
||||
srt_content="\n".join([s['srt_content'] for s in results])
|
||||
text="\n".join([s['text'] for s in results])
|
||||
|
||||
return (results,srt_content,text,total_length_seconds,)
|
||||
|
||||
+1
-1
@@ -234,7 +234,7 @@ class StyleAlignedSampleReferenceLatents:
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS.reverse(), ),
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
|
||||
"denoise": ("FLOAT", {"default": 1, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
|
||||
}
|
||||
|
||||
+53
-5
@@ -624,11 +624,13 @@ class AppInfo:
|
||||
|
||||
im=[]
|
||||
if image:
|
||||
# img=image[0][0]
|
||||
print('AppInfo_image',len(image))
|
||||
images=[image]
|
||||
# batch 的方式需要处理
|
||||
for i in range(len(image)):
|
||||
img=image[i]
|
||||
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,
|
||||
|
||||
@@ -646,6 +648,49 @@ class AppInfo:
|
||||
|
||||
|
||||
|
||||
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),)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class SwitchByIndex:
|
||||
@@ -823,7 +868,7 @@ class TESTNODE_:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"ANY":(any_type,),
|
||||
"ANY":(any_type,),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -838,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)
|
||||
|
||||
+16
-6
@@ -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:
|
||||
@@ -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,173 @@
|
||||
import os,re
|
||||
import sys,time
|
||||
from pathlib import Path
|
||||
import torchaudio
|
||||
import hashlib
|
||||
import torch
|
||||
import folder_paths
|
||||
import comfy.utils
|
||||
|
||||
from faster_whisper import WhisperModel
|
||||
|
||||
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 get_model_dir(m):
|
||||
try:
|
||||
return folder_paths.get_folder_paths(m)[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, m)
|
||||
|
||||
|
||||
|
||||
whisper_model_path=get_model_dir('whisper')
|
||||
|
||||
model_sizes=[
|
||||
d for d in os.listdir(whisper_model_path) if os.path.isdir(
|
||||
os.path.join(whisper_model_path, d)
|
||||
) and os.path.isfile(os.path.join(os.path.join(whisper_model_path, d), "config.json"))
|
||||
]
|
||||
|
||||
|
||||
class LoadWhisperModel:
|
||||
def __init__(self):
|
||||
self.model = None
|
||||
self.device="cuda" if torch.cuda.is_available() else "cpu"
|
||||
self.model_size=model_sizes[0]
|
||||
self.compute_type='float16'
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model_size": (model_sizes,),
|
||||
"device": (["auto","cpu"],),
|
||||
"compute_type": (["float16","int8_float16","int8"],),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WHISPER",)
|
||||
RETURN_NAMES = ("whisper_model",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/Whisper"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self,model_size,device,compute_type):
|
||||
|
||||
if device=="auto" and self.device!='cuda':
|
||||
self.device="cuda" if torch.cuda.is_available() else "cpu"
|
||||
self.model=None
|
||||
|
||||
if device=='cpu' and self.device!='cpu':
|
||||
self.device="cpu"
|
||||
self.model=None
|
||||
|
||||
if model_size!= self.model_size:
|
||||
self.model_size=model_size
|
||||
self.model=None
|
||||
|
||||
if compute_type!=self.compute_type:
|
||||
self.compute_type=compute_type
|
||||
self.model=None
|
||||
|
||||
if self.model==None:
|
||||
self.model = WhisperModel(
|
||||
os.path.join(whisper_model_path, self.model_size),
|
||||
device=self.device,
|
||||
compute_type=self.compute_type
|
||||
)
|
||||
|
||||
return (self.model,)
|
||||
|
||||
|
||||
class WhisperTranscribe:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"whisper_model": ("WHISPER",),
|
||||
"audio": ("AUDIO",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,"STRING","STRING","FLOAT",)
|
||||
RETURN_NAMES = ("result","srt","text","total_seconds",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/Whisper"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
# OUTPUT_IS_LIST = (False,False,False,)
|
||||
|
||||
def run(self,whisper_model,audio):
|
||||
|
||||
if 'audio_path' in audio and (not 'waveform' in audio):
|
||||
waveform, sample_rate = torchaudio.load(audio['audio_path'])
|
||||
waveform=waveform.mean(0)
|
||||
total_length_seconds = waveform.shape[0] / sample_rate
|
||||
waveform=waveform.numpy()
|
||||
|
||||
elif 'waveform' in audio and 'sample_rate' in audio:
|
||||
print("Original shape:", audio["waveform"].shape, isinstance(audio["waveform"], torch.Tensor)) # 打印原始形状
|
||||
waveform = audio["waveform"].squeeze(0) # Remove the added batch dimension
|
||||
sample_rate = audio["sample_rate"]
|
||||
|
||||
# if audio_sf != sampling_rate:
|
||||
# waveform = torchaudio.functional.resample(
|
||||
# waveform, orig_freq=audio_sf, new_freq=sampling_rate
|
||||
# )
|
||||
|
||||
waveform=waveform.mean(0)
|
||||
|
||||
total_length_seconds = waveform.shape[0] / sample_rate
|
||||
|
||||
waveform=waveform.numpy() #whisper_model.transcribe 旧版不支持直接传tensor,先用numpy
|
||||
|
||||
segments, info = whisper_model.transcribe(waveform, beam_size=5)
|
||||
|
||||
print("Detected language '%s' with probability %f" % (info.language, info.language_probability))
|
||||
|
||||
# Function to format time for SRT
|
||||
def format_time(seconds):
|
||||
millis = int((seconds - int(seconds)) * 1000)
|
||||
hours, remainder = divmod(int(seconds), 3600)
|
||||
minutes, seconds = divmod(remainder, 60)
|
||||
return f"{hours:02}:{minutes:02}:{seconds:02},{millis:03}"
|
||||
|
||||
# Prepare SRT content as a string
|
||||
results = []
|
||||
for i, segment in enumerate(segments):
|
||||
start_time = format_time(segment.start)
|
||||
end_time = format_time(segment.end)
|
||||
srt_content = f"{i + 1}\n"
|
||||
srt_content += f"{start_time} --> {end_time}\n"
|
||||
|
||||
text=segment.text.strip()
|
||||
|
||||
srt_content += f"{text}\n\n"
|
||||
|
||||
start_time=segment.start
|
||||
end_time=segment.end
|
||||
|
||||
|
||||
results.append({
|
||||
"srt_content":srt_content,
|
||||
"start_time":start_time,
|
||||
"end_time":end_time,
|
||||
"text":text,
|
||||
"language":[info.language]
|
||||
})
|
||||
|
||||
srt_content="\n".join([s['srt_content'] for s in results])
|
||||
text="\n".join([s['text'] for s in results])
|
||||
|
||||
return (results,srt_content,text,total_length_seconds,)
|
||||
|
||||
+6
-4
@@ -7,6 +7,7 @@ import os
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
import hashlib
|
||||
from uuid import uuid4
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
@@ -26,7 +27,7 @@ def tensor_to_hash(tensor):
|
||||
return hash_value
|
||||
|
||||
|
||||
def create_temp_file(image):
|
||||
def create_temp_file(image, uuid):
|
||||
output_dir = folder_paths.get_temp_directory()
|
||||
|
||||
(
|
||||
@@ -35,7 +36,7 @@ def create_temp_file(image):
|
||||
counter,
|
||||
subfolder,
|
||||
_,
|
||||
) = folder_paths.get_save_image_path('material', output_dir)
|
||||
) = folder_paths.get_save_image_path(f'material_{uuid}', output_dir)
|
||||
|
||||
|
||||
image=tensor2pil(image)
|
||||
@@ -59,6 +60,7 @@ class EditMask:
|
||||
|
||||
def __init__(self):
|
||||
self.image_id = None
|
||||
self.uuid = str(uuid4())
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -117,13 +119,13 @@ class EditMask:
|
||||
image_path = os.path.join(base_dir,subfolder, name)
|
||||
|
||||
if image_path==None:
|
||||
image_path,images=create_temp_file(image)
|
||||
image_path,images=create_temp_file(image, self.uuid)
|
||||
|
||||
print('#image_path',os.path.exists(image_path),image_path)
|
||||
# image_path = folder_paths.get_annotated_filepath(image) #文件名
|
||||
|
||||
if not os.path.exists(image_path):
|
||||
image_path,images=create_temp_file(image)
|
||||
image_path,images=create_temp_file(image, self.uuid)
|
||||
|
||||
|
||||
img = node_helpers.pillow(Image.open, image_path)
|
||||
|
||||
@@ -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
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-mixlab-nodes"
|
||||
description = "3D, ScreenShareNode & FloatingVideoNode, SpeechRecognition & SpeechSynthesis, GPT, LoadImagesFromLocal, Layers, Other Nodes, ..."
|
||||
version = "0.41.0"
|
||||
version = "0.46.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"]
|
||||
|
||||
|
||||
+15
-2
@@ -5,7 +5,7 @@ opencv-python-headless
|
||||
matplotlib
|
||||
openai
|
||||
torchaudio
|
||||
# simple-lama-inpainting
|
||||
|
||||
clip-interrogator==0.6.0
|
||||
transformers>=4.36.0
|
||||
lark-parser
|
||||
@@ -21,4 +21,17 @@ soundfile>=0.12.1
|
||||
json-repair
|
||||
|
||||
bitsandbytes
|
||||
accelerate
|
||||
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
|
||||
|
||||
faster_whisper
|
||||
|
||||
git+https://github.com/openai/swarm.git
|
||||
|
||||
@@ -214,7 +214,8 @@ async function extractInputAndOutputData (
|
||||
node.type === 'ChinesePrompt_Mix' ||
|
||||
node.type === 'Seed_' ||
|
||||
node.type === 'SiliconflowLLM' ||
|
||||
node.type === 'ChatGPTOpenAI'
|
||||
node.type === 'ChatGPTOpenAI' ||
|
||||
node.type === 'SiliconflowTextToImageNode'
|
||||
) {
|
||||
// seed 的类型收集
|
||||
try {
|
||||
@@ -266,7 +267,13 @@ function downloadJsonFile (jsonData, fileName = 'mix_app.json') {
|
||||
}
|
||||
|
||||
async function save (json, download = false, showInfo = true) {
|
||||
let nodesAll = window._nodesAll || (await getObjectInfo())
|
||||
if (!window._nodesAll) {
|
||||
window._nodesAll = await getObjectInfo();
|
||||
}
|
||||
|
||||
let nodesAll = window._nodesAll;
|
||||
|
||||
// let nodesAll = window._nodesAll || (await getObjectInfo())
|
||||
|
||||
console.log('####SAVE', nodesAll, json)
|
||||
|
||||
@@ -352,7 +359,7 @@ async function save (json, download = false, showInfo = true) {
|
||||
const imgurl = images[index]
|
||||
images[index] = await drawImageToCanvas(imgurl)
|
||||
}
|
||||
data.app.idle_animation = images
|
||||
if (idle_animation) data.app.idle_animation = images
|
||||
} catch (error) {}
|
||||
|
||||
// console.log(data.app)
|
||||
@@ -416,9 +423,9 @@ function getInputsAndOutputs () {
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.utils.AppInfo',
|
||||
init () {
|
||||
if (!window._nodesAll) {
|
||||
getObjectInfo().then(r => (window._nodesAll = r))
|
||||
}
|
||||
// if (!window._nodesAll) {
|
||||
// getObjectInfo().then(r => (window._nodesAll = r))
|
||||
// }
|
||||
},
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'AppInfo') {
|
||||
|
||||
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
|
||||
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
|
||||
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
|
||||
|
||||
const version = 'v0.41.0'
|
||||
const version = 'v0.46.0'
|
||||
|
||||
fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
|
||||
.then(response => response.json())
|
||||
|
||||
@@ -3,7 +3,7 @@ import { api } from '../../../scripts/api.js'
|
||||
import { ComfyWidgets } from '../../../scripts/widgets.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
import WaveSurfer from 'https://cdn.jsdelivr.net/npm/wavesurfer.js@7/dist/wavesurfer.esm.js'
|
||||
import WaveSurfer from './wavesurfer.esm.js'
|
||||
|
||||
function get_position_style (ctx, widget_width, y, node_height) {
|
||||
const MARGIN = 4 // the margin around the html element
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -544,7 +544,7 @@ async function getCustomnodeMappings () {
|
||||
const data = (await get_nodes_map()).data
|
||||
window._nodes_maps = data
|
||||
}
|
||||
console.log('#getCustomnodeMappings', window._nodes_maps)
|
||||
// console.log('#getCustomnodeMappings', window._nodes_maps)
|
||||
for (let url in window._nodes_maps) {
|
||||
let n = window._nodes_maps[url]
|
||||
for (let node of n[0]) {
|
||||
@@ -1735,14 +1735,16 @@ app.registerExtension({
|
||||
|
||||
// 把json往里 拖
|
||||
document.addEventListener('drop', async event => {
|
||||
event.preventDefault()
|
||||
event.stopPropagation()
|
||||
|
||||
// Dragging from Chrome->Firefox there is a file but its a bmp, so ignore that
|
||||
// Only intercept the JSON files handled here. Calling preventDefault()
|
||||
// unconditionally swallowed every drop, so ComfyUI's native drag&drop
|
||||
// (guarded by event.defaultPrevented) never loaded dropped workflows
|
||||
// (PNG / JSON / etc.). Keep preventDefault scoped to the handled case.
|
||||
if (
|
||||
event.dataTransfer.files.length &&
|
||||
event.dataTransfer.files[0].type == 'application/json'
|
||||
) {
|
||||
event.preventDefault()
|
||||
event.stopPropagation()
|
||||
const reader = new FileReader()
|
||||
reader.onload = async () => {
|
||||
loadAppJson(reader.result)
|
||||
@@ -2171,6 +2173,10 @@ app.registerExtension({
|
||||
|
||||
fetch('manager/badge_mode').then(r => {
|
||||
if (r.status === 404) {
|
||||
// 已有ComfyUI自带的badge
|
||||
if(node.badges?.[0]?.()){
|
||||
return
|
||||
}
|
||||
// 右上角的badge是否已经绘制
|
||||
if (!node.badge_enabled) {
|
||||
if (!node.getNickname) {
|
||||
|
||||
@@ -41,6 +41,9 @@ class Visualizer {
|
||||
overflow: 'hidden'
|
||||
})
|
||||
this.iframe.src = '/mixlab/app/' + visualSrc + '.html'
|
||||
// this.iframe.width="300";
|
||||
// this.iframe.height="400";
|
||||
|
||||
console.log('#Visualizer', container, this.iframe)
|
||||
container.appendChild(this.iframe)
|
||||
}
|
||||
@@ -73,7 +76,7 @@ function createVisualizer (node, inputName, typeName, inputData, app) {
|
||||
draw: function (ctx, node, widgetWidth, widgetY, widgetHeight) {
|
||||
const margin = 10
|
||||
const top_offset = 5
|
||||
const visible = app.canvas.ds.scale > 0.5 && this.type === typeName
|
||||
const visible = app.canvas.ds.scale > 0.3 && this.type === typeName
|
||||
const w = widgetWidth - margin * 4
|
||||
const clientRectBound = ctx.canvas.getBoundingClientRect()
|
||||
const transform = new DOMMatrix()
|
||||
@@ -85,12 +88,13 @@ function createVisualizer (node, inputName, typeName, inputData, app) {
|
||||
.translateSelf(margin, margin + widgetY)
|
||||
|
||||
Object.assign(this.visualizer.style, {
|
||||
left: `${transform.a * margin + transform.e + 40}px`,
|
||||
left: `${transform.a * margin + transform.e + 0}px`,
|
||||
top: `${transform.d + transform.f + top_offset}px`,
|
||||
width: `${w * transform.a}px`,
|
||||
height: `${
|
||||
w * transform.d - widgetHeight - margin * 15 * transform.d
|
||||
}px`,
|
||||
height: `${(w * transform.a * 4) / 3 - margin * 5 * transform.d}px`,
|
||||
// height: `${
|
||||
// w * transform.d - widgetHeight - margin * 15 * transform.d
|
||||
// }px`,
|
||||
position: 'absolute',
|
||||
overflow: 'hidden',
|
||||
zIndex: app.graph._nodes.indexOf(node)
|
||||
@@ -137,11 +141,11 @@ function createVisualizer (node, inputName, typeName, inputData, app) {
|
||||
// Make sure visualization iframe is always inside the node when resize the node
|
||||
node.onResize = function () {
|
||||
let [w, h] = this.size
|
||||
if (w <= 600) w = 600
|
||||
if (h <= 500) h = 500
|
||||
if (w <= 300) w = 300
|
||||
if (h <= 400) h = 400
|
||||
|
||||
if (w > 600) {
|
||||
h = w - 100
|
||||
if (w > 300) {
|
||||
h = Math.round((w * 4) / 3)
|
||||
}
|
||||
|
||||
this.size = [w, h]
|
||||
@@ -181,14 +185,14 @@ function registerVisualizer (nodeType, nodeData, nodeClassName, typeName) {
|
||||
app
|
||||
])
|
||||
|
||||
this.setSize([600, 500])
|
||||
this.setSize([300, 400])
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
nodeType.prototype.onExecuted = async function (message) {
|
||||
// Check if reference image and depth map are available
|
||||
console.log("#message",message)
|
||||
console.log('#message', message)
|
||||
if (message.reference_image && message.depth_map) {
|
||||
const params = {}
|
||||
params.reference_image = message.reference_image[0]
|
||||
|
||||
File diff suppressed because one or more lines are too long
+4713
File diff suppressed because one or more lines are too long
+5
-1
@@ -223,6 +223,10 @@
|
||||
margin-top: 0;
|
||||
}
|
||||
|
||||
.image-with-grid img {
|
||||
margin: 0 !important;
|
||||
}
|
||||
|
||||
/* .card:hover {
|
||||
box-shadow: 0px 0px 10px 10px #e9fbfa;
|
||||
} */
|
||||
@@ -1983,6 +1987,7 @@
|
||||
let inp = document.createElement('input')
|
||||
inp.type = 'file'
|
||||
inp.setAttribute('accept', "audio/*")
|
||||
|
||||
inp.style.display = 'none'
|
||||
inp.addEventListener('change', async e => {
|
||||
e.preventDefault()
|
||||
@@ -2546,7 +2551,6 @@
|
||||
let submitDiv = document.createElement('div');
|
||||
submitDiv.className = "submit_div"
|
||||
|
||||
|
||||
submitDiv.appendChild(appStatus);
|
||||
|
||||
const promptButton = document.createElement('button');
|
||||
|
||||
@@ -237,8 +237,11 @@ const sleep = (t = 1000) => {
|
||||
// 方法:旋转摄像机并拍摄图片 // 每次旋转的角度增量,转换为弧度
|
||||
async function captureImages (
|
||||
totalFrames = 20,
|
||||
angleIncrement = THREE.MathUtils.degToRad(1.5)
|
||||
angleIncrement = 1.5,
|
||||
scaleFactor = 1 // 添加放大倍数参数,默认为1
|
||||
) {
|
||||
angleIncrement = THREE.MathUtils.degToRad(angleIncrement)
|
||||
|
||||
// 计算场景中所有物体的中心点
|
||||
const box = new THREE.Box3().setFromObject(scene)
|
||||
const center = new THREE.Vector3()
|
||||
@@ -264,6 +267,21 @@ async function captureImages (
|
||||
const startAngle = initialAngle
|
||||
// - (angleIncrement * totalFrames) / 2
|
||||
|
||||
// 保存原始尺寸
|
||||
const originalWidth = renderer.domElement.width
|
||||
const originalHeight = renderer.domElement.height
|
||||
|
||||
// 调整渲染器尺寸
|
||||
renderer.setSize(
|
||||
originalWidth * scaleFactor,
|
||||
originalHeight * scaleFactor,
|
||||
false
|
||||
)
|
||||
|
||||
// 调整相机的视图矩阵(如果需要)
|
||||
camera.aspect = (originalWidth * scaleFactor) / (originalHeight * scaleFactor)
|
||||
camera.updateProjectionMatrix()
|
||||
|
||||
for (let i = 0; i < totalFrames; i++) {
|
||||
const angle = startAngle + i * angleIncrement
|
||||
|
||||
@@ -284,6 +302,13 @@ async function captureImages (
|
||||
await new Promise(resolve => setTimeout(resolve, 500))
|
||||
}
|
||||
|
||||
// 恢复渲染器尺寸
|
||||
renderer.setSize(originalWidth, originalHeight, false)
|
||||
|
||||
// 恢复相机的视图矩阵
|
||||
camera.aspect = originalWidth / originalHeight
|
||||
camera.updateProjectionMatrix()
|
||||
|
||||
// 恢复相机到初始位置和朝向
|
||||
camera.position.copy(initialPosition)
|
||||
camera.lookAt(initialTarget)
|
||||
@@ -294,7 +319,7 @@ async function captureImages (
|
||||
async function takeScreenshot () {
|
||||
// 更新相机的矩阵,以确保其世界矩阵是最新的
|
||||
camera.updateMatrixWorld()
|
||||
const imgs = await captureImages()
|
||||
const imgs = await captureImages(12,3,4)
|
||||
|
||||
// 获取当前网页的 URL
|
||||
const currentUrl = window.location.href
|
||||
|
||||
@@ -0,0 +1,416 @@
|
||||
{
|
||||
"last_node_id": 9,
|
||||
"last_link_id": 8,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 3,
|
||||
"type": "TextInput_",
|
||||
"pos": [
|
||||
137,
|
||||
420
|
||||
],
|
||||
"size": {
|
||||
"0": 400,
|
||||
"1": 200
|
||||
},
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
2
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"title": "使用 Azure OpenAI",
|
||||
"properties": {
|
||||
"Node name for S&R": "TextInput_"
|
||||
},
|
||||
"widgets_values": [
|
||||
"https://mixcopilot.openai.azure.com"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"type": "KeyInput",
|
||||
"pos": [
|
||||
144,
|
||||
257
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 70
|
||||
},
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "key",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
1
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"title": "使用你自己的key",
|
||||
"properties": {
|
||||
"Node name for S&R": "KeyInput"
|
||||
},
|
||||
"widgets_values": [
|
||||
null,
|
||||
null
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "MultiPersonPodcast",
|
||||
"pos": [
|
||||
1099,
|
||||
480
|
||||
],
|
||||
"size": [
|
||||
481.8963185574753,
|
||||
268.61682945154007
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "speaker",
|
||||
"type": "SPEAKER",
|
||||
"link": 7,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"link": 4,
|
||||
"widget": {
|
||||
"name": "text"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "audio_list",
|
||||
"type": "AUDIO",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"links": [
|
||||
8
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 1
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "MultiPersonPodcast"
|
||||
},
|
||||
"widgets_values": [
|
||||
"小明:大家好,欢迎收听本周的《AI新动态》。我是主持人小明,今天我们有两位嘉宾,分别是小李和小王。大家跟听众打个招呼吧!\n小李:大家好,我是小李,很高兴今天能和大家聊聊最新的AI动态。\n小王:大家好,我是小王,也很期待今天的讨论。",
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
false,
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 7,
|
||||
"type": "LoadSpeaker",
|
||||
"pos": [
|
||||
584,
|
||||
567
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 58
|
||||
},
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "speaker",
|
||||
"type": "SPEAKER",
|
||||
"links": [
|
||||
6
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"title": "opus",
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadSpeaker"
|
||||
},
|
||||
"widgets_values": [
|
||||
"opus_00001"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "SimulateDevDesignDiscussions",
|
||||
"pos": [
|
||||
611,
|
||||
201
|
||||
],
|
||||
"size": [
|
||||
391.9864763335838,
|
||||
217.95792637114943
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "api_key",
|
||||
"type": "STRING",
|
||||
"link": 1,
|
||||
"widget": {
|
||||
"name": "api_key"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "custom_model_name",
|
||||
"type": "STRING",
|
||||
"link": null,
|
||||
"widget": {
|
||||
"name": "custom_model_name"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "custom_api_url",
|
||||
"type": "STRING",
|
||||
"link": 2,
|
||||
"widget": {
|
||||
"name": "custom_api_url"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
3,
|
||||
4
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "SimulateDevDesignDiscussions"
|
||||
},
|
||||
"widgets_values": [
|
||||
"数字艺术好看吗?",
|
||||
"gpt-4o",
|
||||
"openai",
|
||||
"",
|
||||
"",
|
||||
""
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "RenameSpeaker",
|
||||
"pos": [
|
||||
593,
|
||||
681
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 58
|
||||
},
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "speaker",
|
||||
"type": "SPEAKER",
|
||||
"link": 6
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "speaker",
|
||||
"type": "SPEAKER",
|
||||
"links": [
|
||||
7
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "RenameSpeaker"
|
||||
},
|
||||
"widgets_values": [
|
||||
"主持人"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "ShowTextForGPT",
|
||||
"pos": [
|
||||
1071,
|
||||
128
|
||||
],
|
||||
"size": [
|
||||
624.2005965936271,
|
||||
279.47889630613906
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"link": 3,
|
||||
"widget": {
|
||||
"name": "text"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "output_dir",
|
||||
"type": "STRING",
|
||||
"link": null,
|
||||
"widget": {
|
||||
"name": "output_dir"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"links": null,
|
||||
"shape": 6
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "ShowTextForGPT"
|
||||
},
|
||||
"widgets_values": [
|
||||
"",
|
||||
"",
|
||||
"* 主持人:作为一名设计师,你如何定义“好看”的数字艺术?\n设计师:好看的数字艺术?就像你在沙漠中看到绿洲的那一刻,它能吸引你的眼球,抓住你的心,它能传达情感,让人产生共鸣。可能是颜色的对撞,也可能是形状的魔法,总之,它让你想多看几眼,还想收藏到你的精神博物馆里。\n* 主持人:程序员,你们在开发支持数字艺术的软件时,如何确保用户体验的直观性和美观性?\n程序员:哎呀,这可是门艺术活啊!这时候我们可不像写代码那样呆板,想象力飞起来。我们会尽量让界面简洁好用,不搞那些让人摸不着头脑的功能。动效啥的也要调校好,太多就变花里胡哨了,太少用户觉得干巴巴。最重要的是,多听设计师的,他们可是颜值担当啊!\n* 主持人:站在设计师的角度,你觉得技术如何影响了数字艺术的表现力?\n设计师:技术啊,那可是我们的魔法棒!有了高端的硬件和软件,我们可以在屏幕上玩出各种花样,大到宇宙,小到细胞,想象力在技术的加持下,才能飞得更高更远。不管是3D渲染,还是AR互动,技术就是让我们的创意从草图变成现实的桥梁,让我们画布上的每一个像素都能发光。\n* 主持人:不知道程序员又是怎么看待数字艺术的后台开发和前端展示关系的呢?\n程序员:后端和前端就像魔法师和舞台演员。后端是幕后默默挥舞魔法杖,搞定数据处理啊、服务器啥的,让那台机器运转得顺溜。前端呢,就是站在舞台中央光彩夺目,把数据和功能打包成美美的界面展示给用户。说白了,后端是灵魂,前端是颜值,两个缺一不可,配合得好才是真正的艺术!\n* 主持人:感谢大家的参与,今天关于数字艺术的讨论让我受益匪浅。\n程序员:不客气,代码和艺术的碰撞总是火花四射!\n\n设计师:没错,灵感和技术结合,才能创作出让人惊艳的作品。期待下次再聊!"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 9,
|
||||
"type": "PreviewAudio",
|
||||
"pos": [
|
||||
1740,
|
||||
453
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 76
|
||||
},
|
||||
"flags": {},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"link": 8
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewAudio"
|
||||
},
|
||||
"widgets_values": [
|
||||
null
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
1,
|
||||
2,
|
||||
0,
|
||||
1,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
2,
|
||||
3,
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
3,
|
||||
1,
|
||||
0,
|
||||
5,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
4,
|
||||
1,
|
||||
0,
|
||||
6,
|
||||
1,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
6,
|
||||
7,
|
||||
0,
|
||||
8,
|
||||
0,
|
||||
"SPEAKER"
|
||||
],
|
||||
[
|
||||
7,
|
||||
8,
|
||||
0,
|
||||
6,
|
||||
0,
|
||||
"SPEAKER"
|
||||
],
|
||||
[
|
||||
8,
|
||||
6,
|
||||
1,
|
||||
9,
|
||||
0,
|
||||
"AUDIO"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 1.3310000000000006,
|
||||
"offset": [
|
||||
-361.773434010461,
|
||||
27.423855709687306
|
||||
]
|
||||
}
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -0,0 +1,775 @@
|
||||
{
|
||||
"last_node_id": 18,
|
||||
"last_link_id": 15,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 6,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
-32,
|
||||
79
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
1,
|
||||
2,
|
||||
3
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"1.jpg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 7,
|
||||
"type": "KeyInput",
|
||||
"pos": [
|
||||
-34,
|
||||
585
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 94
|
||||
},
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "key",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
4,
|
||||
5,
|
||||
6
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "KeyInput"
|
||||
},
|
||||
"widgets_values": [
|
||||
null,
|
||||
null
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 10,
|
||||
"type": "LoadVideoFromURL",
|
||||
"pos": [
|
||||
1038,
|
||||
21
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
266
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "url",
|
||||
"type": "STRING",
|
||||
"link": 7,
|
||||
"widget": {
|
||||
"name": "url"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "frames",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
8
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadVideoFromURL"
|
||||
},
|
||||
"widgets_values": [
|
||||
"https://example.com/video.mp4",
|
||||
0,
|
||||
"Disabled",
|
||||
512,
|
||||
512,
|
||||
0,
|
||||
0,
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 13,
|
||||
"type": "LoadVideoFromURL",
|
||||
"pos": [
|
||||
1032,
|
||||
715
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 266
|
||||
},
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "url",
|
||||
"type": "STRING",
|
||||
"link": 10,
|
||||
"widget": {
|
||||
"name": "url"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "frames",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
12
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadVideoFromURL"
|
||||
},
|
||||
"widgets_values": [
|
||||
"https://example.com/video.mp4",
|
||||
0,
|
||||
"Disabled",
|
||||
512,
|
||||
512,
|
||||
0,
|
||||
0,
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 11,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
1436,
|
||||
6
|
||||
],
|
||||
"size": [
|
||||
210,
|
||||
246
|
||||
],
|
||||
"flags": {},
|
||||
"order": 11,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 8
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewImage"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "VideoGenKlingNode",
|
||||
"pos": [
|
||||
527.4816383216086,
|
||||
19.100152539671797
|
||||
],
|
||||
"size": {
|
||||
"0": 400,
|
||||
"1": 200
|
||||
},
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 1
|
||||
},
|
||||
{
|
||||
"name": "fal_key",
|
||||
"type": "STRING",
|
||||
"link": 4,
|
||||
"widget": {
|
||||
"name": "fal_key"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
7,
|
||||
13
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VideoGenKlingNode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"The man is shaking his head with a wry smile.\n\n",
|
||||
"5",
|
||||
"16:9",
|
||||
"standard",
|
||||
""
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 16,
|
||||
"type": "ShowTextForGPT",
|
||||
"pos": [
|
||||
1702,
|
||||
-85
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
200
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"link": 13,
|
||||
"widget": {
|
||||
"name": "text"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "output_dir",
|
||||
"type": "STRING",
|
||||
"link": null,
|
||||
"widget": {
|
||||
"name": "output_dir"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"links": null,
|
||||
"shape": 6
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "ShowTextForGPT"
|
||||
},
|
||||
"widgets_values": [
|
||||
"",
|
||||
"",
|
||||
"https://v2.fal.media/files/0f9f44093c7442d5b0819616880bca65_output.mp4"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "VideoGenRunwayGen3Node",
|
||||
"pos": [
|
||||
529,
|
||||
297
|
||||
],
|
||||
"size": {
|
||||
"0": 400,
|
||||
"1": 200
|
||||
},
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 2
|
||||
},
|
||||
{
|
||||
"name": "fal_key",
|
||||
"type": "STRING",
|
||||
"link": 5,
|
||||
"widget": {
|
||||
"name": "fal_key"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
9,
|
||||
14
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VideoGenRunwayGen3Node"
|
||||
},
|
||||
"widgets_values": [
|
||||
"The man is shaking his head with a wry smile.\n\n",
|
||||
"5",
|
||||
"16:9",
|
||||
""
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "VideoGenLumaDreamMachineNode",
|
||||
"pos": [
|
||||
534,
|
||||
570
|
||||
],
|
||||
"size": {
|
||||
"0": 400,
|
||||
"1": 200
|
||||
},
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 3
|
||||
},
|
||||
{
|
||||
"name": "fal_key",
|
||||
"type": "STRING",
|
||||
"link": 6,
|
||||
"widget": {
|
||||
"name": "fal_key"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
10,
|
||||
15
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VideoGenLumaDreamMachineNode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"The man is shaking his head with a wry smile.\n\n",
|
||||
"16:9",
|
||||
"",
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 18,
|
||||
"type": "ShowTextForGPT",
|
||||
"pos": [
|
||||
1719,
|
||||
665
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
200
|
||||
],
|
||||
"flags": {},
|
||||
"order": 10,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"link": 15,
|
||||
"widget": {
|
||||
"name": "text"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "output_dir",
|
||||
"type": "STRING",
|
||||
"link": null,
|
||||
"widget": {
|
||||
"name": "output_dir"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"links": null,
|
||||
"shape": 6
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "ShowTextForGPT"
|
||||
},
|
||||
"widgets_values": [
|
||||
"",
|
||||
"",
|
||||
"https://v2.fal.media/files/a622a5aac002452ba0e75f7d8871389d_output.mp4"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 12,
|
||||
"type": "LoadVideoFromURL",
|
||||
"pos": [
|
||||
1049,
|
||||
349
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 266
|
||||
},
|
||||
"flags": {},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "url",
|
||||
"type": "STRING",
|
||||
"link": 9,
|
||||
"widget": {
|
||||
"name": "url"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "frames",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
11
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadVideoFromURL"
|
||||
},
|
||||
"widgets_values": [
|
||||
"https://example.com/video.mp4",
|
||||
0,
|
||||
"Disabled",
|
||||
512,
|
||||
512,
|
||||
0,
|
||||
0,
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 17,
|
||||
"type": "ShowTextForGPT",
|
||||
"pos": [
|
||||
1711,
|
||||
318
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
200
|
||||
],
|
||||
"flags": {},
|
||||
"order": 8,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"link": 14,
|
||||
"widget": {
|
||||
"name": "text"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "output_dir",
|
||||
"type": "STRING",
|
||||
"link": null,
|
||||
"widget": {
|
||||
"name": "output_dir"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"links": null,
|
||||
"shape": 6
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "ShowTextForGPT"
|
||||
},
|
||||
"widgets_values": [
|
||||
"",
|
||||
"",
|
||||
"https://v2.fal.media/files/755d51c8984a4445a852268a209b07e4_output.mp4"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 15,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
1444,
|
||||
650
|
||||
],
|
||||
"size": [
|
||||
210,
|
||||
246
|
||||
],
|
||||
"flags": {},
|
||||
"order": 13,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 12
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewImage"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 14,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
1416,
|
||||
310
|
||||
],
|
||||
"size": [
|
||||
210,
|
||||
246
|
||||
],
|
||||
"flags": {},
|
||||
"order": 12,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 11
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewImage"
|
||||
}
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
1,
|
||||
6,
|
||||
0,
|
||||
3,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
2,
|
||||
6,
|
||||
0,
|
||||
4,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
3,
|
||||
6,
|
||||
0,
|
||||
5,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
4,
|
||||
7,
|
||||
0,
|
||||
3,
|
||||
1,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
5,
|
||||
7,
|
||||
0,
|
||||
4,
|
||||
1,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
6,
|
||||
7,
|
||||
0,
|
||||
5,
|
||||
1,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
7,
|
||||
3,
|
||||
0,
|
||||
10,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
8,
|
||||
10,
|
||||
0,
|
||||
11,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
9,
|
||||
4,
|
||||
0,
|
||||
12,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
10,
|
||||
5,
|
||||
0,
|
||||
13,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
11,
|
||||
12,
|
||||
0,
|
||||
14,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
12,
|
||||
13,
|
||||
0,
|
||||
15,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
13,
|
||||
3,
|
||||
0,
|
||||
16,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
14,
|
||||
4,
|
||||
0,
|
||||
17,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
15,
|
||||
5,
|
||||
0,
|
||||
18,
|
||||
0,
|
||||
"STRING"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.7247295000000012,
|
||||
"offset": [
|
||||
-59.526177135875514,
|
||||
197.10624904779425
|
||||
]
|
||||
}
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
Reference in New Issue
Block a user