Compare commits

..
284 Commits
Author SHA1 Message Date
shadow c5392aa237 Merge pull request #441 from Creepybits/main
Fix: Correct scheduler input on StyleAligned Sample Reference Latents node
2026-06-04 15:37:37 +08:00
shadow 16a2b55fa1 Merge pull request #462 from LordTaylor/fix/dragdrop-preventdefault-swallows-workflow-drop
Fix: drop handler swallows all drops and breaks ComfyUI native drag&drop
2026-06-04 15:36:43 +08:00
LordTaylorandClaude Opus 4.8 b423b09ff3 Fix: don't swallow non-JSON drops, restoring ComfyUI native drag&drop
The document 'drop' listener called event.preventDefault() and
stopPropagation() unconditionally on every drop. ComfyUI's native
drag&drop handler bails when event.defaultPrevented is already true, so
dropping any workflow file (PNG/JSON/webp) onto the canvas silently did
nothing as long as Mixlab was installed.

Scope preventDefault()/stopPropagation() to the application/json case
actually handled here, so all other drops fall through to ComfyUI.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-05-30 23:25:58 +02:00
Creepybits e66add88cb Update Style.py
Fix: 
Correct scheduler input on StyleAligned Sample Reference Latents node

Description:
The scheduler input for the StyleAligned Sample Reference Latents node was defined using a .reverse() method, which caused it to register as an invalid input type that could not accept connections.

This commit changes the input from a broken socket to a dropdown widget, making it consistent with the StyleAligned Reference Sampler node and allowing it to function as intended.
2025-10-03 00:26:06 +02:00
shadow 32b22c39cb Merge pull request #411 from ComfyNodePRs/update-publish-yaml
Update Github Action for Publishing to Comfy Registry
2025-07-22 09:44:28 +08:00
shadow 259baac177 Merge pull request #430 from torzdf/main
Bugfix: Give EditMask temp files unique filenames
2025-07-22 09:44:03 +08:00
torzdf 67ef8c13a8 Bugfix: Give EditMask temp files unique filenames 2025-07-01 00:50:43 +01:00
shadow b2bb1876de Merge pull request #393 from wengxiaoxiong/main
fix: Pillow 10.0.0+ compatibility
2025-02-05 18:24:45 +08:00
wengxiaoxiong cda4e626e7 fix: fix 'FreeTypeFont' object has no attribute 'getsize' 2025-02-01 15:33:24 +08:00
shadow d835aff0cb Merge pull request #391 from TangYanxin/main
Solve the problem that there are two duplicate badges on the node
2025-01-25 10:35:15 +08:00
snomiao 21e1967c5e chore(publish): update GitHub Actions workflow for node publishing
- Add permissions for writing issues
- Update action version to v1 for publish-node-action
- Add condition to run job only for specific repository owner
2025-01-20 21:29:05 +00:00
唐焱鑫 c9b5baf4d9 Update ui_mixlab.js: if ComfyUI already comes with a badge, don't add a new badge. 2025-01-15 01:16:48 +08:00
shadow 67c974c96e Merge pull request #374 from hieuck/fix-module-'PIL.Image'-has-no-attribute-'ANTIALIAS'
- **Exception Message:** module 'PIL.Image' has no attribute 'ANTIALIAS'
2024-11-26 19:59:58 +08:00
Lê Trung Hiếu b46ccb03c9 Update ImageNode.py 2024-11-26 07:10:48 +07:00
shadow 0ecf98e08b Merge pull request #361 from shiertier/main
fix GFW block cdn.jsdelivr.net
2024-11-21 09:40:33 +08:00
shiertier f024034724 do not perform get object in init 2024-11-01 12:48:48 +08:00
shiertier 327a21f009 Merge pull request #1 from dionren/main
fix GFW block cdn.jsdelivr.net
2024-11-01 00:47:50 +08:00
任嘉 cfc51532b8 fix GFW block wavesurfer.esm.js 2024-10-27 22:03:48 +08:00
任嘉 00988f92e4 fix GFW block cdn.jsdelivr.net
Add wavesurfer.esm.js
2024-10-27 22:02:13 +08:00
shadowcz007 868c6085a8 Update extension-node-map.json 2024-10-25 14:24:39 +08:00
shadowcz007 a47a56bda0 Update ImageNode.py 2024-10-21 08:31:05 +08:00
shadowcz007 3667b42b2f Update ImageNode.py 2024-10-21 08:28:11 +08:00
shadowcz007 7d142d7d62 Update extension-node-map.json 2024-10-19 11:53:10 +08:00
shadow 24863e2ed3 Merge pull request #350 from shadowcz007/video-all-in-one-fal
0.46.0
2024-10-14 10:44:52 +08:00
shadowcz007 fe8b526bbb 0.46.0 2024-10-14 10:44:05 +08:00
shadowcz007 6298be393a add workflow# 2024-10-14 09:46:03 +08:00
shadowcz007 3a7853f9cc init 2024-10-14 09:19:55 +08:00
shadowcz007 4a9413c83d Update ChatGPT.py 2024-10-12 20:54:20 +08:00
shadowcz007 21b04d62ae Update README.md 2024-10-12 20:42:28 +08:00
shadowcz007 96929b6d7c Update README.md 2024-10-12 20:40:06 +08:00
shadowcz007 07712d80a5 add SimulateDevDesignDiscussions 多智能体播客节点 2024-10-12 20:39:39 +08:00
shadow 10c9eff16f Merge pull request #348 from shadowcz007/whisper-sensevoice
Whisper sensevoice
2024-10-12 10:46:17 +08:00
shadowcz007 edd7af986d update 2024-10-12 10:44:39 +08:00
shadowcz007 1dc31927e3 Update extension-node-map.json 2024-10-12 10:43:15 +08:00
shadowcz007 36ef7d25ef Update ui_mixlab.js 2024-10-05 12:31:22 +08:00
shadowcz007 b766b8b65d Update SenseVoice.py 2024-10-03 11:13:14 +08:00
shadowcz007 6579ff20b4 json_string2 2024-10-03 11:12:22 +08:00
shadowcz007 2fbee59c3e fixbug 2024-10-03 11:12:10 +08:00
shadowcz007 d3aaa19148 Update ChatGPT.py 2024-10-02 17:07:10 +08:00
shadowcz007 e32a3675fc Update Whisper.py 2024-10-02 17:05:43 +08:00
shadowcz007 b72e7dda08 Update Audio.py 2024-10-02 17:05:35 +08:00
shadowcz007 0f77f28a95 Update Whisper.py 2024-10-02 16:41:46 +08:00
shadowcz007 289f83675b Update SenseVoice.py 2024-10-02 16:41:44 +08:00
shadowcz007 36633b4c72 update 2024-10-02 14:09:08 +08:00
shadowcz007 4f45457811 Update extension-node-map.json 2024-10-02 11:05:33 +08:00
shadowcz007 c39890cd64 MiniCPM_VQA_Simple add extract_keywords 2024-10-02 11:05:20 +08:00
shadowcz007 90f1e49263 Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-10-02 09:34:42 +08:00
shadowcz007 9a1cf205db Update requirements.txt 2024-10-02 09:34:18 +08:00
shadow 45eacb6a50 Merge pull request #342 from shadowcz007/SenseVoice
Update requirements.txt
2024-10-02 09:23:38 +08:00
shadowcz007 6cb2b57463 Update requirements.txt 2024-10-02 09:23:15 +08:00
shadow 59f654fa39 Merge pull request #341 from shadowcz007/SenseVoice
Sense voice
2024-10-01 23:17:37 +08:00
shadowcz007 be8ccc1dc4 新增 SenseVoice 2024-10-01 23:16:02 +08:00
shadowcz007 228e5d9183 Update SenseVoice.py 2024-10-01 23:13:31 +08:00
shadowcz007 8afe6d0383 update 2024-10-01 22:15:35 +08:00
shadowcz007 5f7190b08f Update SenseVoice.py 2024-10-01 21:01:27 +08:00
shadowcz007 a70a9b4bb1 update 2024-10-01 20:41:05 +08:00
shadowcz007 b796e66890 Create SenseVoice.py 2024-10-01 17:57:54 +08:00
shadow f1a663779a Update README.md 2024-09-27 17:49:37 +08:00
shadowcz007 b0aa972326 Update index.html 2024-09-24 15:05:36 +08:00
shadowcz007 ef927a7ed1 loadimage from path ,优化 排序逻辑 ,增加 sort_by_filename 2024-09-23 10:37:42 +08:00
shadowcz007 aa8fc59051 scenedetect & createJSON & PromptImage
- 优化从视频提取片段,并输出json保存
2024-09-22 21:35:00 +08:00
shadowcz007 ce62204392 Update PromptNode.py 2024-09-22 19:30:35 +08:00
shadowcz007 837f28142d Update extension-node-map.json 2024-09-21 21:17:32 +08:00
shadowcz007 60c79c991d Qwen2.5 2024-09-20 12:20:27 +08:00
shadowcz007 078aaeb679 depth viewer 2024-09-18 18:34:55 +08:00
shadowcz007 d9edbd535e 适配不同前端版本 2024-09-18 15:37:16 +08:00
shadowcz007 b4a61b21c3 Update td_background.js 2024-09-18 15:25:01 +08:00
shadowcz007 bdc4193ffe - 新增API调用图像生成节点 TextToImage Siliconflow,可以直接调用Siliconflow提供的flux生成图像
v0.42.0
2024-09-18 09:45:57 +08:00
shadowcz007 74fdd6e396 Compatible with ComfyUI_frontend v1.2.48. 2024-09-18 09:33:44 +08:00
shadowcz007 b2479ebff2 Update td_background.js 2024-09-18 09:28:25 +08:00
shadowcz007 ce2162c764 add SiliconflowTextToImageNode 2024-09-17 22:29:28 +08:00
shadowcz007 16ffd63c80 Update ChatGPT.py 2024-09-17 21:23:43 +08:00
shadowcz007 8faf68348d fixbug 2024-09-17 18:58:45 +08:00
shadowcz007 02dbc72856 fixbug 2024-09-12 18:38:26 +08:00
shadowcz007 da4dcf92dc Update scenedetectNode.py 2024-09-12 13:58:49 +08:00
shadowcz007 49b750abcc Update FishSpeech.py 2024-09-12 13:55:44 +08:00
shadowcz007 4bb4122628 add fishspeech 2024-09-12 13:54:28 +08:00
shadowcz007 cee54f336e 支持设置采样数量 2024-09-12 11:16:22 +08:00
shadowcz007 e95b3813cc fixbug-image batch 2024-09-12 10:39:32 +08:00
shadowcz007 6815cfb05e textImage add fixed_width 2024-09-11 17:24:37 +08:00
shadowcz007 b6acbbce35 add max_characters_per_line 2024-09-10 21:26:31 +08:00
shadowcz007 399e74877d fixbug 2024-09-10 13:35:32 +08:00
shadowcz007 61083e91a6 add scenedetect 2024-09-10 13:19:08 +08:00
shadowcz007 67b4ec3178 Update Video.py 2024-09-08 09:52:28 +08:00
shadowcz007 0fcb725a7a Update __init__.py 2024-09-08 09:47:54 +08:00
shadowcz007 0dbdcdfdc7 Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-09-08 09:45:45 +08:00
shadowcz007 e426d77353 Update ImageNode.py 2024-09-08 09:44:22 +08:00
shadow bd15e29f17 Merge pull request #314 from DropFan/main
fix Error starting the server: [Errno 8] nodename nor servname provid…
2024-09-07 20:53:58 +08:00
shadow b323d29567 Merge branch 'main' into main 2024-09-07 20:52:57 +08:00
shadowcz007 0e54af3356 her 2024-09-07 20:17:01 +08:00
shadowcz007 97f12f3bed add her demo 2024-09-07 20:13:14 +08:00
shadowcz007 5612047b97 fixbug 2024-09-07 17:15:28 +08:00
shadowcz007 2d147a3ae1 fixbug 2024-09-07 16:36:13 +08:00
shadowcz007 1a93c0f8e8 Update Video.py 2024-09-07 16:13:24 +08:00
shadowcz007 0a2b64881a Update video_mixlab.js 2024-09-07 15:14:07 +08:00
shadowcz007 42e7fe4d93 Update extension-node-map.json 2024-09-07 10:29:15 +08:00
Tiger 7ada28258c optimize node class sequence 2024-09-05 23:54:54 +08:00
Tiger bd312afd00 fix Error starting the server: [Errno 8] nodename nor servname provided, or not known 2024-09-05 22:16:55 +08:00
shadowcz007 d94a8af35b Update ImageNode.py 2024-09-01 18:46:00 +08:00
shadowcz007 078fd10147 Update __init__.py 2024-09-01 18:41:43 +08:00
shadowcz007 824e25d77c Update video_mixlab.js 2024-09-01 17:53:35 +08:00
shadowcz007 899b887e47 Update ImageNode.py 2024-09-01 17:43:03 +08:00
shadowcz007 e58981d8a3 Update requirements.txt 2024-09-01 09:59:02 +08:00
shadowcz007 fc41d977a5 Update main_mixlab.js 2024-08-31 22:16:19 +08:00
shadowcz007 ab6210e667 fixbug 2024-08-31 17:13:06 +08:00
shadowcz007 f41805f053 test 2024-08-31 13:01:43 +08:00
shadowcz007 baa809fcd6 Update command.js 2024-08-30 18:47:04 +08:00
shadowcz007 a38d15e495 Update video_mixlab.js 2024-08-29 18:06:50 +08:00
shadowcz007 e97641372a Update ui_mixlab.js 2024-08-29 18:01:32 +08:00
shadowcz007 9aecc2cb08 fixbug 2024-08-29 17:21:36 +08:00
shadowcz007 697667945e fixbug 2024-08-29 16:43:29 +08:00
shadowcz007 d908024577 Update image_mixlab.js 2024-08-29 09:42:52 +08:00
shadowcz007 a5a656d958 Update image_mixlab.js 2024-08-29 09:40:06 +08:00
shadowcz007 ddc3cf05dd Update image_mixlab.js 2024-08-29 09:14:02 +08:00
shadowcz007 66ad4b0abd Update 3d_mixlab.js 2024-08-26 19:14:07 +08:00
shadowcz007 a66023adc6 fixbug 2024-08-26 18:33:37 +08:00
shadowcz007 7277844128 fixbug 2024-08-26 18:08:15 +08:00
shadowcz007 6ef82b1d56 Update requirements.txt 2024-08-26 18:08:08 +08:00
shadowcz007 8ded4829f3 0.40.0 2024-08-23 12:31:16 +08:00
shadowcz007 c4b6acb916 fixbug 2024-08-23 12:27:53 +08:00
shadowcz007 9beb81c303 Update README.md 2024-08-23 12:04:52 +08:00
shadowcz007 f8dd4c6efa node-not-found 2024-08-23 12:00:41 +08:00
shadowcz007 6ce5aa6a3a Enhanced
Enhanced navigation to GitHub
右键菜单支持 text-to-text,方便对 prompt 词补全,支持云LLM或者是本地LLM。
2024-08-23 10:45:23 +08:00
shadowcz007 d2efa8a90a Update requirements.txt 2024-08-22 20:48:43 +08:00
shadowcz007 c6374063e9 Update MiniCPMNode.py 2024-08-22 20:48:17 +08:00
shadowcz007 0846013378 增加 MiniCPM-V 2.6 int4 2024-08-22 14:46:11 +08:00
shadowcz007 c141ba405f fixbug: 自动监听文件夹 2024-08-20 15:06:58 +08:00
shadowcz007 0320f13a9f fixbug 2024-08-19 11:46:55 +08:00
shadowcz007 8adc34be4d fixbug 2024-08-19 10:29:54 +08:00
shadowcz007 cb6810d3c1 Update TextGenerateNode.py 2024-08-18 14:31:22 +08:00
shadowcz007 d384f64abf update :text-to-text 2024-08-17 22:15:31 +08:00
shadowcz007 ef7035f8ee update 2024-08-17 15:11:22 +08:00
shadowcz007 bfcadde5c3 update 2024-08-17 14:22:31 +08:00
shadowcz007 496ff41782 新ui支持,适配后,暂未全面测试 2024-08-16 18:13:05 +08:00
shadowcz007 46f0be5484 add image 2024-08-14 00:27:50 +08:00
shadowcz007 f3db0131c1 fixbug 2024-08-14 00:02:18 +08:00
shadowcz007 83a8d47f51 Update README.md 2024-08-13 23:32:10 +08:00
shadowcz007 fee0222910 v0.37.0 移动端适配、修改app模式的Mask编辑器 2024-08-12 10:10:43 +08:00
shadowcz007 1ed7b5511f mixlab app new mask editor 2024-08-12 00:19:26 +08:00
shadowcz007 b8f7c31537 Update index.html 2024-08-11 17:53:18 +08:00
shadowcz007 164791c257 webui 移动端适配 2024-08-11 17:20:36 +08:00
shadowcz007 8f5e599928 fixbug & ui 2024-08-11 16:29:16 +08:00
shadowcz007 7a7aaeb84d Update index.html 2024-08-10 23:00:58 +08:00
shadowcz007 e2136ab2fc fixbug 2024-08-10 17:13:28 +08:00
shadowcz007 c75cb21946 clean 2024-08-10 10:47:54 +08:00
shadowcz007 bf95218c91 p5-video-workflow 2024-08-10 00:59:44 +08:00
shadowcz007 8cb4507a5f v0.36.0 p5.js 2024-08-10 00:39:27 +08:00
shadowcz007 555890d1ba Update pyproject.toml 2024-08-09 19:08:36 +08:00
shadowcz007 e4f54e83b6 Update Text-to-Image-app.json 2024-08-09 16:30:47 +08:00
shadowcz007 692c4a709e fixbug:web app 2024-08-09 16:27:34 +08:00
shadowcz007 cbd1961459 test 2024-08-08 21:58:44 +08:00
shadowcz007 2e31a33ebf fixbug 2024-08-08 11:37:36 +08:00
shadowcz007 d16c6137d2 update 2024-08-06 23:07:20 +08:00
shadowcz007 0416ab79ec Update 3d_mixlab.js 2024-08-06 21:08:52 +08:00
shadow fc9a1c62b9 Merge pull request #295 from shadowcz007/0.36.0-py5-processing
Lama 改成手动安装,新增JsonRepair
2024-08-06 11:09:27 +08:00
shadowcz007 5d4567b134 Lama 改成手动安装,新增JsonRepair 2024-08-06 11:08:50 +08:00
shadow ae4a17d271 Merge pull request #293 from shadowcz007/0.36.0-py5-processing
0.36.0 py5 processing
2024-08-06 00:24:13 +08:00
shadowcz007 d110a08889 Update __init__.py 2024-08-06 00:23:34 +08:00
shadowcz007 e0157293cb Update P5.py 2024-08-06 00:21:51 +08:00
shadowcz007 0d985b3b65 update 2024-08-06 00:14:05 +08:00
shadowcz007 a65ade9fda updage 2024-08-05 21:30:30 +08:00
shadowcz007 874d6c8cb1 1 2024-08-05 21:16:19 +08:00
shadowcz007 f70ba2afa3 update 2024-08-05 21:08:32 +08:00
shadowcz007 e9f821e578 update 2024-08-05 20:49:06 +08:00
shadowcz007 8e488d4b1d update 2024-08-05 11:55:07 +08:00
shadowcz007 77201a457d 基本打通 2024-08-04 23:48:06 +08:00
shadowcz007 076e3b1178 test 2024-08-04 22:28:16 +08:00
shadowcz007 6b13fa64dc update 2024-08-04 20:44:56 +08:00
shadowcz007 846671a890 preview audio 2024-08-04 18:06:37 +08:00
shadowcz007 05b3088b75 0.35.1 2024-08-04 18:02:13 +08:00
shadowcz007 fe57286959 v0.34.0 2024-08-04 15:28:47 +08:00
shadowcz007 03645bbb33 image batch to list 2024-08-04 13:35:22 +08:00
shadowcz007 93dba9a399 fixbug :load image (base64) 2024-08-04 12:12:41 +08:00
shadowcz007 5627ea8073 Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-08-04 09:40:43 +08:00
shadowcz007 7ba679c9ce fixbug 2024-08-04 09:40:40 +08:00
shadow c7a450e6ce Merge pull request #289 from ComfyNodePRs/licence-update
Update PyProject Toml - License
2024-08-03 17:50:25 +08:00
snomiao beda5156bf chore(licence-update): Update PyProject Toml - License 2024-08-02 23:03:55 +00:00
shadowcz007 76a9da7163 fixbug 2024-08-02 18:32:27 +08:00
shadowcz007 edd0303f59 App模式增加batch prompt,批量提示词,可以把动态提示词批量组成后运行 2024-08-01 21:12:58 +08:00
shadowcz007 be6f47a333 batch prompt :批量提示 2024-08-01 21:03:40 +08:00
shadowcz007 4cd6a072ca Update install.bat 2024-08-01 11:58:23 +08:00
shadowcz007 743a82efe9 fixbug 2024-07-29 18:29:16 +08:00
shadowcz007 9589f28ef7 v0.32.0 2024-07-29 18:11:57 +08:00
shadowcz007 35492c5671 add SiliconflowLLM 2024-07-29 18:06:32 +08:00
shadow db1e695bf3 Merge pull request #284 from cd0304/main
修正text image节点的padding问题
2024-07-29 17:51:29 +08:00
shadowcz007 ecc4aec43b Update ChatGPT.py 2024-07-29 15:17:00 +08:00
shadowcz007 fc063c2205 Update __init__.py 2024-07-29 14:17:57 +08:00
shadowcz007 4d60ce138a Update __init__.py 2024-07-28 21:12:39 +08:00
shadowcz007 2afd24f6e4 fixbug 2024-07-28 20:52:55 +08:00
shadowcz007 437acd023a fixbug 2024-07-28 20:28:34 +08:00
shadowcz007 b00523ae14 优化mixlab app,前端不传workflow,只传输入和输出 2024-07-28 20:21:53 +08:00
shadowcz007 4405a74993 Update Audio.py 2024-07-26 18:56:38 +08:00
cd0304 cb16090868 Update ImageNode.py 2024-07-26 13:04:17 +08:00
cd0304 396e510dce Update ImageNode.py
fix height
2024-07-26 00:32:56 +08:00
shadowcz007 3b9790b969 Update __init__.py 2024-07-25 13:39:41 +08:00
shadowcz007 a35d07a7ac video 2024-07-17 20:49:15 +08:00
shadowcz007 6d004c61fc Update pyproject.toml 2024-07-17 14:41:33 +08:00
shadowcz007 ffdd06da1b Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-07-17 14:41:02 +08:00
shadowcz007 f03f34cacb Update checkVersion_mixlab.js 2024-07-17 14:40:59 +08:00
shadow 0c86ea849e Merge pull request #273 from cd0304/main
textimge节点增加对otf后缀字体支持
2024-07-17 14:37:35 +08:00
cd0304 0efa4c38c0 Update ImageNode.py 2024-07-17 13:59:22 +08:00
cd0304 6092ab7793 Update ImageNode.py 2024-07-17 13:17:40 +08:00
shadowcz007 929def87eb Update ui_mixlab.js 2024-07-17 11:16:43 +08:00
shadowcz007 be074ccff7 Update __init__.py 2024-07-16 22:47:10 +08:00
shadowcz007 3445199393 AUDIO 2024-07-16 21:38:54 +08:00
shadowcz007 216c7e152e 0.30.3 2024-07-08 00:03:33 +08:00
shadowcz007 cc8bc10690 update 2024-07-07 18:41:56 +08:00
shadowcz007 69b4218d60 Update __init__.py 2024-07-07 17:05:06 +08:00
shadowcz007 1dd18dc4f8 fixbug 2024-07-06 20:30:50 +08:00
shadowcz007 4ccbd999d9 fixbug 2024-07-06 00:54:19 +08:00
shadowcz007 fa8d404964 0.30.2 2024-07-06 00:39:02 +08:00
shadowcz007 30086957c9 fixbug 2024-07-06 00:37:52 +08:00
shadowcz007 0e57c620c9 Update Video.py 2024-07-04 18:19:53 +08:00
shadowcz007 3ce1c59a2d Update README.md 2024-07-04 17:37:50 +08:00
shadowcz007 3337e20b9e Math Operation 2024-06-23 16:50:06 +08:00
shadowcz007 e816b3626e update 2024-06-22 21:44:33 +08:00
shadowcz007 3e0cb0f17a Update ui_mixlab.js 2024-06-22 18:42:12 +08:00
shadowcz007 41bc606217 Update 2-screeshare.json 2024-06-22 11:56:32 +08:00
shadowcz007 5a5f4ca49a Update pyproject.toml 2024-06-21 23:08:27 +08:00
shadowcz007 c3a8437cd1 Update ImageNode.py 2024-06-21 22:10:19 +08:00
shadowcz007 8d8a1a392d fixbug 2024-06-21 21:54:41 +08:00
shadowcz007 5f93fb5e55 增加支持的国产大模型 2024-06-21 17:40:15 +08:00
shadowcz007 d05050d7d8 v0.30.1 2024-06-20 20:39:43 +08:00
shadowcz007 8e9744100d 优化composite images节点 2024-06-20 17:46:46 +08:00
shadowcz007 1e4e7e287d Update ImageNode.py 2024-06-20 16:33:58 +08:00
shadowcz007 e8f0c73f08 优化text image节点,更为精准控制空白间距,字体修改为选择方式 2024-06-20 16:32:04 +08:00
shadowcz007 e923e28f8d Canvas Mode 2024-06-20 14:59:51 +08:00
shadowcz007 5cc75bfa7c Merge branch 'main' of https://github.com/shadowcz007/comfyui-mixlab-nodes 2024-06-20 12:04:29 +08:00
shadowcz007 d6701769b8 fixbug:showtext 2024-06-20 12:04:23 +08:00
shadow 0ddc67bdab Create CNAME 2024-06-20 11:13:12 +08:00
shadowcz007 38b62b7a68 Update pyproject.toml 2024-06-19 11:10:27 +08:00
shadowcz007 7e726000c7 v0.30.0 2024-06-18 16:57:14 +08:00
shadowcz007 d8dfb292ec 增加 Edit Mask & SD3 示例 2024-06-18 16:55:59 +08:00
shadowcz007 826975241d Audio Play 2024-06-17 10:50:26 +08:00
shadowcz007 743637ceaf Update Video.py 2024-06-14 11:50:05 +08:00
shadowcz007 e350c7e31e CombineAudioVideo、LoadAndCombinedAudio 2024-06-14 11:38:10 +08:00
shadowcz007 66b1e0ab9f Update __init__.py 2024-06-14 08:17:04 +08:00
shadowcz007 7b0374d110 Update requirements.txt 2024-06-13 09:04:48 +08:00
shadowcz007 e86ef8cbb0 ImageBatchToList、LoadAndCombinedAudio、combine_audio_video、GenerateFramesByCount 2024-06-12 20:54:16 +08:00
shadowcz007 8c901c54bc Update extension-node-map.json 2024-06-08 17:41:40 +08:00
shadowcz007 408d85691e v0.29.0 支持把输出显示到comfyui背景(TouchDesigner 风格) 2024-06-08 16:58:21 +08:00
shadowcz007 c66cd6901b appinfo add performance features
Appinfo supports outputting to the background, enhancing the performance features of ComfyUI.
2024-06-08 16:03:10 +08:00
shadowcz007 aeadbc4f6d fixbug 2024-06-06 15:17:40 +08:00
shadowcz007 224136890e fixbug 2024-06-06 08:02:51 +08:00
shadowcz007 3669a1e86d 0.28.3 2024-06-01 23:21:33 +08:00
shadowcz007 d588b5b327 Update index.html 2024-05-29 22:37:52 +08:00
shadowcz007 b705679098 Update index.html 2024-05-29 21:49:32 +08:00
shadowcz007 f71a0b0da5 Update index.html 2024-05-29 20:20:19 +08:00
shadowcz007 ebc2c76b6b fixbug 2024-05-25 22:50:19 +08:00
shadow 2e3fff278e Merge pull request #240 from audioscavenger/patch-1
Update extension-node-map.json
2024-05-24 11:12:17 +08:00
Eric 1f4bc5e089 Update extension-node-map.json
i'm the new maintainer, thanks
2024-05-23 16:41:33 -07:00
shadowcz007 52c38b10dd v0.28.2 2024-05-23 18:19:30 +08:00
shadowcz007 7047aa5456 add video format 2024-05-23 16:59:04 +08:00
shadowcz007 33fe4019f7 Update ui_mixlab.js 2024-05-23 16:43:50 +08:00
shadowcz007 80b9d97690 Update Video.py 2024-05-23 15:58:24 +08:00
shadowcz007 3c3c92723f Update pyproject.toml 2024-05-23 10:34:36 +08:00
shadowcz007 037bd87006 Update pyproject.toml 2024-05-23 10:26:48 +08:00
shadowcz007 f688310d28 Update Utils.py 2024-05-23 10:14:25 +08:00
shadow c4d65e7a45 Merge pull request #234 from haohaocreates/publish
Add Github Action for Publishing to Comfy Registry
2024-05-22 23:14:33 +08:00
shadow 6f208b710d Merge pull request #235 from haohaocreates/pyproject
Add pyproject.toml for Custom Node Registry
2024-05-22 23:14:17 +08:00
haohaocreates b599faaf85 Update pyproject.toml desc 2024-05-21 15:24:07 -04:00
haohaocreates 6d991d20dc chore(publish): Add Github Action for Publishing to Comfy Registry 2024-05-21 19:19:01 +00:00
haohaocreates c87e0296f6 chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-05-21 19:19:01 +00:00
shadow 16cdb4c5b4 Merge pull request #231 from 295958090/main
修复FloatSlider的bug
2024-05-21 21:57:10 +08:00
Bai Shui 7631b8924d 修复bug 2024-05-21 13:34:29 +08:00
shadowcz007 785d307ff3 Update index.html 2024-05-18 16:13:30 +08:00
shadowcz007 8c713ff35e Update index.html 2024-05-18 16:08:25 +08:00
shadowcz007 7d80493bef Update index.html 2024-05-18 16:05:39 +08:00
shadowcz007 bff2760c3d v0.28.1
修复bug
2024-05-18 11:38:50 +08:00
shadowcz007 a0f8848367 修复 当上传新的图片,编辑mask的bug 2024-05-18 11:38:28 +08:00
shadowcz007 5b1cbcd8d5 修复bug 2024-05-16 13:20:06 +08:00
shadowcz007 05857a92d5 v0.28.0
add rembg api & webapp rembg
2024-05-16 11:49:45 +08:00
shadowcz007 6bdc811286 add rembg api & webapp rembg 2024-05-16 11:49:14 +08:00
shadowcz007 469d50a5b8 Update index.html 2024-05-16 09:02:13 +08:00
shadowcz007 ef86904bfb Update ui_mixlab.js 2024-05-16 09:02:08 +08:00
shadowcz007 d4181ea67c v0.27.1 fixbug 2024-05-16 08:49:30 +08:00
shadowcz007 1c6d17309f Update index.html 2024-05-16 08:49:09 +08:00
shadowcz007 db293ec41d fixbug css 2024-05-16 08:47:17 +08:00
shadowcz007 db8d468f29 0.27.0 增加webapp的mask绘制 2024-05-16 00:11:16 +08:00
shadowcz007 d7d7af7265 add mask edit for webapp 2024-05-16 00:06:42 +08:00
shadowcz007 bd763cadc1 fixbug 2024-05-16 00:05:24 +08:00
shadowcz007 22799fc549 fixbug for mask 2024-05-16 00:05:16 +08:00
shadowcz007 0f231d1271 add minPaint for mask 2024-05-16 00:04:54 +08:00
shadowcz007 8c0c911020 Create LICENSE 2024-05-14 10:02:46 +08:00
173 changed files with 159989 additions and 5958 deletions
+25
View File
@@ -0,0 +1,25 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
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@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+1
View File
@@ -0,0 +1 @@
mixlabnodes.com
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2024 shadow
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+101 -17
View File
@@ -1,26 +1,68 @@
![](https://img.shields.io/github/release/shadowcz007/comfyui-mixlab-nodes)
> 适配了最新版 comfyui 的 py3.11 ,torch 2.1.2+cu121
> 适配了最新版 comfyui 的 py3.11 ,torch 2.3.1+cu121
> [Mixlab nodes discord](https://discord.gg/cXs9vZSqeK)
商务合作请联系 389570357@qq.com
For business cooperation, please contact email 389570357@qq.com
![Her 的DEMO页面](assets/1725710761451.png)
##### `最新`:
ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。模型下载后,放置到 `models/llamafile/`
- 新增[fal.ai](https://fal.ai/dashboard)的视频生成:Kling、RunwayGen3、LumaDreamMachine,[工作流下载](./workflow/video-all-in-one-test-workflow.json)
- 右键菜单支持 text-to-text,方便对 prompt 词补全
- 新增 SimulateDevDesignDiscussions,需要安装[swarm](https://github.com/openai/swarm)和[Comfyui-ChatTTS](https://github.com/shadowcz007/Comfyui-ChatTTS),[工作流下载](./workflow/swarm制作的播客节点workflow.json)
强烈推荐:[Phi-3-mini-4k-instruct-GGUF](https://huggingface.co/lmstudio-community/Phi-3-mini-4k-instruct-GGUF/tree/main),备选:[llama3_if_ai_sdpromptmkr_q2k](https://hf-mirror.com/impactframes/llama3_if_ai_sdpromptmkr_q2k/tree/main)
- 新增 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。
- 增加 MiniCPM-V 2.6 int4
This is the int4 quantized version of MiniCPM-V 2.6.
Running with int4 version would use lower GPU memory (about 7GB).
- 移动端适配、修改 app 模式的 Mask 编辑器
- 增加 p5.js 作为输入节点
[workflow](./workflow/p5workflow.json)
[workflow2](./workflow/p5-video-workflow.json)
- App 模式增加 batch prompt,批量提示词,可以把动态提示词批量组成后运行
![alt text](./assets/1722517810720.png)
- 增加 API Key Input 节点,用于管理 LLM 的 Key,同时优化 LLM 相关节点,为后续 agent 模式做准备
- 增加 SiliconflowLLM,可以使用由 Siliconflow 提供的免费 LLM
<!-- - ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。模型下载后,放置到 `models/llamafile/` -->
<!--
强烈推荐:
[Phi-3-mini-4k-instruct-function-calling-GGUF](https://huggingface.co/nold/Phi-3-mini-4k-instruct-function-calling-GGUF)
[Phi-3-mini-4k-instruct-GGUF](https://huggingface.co/lmstudio-community/Phi-3-mini-4k-instruct-GGUF/tree/main),备选:[llama3_if_ai_sdpromptmkr_q2k](https://hf-mirror.com/impactframes/llama3_if_ai_sdpromptmkr_q2k/tree/main)
- 右键菜单支持 image-to-text,使用多模态模型,多模态使用 [llava-phi-3-mini-gguf](https://huggingface.co/xtuner/llava-phi-3-mini-gguf/tree/main),注意需要把llava-phi-3-mini-mmproj-f16.gguf也下载
![](./assets/prompt_ai_setup.png)
![](./assets/prompt-ai.png)
![](./assets/prompt-ai.png) -->
#### `相关插件推荐`
<!-- [comfyui-sd-prompt-mixlab](https://github.com/shadowcz007/comfyui-sd-prompt-mixlab) -->
[comfyui-liveportrait](https://github.com/shadowcz007/comfyui-liveportrait)
[Comfyui-ChatTTS](https://github.com/shadowcz007/Comfyui-ChatTTS)
[comfyui-sound-lab](https://github.com/shadowcz007/comfyui-sound-lab)
[comfyui-Image-reward](https://github.com/shadowcz007/comfyui-Image-reward)
@@ -37,6 +79,8 @@ ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一
- 发布为 app 的 workflow,可以在右键里再次编辑了
- web app 可以设置分类,在 comfyui 右键菜单可以编辑更新 web app
- 支持动态提示
- 支持把输出显示到 comfyui 背景(TouchDesigner 风格)
- 如果转为 web app 打开是空白的,注意检查下插件目录的名字需要是:comfyui-mixlab-nodes(如果是 zip 包下载会多了个-main 的后缀,需要去掉)
![](./assets/微信图片_20240421205440.png)
@@ -95,15 +139,20 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
[Voice + Real-time Face Swap Workflow](./workflow/语音+实时换脸workflow.json)
- Preview Audio
[text-to-audio](./workflow/text-to-audio-base-workflow.json)
### GPT
> Support for calling multiple GPTs.Local LLM(llama.cpp)、 ChatGPT、ChatGLM3 、ChatGLM4 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1 . Azure OpenAI:https://xxxx.openai.azure.com
> Support for calling multiple GPTs.Local LLM 、 ChatGPT、ChatGLM3 、ChatGLM4 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1 . Azure OpenAI:https://xxxx.openai.azure.com
![gpt-workflow.svg](./assets/gpt-workflow.svg)
[LLM_base_workflow](./workflow/LLM_base_workflow.json)
[workflow-5](./workflow/5-gpt-workflow.json)
- SiliconflowLLM
- ChatGPTOpenAI
最新:ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
<!-- 最新:ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
Model download,move to :`models/llamafile/`
@@ -131,7 +180,7 @@ pip install 'llama-cpp-python[server]'
```
pip install llama-cpp-python \
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/metal
```
``` -->
## Prompt
@@ -158,6 +207,8 @@ pip install llama-cpp-python \
> A new layer class node has been added, allowing you to separate the image into layers. After merging the images, you can input the controlnet for further processing.
> The composite images node overlays a foreground image onto a background image at specified positions and scales, with optional blending modes and masking capabilities. position : 'overall',"center_center","left_bottom","center_bottom","right_bottom","left_top","center_top","right_top"
![layers](./assets/layers-workflow.svg)
![poster](./assets/poster-workflow.svg)
@@ -189,6 +240,19 @@ pip install llama-cpp-python \
> Conveniently load images from a fixed address on the internet to ensure that default images in the workflow can be executed.
#### TextImage
> [下载字体](https://drxie.github.io/OSFCC/)放到 `custom_nodes/comfyui-mixlab-nodes/assets/fonts`
#### MiniCPM-VQA Simple
This is the int4 quantized version of MiniCPM-V 2.6.
Running with int4 version would use lower GPU memory (about 7GB).
[模型](https://huggingface.co/openbmb/MiniCPM-V-2_6-int4)
![alt text](assets/1724308322276.png)
### Style
> Apply VisualStyle Prompting , Modified from [ComfyUI_VisualStylePrompting](https://github.com/ExponentialML/ComfyUI_VisualStylePrompting)
@@ -211,6 +275,8 @@ pip install llama-cpp-python \
### Other Nodes
- 增加 Edit Mask,方便在生成的时候手动绘制 mask [workflow](./workflow/edit-mask-workflow.json)
![main](./assets/all-workflow.svg)
![main2](./assets/detect-face-all.png)
@@ -226,27 +292,45 @@ Add edges to an image.
![FeatheredMask](./assets/FlVou_Y6kaGWYoEj1Tn0aTd4AjMI.jpg)
> LaMaInpainting
> LaMaInpainting(需要手动安装)
- simple-lama-inpainting 里的 pillow 造成冲突,暂时从依赖里移除,如果有安装 simple-lama-inpainting ,节点会自动添加,没有,则不会自动添加。
from [simple-lama-inpainting](https://github.com/enesmsahin/simple-lama-inpainting)
- [问题汇总](https://github.com/shadowcz007/comfyui-mixlab-nodes/issues/294)
> rembgNode
"briarmbg","u2net","u2netp","u2net_human_seg","u2net_cloth_seg","silueta","isnet-general-use","isnet-anime"
**_ briarmbg _** model was developed by BRlA Al and can be used as an open-source model for non-commercial purposes
### Improvement
### Enhancement
- Add "help" option to the context menu for each node.
- Add "Nodes Map" option to the global context menu.
- Direct "Help" option accessible through node context menu.
An improvement has been made to directly redirect to GitHub to search for missing nodes when loading the graph.
- "Nodes Map" feature added to global context menu.
- An improvement has been made to directly redirect to GitHub to search for missing nodes when loading the graph.
*** If not needed, you can comment out ```app.showMissingNodesError``` in the ```ui_mixlab.js``` file.
![help](./assets/help.png)
![node-not-found](./assets/node-not-found.png)
- Right-click shortcut
右键菜单支持 text-to-text,方便对 prompt 词补全,支持云LLM或者是本地LLM。
The right-click menu supports text-to-text conversion, facilitating prompt word completion, and supports cloud LLMs or local LLMs.
Local LLM API example:```http://localhost:1234/v1```
![alt text](./assets/1724380841822.png)
### Models
- [Download TripoSR](https://huggingface.co/stabilityai/TripoSR/blob/main/model.ckpt) and place it in `models/triposr`
+775 -282
View File
File diff suppressed because it is too large Load Diff
Binary file not shown.

After

Width:  |  Height:  |  Size: 537 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 340 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 29 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 366 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 2.2 MiB

Binary file not shown.
Binary file not shown.

Before

Width:  |  Height:  |  Size: 35 KiB

After

Width:  |  Height:  |  Size: 94 KiB

+17199 -582
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+2 -2
View File
@@ -11,9 +11,9 @@ if exist "%python_exec%" (
%python_exec% -s -m pip install "%%i" -i https://pypi.tuna.tsinghua.edu.cn/simple
)
%python_exec% -s -m pip install --upgrade --force llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
%python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
) else (
+150 -35
View File
@@ -1,4 +1,100 @@
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,)
@@ -55,46 +151,65 @@ class SpeechSynthesis:
return {"ui": {"text": text}, "result": (text,)}
#
class GamePal:
class AudioPlayNode:
def __init__(self):
self.output_dir = folder_paths.get_temp_directory()
self.type = "temp"
self.prefix_append = ""
self.compress_level = 4
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"input_text": ("STRING",{"multiline": True,"default": ""}),
},
"optional": {
"input_num": ("INT",{
"default":100,
"min": -1, #Minimum value
"max": 0xffffffffffffffff, #Maximum value
"step": 1, #Slider's step
"display": "slider" # Cosmetic only: display as "number" or "slider"
}),
"python_code": ("STRING",{"multiline": True,"default": "result= 1 if 'Mixlab' in input_text else 0"}),
}
}
INPUT_IS_LIST = False
RETURN_TYPES = ("INT",)
return {"required": {
"audio": ("AUDIO",),
},
}
RETURN_TYPES = ()
FUNCTION = "run"
OUTPUT_NODE = True
OUTPUT_IS_LIST = (False,)
CATEGORY = "♾️Mixlab/Audio"
def run(self, input_text,input_num,python_code):
exec(python_code)
res=None
try:
# 可能会引发异常的代码
res=result
except:
# 处理异常的代码
print('')
INPUT_IS_LIST = False
OUTPUT_IS_LIST = ()
print(res)
OUTPUT_NODE = True
def run(self,audio):
# print(session_history)
return {"ui": {"text": [input_text],"num":[input_num]}, "result": (res,)}
# 判断是否是 Tensor 类型
is_tensor = not isinstance(audio, dict)
# print('#判断是否是 Tensor 类型',is_tensor,audio)
if not is_tensor and 'waveform' in audio and 'sample_rate' in audio:
# {'waveform': tensor([], size=(1, 1, 0)), 'sample_rate': 44100}
is_tensor=True
if is_tensor and (not 'audio_path' in audio):
filename_prefix=""
# 保存
filename_prefix += self.prefix_append
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
results = list()
filename_with_batch_num = filename.replace("%batch_num%", str(1))
file = f"{filename_with_batch_num}_{counter:05}_.wav"
torchaudio.save(os.path.join(full_output_folder, file), audio['waveform'].squeeze(0), audio["sample_rate"])
results.append({
"filename": file,
"subfolder": subfolder,
"type": self.type
})
else:
results=[{
"filename": audio['filename'],
"subfolder":audio['subfolder'],
"type": audio['type'],
"audio_path":audio['audio_path']
}]
# print(audio)
return {"ui": {"audio":results}}
+699 -98
View File
@@ -1,4 +1,6 @@
import openai
from swarm import Swarm, Agent
import time
import urllib.error
import re,json,os,string,random
@@ -6,14 +8,79 @@ import folder_paths
import hashlib
import codecs,sys
import importlib.util
import subprocess
import requests
from PIL import Image
from io import BytesIO
import torch
import numpy as np
python = sys.executable
# Convert PIL to Tensor
def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
def is_installed(package):
# 从文本中提取json
def extract_json_strings(text):
json_strings = []
brace_level = 0
json_str = ''
in_json = False
for char in text:
if char == '{':
brace_level += 1
in_json = True
if in_json:
json_str += char
if char == '}':
brace_level -= 1
if in_json and brace_level == 0:
json_strings.append(json_str)
json_str = ''
in_json = False
return json_strings[0] if len(json_strings)>0 else "{}"
def is_installed(package, package_overwrite=None,auto_install=True):
is_has=False
try:
spec = importlib.util.find_spec(package)
is_has=spec is not None
except ModuleNotFoundError:
return False
return spec is not None
pass
package = package_overwrite or package
if spec is None:
if auto_install==True:
print(f"Installing {package}...")
# 清华源 -i https://pypi.tuna.tsinghua.edu.cn/simple
command = f'"{python}" -m pip install {package}'
result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=True, env=os.environ)
is_has=True
if result.returncode != 0:
print(f"Couldn't install\nCommand: {command}\nError code: {result.returncode}")
is_has=False
else:
print(package+'## OK')
return is_has
# def is_installed(package):
# try:
# spec = importlib.util.find_spec(package)
# except ModuleNotFoundError:
# return False
# return spec is not None
def get_unique_hash(string):
@@ -53,30 +120,14 @@ def azure_client(key,url):
def openai_client(key,url):
client = openai.OpenAI(
api_key=key,
base_url=url
api_key=key,
base_url=url
)
return client
def ZhipuAI_client(key):
try:
if is_installed('zhipuai')==False:
import subprocess
# 安装
print('#pip install zhipuai')
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'zhipuai'], capture_output=True, text=True)
#检查命令执行结果
if result.returncode == 0:
print("#install success")
from zhipuai import ZhipuAI
else:
print("#install error")
else:
if is_installed('zhipuai')==True:
from zhipuai import ZhipuAI
except:
print("#install zhipuai error")
@@ -97,73 +148,76 @@ def get_llama_path():
except:
return os.path.join(folder_paths.models_dir, "llamafile")
def get_llama_models():
res=[]
# def get_llama_models():
# res=[]
model_path=get_llama_path()
if os.path.exists(model_path):
files = os.listdir(model_path)
for file in files:
if os.path.isfile(os.path.join(model_path, file)):
res.append(file)
res=phi_sort(res)
return res
# model_path=get_llama_path()
# if os.path.exists(model_path):
# files = os.listdir(model_path)
# for file in files:
# if os.path.isfile(os.path.join(model_path, file)):
# res.append(file)
# res=phi_sort(res)
# return res
llama_modes_list=get_llama_models()
# llama_modes_list=get_llama_models()
# llama_modes_list=[]
def get_llama_model_path(file_name):
model_path=get_llama_path()
mp=os.path.join(model_path,file_name)
return mp
# def get_llama_model_path(file_name):
# model_path=get_llama_path()
# mp=os.path.join(model_path,file_name)
# return mp
def llama_cpp_client(file_name):
try:
if is_installed('llama_cpp')==False:
import subprocess
# def llama_cpp_client(file_name):
# try:
# if is_installed('llama_cpp')==False:
# import subprocess
# 安装
print('#pip install llama-cpp-python')
# # 安装
# print('#pip install llama-cpp-python')
result = subprocess.run([sys.executable, '-s', '-m', 'pip',
'install',
'llama-cpp-python',
'--extra-index-url',
'https://abetlen.github.io/llama-cpp-python/whl/cu121'
], capture_output=True, text=True)
# result = subprocess.run([sys.executable, '-s', '-m', 'pip',
# 'install',
# 'llama-cpp-python',
# '--extra-index-url',
# 'https://abetlen.github.io/llama-cpp-python/whl/cu121'
# ], capture_output=True, text=True)
#检查命令执行结果
if result.returncode == 0:
print("#install success")
from llama_cpp import Llama
# #检查命令执行结果
# if result.returncode == 0:
# print("#install success")
# from llama_cpp import Llama
subprocess.run([sys.executable, '-s', '-m', 'pip',
'install',
'llama-cpp-python[server]'
], capture_output=True, text=True)
# subprocess.run([sys.executable, '-s', '-m', 'pip',
# 'install',
# 'llama-cpp-python[server]'
# ], capture_output=True, text=True)
else:
print("#install error")
# else:
# print("#install error")
else:
from llama_cpp import Llama
except:
print("#install llama-cpp-python error")
# else:
# from llama_cpp import Llama
# except:
# print("#install llama-cpp-python error")
if file_name:
mp=get_llama_model_path(file_name)
# file_name=get_llama_models()[0]
# model_path=os.path.join(folder_paths.models_dir, "llamafile")
# mp=os.path.join(model_path,file_name)
# if file_name:
# mp=get_llama_model_path(file_name)
# # file_name=get_llama_models()[0]
# # model_path=os.path.join(folder_paths.models_dir, "llamafile")
# # mp=os.path.join(model_path,file_name)
llm = Llama(model_path=mp, chat_format="chatml",n_gpu_layers=-1,n_ctx=512)
# llm = Llama(model_path=mp, chat_format="chatml",n_gpu_layers=-1,n_ctx=512)
return llm
# return llm
def chat(client, model_name,messages ):
if is_installed('json_repair'):
from json_repair import repair_json
def chat(client, model_name,messages,max_tokens=4096,temperature=0.6 ):
print('#chat',model_name,messages)
try_count = 0
while True:
try_count += 1
@@ -171,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的
@@ -206,6 +262,36 @@ def chat(client, model_name,messages ):
return content
llm_apis=[
{
"value": "https://api.openai.com/v1",
"label": "openai"
},
{
"value": "https://openai.api2d.net/v1",
"label": "api2d"
},
# {
# "value": "https://docs-test-001.openai.azure.com",
# "label": "https://docs-test-001.openai.azure.com"
# },
{
"value": "https://api.moonshot.cn/v1",
"label": "Kimi"
},
{
"value": "https://api.deepseek.com/v1",
"label": "DeepSeek-V2"
},
{
"value": "https://api.siliconflow.cn/v1",
"label": "SiliconCloud"
}]
llm_apis_dict = {api["label"]: api["value"] for api in llm_apis}
class ChatGPTNode:
def __init__(self):
# self.__client = OpenAI()
@@ -215,35 +301,60 @@ class ChatGPTNode:
@classmethod
def INPUT_TYPES(cls):
model_list=llama_modes_list+[
"gpt-3.5-turbo",
"gpt-3.5-turbo-0125",
"gpt-35-turbo",
"gpt-3.5-turbo-16k",
"gpt-3.5-turbo-16k-0613",
"gpt-4-0613",
"gpt-4-1106-preview",
"glm-4"
model_list=[
"gpt-3.5-turbo",
"gpt-3.5-turbo-16k",
"gpt-4o",
"gpt-4o-2024-05-13",
"gpt-4",
"gpt-4-0314",
"gpt-4-0613",
"gpt-3.5-turbo-0301",
"gpt-3.5-turbo-0613",
"gpt-3.5-turbo-16k-0613",
"qwen-turbo",
"qwen-plus",
"qwen-long",
"qwen-max",
"qwen-max-longcontext",
"glm-4",
"glm-3-turbo",
"moonshot-v1-8k",
"moonshot-v1-32k",
"moonshot-v1-128k",
"deepseek-chat",
"Qwen/Qwen2-7B-Instruct",
"THUDM/glm-4-9b-chat",
"01-ai/Yi-1.5-9B-Chat-16K",
"meta-llama/Meta-Llama-3.1-8B-Instruct"
]
return {
"required": {
"api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
"api_url":("URL", {"default": "", "multiline": True,"dynamicPrompts": False}),
# "api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
# "api_key":("STRING", {"forceInput": True,}),
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"system_content": ("STRING",
{
"default": "You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"multiline": True,"dynamicPrompts": False
}),
"model": ( model_list,
{"default": model_list[0]}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
"api_url":(list(llm_apis_dict.keys()),
{"default": list(llm_apis_dict.keys())[0]}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
"extra_pnginfo": "EXTRA_PNGINFO",
},
"optional":{
"api_key":("STRING", {"forceInput": True,}),
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
"custom_api_url":("STRING", {"forceInput": True,}), #适合自定义model
},
}
RETURN_TYPES = ("STRING","STRING","STRING",)
@@ -255,12 +366,29 @@ class ChatGPTNode:
def generate_contextual_text(self,
api_key,
api_url,
# api_key,
prompt,
system_content,
model,
seed,context_size,unique_id = None, extra_pnginfo=None):
model,
seed,
context_size,
api_url,
api_key=None,
custom_model_name=None,
custom_api_url=None,
):
if custom_model_name!=None:
model=custom_model_name
api_url=llm_apis_dict[api_url] if api_url in llm_apis_dict else ""
if custom_api_url!=None:
api_url=custom_api_url
if api_key==None:
api_key="lm_studio"
# print(api_key!='',api_url,prompt,system_content,model,seed)
# 可以选择保留会话历史以维持上下文记忆
# 或者在此处清除会话历史 self.session_history.clear()
@@ -273,7 +401,7 @@ class ChatGPTNode:
self.system_content=system_content
# self.session_history=[]
# self.session_history.append({"role": "system", "content": system_content})
print("api_key,api_url",api_key,api_url)
#
if is_azure_url(api_url):
client=azure_client(api_key,api_url)
@@ -282,12 +410,12 @@ class ChatGPTNode:
if model == "glm-4" :
client = ZhipuAI_client(api_key) # 使用 Zhipuai 的接口
print('using Zhipuai interface')
elif model in llama_modes_list:
#
client=llama_cpp_client(model)
# elif model in llama_modes_list:
# #
# client=llama_cpp_client(model)
else :
client = openai_client(api_key,api_url) # 使用 ChatGPT 的接口
print('using ChatGPT interface')
# print('using ChatGPT interface',api_key,api_url)
# 把用户的提示添加到会话历史中
# 调用API时传递整个会话历史
@@ -303,6 +431,7 @@ class ChatGPTNode:
session_history=crop_list_tail(self.session_history,context_size)
messages=[{"role": "system", "content": self.system_content}]+session_history+[{"role": "user", "content": prompt}]
response_content = chat(client,model,messages)
self.session_history=self.session_history+[{"role": "user", "content": prompt}]+[{'role':'assistant',"content":response_content}]
@@ -323,6 +452,174 @@ class ChatGPTNode:
return (response_content,json.dumps(messages, indent=4),json.dumps(self.session_history, indent=4),)
class SiliconflowFreeNode:
def __init__(self):
# self.__client = OpenAI()
self.session_history = [] # 用于存储会话历史的列表
# self.seed=0
self.system_content="You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible."
@classmethod
def INPUT_TYPES(cls):
model_list= [
"Qwen/Qwen2.5-7B-Instruct",
"Qwen/Qwen2-7B-Instruct",
"THUDM/glm-4-9b-chat",
"01-ai/Yi-1.5-9B-Chat-16K",
"meta-llama/Meta-Llama-3.1-8B-Instruct"
]
return {
"required": {
"api_key":("STRING", {"forceInput": True,}),
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"system_content": ("STRING",
{
"default": "You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"multiline": True,"dynamicPrompts": False
}),
"model": ( model_list,
{"default": model_list[0]}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
"max_tokens":("INT", {"default": 512, "min": 512, "max":200000, "step": 1}),
},
"optional":{
"custom_model_name":("STRING", {"forceInput": True,}), #适合自定义model
},
}
RETURN_TYPES = ("STRING","STRING","STRING",)
RETURN_NAMES = ("text","messages","session_history",)
FUNCTION = "generate_contextual_text"
CATEGORY = "♾️Mixlab/GPT"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,False,False,)
def generate_contextual_text(self,
api_key,
prompt,
system_content,
model,
seed,
context_size,
max_tokens,
custom_model_name=None):
if custom_model_name!=None:
model=custom_model_name
api_url="https://api.siliconflow.cn/v1"
# 把系统信息和初始信息添加到会话历史中
if system_content:
self.system_content=system_content
# self.session_history=[]
# self.session_history.append({"role": "system", "content": system_content})
#
client = openai_client(api_key,api_url) # 使用 ChatGPT 的接口
# print('using ChatGPT interface',api_key,api_url)
# 把用户的提示添加到会话历史中
# 调用API时传递整个会话历史
def crop_list_tail(lst, size):
if size >= len(lst):
return lst
elif size==0:
return []
else:
return lst[-size:]
session_history=crop_list_tail(self.session_history,context_size)
messages=[{"role": "system", "content": self.system_content}]+session_history+[{"role": "user", "content": prompt}]
response_content = chat(client,model,messages,max_tokens)
self.session_history=self.session_history+[{"role": "user", "content": prompt}]+[{'role':'assistant',"content":response_content}]
return (response_content,json.dumps(messages, indent=4),json.dumps(self.session_history, indent=4),)
class 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
@@ -484,3 +781,307 @@ class TextSplitByDelimiter:
arr= arr[start_index:start_index + max_count * (skip_every+1):(skip_every+1)]
return (arr,)
class JsonRepair:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"json_string":("STRING", {"forceInput": True,}),
"key":("STRING", {"multiline": False,"dynamicPrompts": False,"default": ""}),
},
"optional":{
"json_string2":("STRING", {"forceInput": True,})
},
}
INPUT_IS_LIST = False
RETURN_TYPES = ("STRING","STRING",)
RETURN_NAMES = ("json_string","value",)
FUNCTION = "run"
# OUTPUT_NODE = True
OUTPUT_IS_LIST = (False,False,)
CATEGORY = "♾️Mixlab/GPT"
def run(self, json_string,key="",json_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)
# 将 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]
# 将 Python 对象转换回 JSON 字符串,确保中文字符不被转义
json_str_with_chinese = json.dumps(data, ensure_ascii=False)
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),)
+1 -1
View File
@@ -79,7 +79,7 @@ def get_clip_interrogator_path():
cache_path=get_clip_interrogator_path()
caption_model_path=os.path.join(cache_path, "Salesforce/blip-image-captioning-base")
caption_model_path=os.path.join(cache_path, "Salesforce","blip-image-captioning-base")
if not os.path.exists(caption_model_path):
print(f"## clip_interrogator_model not found: {caption_model_path}, pls download from https://huggingface.co/Salesforce/blip-image-captioning-base")
caption_model_path='Salesforce/blip-image-captioning-base'
+332
View File
@@ -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)
+250
View File
@@ -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)
+565 -282
View File
File diff suppressed because it is too large Load Diff
-2
View File
@@ -85,8 +85,6 @@ class LaMaInpainting:
"image": ("IMAGE",),
"mask": ("MASK",),
},
}
RETURN_TYPES = ("IMAGE",)
+157
View File
@@ -0,0 +1,157 @@
# 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
import folder_paths
from transformers import AutoTokenizer, AutoModel
from torchvision.transforms.v2 import ToPILImage
# from decord import VideoReader, cpu # pip install decord
# from PIL import Image
def get_model_path(n=""):
try:
return folder_paths.get_folder_paths(n)[0]
except:
return os.path.join(folder_paths.models_dir, n)
class MiniCPM_VQA_Simple:
def __init__(self):
self.model_checkpoint = None
self.tokenizer = None
self.model = None
self.device = (
torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
)
self.bf16_support = (
torch.cuda.is_available()
and torch.cuda.get_device_capability(self.device)[0] >= 8
)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
"text": ("STRING", {"default": "", "multiline": True}),
"seed": ("INT", {"default": -1}), # add seed parameter, default is -1
"extract_keywords":("BOOLEAN", {"default": False}),
"temperature": (
"FLOAT",
{
"default": 0.7,
},
),
"keep_model_loaded": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("STRING","STRING",)
RETURN_NAMES = ("result","keywords",)
FUNCTION = "inference"
CATEGORY = "♾️Mixlab/Image"
def inference(
self,
images,
text,
seed, # add seed parameter, default is -1
extract_keywords,
temperature,
keep_model_loaded,
):
if seed != -1:
torch.manual_seed(seed)
model_id = "openbmb/MiniCPM-V-2_6-int4"
self.model_checkpoint = os.path.join( get_model_path("prompt_generator"), os.path.basename(model_id))
if not os.path.exists(self.model_checkpoint):
from huggingface_hub import snapshot_download
snapshot_download(
repo_id=model_id,
local_dir=self.model_checkpoint,
local_dir_use_symlinks=False,
endpoint='https://hf-mirror.com'
)
if self.tokenizer is None:
self.tokenizer = AutoTokenizer.from_pretrained(
self.model_checkpoint,
trust_remote_code=True,
low_cpu_mem_usage=True,
)
if self.model is None:
self.model = AutoModel.from_pretrained(
self.model_checkpoint,
trust_remote_code=True,
low_cpu_mem_usage=True,
attn_implementation="sdpa",
torch_dtype=torch.bfloat16 if self.bf16_support else torch.float16,
)
with torch.no_grad():
images = images.permute([0, 3, 1, 2])
images = [ToPILImage()(img).convert("RGB") for img in images]
msgs = [{"role": "user", "content": images + [text]}]
params = {"use_image_id": False, }
# offload model to CPU
# self.model = self.model.to(torch.device("cpu"))
# self.model.eval()
result = self.model.chat(
image=None,
msgs=msgs,
tokenizer=self.tokenizer,
sampling=True,
# top_k=top_k,
# top_p=top_p,
temperature=temperature,
# repetition_penalty=repetition_penalty,
# max_new_tokens=max_new_tokens,
**params,
)
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()
if not keep_model_loaded:
del self.tokenizer # release tokenizer memory
del self.model # release model memory
self.tokenizer = None # set tokenizer to None
self.model = None # set model to None
torch.cuda.empty_cache() # release GPU memory
torch.cuda.ipc_collect()
# print(result)
return (result,keyword_result,)
+104
View File
@@ -0,0 +1,104 @@
import torch
import numpy as np
from PIL import Image,ImageSequence,ImageOps
import base64
import io
import comfy.utils
import folder_paths
import node_helpers
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
# Convert PIL to Tensor
def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
def load_image_to_tensor( image):
image_path = folder_paths.get_annotated_filepath(image)
img = node_helpers.pillow(Image.open, image_path)
output_images = []
output_masks = []
w, h = None, None
excluded_formats = ['MPO']
for i in ImageSequence.Iterator(img):
i = node_helpers.pillow(ImageOps.exif_transpose, i)
if i.mode == 'I':
i = i.point(lambda i: i * (1 / 255))
image = i.convert("RGB")
if len(output_images) == 0:
w = image.size[0]
h = image.size[1]
if image.size[0] != w or image.size[1] != h:
continue
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
if 'A' in i.getbands():
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
mask = torch.zeros((64,64), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 1 and img.format not in excluded_formats:
output_image = torch.cat(output_images, dim=0)
output_mask = torch.cat(output_masks, dim=0)
else:
output_image = output_images[0]
output_mask = output_masks[0]
return (output_image, output_mask)
class P5Input:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"frames":("IMAGEBASE64",),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("frames",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
OUTPUT_NODE = True
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self, frames):
ims=[]
for im in frames['images']:
# print(im)
if 'type' in im and (not f"[{im['type']}]" in im['name']):
im['name']=im['name']+" "+f"[{im['type']}]"
output_image, output_mask = load_image_to_tensor(im['name'])
ims.append(output_image)
if len(ims)==0:
image1 = Image.new('RGB', (512, 512), color='black')
return (pil2tensor(image1),)
image1 = ims[0]
for image2 in ims[1:]:
if image1.shape[1:] != image2.shape[1:]:
image2 = comfy.utils.common_upscale(image2.movedim(-1, 1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1, -1)
image1 = torch.cat((image1, image2), dim=0)
# 用于节点提示:p5节点提示有多少帧
return {"ui": {"_info": [len(frames['images'])]}, "result": (image1,)}
+33 -7
View File
@@ -18,7 +18,13 @@ import json
# req = request.Request("http://127.0.0.1:8188/prompt", data=data)
# request.urlopen(req)
embeddings_path=os.path.join(folder_paths.models_dir, "embeddings")
def get_model_path(n=""):
try:
return folder_paths.get_folder_paths(n)[0]
except:
return os.path.join(folder_paths.models_dir, n)
embeddings_path=get_model_path("embeddings")
def get_files_with_extension(directory, extension):
@@ -181,7 +187,8 @@ class PromptImage:
}
}
RETURN_TYPES = ()
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json_str",)
OUTPUT_NODE = True
@@ -196,12 +203,19 @@ class PromptImage:
filename_prefix="mixlab_"
filename_prefix += self.prefix_append
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
filename_prefix, self.output_dir, images[0].shape[1], images[0].shape[0])
filename_prefix,self.output_dir, images[0].shape[1], images[0].shape[0])
full_output_folder=os.path.join(full_output_folder,'PromptImage')
subfolder='PromptImage'
results = list()
save_to_image=save_to_image[0]=='enable'
#保存到本地的json文件,记录图片和prompt的对应关系
output_images=[]
output_prompt=[]
for index in range(len(images)):
res=[]
imgs=images[index]
@@ -209,24 +223,36 @@ class PromptImage:
for image in imgs:
img=tensor2pil(image)
prompt_text=prompts[index]
metadata = None
if save_to_image:
metadata = PngInfo()
prompt_text=prompts[index]
if prompt_text is not None:
metadata.add_text("prompt_text", prompt_text)
file = f"{filename}_{index}_{counter:05}_.png"
img.save(os.path.join(full_output_folder, file), pnginfo=metadata, compress_level=self.compress_level)
fp=os.path.join(full_output_folder,file)
img.save(fp, pnginfo=metadata, compress_level=self.compress_level)
res.append({
"filename": file,
"subfolder": subfolder,
"type": self.type
})
output_images.append(fp)
output_prompt.append(prompt_text)
counter += 1
results.append(res)
return { "ui": { "_images": results,"prompts":prompts } }
# if save_to_image:
# # 保存为本地文件
# with open(os.path.join(full_output_folder,'PromptImage.json'), 'w') as file:
# json.dump(output_dict, file, ensure_ascii=False, indent=4)
return { "ui": { "_images": results,"prompts":prompts },"result":(json.dumps({
"images":output_images,
"prompts":output_prompt
}),) }
+36 -22
View File
@@ -467,15 +467,37 @@ class BriaRMBG(nn.Module):
def get_U2NET_model_path():
try:
return folder_paths.get_folder_paths('rembg')[0]
except:
return os.path.join(folder_paths.models_dir, "rembg")
U2NET_HOME=os.path.join(folder_paths.models_dir, "rembg")
U2NET_HOME=get_U2NET_model_path()
os.environ["U2NET_HOME"] = U2NET_HOME
global _available
_available=False
def get_rembg_models(path):
"""从目录中获取文件并提取文件名
Args:
path: 目录路径
Returns:
文件名列表
"""
filenames = []
for root, _, files in os.walk(path):
for filename in files:
# 过滤隐藏文件
if not filename.startswith('.'):
name, ext = os.path.splitext(os.path.basename(filename))
filenames.append(name)
return filenames
def is_installed(package):
try:
spec = importlib.util.find_spec(package)
@@ -509,8 +531,8 @@ except:
_available=False
def briarmbg_run(images=[]):
mroot=os.path.join(folder_paths.models_dir, "rembg")
def run_briarmbg(images=[]):
mroot=U2NET_HOME
m=os.path.join(mroot,'briarmbg.pth')
if os.path.exists(m)==False:
# 下载
@@ -573,14 +595,15 @@ def briarmbg_run(images=[]):
return (masks,rgba_images,rgb_images)
def run_bg(model_name= "unet",images=[]):
def run_rembg(model_name= "unet",images=[],callback=None):
# model_name = "unet" # "isnet-general-use"
# print('#run_rembg',model_name)
rembg_session = new_session(model_name)
masks=[]
rgba_images=[]
rgb_images=[]
# 进度条
pbar = comfy.utils.ProgressBar(len(images) )
pbar=callback
for img in images:
# use the post_process_mask argument to post process the mask to get better results.
mask = remove(img, session=rembg_session,only_mask=True,post_process_mask=True)
@@ -620,8 +643,9 @@ def run_bg(model_name= "unet",images=[]):
rgb_image = Image.new("RGB", image_rgba.size, (0, 0, 0))
rgb_image.paste(image_rgba, mask=image_rgba.split()[3])
rgb_images.append(rgb_image)
pbar.update(1)
if pbar:
pbar.update(1)
return (masks,rgba_images,rgb_images)
@@ -643,17 +667,7 @@ class RembgNode_:
def INPUT_TYPES(s):
return {"required": {
"image": ("IMAGE",),
"model_name": ([
"briarmbg",
"u2net",
"u2netp",
"u2net_human_seg",
"u2net_cloth_seg",
"silueta",
"isnet-general-use",
"isnet-anime",
],),
"model_name": (get_rembg_models(U2NET_HOME),),
},
}
@@ -681,9 +695,9 @@ class RembgNode_:
images.append(im)
if model_name=='briarmbg':
masks,rgba_images,rgb_images=briarmbg_run(images)
masks,rgba_images,rgb_images=run_briarmbg(images)
else:
masks,rgba_images,rgb_images=run_bg(model_name,images)
masks,rgba_images,rgb_images=run_rembg(model_name,images, comfy.utils.ProgressBar(len(images) ))
masks=[pil2tensor(m) for m in masks]
+6 -6
View File
@@ -90,7 +90,7 @@ class ScreenShareNode:
} }
RETURN_TYPES = ('IMAGE','STRING','FLOAT',"INT")
RETURN_NAMES = ("IMAGE","PROMPT","FLOAT","INT")
RETURN_NAMES = ("current frame (image)","prompt","denoise (float)","seed (int)")
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Screen"
@@ -109,7 +109,7 @@ class FloatingVideo:
@classmethod
def INPUT_TYPES(s):
return { "required":{
"images": ("IMAGE",)
"image": ("IMAGE",)
}, }
# RETURN_TYPES = ('IMAGE','MASK')
@@ -124,16 +124,16 @@ class FloatingVideo:
# OUTPUT_IS_LIST = (False,False,)
# 运行的函数
def run(self,images):
def run(self,image):
results = list()
for image in images:
image=tensor2pil(image)
for im in image:
im=tensor2pil(im)
# image_base64 = base64.b64encode(image.tobytes())
buffered = BytesIO()
image.save(buffered, format="JPEG")
im.save(buffered, format="JPEG")
image_base64 = base64.b64encode(buffered.getvalue()).decode("utf-8")
results.append(image_base64)
+226
View File
@@ -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
View File
@@ -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}),
}
+10 -8
View File
@@ -280,7 +280,7 @@ class ChinesePrompt:
},
"optional":{
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
"seed":("INT", {"default": 100, "min": 100, "max": 0xffffffffffffffff}),
},
@@ -331,13 +331,15 @@ class ChinesePrompt:
for t in texts:
if t:
# translated_text = translated_word = translate(zh_en_tokenizer,zh_en_model,str(t))
parser = Lark(grammar, start="start", parser="lalr", transformer=ChinesePromptTranslate())
# print('t',t)
result = parser.parse(t).children
# print('en_result',result)
# en_text=translate(zh_en_tokenizer,zh_en_model,text_without_syntax)
en_texts.append(result[0])
try:
result = parser.parse(t).children
en_texts.append(result[0])
except:
print(f"Error parsing '{t}'")
t = translate(str(t))
en_texts.append(t)
zh_en_model.to('cpu')
print("test en_text",en_texts)
@@ -384,7 +386,7 @@ class PromptGenerate:
"optional":{
"multiple": (["off","on"],),
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
"seed":("INT", {"default": 100, "min": 100, "max": 0xffffffffffffffff}),
},
}
+9 -2
View File
@@ -5,13 +5,20 @@ from PIL import Image
import numpy as np
import torch
from folder_paths import get_filename_list, get_full_path, get_save_image_path, get_output_directory,models_dir
from folder_paths import get_folder_paths, get_full_path, get_save_image_path, get_output_directory,models_dir
from comfy.model_management import get_torch_device
from .tsr.system import TSR
import comfy.utils
triposr_model_path=path.join(models_dir,'triposr/model.ckpt')
def get_triposr_model_path():
try:
return path.join(get_folder_paths('triposr')[0],'model.ckpt')
except:
return path.join(path.join(models_dir, "triposr"),'model.ckpt')
triposr_model_path=get_triposr_model_path()
# Tensor to PIL
+90 -16
View File
@@ -82,13 +82,13 @@ def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
def create_temp_file(image):
def create_temp_file(image,counter=1):
output_dir = folder_paths.get_temp_directory()
(
full_output_folder,
filename,
counter,
_,
subfolder,
_,
) = folder_paths.get_save_image_path('tmp', output_dir)
@@ -133,7 +133,7 @@ def get_font_files(directory):
return font_files
r_directory = os.path.join(os.path.dirname(__file__), '../assets/')
r_directory = os.path.join(os.path.dirname(__file__), '..','assets','/')
font_files = get_font_files(r_directory)
# print(font_files)
@@ -181,6 +181,28 @@ class ColorInput:
return (h,r,g,b,a,)
class KeyInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"key":("KEY",),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("key",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self,key):
return (key,)
class FontInput:
@classmethod
@@ -284,7 +306,7 @@ class FloatSlider:
}
RETURN_TYPES = ("FLOAT",)
RETURN_NAMES = ('weight(0-1)',)
RETURN_NAMES = ('FLOAT',)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
@@ -297,9 +319,7 @@ class FloatSlider:
number = min_value
elif number > max_value:
number = max_value
scaled_number = (number - min_value) / (max_value - min_value)
return (scaled_number,)
return (number,)
class IntNumber:
@classmethod
@@ -568,7 +588,7 @@ class AppInfo:
},
"optional":{
"IMAGE": ("IMAGE",),
"image": ("IMAGE",),
"description":("STRING",{"multiline": True,"default": "","dynamicPrompts": False}),
"version":("INT", {
"default": 1,
@@ -581,6 +601,7 @@ class AppInfo:
"link":("STRING",{"multiline": False,"default": "https://","dynamicPrompts": False}),
"category":("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
"auto_save": (["enable","disable"],),
"idle_animation": ("BOOLEAN", {"default": False},),
}
}
@@ -596,14 +617,21 @@ class AppInfo:
INPUT_IS_LIST = True
# OUTPUT_IS_LIST = (True,)
def run(self,name,input_ids,output_ids,IMAGE,description,version,share_prefix,link,category,auto_save):
def run(self,name,input_ids,output_ids,image,description,version,share_prefix,link,category,auto_save,idle_animation):
name=name[0]
idle_animation=idle_animation[0]
im=None
if IMAGE:
im=IMAGE[0][0]
#TODO batch 的方式需要处理
im=create_temp_file(im)
im=[]
if image:
images=[image]
# batch 的方式需要处理
images=flatten_list(images)
# img=image[0][0]
print('AppInfo_image',len(images))
for i in range(len(images)):
img=images[i]
im.append(create_temp_file(img,i+1)[0])
# image [img,] img[batch,w,h,a] 列表里面是batch,
input_ids=input_ids[0]
@@ -616,7 +644,50 @@ class AppInfo:
# id=get_json_hash([name,im,input_ids,output_ids,description,version])
return {"ui": {"json": [name,im,input_ids,output_ids,description,version,share_prefix,link,category]}, "result": ()}
return {"ui": {"json": [name,im,input_ids,output_ids,description,version,share_prefix,link,category,idle_animation]}, "result": ()}
class CreateJsonNode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"key": ("STRING",{"multiline": False,"default": "data","dynamicPrompts": False}),
"value":(any_type,),
"save":("BOOLEAN", {"default": True},),
},
"optional":{
"json_str":("STRING", {"forceInput": True,}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json_str",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Output"
OUTPUT_NODE = True
INPUT_IS_LIST = False
# OUTPUT_IS_LIST = (True,)
def run(self,key,value,save,json_str=None):
data={}
data[key]=value
if json_str:
json_obj = json.loads(json_str)
data.update(json_obj)
if save:
# 保存为本地文件
with open(os.path.join(folder_paths.get_output_directory(),'data.json'), 'w') as file:
json.dump(data, file, ensure_ascii=False, indent=4)
return (json.dumps(data),)
@@ -797,7 +868,7 @@ class TESTNODE_:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"ANY":(any_type,),
"ANY":(any_type,),
},
}
@@ -812,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)
+389 -70
View File
@@ -17,9 +17,136 @@ import folder_paths
from comfy.k_diffusion.utils import FolderOfImages
from comfy.utils import common_upscale
import torchaudio
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:
if frame_count <= len(frames):
return frames[:frame_count]
else:
return [frames[i % len(frames)] for i in range(frame_count)]
else:
extended_frames = frames + frames[-2:0:-1] # 正向加反向中间部分
if frame_count <= len(extended_frames):
return extended_frames[:frame_count]
else:
return [extended_frames[i % len(extended_frames)] for i in range(frame_count)]
# # 示例用法
# frames = ["frame1", "frame2", "frame3"]
# frame_count = 2
# result = get_frames(frame_count, frames, revert=False)
# print(result) # 输出: ['frame1', 'frame2', 'frame3', 'frame1', 'frame2', 'frame3', 'frame1']
# result = get_frames(frame_count, frames, revert=True)
# print(result) # 输出: ['frame1', 'frame2', 'frame3', 'frame2', 'frame1', 'frame2', 'frame3']
def get_mime_type(file_path):
# 获取文件的 MIME 类型
mime_type, _ = mimetypes.guess_type(file_path)
# 如果无法猜测类型,返回默认类型
if mime_type is None:
return 'application/octet-stream'
return mime_type
# import subprocess
# from imageio_ffmpeg import get_ffmpeg_exe
def save_audio_base64s_to_file(base64_audios, output_folder, file_name):
# Ensure the output folder exists
if not os.path.exists(output_folder):
os.makedirs(output_folder)
decoded_audios=[]
for a in base64_audios:
# If the base64 string contains a header, remove it
if ',' in a:
a = a.split(',')[1]
# 解码 base64 数据
a=base64.b64decode(a)
decoded_audios.append(a)
# 拼接音频数据
combined_audio = b''.join(decoded_audios)
# Create the full file path
file_path = os.path.join(output_folder, file_name)
# Write the decoded audio to the file
with open(file_path, 'wb') as audio_file:
audio_file.write(combined_audio)
return file_path
# Example usage
# base64_audio = "data:audio/wav;base64,UklGRiQAAABXQVZFZm10IBAAAAABAAEAIlYAAESsAAACABAAZGF0YQAAAAA="
# output_folder = "audio_files"
# file_name = "output.wav"
# file_path = save_audio_base64_to_file(base64_audio, output_folder, file_name)
# print(f"Audio saved to: {file_path}")
# 写一个python文件,用来 判断文件夹内命名为 所有chat_tts开头的文件数量(chat_tts_00001),并输出新的编号
def get_new_counter(full_output_folder, filename_prefix):
# 获取目录中的所有文件
files = os.listdir(full_output_folder)
# 过滤出以 filename_prefix 开头并且后续部分为数字的文件
filtered_files = []
for f in files:
if f.startswith(filename_prefix):
# 去掉文件名中的前缀和后缀,只保留中间的数字部分
base_name = f[len(filename_prefix)+1:]
number_part = base_name.split('.')[0] # 假设文件名中只有一个点,即扩展名
if number_part.isdigit():
filtered_files.append(int(number_part))
if not filtered_files:
return 1
# 获取最大的编号
max_number = max(filtered_files)
# 新的编号
return max_number + 1
def crop_audio(input_file, start_time, duration):
# Load the audio file
audio_tensor, sample_rate = torchaudio.load(input_file)
# Convert start_time and duration from seconds to sample indices
start_sample = int(start_time * sample_rate)
end_sample = start_sample + int(duration * sample_rate)
# Perform the slicing
cropped_audio_tensor = audio_tensor[:, start_sample:end_sample]
# Save the cropped audio to a new file
torchaudio.save(input_file, cropped_audio_tensor, sample_rate)
return input_file
def generate_folder_name(directory,video_path):
# Get the directory and filename from the video path
_, filename = os.path.split(video_path)
@@ -60,6 +187,9 @@ def split_video(video_path, video_segment_frames, transition_frames, output_dir)
# 打印当前片段的起始帧和结束帧
print(f"Segment {i+1}: Start Frame {start_frame}, End Frame {end_frame}")
if end_frame<start_frame:
break
# 保存当前片段为一个视频文件
segment_video_path = f"{output_dir}/segment_{i+1}.avi"
@@ -68,6 +198,7 @@ def split_video(video_path, video_segment_frames, transition_frames, output_dir)
segment_video = cv2.VideoWriter(segment_video_path, fourcc, fps, (int(video_capture.get(cv2.CAP_PROP_FRAME_WIDTH)),
int(video_capture.get(cv2.CAP_PROP_FRAME_HEIGHT))))
for frame_num in range(start_frame, end_frame):
ret, frame = video_capture.read()
if ret:
@@ -87,7 +218,7 @@ def split_video(video_path, video_segment_frames, transition_frames, output_dir)
folder_paths.folder_names_and_paths["video_formats"] = (
[
os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "video_formats"),
os.path.join(os.path.dirname(os.path.abspath(__file__)), ".", "video_formats"),
],
[".json"]
)
@@ -101,6 +232,25 @@ if ffmpeg_path is None:
except:
print("ffmpeg could not be found. Outputs that require it have been disabled")
def combine_audio_video(audio_path, video_path, output_path):
command = [
ffmpeg_path,
'-i', video_path,
'-i', audio_path,
'-c:v', 'copy',
'-c:a', 'aac',
'-shortest',
output_path
]
subprocess.run(command, check=True)
return output_path
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
@@ -262,7 +412,7 @@ class LoadVideoAndSegment:
files.append(f)
return {"required": {
"video": (sorted(files), {"video_upload": True}),
"video_segment_frames": ("INT", {"default": 10, "min": 1, "step": 1}),
"video_segment_frames": ("INT", {"default": 10, "min": -1, "step": 1}),
"transition_frames": ("INT", {"default": 0, "min": 0, "step": 1}),
},}
@@ -332,63 +482,6 @@ class LoadVideoAndSegment:
video_path = folder_paths.get_annotated_filepath(video)
# check if video is a gif - will need to use cv fallback to read frames
# use cv fallback if ffmpeg not installed or gif
# if ffmpeg_path is None:
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# otherwise, continue with ffmpeg
# args_dummy = [ffmpeg_path, "-i", video_path, "-f", "null", "-"]
# try:
# with subprocess.Popen(args_dummy, stdout=subprocess.DEVNULL, stderr=subprocess.PIPE) as proc:
# for line in proc.stderr.readlines():
# match = re.search(", ([1-9]|\\d{2,})x(\\d+)",line.decode('utf-8'))
# if match is not None:
# size = [int(match.group(1)), int(match.group(2))]
# break
# except Exception as e:
# print(f"Retrying with opencv due to ffmpeg error: {e}")
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# args_all_frames = [ffmpeg_path, "-i", video_path, "-v", "error",
# "-pix_fmt", "rgb24"]
# vfilters = []
# if skip_first_frames > 0:
# vfilters.append(f"select=gt(n\\,{skip_first_frames-1})")
# if frame_load_cap > 0:
# vfilters.append(f"select=gt({frame_load_cap}\\,n)")
# #manually calculate aspect ratio to ensure reads remain aligned
# if len(vfilters) > 0:
# args_all_frames += ["-vf", ",".join(vfilters)]
# args_all_frames += ["-f", "rawvideo", "-"]
# images = []
# try:
# with subprocess.Popen(args_all_frames, stdout=subprocess.PIPE) as proc:
# #Manually buffer enough bytes for an image
# bpi = size[0]*size[1]*3
# current_bytes = bytearray(bpi)
# current_offset=0
# while True:
# bytes_read = proc.stdout.read(bpi - current_offset)
# if bytes_read is None:#sleep to wait for more data
# time.sleep(.2)
# continue
# if len(bytes_read) == 0:#EOF
# break
# current_bytes[current_offset:len(bytes_read)] = bytes_read
# current_offset+=len(bytes_read)
# if current_offset == bpi:
# images.append(np.array(current_bytes, dtype=np.float32).reshape(size[1], size[0], 3) / 255.0)
# current_offset = 0
# except Exception as e:
# print(f"Retrying with opencv due to ffmpeg error: {e}")
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# imgs=split_list(images,video_segment_frames,transition_frames)
# temp path
tp=folder_paths.get_temp_directory()
basename = os.path.basename(video_path) # 获取文件名
@@ -396,15 +489,22 @@ class LoadVideoAndSegment:
folder_path = create_folder(tp,name_without_extension)
# 导出的数据
scenes_video,total_frames,fps=split_video(video_path,video_segment_frames,
transition_frames,folder_path)
if video_segment_frames==-1:
# 不切割视频
scenes_video=[video_path]
# 读取视频文件
video_capture = cv2.VideoCapture(video_path)
# 获取视频的总帧数和帧率
total_frames = int(video_capture.get(cv2.CAP_PROP_FRAME_COUNT))
fps = video_capture.get(cv2.CAP_PROP_FPS)
else:
# 导出的数据
scenes_video,total_frames,fps=split_video(video_path,video_segment_frames,
transition_frames,folder_path)
# imgs=[torch.from_numpy(np.stack(im)) for im in imgs]
# images = torch.from_numpy(np.stack(images))
return (scenes_video,len(scenes_video), total_frames,fps,)
@@ -422,7 +522,113 @@ class LoadVideoAndSegment:
return "Invalid image file: {}".format(video)
return True
class LoadAndCombinedAudio_:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"audios": ("AUDIOBASE64",),
"start_time": ("FLOAT" , {"default": 0, "min": 0, "max": 10000000, "step": 0.01}),
"duration": ("FLOAT" , {"default": 10, "min": -1, "max": 10000000, "step": 0.01}),
},
}
CATEGORY = "♾️Mixlab/Audio"
RETURN_TYPES = ("STRING","AUDIO",)
RETURN_NAMES = ("audio_file_path","audio",)
FUNCTION = "run"
def run(self,audios, start_time, duration):
output_dir = folder_paths.get_output_directory()
counter=get_new_counter(output_dir,'audio_')
audio_file_name = f"audio_{counter:05}.wav"
audio_file=save_audio_base64s_to_file(audios['base64'],output_dir,audio_file_name)
# duration == -1 则不裁切
if duration > -1:
crop_audio(audio_file, start_time, duration)
waveform, sample_rate = torchaudio.load(audio_file)
audio = {
"filename": audio_file_name,
"subfolder": "",
"type": "output",
"audio_path":audio_file,
"waveform": waveform.unsqueeze(0),
"sample_rate": sample_rate}
return (audio_file,audio ,)
class CombineAudioVideo:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"video": ("SCENE_VIDEO",),
"audio": ("AUDIO", ),
},
}
CATEGORY = "♾️Mixlab/Video"
OUTPUT_NODE = True
FUNCTION = "run"
RETURN_TYPES = ("SCENE_VIDEO",)
RETURN_NAMES = ("SCENE_VIDEO",)
def run(self,video, audio):
output_dir = folder_paths.get_output_directory()
# 判断是否是 Tensor 类型
is_tensor = not isinstance(audio, dict)
# print('#判断是否是 Tensor 类型',is_tensor,audio)
if not is_tensor and 'waveform' in audio and 'sample_rate' in audio:
# {'waveform': tensor([], size=(1, 1, 0)), 'sample_rate': 44100}
is_tensor=True
if "audio_path" in audio:
is_tensor=False
audio_file_path=audio["audio_path"]
if is_tensor:
filename_prefix="audio_tmp"
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
filename_prefix,
folder_paths.get_temp_directory())
filename_with_batch_num = filename.replace("%batch_num%", str(1))
file = f"{filename_with_batch_num}_{counter:05}_.wav"
audio_file_path=os.path.join(full_output_folder, file)
torchaudio.save(audio_file_path, audio['waveform'].squeeze(0), audio["sample_rate"])
# 获取文件名和扩展名
base, ext = os.path.splitext(video)
counter=get_new_counter(output_dir,'video_final_')
v_file = f"video_final_{counter:05}{ext}"
v_file_path=os.path.join(output_dir, v_file)
combine_audio_video(audio_file_path,video,v_file_path)
previews = [
{
"filename": v_file,
"subfolder": "",
"type": "output",
"format": get_mime_type(v_file),
}
]
return {"ui": {"gifs": previews},"result":(v_file_path,)}
# The code is based on ComfyUI-VideoHelperSuite modification.
class VideoCombine_Adv:
@@ -433,6 +639,7 @@ class VideoCombine_Adv:
ffmpeg_formats = ["video/"+x[:-5] for x in folder_paths.get_filename_list("video_formats")]
else:
ffmpeg_formats = []
# ffmpeg_formats =["video/"+x for x in ['webm', 'mp4', 'mkv']]
return {
"required": {
"image_batch": ("IMAGE",),
@@ -453,7 +660,8 @@ class VideoCombine_Adv:
},
}
RETURN_TYPES = ()
RETURN_TYPES = ("SCENE_VIDEO",)
RETURN_NAMES = ("scenes_video",)
OUTPUT_NODE = True
CATEGORY = "♾️Mixlab/Video"
FUNCTION = "run"
@@ -622,7 +830,7 @@ class VideoCombine_Adv:
"format": format,
}
]
return {"ui": {"gifs": previews}}
return {"ui": {"gifs": previews},"result":(file_path,)}
class VAEEncodeForInpaint_Frames:
@@ -689,4 +897,115 @@ class VAEEncodeForInpaint_Frames:
result.append({"samples":t, "noise_mask": (mask_erosion[:,:,:x,:y].round())})
return (result, )
return (result, )
class GenerateFramesByCount:
@classmethod
def INPUT_TYPES(cls):
return {"required": {
"frames": ('IMAGE',),
"frame_count": ("INT", {"default": 72, "min": 1, "step": 1}),
"revert" :("BOOLEAN", {"default": True},),
},}
RETURN_TYPES = ('IMAGE',)
RETURN_NAMES = ("frames",)
FUNCTION = "r"
CATEGORY = "♾️Mixlab/Video"
# INPUT_IS_LIST = True
def r(self, frames, frame_count, revert):
image_list = [frames[i:i + 1, ...] for i in range(frames.shape[0])]
print('#image_list',len(image_list),frame_count)
image_list=get_frames(frame_count,image_list,revert)
images = torch.cat(image_list, dim=0)
return (images,)
class scenesNode_:
@classmethod
def INPUT_TYPES(cls):
return {"required": {
"scenes_video": ('SCENE_VIDEO',),
"index": ("INT", {"default": 0, "min": 0, "step": 1}),
},}
RETURN_TYPES = ('IMAGE','INT',)
RETURN_NAMES = ("video frames (batch)","count",)
# OUTPUT_IS_LIST = (False,)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Video"
INPUT_IS_LIST = True
def load_video_cv_fallback(self, video, frame_load_cap, skip_first_frames):
images = []
total_frame_count = 0
video_cap = cv2.VideoCapture(video)
try:
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
frames_added = 0
base_frame_time = 1/video_cap.get(cv2.CAP_PROP_FPS)
target_frame_time = base_frame_time
time_offset=0.0
while video_cap.isOpened():
if time_offset < target_frame_time:
is_returned, frame = video_cap.read()
# if didn't return frame, video has ended
if not is_returned:
break
time_offset += base_frame_time
if time_offset < target_frame_time:
continue
time_offset -= target_frame_time
# if not at start_index, skip doing anything with frame
total_frame_count += 1
if total_frame_count <= skip_first_frames:
continue
# TODO: do whatever operations need to happen, like force_size, etc
# opencv loads images in BGR format (yuck), so need to convert to RGB for ComfyUI use
# follow up: can videos ever have an alpha channel?
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
# convert frame to comfyui's expected format (taken from comfy's load image code)
image = Image.fromarray(frame)
image = ImageOps.exif_transpose(image)
image = np.array(image, dtype=np.float32) / 255.0
image = torch.from_numpy(image)[None,]
images.append(image)
frames_added += 1
# if cap exists and we've reached it, stop processing frames
if frame_load_cap > 0 and frames_added >= frame_load_cap:
break
finally:
video_cap.release()
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:
vp=scenes_video[index]
else:
vp=scenes_video[-1]
return self.load_video_cv_fallback(vp,0,0)
+173
View File
@@ -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,)
+174
View File
@@ -0,0 +1,174 @@
import torch
from PIL import Image, ImageOps, ImageSequence, ImageFile
from PIL.PngImagePlugin import PngInfo
import numpy as np
import os
import folder_paths
import node_helpers
import hashlib
from uuid import uuid4
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
# tensor 取hash值
def tensor_to_hash(tensor):
# 将 Tensor 转换为 NumPy 数组
np_array = tensor.cpu().numpy()
# 将 NumPy 数组转换为字节数据
byte_data = np_array.tobytes()
# 计算哈希值
hash_value = hashlib.md5(byte_data).hexdigest()
return hash_value
def create_temp_file(image, uuid):
output_dir = folder_paths.get_temp_directory()
(
full_output_folder,
filename,
counter,
subfolder,
_,
) = folder_paths.get_save_image_path(f'material_{uuid}', output_dir)
image=tensor2pil(image)
image_file = f"{filename}_{counter:05}.png"
image_path=os.path.join(full_output_folder, image_file)
image.save(image_path,compress_level=4)
return (image_path,[{
"filename": image_file,
"subfolder": subfolder,
"type": "temp"
}])
# image - tensor - 文件路径
# loadImage的方法( 文件路径 - image-mask )
class EditMask:
def __init__(self):
self.image_id = None
self.uuid = str(uuid4())
@classmethod
def INPUT_TYPES(s):
return {"required":
{"image": ("IMAGE",), # 表示一个张量
},
"optional":{
"image_update": ("IMAGE_FILE",)
},
}
CATEGORY = "♾️Mixlab/Mask"
RETURN_TYPES = ("IMAGE", "MASK")
RETURN_NAMES = ("image", "mask")
FUNCTION = "edit"
OUTPUT_NODE = True
def edit(self, image,image_update=None):
# 根据image输入来判断是否是新的图片
if self.image_id==None:
self.image_id=tensor_to_hash(image)
image_update=None
else:
image_id=tensor_to_hash(image)
if image_id!=self.image_id:
image_update=None
self.image_id=image_id
image_path=None
# print('#image_update',self.image_id,image_update)
if image_update==None:
print('--')
else:
if 'images' in image_update:
images=image_update['images']
filename=images[0]['filename']
subfolder=images[0]['subfolder']
type=images[0]['type']
name, base_dir=folder_paths.annotated_filepath(filename)
if type.endswith("output"):
base_dir = folder_paths.get_output_directory()
elif type.endswith("input"):
base_dir = folder_paths.get_input_directory()
elif type.endswith("temp"):
base_dir = folder_paths.get_temp_directory()
#base_dir = folder_paths.get_input_directory()
# print(base_dir,subfolder, name)
image_path = os.path.join(base_dir,subfolder, name)
if image_path==None:
image_path,images=create_temp_file(image, 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, self.uuid)
img = node_helpers.pillow(Image.open, image_path)
output_images = []
output_masks = []
w, h = None, None
excluded_formats = ['MPO']
for i in ImageSequence.Iterator(img):
i = node_helpers.pillow(ImageOps.exif_transpose, i)
if i.mode == 'I':
i = i.point(lambda i: i * (1 / 255))
image = i.convert("RGB")
if len(output_images) == 0:
w = image.size[0]
h = image.size[1]
if image.size[0] != w or image.size[1] != h:
continue
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
if 'A' in i.getbands():
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
# 尺寸不对,需要按照image来
mask = torch.zeros((h, w), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 1 and img.format not in excluded_formats:
output_image = torch.cat(output_images, dim=0)
output_mask = torch.cat(output_masks, dim=0)
else:
output_image = output_images[0]
output_mask = output_masks[0]
return {"ui":{"images": images},"result": (output_image, output_mask)}
# return (output_image, output_mask)
+69
View File
@@ -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
+87
View File
@@ -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}
+2
View File
@@ -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
+496
View File
@@ -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
+147
View File
@@ -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
+104
View File
@@ -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
+94
View File
@@ -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()
+130
View File
@@ -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) == [
"这是一段很长的中文文本,",
"而且没有句号,也没有感叹号,",
"也没有问号,也没有换行符.",
]
+4
View File
@@ -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())
+31
View File
@@ -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
+130
View File
@@ -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()
+699
View File
@@ -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()
+497
View File
@@ -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)
+23
View File
@@ -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",
]
+217
View File
@@ -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)
+13
View File
@@ -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)
)
+16
View File
@@ -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]
+50
View File
@@ -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
+56
View File
@@ -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)
+48
View File
@@ -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)
+100
View File
@@ -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)
+122
View 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
+114
View File
@@ -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
+98
View File
@@ -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, )
+321
View File
@@ -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,)
+10
View File
@@ -0,0 +1,10 @@
{
"main_pass":
[
"-n", "-c:v", "libsvtav1",
"-pix_fmt", "yuv420p10le",
"-crf", "23"
],
"extension": "webm",
"environment": {"SVT_LOG": "1"}
}
+9
View File
@@ -0,0 +1,9 @@
{
"main_pass":
[
"-n", "-c:v", "libx264",
"-pix_fmt", "yuv420p",
"-crf", "19"
],
"extension": "mp4"
}
+11
View File
@@ -0,0 +1,11 @@
{
"main_pass":
[
"-n", "-c:v", "libx265",
"-pix_fmt", "yuv420p10le",
"-preset", "medium",
"-crf", "22",
"-x265-params", "log-level=quiet"
],
"extension": "mp4"
}

Some files were not shown because too many files have changed in this diff Show More