Compare commits

..
205 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
205 changed files with 147304 additions and 7863 deletions
+6 -2
View File
@@ -7,15 +7,19 @@ on:
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'shadowcz007' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+83 -19
View File
@@ -1,18 +1,51 @@
![](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)
##### `最新`:
- 增加 Edit Mask,方便在生成的时候手动绘制 mask [workflow](./workflow/edit-mask-workflow.json)
- 新增[fal.ai](https://fal.ai/dashboard)的视频生成:Kling、RunwayGen3、LumaDreamMachine,[工作流下载](./workflow/video-all-in-one-test-workflow.json)
- 新增 SimulateDevDesignDiscussions,需要安装[swarm](https://github.com/openai/swarm)和[Comfyui-ChatTTS](https://github.com/shadowcz007/Comfyui-ChatTTS),[工作流下载](./workflow/swarm制作的播客节点workflow.json)
- ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。模型下载后,放置到 `models/llamafile/`
- 新增 SenseVoice
- 右键菜单支持 text-to-text,方便对 prompt 词补全
- [新增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)
@@ -21,8 +54,7 @@
- 右键菜单支持 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) -->
#### `相关插件推荐`
@@ -47,7 +79,8 @@
- 发布为 app 的 workflow,可以在右键里再次编辑了
- web app 可以设置分类,在 comfyui 右键菜单可以编辑更新 web app
- 支持动态提示
- 支持把输出显示到comfyui背景(TouchDesigner 风格)
- 支持把输出显示到 comfyui 背景(TouchDesigner 风格)
- 如果转为 web app 打开是空白的,注意检查下插件目录的名字需要是:comfyui-mixlab-nodes(如果是 zip 包下载会多了个-main 的后缀,需要去掉)
![](./assets/微信图片_20240421205440.png)
@@ -106,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/`
@@ -142,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
@@ -171,7 +209,6 @@ pip install llama-cpp-python \
> The composite images node overlays a foreground image onto a background image at specified positions and scales, with optional blending modes and masking capabilities. position : 'overall',"center_center","left_bottom","center_bottom","right_bottom","left_top","center_top","right_top"
![layers](./assets/layers-workflow.svg)
![poster](./assets/poster-workflow.svg)
@@ -205,9 +242,16 @@ pip install llama-cpp-python \
#### TextImage
> [下载字体](https://drxie.github.io/OSFCC/)放到 ```custom_nodes/comfyui-mixlab-nodes/assets/fonts```
> [下载字体](https://drxie.github.io/OSFCC/)放到 `custom_nodes/comfyui-mixlab-nodes/assets/fonts`
#### MiniCPM-VQA Simple
This is the int4 quantized version of MiniCPM-V 2.6.
Running with int4 version would use lower GPU memory (about 7GB).
[模型](https://huggingface.co/openbmb/MiniCPM-V-2_6-int4)
![alt text](assets/1724308322276.png)
### Style
@@ -231,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)
@@ -246,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`
+630 -287
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

+8473 -567
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 (
+102 -2
View File
@@ -3,6 +3,101 @@ import os
import folder_paths
import torchaudio
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
def __ne__(self, __value: object) -> bool:
return False
any_type = AnyType("*")
def analyze_audio_data(audio_data):
total_duration = 0
total_gap_duration = 0
emotion_counts = {}
audio_types = set()
languages = set()
for i, entry in enumerate(audio_data):
# Calculate the duration of each audio segment
start_time = entry['start_time']
end_time = entry['end_time']
duration = end_time - start_time
total_duration += duration
# Count the emotions
if "emotion" in entry:
emotion = entry['emotion']
if emotion in emotion_counts:
emotion_counts[emotion] += 1
else:
emotion_counts[emotion] = 1
# Collect the audio types
if "audio_type" in entry:
audio_types.add(entry['audio_type'])
if "language" in entry:
languages.add(entry['language'])
# Calculate gap duration if not the last entry
if i < len(audio_data) - 1:
next_start_time = audio_data[i + 1]['start_time']
gap_duration = next_start_time - end_time
if gap_duration > 0:
total_gap_duration += gap_duration
# Get the most frequent emotion
if len(emotion_counts.keys())>0:
most_frequent_emotion = max(emotion_counts, key=emotion_counts.get)
else:
most_frequent_emotion=None
# Convert audio_types set to list for better readability
audio_types = list(audio_types)
languages=list(languages)
# Print the results
print(f"Total Effective Duration: {total_duration:.2f} seconds")
print(f"Total Gap Duration: {total_gap_duration:.2f} seconds")
print(f"Emotion Changes: {emotion_counts}")
print(f"Most Frequent Emotion: {most_frequent_emotion}")
print(f"Audio Types: {audio_types}")
return {
"total_duration": total_duration,
"total_gap_duration": total_gap_duration,
"emotion_changes": emotion_counts,
"most_frequent_emotion": most_frequent_emotion,
"audio_types": audio_types,
"languages":languages
}
# 分析音频数据
class AnalyzeAudioNone:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"json":(any_type,),},
}
RETURN_TYPES = (any_type,)
RETURN_NAMES = ("result",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Audio"
def run(self,json):
result=analyze_audio_data(json)
return (result,)
class SpeechRecognition:
@classmethod
def INPUT_TYPES(s):
@@ -90,7 +185,7 @@ class AudioPlayNode:
# {'waveform': tensor([], size=(1, 1, 0)), 'sample_rate': 44100}
is_tensor=True
if is_tensor:
if is_tensor and (not 'audio_path' in audio):
filename_prefix=""
# 保存
filename_prefix += self.prefix_append
@@ -108,7 +203,12 @@ class AudioPlayNode:
})
else:
results=[audio]
results=[{
"filename": audio['filename'],
"subfolder":audio['subfolder'],
"type": audio['type'],
"audio_path":audio['audio_path']
}]
# print(audio)
+680 -92
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,8 +301,9 @@ class ChatGPTNode:
@classmethod
def INPUT_TYPES(cls):
model_list=llama_modes_list+[
"gpt-3.5-turbo",
model_list=[
"gpt-3.5-turbo",
"gpt-3.5-turbo-16k",
"gpt-4o",
"gpt-4o-2024-05-13",
@@ -236,27 +323,38 @@ class ChatGPTNode:
"moonshot-v1-8k",
"moonshot-v1-32k",
"moonshot-v1-128k",
"deepseek-chat"
"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",)
@@ -268,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()
@@ -286,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)
@@ -295,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时传递整个会话历史
@@ -316,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}]
@@ -336,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
@@ -497,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),)
+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)
+375 -87
View File
@@ -8,17 +8,31 @@ from PIL.PngImagePlugin import PngInfo
import base64,os,random
from io import BytesIO
import folder_paths
import node_helpers
import json,io
import comfy.utils
from comfy.cli_args import args
import cv2
import string
import string,re
import math,glob
from .Watcher import FolderWatcher
from itertools import product
# 文件名排序
def sort_by_filename(items):
def extract_parts(filename):
# 使用正则表达式将文件名拆分为数字和非数字部分
parts = re.split(r'(\d+)', filename)
# 将数字部分转换为整数以便正确排序,同时保留非数字部分
parts = [int(part) if part.isdigit() else part for part in parts]
return parts
# 按照 file_name 的拆分部分进行排序
sorted_items = sorted(items, key=lambda x: extract_parts(x['file_name']))
return sorted_items
# 将PIL图片转换为OpenCV格式
def pil_to_opencv(image):
open_cv_image = cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
@@ -30,24 +44,26 @@ def opencv_to_pil(image):
return pil_image
# 列出目录下面的所有文件
def get_files_with_extension(directory, extension):
def get_files_with_extension(directory, extensions):
file_list = []
# 确保extensions参数是一个list,即使只有一个元素
if not isinstance(extensions, (tuple, list)):
extensions = [extensions]
for root, dirs, files in os.walk(directory):
# print(f"Files at {root}: {files}") # 确认files是一个字符串列表
for file in files:
if 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)
# 检查文件是否以任何一个提供的扩展名结尾
if any(file.endswith(ext) for ext in extensions):
# 直接将文件名添加到列表中
file_list.append(file)
return file_list
def composite_images(foreground, background, mask, is_multiply_blend=False, position="overall", scale=0.25):
width, height = foreground.size
bg_image = background
bwidth, bheight = bg_image.size
scale=max(scale,1/bwidth)
scale=max(scale,1/bheight)
scale = max(scale, 1 / bwidth)
scale = max(scale, 1 / bheight)
def determine_scale_option(width, height):
return 'height' if height > width else 'width'
@@ -66,9 +82,9 @@ def composite_images(foreground, background, mask, is_multiply_blend=False, posi
else:
scale_option = determine_scale_option(width, height)
if scale_option == 'height':
scale = int(bheight * scale) / height
scale = bheight * scale / height
else:
scale = int(bwidth * scale) / width
scale = bwidth * scale / width
new_width = int(width * scale)
new_height = int(height * scale)
@@ -106,22 +122,18 @@ def composite_images(foreground, background, mask, is_multiply_blend=False, posi
"mask": mask
}
layer_image = layer['image']
layer_mask = layer['mask']
# Resize the foreground image with antialiasing
try:
resampling_method = Image.Resampling.LANCZOS
except AttributeError:
resampling_method = Image.ANTIALIAS
bg_image = merge_images(bg_image,
layer_image,
layer_mask,
layer['x'],
layer['y'],
layer['width'],
layer['height'],
layer['scale_option'],
is_multiply_blend)
layer_image = layer['image'].resize((layer['width'], layer['height']), resampling_method)
layer_mask = layer['mask'].resize((layer['width'], layer['height']), resampling_method)
bg_image = bg_image.convert('RGB')
bg_image.paste(layer_image, (layer['x'], layer['y']), layer_mask)
return bg_image
return bg_image.convert('RGB')
@@ -488,6 +500,53 @@ def load_image(fp,white_bg=False):
return images
# 读取图片数据,转成tensor
def load_image_to_tensor( image):
image_path = folder_paths.get_annotated_filepath(image)
img = node_helpers.pillow(Image.open, image_path)
output_images = []
output_masks = []
w, h = None, None
excluded_formats = ['MPO']
for i in ImageSequence.Iterator(img):
i = node_helpers.pillow(ImageOps.exif_transpose, i)
if i.mode == 'I':
i = i.point(lambda i: i * (1 / 255))
image = i.convert("RGB")
if len(output_images) == 0:
w = image.size[0]
h = image.size[1]
if image.size[0] != w or image.size[1] != h:
continue
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
if 'A' in i.getbands():
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
mask = torch.zeros((64,64), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 1 and img.format not in excluded_formats:
output_image = torch.cat(output_images, dim=0)
output_mask = torch.cat(output_masks, dim=0)
else:
output_image = output_images[0]
output_mask = output_masks[0]
return (output_image, output_mask)
def load_image_and_mask_from_url(url, timeout=10):
# Load the image from the URL
response = requests.get(url, timeout=timeout)
@@ -728,8 +787,7 @@ def areaToMask(x,y,w,h,image):
# return bg_image
import cv2
import numpy as np
# ps的正片叠底
# 可以基于https://www.cnblogs.com/jsxyhelu/p/16947810.html ,用gpt写python代码
@@ -909,9 +967,68 @@ def resize_image(layer_image, scale_option, width, height,color="white"):
return layer_image
def generate_text_image(text, font_path, font_size, text_color, vertical=True, stroke=False, stroke_color=(0, 0, 0), stroke_width=1, spacing=0, padding=4):
# Split text into lines based on line breaks
lines = text.split("\n")
def generate_text_image(text,
font_path,
font_size,
text_color,
vertical=True,
stroke=False,
stroke_color=(0, 0, 0),
stroke_width=1,
spacing=0,
line_spacing=0,
padding=4,
max_characters_per_line=48,
fixed_width=None):
def split_text(text, max_chars, fixed_width=False):
lines = []
current_line = ""
current_length = 0
for char in text:
if char == '\n':
lines.append(current_line)
current_line = ""
current_length = 0
elif '\u4e00' <= char <= '\u9fff': # Chinese character
if current_length + 1 <= max_chars:
current_line += char
current_length += 1
else:
lines.append(current_line)
current_line = char
current_length = 1
else: # English character or other
if char == ' ':
space_length = 1
else:
space_length = 1
if current_length + space_length <= max_chars:
current_line += char
current_length += space_length
else:
lines.append(current_line)
current_line = char
current_length = space_length
if current_line:
lines.append(current_line)
# Pad lines to max_chars if fixed_width is provided
if fixed_width:
lines = [line.ljust(max_chars) for line in lines]
# If there's only one line and fixed_width is True, pad it
if fixed_width and len(lines) == 1:
lines[0] = lines[0].ljust(max_chars)
return lines
# lines = text.split("\n")
# Split text into lines based on max_characters_per_line
lines = split_text(text, max_characters_per_line,fixed_width)
# Load font
font = ImageFont.truetype(font_path, font_size)
@@ -929,26 +1046,35 @@ def generate_text_image(text, font_path, font_size, text_color, vertical=True, s
if layout == "vertical":
for line in lines:
max_char_width = max(font.getsize(char)[0] for char in line)
max_char_width = max(font.getbbox(char)[2] - font.getbbox(char)[0] for char in line)
for char in line:
char_width, char_height = font.getsize(char)
left, top, right, bottom = font.getbbox(char)
char_width = right - left
char_height = bottom - top
char_coordinates.append((x, y))
y += char_height + spacing
max_height = max(max_height, y + padding)
x += max_char_width + spacing
x += max_char_width + line_spacing
y = padding
max_width = x
total_line_width = sum(font.getbbox(line)[2] - font.getbbox(line)[0] for line in lines)
total_spacing = line_spacing * (len(lines) - 1)
max_width = total_line_width + total_spacing + padding * 2
else:
for line in lines:
line_width, line_height = font.getsize(line)
line_width, line_height = font.getbbox(line)[2] - font.getbbox(line)[0], font.getbbox(line)[3] - font.getbbox(line)[1]
for char in line:
char_width, char_height = font.getsize(char)
left, top, right, bottom = font.getbbox(char)
char_width = right - left
char_height = bottom - top
char_coordinates.append((x, y))
x += char_width + spacing
max_width = max(max_width, x + padding)
y += line_height + spacing
y += line_height + line_spacing
x = padding
max_height = y
total_line_heights = sum(font.getbbox(line)[3] - font.getbbox(line)[1] for line in lines)
total_spacing = line_spacing * (len(lines) - 1)
max_height = total_line_heights + total_spacing + padding * 2
# 3. Create image with calculated width and height
image = Image.new('RGBA', (max_width, max_height), (255, 255, 255, 0))
@@ -960,10 +1086,10 @@ def generate_text_image(text, font_path, font_size, text_color, vertical=True, s
for char in line:
x, y = char_coordinates[index]
if stroke:
draw.text((x-stroke_width, y), char, font=font, fill=text_color)
draw.text((x+stroke_width, y), char, font=font, fill=text_color)
draw.text((x, y-stroke_width), char, font=font, fill=text_color)
draw.text((x, y+stroke_width), char, font=font, fill=text_color)
draw.text((x-stroke_width, y), char, font=font, fill=stroke_color)
draw.text((x+stroke_width, y), char, font=font, fill=stroke_color)
draw.text((x, y-stroke_width), char, font=font, fill=stroke_color)
draw.text((x, y+stroke_width), char, font=font, fill=stroke_color)
draw.text((x, y), char, font=font, fill=text_color)
index += 1
@@ -977,10 +1103,18 @@ def generate_text_image(text, font_path, font_size, text_color, vertical=True, s
image = image.convert('RGB')
# 5. Scale the image if fixed_width is specified
if fixed_width and fixed_width < max_width:
scaling_factor = fixed_width / max_width
new_height = int(max_height * scaling_factor)
image = image.resize((fixed_width, new_height), Image.ANTIALIAS)
alpha_image = alpha_image.resize((fixed_width, new_height), Image.ANTIALIAS)
return (image, alpha_image)
def base64_to_image(base64_string):
# 去除前缀
prefix, base64_data = base64_string.split(",", 1)
@@ -1288,6 +1422,9 @@ class LoadImages_:
image=pil2tensor(image)
ims.append(image)
if len(ims)==0:
image1 = Image.new('RGB', (512, 512), color='black')
return (pil2tensor(image1),)
image1 = ims[0]
for image2 in ims[1:]:
if image1.shape[1:] != image2.shape[1:]:
@@ -1318,7 +1455,7 @@ class LoadImagesFromPath:
},
"optional":{
"white_bg": (["disable","enable"],),
"newest_files": (["enable", "disable"],),
"sort_by": (["file_name", "newest"],),#根据文件名来排序,还是按照最新创建时间
"index_variable":("INT", {
"default": 0,
"min": -1, #Minimum value
@@ -1328,13 +1465,13 @@ class LoadImagesFromPath:
}),
"watcher":(["disable","enable"],),
"result": ("WATCHER",),#为了激活本节点运行
"prompt": ("PROMPT",),
# "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"prompt": ("PROMPT",),
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ('IMAGE','MASK','STRING','STRING',)
RETURN_NAMES = ("IMAGE","MASK","prompt_for_FloatingVideo","filepaths",)
RETURN_NAMES = ("image list","MASK","prompt_for_FloatingVideo","filepaths",)
FUNCTION = "run"
@@ -1347,7 +1484,7 @@ class LoadImagesFromPath:
watcher_folder=None
# 运行的函数
def run(self,file_path,white_bg,newest_files,index_variable,watcher,result,prompt):
def run(self,file_path,white_bg,sort_by,index_variable,watcher,result,prompt,seed=1):
global watcher_folder
# print('###监听:',watcher_folder,watcher,file_path,result)
@@ -1370,19 +1507,23 @@ class LoadImagesFromPath:
# 当开启了监听,则取最新的,第一个文件
if watcher=='enable':
index_variable=0
newest_files='enable'
sort_by='newest'
# 排序
sorted_files = sorted(images, key=lambda x: os.path.getmtime(x['file_path']), reverse=(newest_files=='enable'))
if sort_by=='newest':
sorted_files = sorted(images, key=lambda x: os.path.getmtime(x['file_path']), reverse=True)
elif sort_by=='file_name':
# 根据文件名排序
sorted_files = sort_by_filename(images)
imgs=[]
masks=[]
file_names=[]
file_paths=[]
for im in sorted_files:
imgs.append(im['image'])
masks.append(im['mask'])
file_names.append(im['file_name'])
file_paths.append(im['file_path'])
# print('index_variable',index_variable)
@@ -1390,12 +1531,13 @@ class LoadImagesFromPath:
if index_variable!=-1:
imgs=[imgs[index_variable]] if index_variable < len(imgs) else None
masks=[masks[index_variable]] if index_variable < len(masks) else None
file_names=[file_names[index_variable]] if index_variable < len(file_names) else None
file_paths=[file_paths[index_variable]] if index_variable < len(file_paths) else None
except Exception as e:
print("发生了一个未知的错误:", str(e))
# print('#prompt::::',prompt)
return {"ui": {"seed": [1]}, "result":(imgs,masks,prompt,file_names,)}
# return {"ui": {"seed": [1]}, "result":(imgs,masks,prompt,file_names,)}
return (imgs,masks,prompt,file_paths,)
# TODO 扩大选区的功能,重新输出mask
@@ -1498,31 +1640,52 @@ class TextImage:
return {"required": {
"text": ("STRING",{"multiline": True,"default": "龍馬精神迎新歲","dynamicPrompts": False}),
"font": (get_files_with_extension(FONT_PATH,'.ttf'),),#后缀为 ttf
"font": (get_files_with_extension(FONT_PATH,['.ttf','.otf']),),#后缀为 ttf
"font_size": ("INT",{
"default":100,
"min": 100, #Minimum value
"max": 1000, #Maximum value
"min": 1, #Minimum value
"max": 10000000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"spacing": ("INT",{
"default":12,
"min": -200, #Minimum value
"max": 200, #Maximum value
"min": -2000000000, #Minimum value
"max": 2000000000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"line_spacing": ("INT",{
"default":12,
"min": -2000000000, #Minimum value
"max": 2000000000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"padding": ("INT",{
"default":8,
"min": 0, #Minimum value
"max": 200, #Maximum value
"max": 2000000000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"text_color":("STRING",{"multiline": False,"default": "#000000","dynamicPrompts": False}),
"vertical":("BOOLEAN", {"default": True},),
"stroke":("BOOLEAN", {"default": False},),
"max_characters_per_line": ("INT",{
"default":44,
"min": 1, #Minimum value
"max": 2000000000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"fixed_width":("INT",{
"default":0,
"min": 0, #Minimum value
"max": 2000000000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
},
}
@@ -1536,14 +1699,20 @@ class TextImage:
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,False,)
def run(self,text,font,font_size,spacing,padding,text_color,vertical,stroke):
def run(self,text,font,font_size,spacing,line_spacing,padding,text_color,vertical,stroke,max_characters_per_line,fixed_width):
font_path=os.path.join(FONT_PATH,font+'.ttf')
font_path=os.path.join(FONT_PATH,font)
if text=="":
text=" "
# stroke=False, stroke_color=(0, 0, 0), stroke_width=1, spacing=0
img,mask=generate_text_image(text,font_path,font_size,text_color,vertical,stroke,(0, 0, 0),1,spacing,padding)
# max_characters_per_line 英文字按照空格计算1个,中文按照字数计算
if fixed_width==0:
fixed_width=None
img,mask=generate_text_image(text,font_path,font_size,text_color,vertical,stroke,(0, 0, 0),1,
spacing,line_spacing,padding,max_characters_per_line,
fixed_width
)
img=pil2tensor(img)
mask=pil2tensor(mask)
@@ -1577,7 +1746,7 @@ class LoadImagesFromURL:
def run(self,url,seed=0):
global urls_image
print(urls_image)
# print(urls_image)
def filter_http_urls(urls):
filtered_urls = []
for url in urls.split('\n'):
@@ -1666,28 +1835,51 @@ class Image3D:
def run(self,upload,material=None):
# print('material',material)
# print(upload )
image = base64_to_image(upload['image'])
mat=None
if 'material' in upload and upload['material']:
mat=base64_to_image(upload['material'])
mat=mat.convert('RGB')
mat=pil2tensor(mat)
# 截取的系列角度截图
images=upload['images'] if "images" in upload else []
mask = image.split()[3]
image=image.convert('RGB')
ims=[]
for im in images:
if 'type' in im and (not f"[{im['type']}]" in im['name']):
im['name']=im['name']+" "+f"[{im['type']}]"
output_image, output_mask = load_image_to_tensor(im['name'])
ims.append(output_image)
mask=mask.convert('L')
mask=None
bg_image=None
if 'bg_image' in upload and upload['bg_image']:
bg_image = base64_to_image(upload['bg_image'])
bg_image=bg_image.convert('RGB')
bg_image=pil2tensor(bg_image)
mat=None
# 如果没有系列截图
if len(ims)==0:
# 这个是3d模型当前截图
image = base64_to_image(upload['image'])
if 'material' in upload and upload['material']:
mat=base64_to_image(upload['material'])
mat=mat.convert('RGB')
mat=pil2tensor(mat)
mask = image.split()[3]
image=image.convert('RGB')
mask=mask.convert('L')
if 'bg_image' in upload and upload['bg_image']:
bg_image = base64_to_image(upload['bg_image'])
bg_image=bg_image.convert('RGB')
bg_image=pil2tensor(bg_image)
mask=pil2tensor(mask)
image=pil2tensor(image)
mask=pil2tensor(mask)
image=pil2tensor(image)
else:
image = torch.cat(ims, dim=0)
m=[]
if not material is None:
@@ -1808,12 +2000,9 @@ class CompositeImages:
def run(self, foreground,mask,background, is_multiply_blend, position, scale):
results = []
f1=[]
for fg, mask in zip(foreground, mask ):
f1.append([fg,mask])
for f, bg in product(f1, background):
[fg,mask]=f
fg_pil = tensor2pil(fg)
@@ -2710,14 +2899,14 @@ class ResizeImage:
"default": 512,
"min": 1, #Minimum value
"max": 8192, #Maximum value
"step": 8, #Slider's step
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"height": ("INT",{
"default": 512,
"min": 1, #Minimum value
"max": 8192, #Maximum value
"step": 8, #Slider's step
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"scale_option": (["width","height",'overall','center'],),
@@ -2733,7 +2922,7 @@ class ResizeImage:
}
RETURN_TYPES = ("IMAGE","IMAGE","STRING","MASK",)
RETURN_NAMES = ("image","average_image","average_hex","mask",)
RETURN_NAMES = ("image list","average_image","average_hex","mask",)
FUNCTION = "run"
@@ -2771,11 +2960,14 @@ class ResizeImage:
im=tensor2pil(im)
im=im.convert('RGB')
a_im,hex=get_average_color_image(im)
a_im,hex=get_average_color_image(im)
if average_color=='on':
fill_color=hex
a_im=resize_image(a_im,scale_option,w,h,fill_color)
im=resize_image(im,scale_option,w,h,fill_color)
im=pil2tensor(im)
@@ -3201,4 +3393,100 @@ class ImageListToBatch_:
out = torch.cat(out, dim=0)
return (out,)
return (out,)
# https://github.com/gokayfem/ComfyUI-Depth-Visualization?tab=readme-ov-file
class DepthViewer_:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"depth_map": ("IMAGE",),
},
"optional":{
"frames":("IMAGEBASE64",),
},
}
def __init__(self):
self.saved_reference = []
self.saved_depth = []
self.full_output_folder,self.filename,self.counter, self.subfolder, self.filename_prefix = folder_paths.get_save_image_path(
"imagesave",
folder_paths.get_output_directory())
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("frames",)
OUTPUT_NODE = True
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/3D"
def run(self, image, depth_map,frames=None):
self.saved_reference.clear()
self.saved_depth.clear()
image = image[0].detach().cpu().numpy()
depth = depth_map[0].detach().cpu().numpy()
image = Image.fromarray(np.clip(255. * image, 0, 255).astype(np.uint8)).convert('RGB')
depth = Image.fromarray(np.clip(255. * depth, 0, 255).astype(np.uint8))
return self.display([image], [depth],frames)
def display(self, reference_image, depth_map,frames):
for (batch_number, (single_image, single_depth)) in enumerate(zip(reference_image, depth_map)):
filename_with_batch_num = self.filename.replace("%batch_num%", str(batch_number))
image_file = f"{filename_with_batch_num}_{self.counter:05}_reference.png"
single_image.save(os.path.join(self.full_output_folder, image_file))
depth_file = f"{filename_with_batch_num}_{self.counter:05}_depth.png"
single_depth.save(os.path.join(self.full_output_folder, depth_file))
self.saved_reference.append({
"filename": image_file,
"subfolder": self.subfolder,
"type": "output"
})
self.saved_depth.append({
"filename": depth_file,
"subfolder": self.subfolder,
"type": "output"
})
self.counter += 1
ims=[]
image1 = Image.new('RGB', (512, 512), color='black')
image1=pil2tensor(image1)
# print('frames',frames)
if frames!=None and "images" in frames:
for im in frames['images']:
# print(im)
if 'type' in im and (not f"[{im['type']}]" in im['name']):
im['name']=im['name']+" "+f"[{im['type']}]"
try:
output_image, output_mask = load_image_to_tensor(im['name'])
ims.append(output_image)
except:
print("no")
if len(ims)>0:
image1 = ims[0]
for image2 in ims[1:]:
if image1.shape[1:] != image2.shape[1:]:
image2 = comfy.utils.common_upscale(image2.movedim(-1, 1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1, -1)
image1 = torch.cat((image1, image2), dim=0)
return {"ui": {"reference_image": self.saved_reference, "depth_map": self.saved_depth}, "result": (image1,)}
+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
}),) }
+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}),
},
}
+85 -9
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)
@@ -181,6 +181,28 @@ class ColorInput:
return (h,r,g,b,a,)
class KeyInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"key":("KEY",),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("key",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self,key):
return (key,)
class FontInput:
@classmethod
@@ -579,6 +601,7 @@ class AppInfo:
"link":("STRING",{"multiline": False,"default": "https://","dynamicPrompts": False}),
"category":("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
"auto_save": (["enable","disable"],),
"idle_animation": ("BOOLEAN", {"default": False},),
}
}
@@ -594,14 +617,21 @@ class AppInfo:
INPUT_IS_LIST = True
# OUTPUT_IS_LIST = (True,)
def run(self,name,input_ids,output_ids,image,description,version,share_prefix,link,category,auto_save):
def run(self,name,input_ids,output_ids,image,description,version,share_prefix,link,category,auto_save,idle_animation):
name=name[0]
idle_animation=idle_animation[0]
im=None
im=[]
if image:
im=image[0][0]
#TODO batch 的方式需要处理
im=create_temp_file(im)
images=[image]
# batch 的方式需要处理
images=flatten_list(images)
# img=image[0][0]
print('AppInfo_image',len(images))
for i in range(len(images)):
img=images[i]
im.append(create_temp_file(img,i+1)[0])
# image [img,] img[batch,w,h,a] 列表里面是batch,
input_ids=input_ids[0]
@@ -614,7 +644,50 @@ class AppInfo:
# id=get_json_hash([name,im,input_ids,output_ids,description,version])
return {"ui": {"json": [name,im,input_ids,output_ids,description,version,share_prefix,link,category]}, "result": ()}
return {"ui": {"json": [name,im,input_ids,output_ids,description,version,share_prefix,link,category,idle_animation]}, "result": ()}
class CreateJsonNode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"key": ("STRING",{"multiline": False,"default": "data","dynamicPrompts": False}),
"value":(any_type,),
"save":("BOOLEAN", {"default": True},),
},
"optional":{
"json_str":("STRING", {"forceInput": True,}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json_str",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Output"
OUTPUT_NODE = True
INPUT_IS_LIST = False
# OUTPUT_IS_LIST = (True,)
def run(self,key,value,save,json_str=None):
data={}
data[key]=value
if json_str:
json_obj = json.loads(json_str)
data.update(json_obj)
if save:
# 保存为本地文件
with open(os.path.join(folder_paths.get_output_directory(),'data.json'), 'w') as file:
json.dump(data, file, ensure_ascii=False, indent=4)
return (json.dumps(data),)
@@ -795,7 +868,7 @@ class TESTNODE_:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"ANY":(any_type,),
"ANY":(any_type,),
},
}
@@ -810,6 +883,9 @@ class TESTNODE_:
OUTPUT_IS_LIST = (True,)
def run(self,ANY):
print('#TESTNODE_',len(ANY))
print(type(ANY))
try:
print(ANY[0].shape)
+61 -23
View File
@@ -22,7 +22,15 @@ import base64
import mimetypes
# 使用递归的方法将嵌套的列表展平为一维列表
def flatten_list(nested_list):
flat_list = []
for item in nested_list:
if isinstance(item, list):
flat_list.extend(flatten_list(item))
else:
flat_list.append(item)
return flat_list
def get_frames(frame_count, frames, revert=False):
if not revert:
@@ -545,20 +553,24 @@ class LoadAndCombinedAudio_:
if duration > -1:
crop_audio(audio_file, start_time, duration)
return (audio_file, {
"filename": audio_file_name,
"subfolder": "",
"type": "output",
"audio_path":audio_file
} ,)
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_file_path": ("STRING", {"forceInput": True}),
"audio_file_path": ("STRING", {"forceInput": True}),
"video": ("SCENE_VIDEO",),
"audio": ("AUDIO", ),
},
}
@@ -566,23 +578,46 @@ class CombineAudioVideo:
OUTPUT_NODE = True
FUNCTION = "run"
RETURN_TYPES = ()
RETURN_NAMES = ()
RETURN_TYPES = ("SCENE_VIDEO",)
RETURN_NAMES = ("SCENE_VIDEO",)
def run(self,video_file_path, audio_file_path):
def run(self,video, audio):
output_dir = folder_paths.get_output_directory()
counter=get_new_counter(output_dir,'video_final_')
# 判断是否是 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_file_path)
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_file_path,v_file_path)
combine_audio_video(audio_file_path,video,v_file_path)
previews = [
{
@@ -592,7 +627,8 @@ class CombineAudioVideo:
"format": get_mime_type(v_file),
}
]
return {"ui": {"gifs": previews}}
return {"ui": {"gifs": previews},"result":(v_file_path,)}
# The code is based on ComfyUI-VideoHelperSuite modification.
class VideoCombine_Adv:
@@ -885,7 +921,7 @@ class GenerateFramesByCount:
def r(self, frames, frame_count, revert):
image_list = [frames[i:i + 1, ...] for i in range(frames.shape[0])]
print('#image_list',len(image_list),frame_count)
image_list=get_frames(frame_count,image_list,revert)
images = torch.cat(image_list, dim=0)
@@ -900,7 +936,6 @@ class scenesNode_:
return {"required": {
"scenes_video": ('SCENE_VIDEO',),
"index": ("INT", {"default": 0, "min": 0, "step": 1}),
},}
RETURN_TYPES = ('IMAGE','INT',)
@@ -912,14 +947,15 @@ class scenesNode_:
INPUT_IS_LIST = True
def load_video_cv_fallback(self, video, frame_load_cap, skip_first_frames):
# print('#video',video)
images = []
total_frame_count = 0
video_cap = cv2.VideoCapture(video)
try:
video_cap = cv2.VideoCapture(video)
if not video_cap.isOpened():
raise ValueError(f"{video} could not be loaded with cv fallback.")
# set video_cap to look at start_index frame
images = []
total_frame_count = 0
frames_added = 0
base_frame_time = 1/video_cap.get(cv2.CAP_PROP_FPS)
@@ -958,11 +994,13 @@ class scenesNode_:
finally:
video_cap.release()
print("total_frame_count",total_frame_count)
images = torch.cat(images, dim=0)
return (images, frames_added,)
def run(self, scenes_video,index):
scenes_video=flatten_list(scenes_video)
print('#scenes_video',index,scenes_video)
index=index[0]
if len(scenes_video) > index:
+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,)
+6 -4
View File
@@ -7,6 +7,7 @@ import os
import folder_paths
import node_helpers
import hashlib
from uuid import uuid4
# Tensor to PIL
def tensor2pil(image):
@@ -26,7 +27,7 @@ def tensor_to_hash(tensor):
return hash_value
def create_temp_file(image):
def create_temp_file(image, uuid):
output_dir = folder_paths.get_temp_directory()
(
@@ -35,7 +36,7 @@ def create_temp_file(image):
counter,
subfolder,
_,
) = folder_paths.get_save_image_path('material', output_dir)
) = folder_paths.get_save_image_path(f'material_{uuid}', output_dir)
image=tensor2pil(image)
@@ -59,6 +60,7 @@ class EditMask:
def __init__(self):
self.image_id = None
self.uuid = str(uuid4())
@classmethod
def INPUT_TYPES(s):
@@ -117,13 +119,13 @@ class EditMask:
image_path = os.path.join(base_dir,subfolder, name)
if image_path==None:
image_path,images=create_temp_file(image)
image_path,images=create_temp_file(image, self.uuid)
print('#image_path',os.path.exists(image_path),image_path)
# image_path = folder_paths.get_annotated_filepath(image) #文件名
if not os.path.exists(image_path):
image_path,images=create_temp_file(image)
image_path,images=create_temp_file(image, self.uuid)
img = node_helpers.pillow(Image.open, image_path)
+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,)
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-mixlab-nodes"
description = "3D, ScreenShareNode & FloatingVideoNode, SpeechRecognition & SpeechSynthesis, GPT, LoadImagesFromLocal, Layers, Other Nodes, ..."
version = "0.30.3"
version = "0.46.0"
license = "MIT"
dependencies = ["numpy", "pyOpenSSL", "watchdog", "opencv-python-headless", "matplotlib", "openai", "simple-lama-inpainting", "clip-interrogator==0.6.0", "transformers>=4.36.0", "lark-parser", "imageio-ffmpeg", "rembg[gpu]", "omegaconf==2.3.0", "Pillow>=9.5.0", "einops==0.7.0", "trimesh>=4.0.5", "huggingface-hub", "scikit-image"]
+22 -5
View File
@@ -4,17 +4,34 @@ watchdog
opencv-python-headless
matplotlib
openai
simple-lama-inpainting
torchaudio
clip-interrogator==0.6.0
transformers>=4.36.0
lark-parser
imageio-ffmpeg
rembg[gpu]
omegaconf==2.3.0
omegaconf>=2.3.0
Pillow>=9.5.0
einops==0.7.0
einops>=0.7.0
trimesh>=4.0.5
huggingface-hub
scikit-image
torchaudio
soundfile>=0.12.1
soundfile>=0.12.1
json-repair
bitsandbytes
accelerate
scenedetect[opencv-headless]
hydra-core>=1.3.2
loralib>=0.1.2
natsort>=8.4.0
#simple-lama-inpainting
git+https://github.com/shadowcz007/SenseVoice-python.git
faster_whisper
git+https://github.com/openai/swarm.git
+435 -279
View File
@@ -2,6 +2,62 @@ import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { $el } from '../../../scripts/ui.js'
import { loadExternalScript, get_position_style } from './common.js'
function setCameraOrbit (modelview, distant, angles, screenNumber) {
//2.1 20
// const angles = {
// 1: -20.0,
// 2: -17.9,
// 3: -15.8,
// 4: -13.7,
// 5: -11.6,
// 6: -9.5,
// 7: -7.4,
// 8: -5.3,
// 9: -3.2,
// 10: -1.1,
// 11: 1.1,
// 12: 3.2,
// 13: 5.3,
// 14: 7.4,
// 15: 9.5,
// 16: 11.6,
// 17: 13.7,
// 18: 15.8,
// 19: 17.9,
// 20: 20.0
// };
// 12 3.6
// const angles = {
// 1: -20.0,
// 2: -16.4,
// 3: -12.7,
// 4: -9.1,
// 5: -5.5,
// 6: -1.8,
// 7: 1.8,
// 8: 5.5,
// 9: 9.1,
// 10: 12.7,
// 11: 16.4,
// 12: 20.0
// }
const angle = angles[screenNumber]
let co=modelview.cameraOrbit.split(" ")
if (angle !== undefined) {
modelview.cameraOrbit = `${angle}deg ${co[1]} ${distant}m`
console.log(screenNumber, angle)
} else {
console.error('Invalid screen number')
}
}
const getLocalData = key => {
let data = {}
try {
@@ -26,7 +82,8 @@ const setLocalDataOfWin = (key, value) => {
localStorage.setItem(key, JSON.stringify(value))
// window[key] = value
}
async function uploadImage (blob, fileType = '.svg', filename) {
async function uploadImage_ (blob, fileType = '.svg', filename) {
// const blob = await (await fetch(src)).blob();
const body = new FormData()
body.append(
@@ -41,13 +98,17 @@ async function uploadImage (blob, fileType = '.svg', filename) {
// console.log(resp)
let data = await resp.json()
return data
}
async function uploadImage (blob, fileType = '.svg', filename) {
let data = await uploadImage_(blob, fileType, filename)
let { name, subfolder } = data
let src = api.apiURL(
`/view?filename=${encodeURIComponent(
name
)}&type=input&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
)
return src
}
@@ -78,38 +139,6 @@ const parseImage = url => {
})
}
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
async function extractMaterial (
modelViewerVariants,
selectMaterial,
@@ -171,6 +200,42 @@ async function changeMaterial (
targetMaterial.pbrMetallicRoughness.baseColorTexture.setTexture(targetTexture)
}
function inputFileClick (isFileURL = false, isGlb = false) {
return new Promise((res, rej) => {
// 创建一个input元素
var input = document.createElement('input')
input.type = 'file'
input.accept = isGlb ? '.glb' : 'image/*'
// 监听input的change事件
input.addEventListener('change', function () {
// 获取上传的文件
var file = input.files[0]
if (isFileURL) {
res(URL.createObjectURL(file))
return
}
// 创建一个FileReader对象来读取文件
var reader = new FileReader()
// 监听FileReader的load事件
reader.addEventListener('load', async () => {
let base64 = reader.result
input.remove()
res(base64)
})
// 读取文件
reader.readAsDataURL(file)
})
// 触发input的点击事件
input.click()
})
}
app.registerExtension({
name: 'Mixlab.3D.3DImage',
async getCustomWidgets (app) {
@@ -189,7 +254,7 @@ app.registerExtension({
let d = getLocalData('_mixlab_3d_image')
// console.log('serializeValue', node)
if (d && d[node.id]) {
let { url, bg, material } = d[node.id]
let { url, bg, material, images } = d[node.id]
let data = {}
if (url) {
data.image = await parseImage(url)
@@ -205,6 +270,10 @@ app.registerExtension({
data.material = await parseImage(material)
}
if (images) {
data.images = images
}
return JSON.parse(JSON.stringify(data))
} else {
return {}
@@ -217,66 +286,62 @@ app.registerExtension({
}
},
async init () {
await loadExternalScript('/mixlab/app/lib/model-viewer.min.js', 'module')
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == '3DImage') {
console.log('nodeType.comfyClass', nodeType.comfyClass)
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
orig_nodeCreated?.apply(this, arguments)
const uploadWidget = this.widgets.filter(w => w.name == 'upload')[0]
const widget = {
type: 'div',
name: 'upload-preview',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 88, node.size[1])
get_position_style(ctx, widget_width - 122, 88, node.size[1], 44)
)
}
}
widget.div = $el('div', {})
widget.div.style.width = `120px`
document.body.appendChild(widget.div)
const inputDiv = (key, placeholder, preview) => {
let div = document.createElement('div')
const ip = document.createElement('input')
ip.type = 'file'
const ip = document.createElement('button')
ip.className = `${'comfy-multiline-input'} ${placeholder}`
div.style = `display: flex;
align-items: center;
margin: 6px 8px;
margin-top: 0;`
ip.placeholder = placeholder
// ip.value = value
ip.style = `outline: none;
border: none;
padding: 4px;
width: 60%;cursor: pointer;
width: 100px;cursor: pointer;
height: 32px;`
const label = document.createElement('label')
label.style = 'font-size: 10px;min-width:32px'
label.innerText = placeholder
div.appendChild(label)
ip.innerText = placeholder
div.appendChild(ip)
let that = this,
filename = new Date().getTime()
let that = this
ip.addEventListener('change', async event => {
const file = event.target.files[0]
const reader = new FileReader()
filename = new Date().getTime()
// 读取文件内容
reader.onload = async e => {
const fileURL = URL.createObjectURL(file)
// console.log('文件URL: ', fileURL)
let html = `<model-viewer src="${fileURL}"
min-field-of-view="0deg" max-field-of-view="180deg"
ip.addEventListener('click', async event => {
let fileURL = await inputFileClick(true, true)
// console.log('文件URL: ', fileURL)
let html = `<model-viewer src="${fileURL}"
oncontextmenu="return false;"
style="outline:1px solid white"
min-field-of-view="0deg"
max-field-of-view="180deg"
min-camera-orbit="auto auto 0m"
max-camera-orbit="auto auto 1000m"
shadow-intensity="1"
camera-controls
touch-action="pan-y">
@@ -285,230 +350,314 @@ app.registerExtension({
<div>Variant: <select class="variant"></select></div>
<div>Material: <select class="material"></select></div>
<div>Material: <div class="material_img"> </div></div>
<div><button class="bg">BG</button></div>
<div>
<button class="bg">BG</button>
</div>
<div>
<input class="ddcap_distant" type="number" min="1" step="1" value="55">
<input class="total_images" type="number" min="1" max="180" step="1" value="20">
<input class="ddcap_range" type="number" min="0" max="20" step="0.1" value="2.1">
<button class="ddcap">Capture Rotational Screenshots</button></div>
<div><button class="export">Export GLB</button></div>
</div></model-viewer>`
preview.innerHTML = html
if (that.size[1] < 400) {
that.setSize([that.size[0], that.size[1] + 300])
app.canvas.draw(true, true)
}
const modelViewerVariants = preview.querySelector('model-viewer')
const select = preview.querySelector('.variant')
const selectMaterial = preview.querySelector('.material')
const material_img = preview.querySelector('.material_img')
const bg = preview.querySelector('.bg')
const exportGLB = preview.querySelector('.export')
if (modelViewerVariants) {
modelViewerVariants.style.width = `${that.size[0] - 24}px`
modelViewerVariants.style.height = `${that.size[1] - 48}px`
}
modelViewerVariants.addEventListener('load', async () => {
const names = modelViewerVariants.availableVariants
// 变量
for (const name of names) {
const option = document.createElement('option')
option.value = name
option.textContent = name
select.appendChild(option)
}
// Adds a default option.
if (names.length === 0) {
const option = document.createElement('option')
option.value = 'default'
option.textContent = 'Default'
select.appendChild(option)
}
// 材质
extractMaterial(
modelViewerVariants,
selectMaterial,
material_img
)
})
let timer = null
const delay = 500 // 延迟时间,单位为毫秒
async function checkCameraChange () {
let dd = getLocalData(key)
let base64Data = modelViewerVariants.toDataURL()
const contentType = getContentTypeFromBase64(base64Data)
const blob = await base64ToBlobFromURL(base64Data, contentType)
// const fileBlob = new Blob([e.target.result], { type: file.type });
let url = await uploadImage(blob, '.png')
// console.log(url)
let bg_blob = await base64ToBlobFromURL(
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mN88uXrPQAFwwK/6xJ6CQAAAABJRU5ErkJggg=='
)
let url_bg = await uploadImage(bg_blob, '.png')
// console.log('url_bg',url_bg)
if (!dd[that.id]) {
dd[that.id] = { url, bg: url_bg }
} else {
dd[that.id] = { ...dd[that.id], url }
}
// 材质贴图
let thumbUrl = material_img.getAttribute('src')
if (thumbUrl) {
let tb = await base64ToBlobFromURL(thumbUrl)
let tUrl = await uploadImage(tb, '.png')
// console.log('材质贴图', tUrl, thumbUrl)
dd[that.id].material = tUrl
}
setLocalDataOfWin(key, dd)
}
function startTimer () {
if (timer) clearTimeout(timer)
timer = setTimeout(checkCameraChange, delay)
}
modelViewerVariants.addEventListener('camera-change', startTimer)
select.addEventListener('input', async event => {
modelViewerVariants.variantName =
event.target.value === 'default' ? null : event.target.value
// 材质
await extractMaterial(
modelViewerVariants,
selectMaterial,
material_img
)
checkCameraChange()
})
selectMaterial.addEventListener('input', event => {
// console.log(selectMaterial.value)
material_img.setAttribute('src', selectMaterial.value)
if (selectMaterial.getAttribute('data-new-material')) {
let index =
~~selectMaterial.selectedOptions[0].getAttribute(
'data-index'
)
changeMaterial(
modelViewerVariants,
modelViewerVariants.model.materials[index],
selectMaterial.getAttribute('data-new-material')
)
}
checkCameraChange()
})
bg.addEventListener('click', () => {
// 创建一个input元素
var input = document.createElement('input')
input.type = 'file'
// 监听input的change事件
input.addEventListener('change', function () {
// 获取上传的文件
var file = input.files[0]
// 创建一个FileReader对象来读取文件
var reader = new FileReader()
// 监听FileReader的load事件
reader.addEventListener('load', async () => {
let base64 = reader.result
// 将读取的文件内容设置为div的背景
preview.style.backgroundImage = 'url(' + base64 + ')'
const contentType = getContentTypeFromBase64(base64)
const blob = await base64ToBlobFromURL(base64, contentType)
// const fileBlob = new Blob([e.target.result], { type: file.type });
let bg_url = await uploadImage(blob, '.png')
let bg_img = await createImage(base64)
let dd = getLocalData(key)
// console.log(dd[that.id],bg_url)
if (!dd[that.id]) dd[that.id] = { url: '', bg: bg_url }
dd[that.id] = {
...dd[that.id],
bg: bg_url,
bg_w: bg_img.naturalWidth,
bg_h: bg_img.naturalHeight
}
setLocalDataOfWin(key, dd)
// 更新尺寸
let w = that.size[0] - 24,
h = (w * bg_img.naturalHeight) / bg_img.naturalWidth
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${h}px`
}
preview.style.width = `${w}px`
})
// 读取文件
reader.readAsDataURL(file)
})
// 触发input的点击事件
input.click()
})
exportGLB.addEventListener('click', async () => {
const glTF = await modelViewerVariants.exportScene()
const file = new File([glTF], 'export.glb')
const link = document.createElement('a')
link.download = file.name
link.href = URL.createObjectURL(file)
link.click()
})
uploadWidget.value = await uploadWidget.serializeValue()
// 更新尺寸
let dd = getLocalData(key)
// console.log(dd[that.id],bg_url)
if (dd[that.id]) {
const { bg_w, bg_h } = dd[that.id]
if (bg_h && bg_w) {
let w = that.size[0] - 24,
h = (w * bg_h) / bg_w
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${h}px`
}
preview.style.width = `${w}px`
}
}
preview.innerHTML = html
if (that.size[1] < 400) {
that.setSize([that.size[0], that.size[1] + 300])
app.canvas.draw(true, true)
}
// 以文本形式读取文件
reader.readAsDataURL(file)
const modelViewerVariants = preview.querySelector('model-viewer')
const select = preview.querySelector('.variant')
const selectMaterial = preview.querySelector('.material')
const material_img = preview.querySelector('.material_img')
const bg = preview.querySelector('.bg')
const exportGLB = preview.querySelector('.export')
const ddcap_distant = preview.querySelector('.ddcap_distant')
const total_images = preview.querySelector('.total_images')
const ddcap_range = preview.querySelector('.ddcap_range')
const ddCap = preview.querySelector('.ddcap')
const sleep = (t = 1000) => {
return new Promise((res, rej) => {
return setTimeout(() => {
res(t)
}, t)
})
}
async function captureImage (isUrl = true) {
let base64Data = modelViewerVariants.toDataURL()
const contentType = getContentTypeFromBase64(base64Data)
const blob = await base64ToBlobFromURL(base64Data, contentType)
if (isUrl) return await uploadImage(blob, '.png')
return await uploadImage_(blob, '.png')
}
async function captureImages (
ddcap_range = 1,
total_images = 12,
distant = 0.23
) {
// 初始 角度
var center = modelViewerVariants.getBoundingBoxCenter().toString()
modelViewerVariants.cameraTarget = center
const startAngle = -((total_images - 1) / 2) * ddcap_range
const angles = {}
for (let i = 0; i < total_images; i++) {
angles[i + 1] = startAngle + i * ddcap_range
}
console.log(angles)
let frames = []
modelViewerVariants.removeAttribute('camera-controls')
for (let i = 0; i < total_images; i++) {
setCameraOrbit(modelViewerVariants, distant, angles, i + 1)
// modelViewerVariants.cameraOrbit = `${currentAngle}deg ${initialCameraOrbit[1]} ${initialCameraOrbit[2]}`
await sleep(1000)
// console.log(`Capturing image at angle: ${currentAngle}deg`)
let file = await captureImage(false)
frames.push(file)
// currentAngle += angleIncrement
}
await sleep(1000)
// 恢复到初始旋转角度
// modelViewerVariants.cameraOrbit = initialCameraOrbit.join(' ')
modelViewerVariants.setAttribute('camera-controls', '')
return frames
}
ddCap.addEventListener('click', async e => {
const distant = Number(ddcap_distant.value), // 23m
totalImages = Number(total_images.value),
angleIncrement = Number(ddcap_range.value)
console.log(angleIncrement, totalImages)
let images = await captureImages(
angleIncrement,
totalImages,
distant
)
let dd = getLocalData(key)
dd[that.id].images = images
setLocalDataOfWin(key, dd)
})
ddcap_distant.addEventListener('input', async e => {
// console.log(ddcap_distant.value)
const center = modelViewerVariants.getBoundingBoxCenter().toString()
modelViewerVariants.cameraTarget = center;
const initialCameraOrbit =
modelViewerVariants.cameraOrbit.split(' ')
modelViewerVariants.cameraOrbit = `${initialCameraOrbit[2]} ${initialCameraOrbit[1]} ${ddcap_distant.value}m`
modelViewerVariants.setAttribute('camera-controls', '')
})
// ddcap_range_top.addEventListener('input', async e => {
// // console.log(ddcap_range.value)
// const initialCameraOrbit =
// modelViewerVariants.cameraOrbit.split(' ')
// modelViewerVariants.cameraOrbit = `${initialCameraOrbit[0]} ${ddcap_range_top.value}deg ${initialCameraOrbit[2]}`
// modelViewerVariants.setAttribute('camera-controls', '')
// })
if (modelViewerVariants) {
modelViewerVariants.style.width = `${that.size[0] - 48}px`
modelViewerVariants.style.height = `${that.size[1] - 48}px`
}
modelViewerVariants.addEventListener('load', async () => {
const names = modelViewerVariants.availableVariants
// 变量
for (const name of names) {
const option = document.createElement('option')
option.value = name
option.textContent = name
select.appendChild(option)
}
// Adds a default option.
if (names.length === 0) {
const option = document.createElement('option')
option.value = 'default'
option.textContent = 'Default'
select.appendChild(option)
}
// 材质
extractMaterial(modelViewerVariants, selectMaterial, material_img)
})
let timer = null
const delay = 500 // 延迟时间,单位为毫秒
async function checkCameraChange () {
let dd = getLocalData(key)
// const fileBlob = new Blob([e.target.result], { type: file.type });
let url = await captureImage()
let bg_blob = await base64ToBlobFromURL(
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mN88uXrPQAFwwK/6xJ6CQAAAABJRU5ErkJggg=='
)
let url_bg = await uploadImage(bg_blob, '.png')
// console.log('url_bg',url_bg)
if (!dd[that.id]) {
dd[that.id] = { url, bg: url_bg }
} else {
dd[that.id] = { ...dd[that.id], url }
}
// 材质贴图
let thumbUrl = material_img.getAttribute('src')
if (thumbUrl) {
let tb = await base64ToBlobFromURL(thumbUrl)
let tUrl = await uploadImage(tb, '.png')
// console.log('材质贴图', tUrl, thumbUrl)
dd[that.id].material = tUrl
}
setLocalDataOfWin(key, dd)
}
function startTimer () {
if (timer) clearTimeout(timer)
timer = setTimeout(checkCameraChange, delay)
}
modelViewerVariants.addEventListener('camera-change', startTimer)
select.addEventListener('input', async event => {
modelViewerVariants.variantName =
event.target.value === 'default' ? null : event.target.value
// 材质
await extractMaterial(
modelViewerVariants,
selectMaterial,
material_img
)
checkCameraChange()
})
selectMaterial.addEventListener('input', event => {
// console.log(selectMaterial.value)
material_img.setAttribute('src', selectMaterial.value)
if (selectMaterial.getAttribute('data-new-material')) {
let index =
~~selectMaterial.selectedOptions[0].getAttribute('data-index')
changeMaterial(
modelViewerVariants,
modelViewerVariants.model.materials[index],
selectMaterial.getAttribute('data-new-material')
)
}
checkCameraChange()
})
//更新bg
const updateBgData = (id, key, url, w, h) => {
let dd = getLocalData(key)
// console.log(dd[that.id],url)
if (!dd[id]) dd[id] = { url: '', bg: url }
dd[id] = {
...dd[id],
bg: url,
bg_w: w,
bg_h: h
}
setLocalDataOfWin(key, dd)
}
bg.addEventListener('click', async () => {
//更新bg
updateBgData(that.id, key, '', 0, 0)
preview.style.backgroundImage = 'none'
let base64 = await inputFileClick(false, false)
// 将读取的文件内容设置为div的背景
preview.style.backgroundImage = 'url(' + base64 + ')'
const contentType = getContentTypeFromBase64(base64)
const blob = await base64ToBlobFromURL(base64, contentType)
// const fileBlob = new Blob([e.target.result], { type: file.type });
let bg_url = await uploadImage(blob, '.png')
let bg_img = await createImage(base64)
//更新bg
updateBgData(
that.id,
key,
bg_url,
bg_img.naturalWidth,
bg_img.naturalHeight
)
// 更新尺寸
let w = that.size[0] - 128,
h = (w * bg_img.naturalHeight) / bg_img.naturalWidth
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${h}px`
}
preview.style.width = `${w}px`
})
exportGLB.addEventListener('click', async () => {
const glTF = await modelViewerVariants.exportScene()
const file = new File([glTF], 'export.glb')
const link = document.createElement('a')
link.download = file.name
link.href = URL.createObjectURL(file)
link.click()
})
uploadWidget.value = await uploadWidget.serializeValue()
// 更新尺寸
let dd = getLocalData(key)
// console.log(dd[that.id],bg_url)
if (dd[that.id]) {
const { bg_w, bg_h } = dd[that.id]
if (bg_h && bg_w) {
let w = that.size[0] - 48,
h = (w * bg_h) / bg_w
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${h}px`
}
preview.style.width = `${w}px`
}
}
})
return div
}
let preview = document.createElement('div')
preview.className = 'preview'
preview.style = `margin-top: 12px;display: flex;
preview.style = `margin-top: 12px;
display: flex;
justify-content: center;
align-items: center;background-repeat: no-repeat;background-size: contain;`
align-items: center;background-repeat: no-repeat;
background-size: contain;`
let upload = inputDiv('_mixlab_3d_image', '3D Model', preview)
@@ -523,18 +672,25 @@ app.registerExtension({
// 更新尺寸
let dd = getLocalData('_mixlab_3d_image')
// console.log(dd[that.id],bg_url)
if (dd[that.id]) {
const { bg_w, bg_h } = dd[that.id]
if (bg_h && bg_w) {
let w = that.size[0] - 24,
h = (w * bg_h) / bg_w
let w = that.size[0] - 128
preview.style.width = `${w}px`
console.log('更新尺寸', w)
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${Math.round(
that.size[1] * 0.8
)}px`
}
if (bg_h && bg_w) {
let h = (w * bg_h) / bg_w
if (modelViewerVariants) {
modelViewerVariants.style.width = `${w}px`
modelViewerVariants.style.height = `${h}px`
}
preview.style.width = `${w}px`
}
}
@@ -561,7 +717,7 @@ app.registerExtension({
const r = onExecuted?.apply?.(this, arguments)
let div = this.widgets.filter(d => d.div)[0]?.div
console.log('Test', this.widgets)
// console.log('Test', this.widgets)
let material = message.material[0]
if (material) {
@@ -617,7 +773,7 @@ app.registerExtension({
// let base64 = await parseImage(url)
let pre = widget.div.querySelector('.preview')
pre.style.width = `${node.size[0]}px`
pre.style.width = `${node.size[0] - 24}px`
pre.innerHTML = `
${url ? `<img src="${url}" style="width:100%"/>` : ''}
`
+57 -84
View File
@@ -3,28 +3,17 @@ import { $el } from '../../../scripts/ui.js'
import { api } from '../../../scripts/api.js'
import { td_bg } from './td_background.js'
console.log('td_bg', td_bg)
// console.log('td_bg', td_bg)
import {
getUrl,
base64Df,
get_position_style,
getObjectInfo
} from './common.js'
//本机安装的插件节点全集
window._nodesAll = null
//获取当前系统的插件,节点清单
function getObjectInfo () {
return new Promise(async (resolve, reject) => {
let url = getUrl()
try {
const response = await fetch(`${url}/object_info`)
const data = await response.json()
resolve(data)
} catch (error) {
reject(error)
}
})
}
const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
const parseImageToBase64 = url => {
return new Promise((res, rej) => {
fetch(url)
@@ -44,39 +33,6 @@ const parseImageToBase64 = url => {
})
}
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 12 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'flex-start',
zIndex: 9999999
}
}
async function drawImageToCanvas (imageUrl, sFactor = 320) {
var canvas = document.createElement('canvas')
var ctx = canvas.getContext('2d')
@@ -256,7 +212,10 @@ async function extractInputAndOutputData (
node.type === 'KSampler' ||
node.type == 'SamplerCustom' ||
node.type === 'ChinesePrompt_Mix' ||
node.type === 'Seed_'
node.type === 'Seed_' ||
node.type === 'SiliconflowLLM' ||
node.type === 'ChatGPTOpenAI' ||
node.type === 'SiliconflowTextToImageNode'
) {
// seed 的类型收集
try {
@@ -276,23 +235,6 @@ async function extractInputAndOutputData (
return { input, output, seed, seedTitle }
}
function getUrl () {
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
return url
}
const getLocalData = key => {
let data = {}
try {
data = JSON.parse(localStorage.getItem(key)) || {}
} catch (error) {
return {}
}
return data
}
async function save_app (json) {
let url = getUrl()
@@ -325,20 +267,28 @@ function downloadJsonFile (jsonData, fileName = 'mix_app.json') {
}
async function save (json, download = false, showInfo = true) {
let nodesAll = window._nodesAll || (await getObjectInfo())
if (!window._nodesAll) {
window._nodesAll = await getObjectInfo();
}
console.log('####SAVE', nodesAll, json[0])
let nodesAll = window._nodesAll;
// let nodesAll = window._nodesAll || (await getObjectInfo())
console.log('####SAVE', nodesAll, json)
const name = json[0],
version = json[5],
share_prefix = json[6], //用于分享的功能扩展
link = json[7], //用于创建界面上的跳转链接
category = json[8] || '', //用于分类
idle_animation = json[9], //用于动画,比如数字人her
description = json[4],
inputIds = json[2].split('\n').filter(f => f),
outputIds = json[3].split('\n').filter(f => f)
const iconData = json[1][0]
let { filename, subfolder, type } = iconData
let iconUrl = api.apiURL(
`/view?filename=${encodeURIComponent(
@@ -391,6 +341,27 @@ async function save (json, download = false, showInfo = true) {
try {
data.app.icon = await drawImageToCanvas(iconUrl)
} catch (error) {}
let images = []
if (json[1].length > 1 && idle_animation) {
images = Array.from(json[1], j => {
let { filename, subfolder, type } = j
return api.apiURL(
`/view?filename=${encodeURIComponent(
filename
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
)
})
}
try {
for (let index = 0; index < images.length; index++) {
const imgurl = images[index]
images[index] = await drawImageToCanvas(imgurl)
}
if (idle_animation) data.app.idle_animation = images
} catch (error) {}
// console.log(data.app)
// let http_workflow = app.graph.serialize()
await save_app(data)
@@ -400,13 +371,17 @@ async function save (json, download = false, showInfo = true) {
if (showInfo) {
let open = window.confirm(
`You can now access the standalone application on a new page!\n${getUrl()}/mixlab/app?filename=${encodeURIComponent(
`You can now access the standalone application on a new page!\n${getUrl()}/mixlab/app${
data.app.idle_animation ? '/her.html' : ''
}?filename=${encodeURIComponent(
data.app.filename
)}&category=${encodeURIComponent(data.app.category)}`
)
if (open)
window.open(
`${getUrl()}/mixlab/app?filename=${encodeURIComponent(
`${getUrl()}/mixlab/app${
data.app.idle_animation ? '/her.html' : ''
}?filename=${encodeURIComponent(
data.app.filename
)}&category=${encodeURIComponent(data.app.category)}`
)
@@ -448,9 +423,9 @@ function getInputsAndOutputs () {
app.registerExtension({
name: 'Mixlab.utils.AppInfo',
init () {
if (!window._nodesAll) {
getObjectInfo().then(r => (window._nodesAll = r))
}
// if (!window._nodesAll) {
// getObjectInfo().then(r => (window._nodesAll = r))
// }
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'AppInfo') {
@@ -466,20 +441,19 @@ app.registerExtension({
const { input, output } = getInputsAndOutputs()
input_ids.value = input.join('\n')
output_ids.value = output.join('\n')
const widget = {
type: 'div',
name: 'AppInfoRun',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(
Object.assign(this.div.style, {
...get_position_style(
ctx,
widget_width,
node.size[1] - widget_height,
node.size[1]
)
)
),
zIndex: 1
})
}
}
@@ -709,7 +683,6 @@ app.registerExtension({
this.serialize_widgets = true //需要保存参数
window._mixlab_app_json = null
}
const onExecuted = nodeType.prototype.onExecuted
+5 -1
View File
@@ -19,7 +19,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
@@ -555,6 +558,7 @@ app.registerExtension({
e.preventDefault()
let inputAudio = document.createElement('input')
inputAudio.type = 'file'
inputAudio.accept = "audio/*"
inputAudio.style.display = 'none'
inputAudio.addEventListener('change', async e => {
e.preventDefault()
+112 -10
View File
@@ -1,3 +1,5 @@
import { getUrl } from './common.js'
async function* completion (url, messages, controller) {
let data = {
model: 'gpt-3.5-turbo-16k',
@@ -91,16 +93,116 @@ async function* completion (url, messages, controller) {
return content
// return (await response.json()).content
}
export async function completion_ (url, messages, controller, callback) {
let request = await completion(url, messages, controller)
export async function completion_ (
apiKey,
url,
model_name,
messages,
controller,
callback
) {
let request = await chatCompletion(
apiKey,
url,
model_name,
messages,
controller
)
for await (const chunk of request) {
let content = chunk.data.choices[0].delta.content || ''
if (chunk.data.choices[0].role == 'assistant') {
//开始
content = ''
}
if (callback) callback(content)
if (callback) callback(chunk)
}
}
export async function* chatCompletion (
apiKey,
api_url,
model_name,
messages,
controller
) {
const mixlabAPI = `${getUrl()}/chat/completions`
const requestBody = {
messages: messages,
stream: true,
key: apiKey,
model_name: model_name,
api_url
}
let response = await fetch(mixlabAPI, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
// Authorization: `Bearer ${apiKey}`
},
body: JSON.stringify(requestBody),
mode: 'cors', // This is to ensure the request is made with CORS
signal: controller.signal
})
const reader = response.body.getReader()
const decoder = new TextDecoder()
let content = ''
let leftover = '' // Buffer for partially read lines
try {
let cont = true
while (cont) {
let result = await reader.read()
if (result.done) {
break
}
const text = leftover + decoder.decode(result.value)
// Check if the last character is a line break
const endsWithLineBreak = text.endsWith('\n')
// Split the text into lines
let lines = text.split('\n')
// If the text doesn't end with a line break, then the last line is incomplete
// Store it in leftover to be added to the next chunk of data
if (!endsWithLineBreak) {
leftover = lines.pop()
} else {
leftover = '' // Reset leftover if we have a line break at the end
}
// Parse all sse events and add them to result
const regex = /^(\S+):\s(.*)$/gm
for (const line of lines) {
const match = regex.exec(line)
if (match) {
result[match[1]] = match[2]
// since we know this is llama.cpp, let's just decode the json in data
if (result.data) {
result.data = JSON.parse(result.data)
content += result.data.choices[0].delta?.content || ''
// console.log('#result.content',content)
// yield
yield result
// if we got a stop token from server, we will break here
if (result.data.choices[0].finish_reason == 'stop') {
if (result.data.generation_settings) {
// generation_settings = result.data.generation_settings;
}
cont = false
break
}
}
}
}
}
} catch (e) {
console.error('chat error: ', e)
throw e
} finally {
controller.abort()
}
return content
}
+1 -1
View File
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
const version = 'v0.30.3'
const version = 'v0.46.0'
fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
.then(response => response.json())
+227
View File
@@ -0,0 +1,227 @@
export const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
export function getUrl () {
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
return url
}
// 获得插件/节点的索引数据
export async function get_nodes_map () {
let url = getUrl()
const res = await fetch(`${url}/mixlab/nodes_map`, {
method: 'POST',
body: JSON.stringify({
data: 'json'
})
})
return await res.json()
}
// 更新或者获取key
export const updateLLMAPIKey = async key => {
try {
const res = await fetch(`${getUrl()}/mixlab/llm_api_key`, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
key: key || null
})
})
const data = await res.json()
if (!res.ok) {
console.error('Error:', data.error)
return
}
if (key) {
console.log('API key saved successfully:', data.message)
return key
} else {
console.log('Retrieved API key:', data.key)
return data.key
}
} catch (error) {
console.error('Request failed:', error)
}
}
//获取当前系统的插件,节点清单
export function getObjectInfo () {
return new Promise(async (resolve, reject) => {
let url = getUrl()
try {
const response = await fetch(`${url}/object_info`)
const data = await response.json()
resolve(data)
} catch (error) {
reject(error)
}
})
}
export function get_position_style (
ctx,
widget_width,
y,
node_height,
left = 44
) {
const MARGIN = 0 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const scaleX = elRect.width / ctx.canvas.width
const scaleY = elRect.height / ctx.canvas.height
const transform = new DOMMatrix()
.scaleSelf(scaleX, scaleY)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `${left}px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'flex-start',
zIndex: 99
}
}
export function loadCSS (url) {
var link = document.createElement('link')
link.rel = 'stylesheet'
link.type = 'text/css'
link.href = url
document.getElementsByTagName('head')[0].appendChild(link)
}
export function injectCSS (css) {
// 检查页面中是否已经存在具有相同内容的style标签
const existingStyle = document.querySelector('style')
if (existingStyle && existingStyle.textContent === css) {
return // 如果已经存在相同的样式,则不进行注入
}
// 创建一个新的style标签,并将CSS内容注入其中
const style = document.createElement('style')
style.textContent = css
// 将style标签插入到页面的head元素中
const head = document.querySelector('head')
head.appendChild(style)
}
export function loadExternalScript (url, type) {
return new Promise((resolve, reject) => {
const existingScript = document.querySelector(`script[src="${url}"]`)
if (existingScript) {
existingScript.onload = () => {
resolve()
}
existingScript.onerror = reject
return
}
const script = document.createElement('script')
script.src = url
if (type) script.type = type // Add this line to load the script as an ES module
script.onload = () => {
resolve()
}
script.onerror = reject
document.head.appendChild(script)
})
}
export async function getQueue () {
try {
const res = await fetch(`${getUrl()}/queue`)
const data = await res.json()
// console.log(data.queue_running,data.queue_pending)
return {
// Running action uses a different endpoint for cancelling
Running: data.queue_running.length,
Pending: data.queue_pending.length
}
} catch (error) {
console.error(error)
return { Running: 0, Pending: 0 }
}
}
export async function interrupt () {
const resp = await fetch(`${getUrl()}/interrupt`, {
method: 'POST'
})
}
export async function sleep (t = 200) {
return new Promise((res, rej) => {
setTimeout(() => {
res(true)
}, t)
})
}
export function createImage (url) {
let im = new Image()
return new Promise((res, rej) => {
im.onload = () => res(im)
im.src = url
})
}
export function convertImageUrlToBase64 (imageUrl) {
return fetch(imageUrl)
.then(response => response.blob())
.then(blob => {
return new Promise((resolve, reject) => {
const reader = new FileReader()
reader.onloadend = () => resolve(reader.result)
reader.onerror = reject
reader.readAsDataURL(blob)
})
})
}
export const getLocalData = key => {
let data = {}
try {
data = JSON.parse(localStorage.getItem(key)) || {}
} catch (error) {
return {}
}
return data
}
export const saveLocalData = (key, id, val) => {
let data = getLocalData(key)
data[id] = val
localStorage.setItem(key, JSON.stringify(data))
}
+1 -220
View File
@@ -1,205 +1,5 @@
import { app } from '../../../scripts/app.js'
// import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
async function getConfig () {
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
const res = await fetch(`${url}/mixlab`, {
method: 'POST'
})
return await res.json()
}
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
const getLocalData = key => {
let data = {}
try {
data = JSON.parse(localStorage.getItem(key)) || {}
} catch (error) {
return {}
}
return data
}
app.registerExtension({
name: 'Mixlab.GPT.ChatGPTOpenAI',
async getCustomWidgets (app) {
return {
KEY (node, inputName, inputData, app) {
// console.log('##inputData', inputData)
const widget = {
type: inputData[0], // the type, CHEESE
name: inputName, // the name, slice
size: [128, 32], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_api_key')
return data[node.id] || 'by Mixlab'
}
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
},
URL (node, inputName, inputData, app) {
// console.log('node', inputName, inputData[0])
const widget = {
type: inputData[0], // the type, CHEESE
name: inputName, // the name, slice
size: [128, 32], // a default size
draw (ctx, node, width, y) {
// a method to draw the widget (ctx is a CanvasRenderingContext2D)
},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_api_url')
return data[node.id] || 'https://api.openai.com/v1'
}
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'ChatGPTOpenAI') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const api_key = this.widgets.filter(w => w.name == 'api_key')[0]
const api_url = this.widgets.filter(w => w.name == 'api_url')[0]
console.log('ChatGPTOpenAI nodeData', this.widgets)
const widget = {
type: 'div',
name: 'chatgptdiv',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, api_key.y, node.size[1])
)
}
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
const inputDiv = (key, placeholder) => {
let div = document.createElement('div')
const ip = document.createElement('input')
ip.type = placeholder === 'Key' ? 'password' : 'text'
ip.className = `${'comfy-multiline-input'} ${placeholder}`
div.style = `display: flex;
align-items: center;
margin: 6px 8px;
margin-top: 0;`
ip.placeholder = placeholder
ip.value = placeholder
ip.style = `margin-left: 24px;
outline: none;
border: none;
padding: 4px;width: 100%;`
const label = document.createElement('label')
label.style = 'font-size: 10px;min-width:32px'
label.innerText = placeholder
div.appendChild(label)
div.appendChild(ip)
ip.addEventListener('change', () => {
let data = getLocalData(key)
data[this.id] = ip.value.trim()
localStorage.setItem(key, JSON.stringify(data))
console.log(this.id, key)
})
return div
}
let inputKey = inputDiv('_mixlab_api_key', 'Key')
let inputUrl = inputDiv('_mixlab_api_url', 'URL')
widget.div.appendChild(inputKey)
widget.div.appendChild(inputUrl)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
inputUrl.remove()
inputKey.remove()
widget.div.remove()
return onRemoved?.()
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
// Fires every time a node is constructed
// You can modify widgets/add handlers/etc here
if (node.type === 'ChatGPTOpenAI') {
let widget = node.widgets.filter(w => w.div)[0]
let apiKey = getLocalData('_mixlab_api_key'),
url = getLocalData('_mixlab_api_url')
let id = node.id
// console.log('ChatGPTOpenAI serialize_widgets', this)
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
widget.div.querySelector('.URL').value =
url[id] || 'https://api.openai.com/v1'
}
}
})
app.registerExtension({
name: 'Mixlab.GPT.ShowTextForGPT',
@@ -214,7 +14,7 @@ app.registerExtension({
for (let i = 0; i < this.widgets.length; i++) {
if (this.widgets[i].name == 'show_text')
this.widgets[i].onRemove?.()
console.log('#ShowTextForGPT', this.widgets[i])
}
this.widgets.length = 2
}
@@ -285,24 +85,5 @@ app.registerExtension({
this.serialize_widgets = true //需要保存参数
}
},
async loadedGraphNode (node, app) {
if (node.type === 'ShowTextForGPT') {
let widget = node.widgets.filter(w => w.name == 'show_text')[0]
// if (widget.value) {
// let [url, prompt] = widget.value
// this[`wavesurfer_${node.id}`] = updateWaveWidgetValue(
// node.widgets,
// node.id,
// url,
// prompt,
// this[`wavesurfer_${node.id}`]
// )
// }
console.log('#loadedGraphNode', node)
}
}
})
+93 -58
View File
@@ -4,6 +4,8 @@ import { api } from '../../../scripts/api.js'
import { $el } from '../../../scripts/ui.js'
import { applyTextReplacements } from '../../../scripts/utils.js'
import { loadExternalScript, get_position_style } from './common.js'
function loadImageToCanvas (base64Image) {
var img = new Image()
var canvas = document.createElement('canvas')
@@ -88,37 +90,40 @@ function getContentTypeFromBase64 (base64Data) {
// const blob = base64ToBlob(base64Data, contentType);
// console.log(blob);
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
// function get_position_style (ctx, widget_width, y, node_height) {
// const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
// /* Create a transform that deals with all the scrolling and zooming */
// const elRect = ctx.canvas.getBoundingClientRect()
// const transform = new DOMMatrix()
// .scaleSelf(
// elRect.width / ctx.canvas.width,
// elRect.height / ctx.canvas.height
// )
// .multiplySelf(ctx.getTransform())
// .translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
// return {
// transformOrigin: '0 0',
// transform: transform,
// left:
// document.querySelector('.comfy-menu').style.display === 'none'
// ? `60px`
// : `0`,
// top: `0`,
// cursor: 'pointer',
// position: 'absolute',
// maxWidth: `${widget_width - MARGIN * 2}px`,
// // maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
// width: `${widget_width - MARGIN * 2}px`,
// // height: `${node_height * 0.3 - MARGIN * 2}px`,
// // background: '#EEEEEE',
// display: 'flex',
// flexDirection: 'column',
// // alignItems: 'center',
// justifyContent: 'space-around'
// }
// }
const getLocalData = key => {
let data = {}
@@ -344,7 +349,7 @@ app.registerExtension({
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 44, node.size[1])
get_position_style(ctx, widget_width, 44, node.size[1], 60)
)
}
}
@@ -530,7 +535,7 @@ app.registerExtension({
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1])
get_position_style(ctx, widget_width, y, node.size[1], 36)
)
}
}
@@ -675,6 +680,18 @@ const createInputImageForBatch = (base64, widget) => {
return im
}
// 添加新图片
const addBase64ToWidgetForLoadImagesToBatch = (
base64,
imagesWidget,
imagesDiv
) => {
if (!imagesWidget.value.base64) imagesWidget.value.base64 = []
imagesWidget.value.base64.push(base64)
let im = createInputImageForBatch(base64, imagesWidget)
imagesDiv.appendChild(im)
}
app.registerExtension({
name: 'Mixlab.Comfy.LoadImagesToBatch',
async getCustomWidgets (app) {
@@ -705,7 +722,6 @@ app.registerExtension({
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'LoadImagesToBatch') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
@@ -717,10 +733,10 @@ app.registerExtension({
type: 'div',
name: 'image_base64',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 44, node.size[1])
)
Object.assign(this.div.style, {
...get_position_style(ctx, widget_width, y, node.size[1], 72),
top: `${widget_height}px`
})
},
serialize: false
}
@@ -751,13 +767,18 @@ app.registerExtension({
base64 = await loadImageToCanvas(base64)
// console.log(base64)
if (!imagesWidget.value) imagesWidget.value = { base64: [] }
imagesWidget.value.base64.push(base64)
let im = createInputImageForBatch(base64, imagesWidget)
imagesDiv.appendChild(im)
addBase64ToWidgetForLoadImagesToBatch(
base64,
imagesWidget,
imagesDiv
)
}
reader.readAsDataURL(file)
})
// 如果是复制的,有数据 , 这个不生效,取不到数据, 需要在nodeCreated里获取
// console.log('#LoadImagesToBatch', imagesWidget.value?.base64)
const btn = document.createElement('button')
btn.innerText = 'Upload Image'
@@ -829,18 +850,36 @@ app.registerExtension({
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'LoadImagesToBatch') {
// await sleep(0)
let imagesWidget = node.widgets.filter(w => w.name === 'images')[0]
let imagePreview = node.widgets.filter(w => w.name == 'image_base64')[0]
let pre = imagePreview.div.querySelector('.images_preview')
// console.log('#LoadImagesToBatch', imagesWidget.value?.base64)
let imagesDiv = imagePreview.div.querySelector('.images_preview')
imagesDiv.innerHTML = ''
for (const d of imagesWidget.value?.base64 || []) {
let im = createInputImageForBatch(d, imagesWidget)
pre.appendChild(im)
imagesDiv.appendChild(im)
}
}
},
nodeCreated (node, app) {
//数据延迟??
setTimeout(() => {
// console.log('#LoadImagesToBatch', node.type)
if (node.type === 'LoadImagesToBatch') {
let imagesWidget = node.widgets.filter(w => w.name === 'images')[0]
let imagePreview = node.widgets.filter(w => w.name == 'image_base64')[0]
let imagesDiv = imagePreview?.div?.querySelector('.images_preview')
imagesDiv.innerHTML = ''
for (const d of imagesWidget.value?.base64 || []) {
let im = createInputImageForBatch(d, imagesWidget)
imagesDiv.appendChild(im)
}
}
}, 1000)
}
})
@@ -848,9 +887,11 @@ app.registerExtension({
app.registerExtension({
name: 'Mixlab.output.ComparingTwoFrames_',
init () {
loadExternalScript('/mixlab/app/lib/juxtapose.min.js')
$el('link', {
rel: 'stylesheet',
href: '/extensions/comfyui-mixlab-nodes/lib/juxtapose.css',
href: '/mixlab/app/lib/juxtapose.css',
parent: document.head
})
@@ -868,8 +909,8 @@ app.registerExtension({
const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated
? onNodeCreated.apply(this, arguments)
: undefined
? onNodeCreated.apply(this, arguments)
: undefined
this.size = [400, this.size[1]]
console.log('##onNodeCreated', this)
@@ -877,10 +918,10 @@ app.registerExtension({
type: 'div',
name: 'preview',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, 400, 44, node.size[1])
)
let s = get_position_style(ctx, widget_width, 44, node.size[1], 36)
delete s.height
Object.assign(this.div.style, s)
},
serialize: false
}
@@ -891,20 +932,15 @@ app.registerExtension({
this.addCustomWidget(widget)
this.serialize_widgets = true //需要保存参数
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
return r
return r
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
@@ -964,7 +1000,7 @@ app.registerExtension({
label: 'After'
}
]
this.size=[this.size[0],300]
this.size = [this.size[0], 300]
}
}
},
@@ -974,7 +1010,6 @@ app.registerExtension({
// node.widgets[0].div.id = 'mix_comparingtowframes_' + node.id
// if (node.widgets_values && node.widgets_values[0]) {
// node.widgets[0].div.innerHTML = ''
// let slider = new juxtapose.JXSlider(
// '#mix_comparingtowframes_' + node.id,
// node.widgets_values,
+4 -1
View File
@@ -81,7 +81,10 @@ function get_position_style (ctx, widget_width, y, node_height) {
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
left:
document.querySelector('.comfy-menu').style.display === 'none'
? `60px`
: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
+19 -91
View File
@@ -3,31 +3,17 @@ import { app } from '../../../scripts/app.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
let api_host = `${window.location.hostname}:${window.location.port}`
let api_base = ''
let url = `${window.location.protocol}//${api_host}${api_base}`
import {
getQueue,
interrupt,
get_position_style,
base64Df,
getUrl,
createImage,
sleep
} from './common.js'
async function getQueue () {
try {
const res = await fetch(`${url}/queue`)
const data = await res.json()
// console.log(data.queue_running,data.queue_pending)
return {
// Running action uses a different endpoint for cancelling
Running: data.queue_running.length,
Pending: data.queue_pending.length
}
} catch (error) {
console.error(error)
return { Running: 0, Pending: 0 }
}
}
async function interrupt () {
const resp = await fetch(`${url}/interrupt`, {
method: 'POST'
})
}
// let url = getUrl()
async function clipboardWriteImage (win, url) {
const canvas = document.createElement('canvas')
@@ -208,22 +194,6 @@ async function shareScreen (
}
}
async function sleep (t = 200) {
return new Promise((res, rej) => {
setTimeout(() => {
res(true)
}, t)
})
}
function createImage (url) {
let im = new Image()
return new Promise((res, rej) => {
im.onload = () => res(im)
im.src = url
})
}
async function compareImages (threshold, previousImage, currentImage) {
// 将 base64 转换为 Image 对象
var previousImg = await createImage(previousImage)
@@ -458,44 +428,6 @@ async function requestCamera () {
return false
}
/*
A method that returns the required style for the html
*/
function get_position_style (ctx, widget_width, y, node_height, top) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `${top}px`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
app.registerExtension({
name: 'Mixlab.image.ScreenShareNode',
async getCustomWidgets (app) {
@@ -593,17 +525,12 @@ app.registerExtension({
type: 'HTML', // whatever
name: 'sreen_share', // whatever
draw (ctx, node, widget_width, y, widget_height) {
// console.log('ScreenSHare', y, widget_height)
// console.log('ScreenSHare', node)
Object.assign(
this.card.style,
get_position_style(
ctx,
widget_width,
widget_height * 5,
node.size[1],
40
)
get_position_style(ctx, widget_width, y, node.size[1], 40)
)
}
}
@@ -1043,12 +970,13 @@ async function setArea (src) {
div.innerHTML = `
<div id='ml_overlay' style='position: absolute;top:0;background: #251f1fc4;
height: 100vh;
z-index:999999;
z-index:99999999999999;
width: 100%;'>
<img id='ml_video' style='position: absolute;
height: ${displayHeight}px;user-select: none;
-webkit-user-drag: none;
outline: 2px solid #eaeaea;
left: 0;
box-shadow: 8px 9px 17px #575757;' />
<div id='ml_selection' style='position: absolute;
border: 2px dashed red;
@@ -1219,10 +1147,10 @@ app.registerExtension({
type: 'video',
name: 'FloatingVideo',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.card.style,
get_position_style(ctx, widget_width, y, node.size[1], 0)
)
Object.assign(this.card.style, {
...get_position_style(ctx, widget_width, y, node.size[1], 40),
top: `${widget_height}px`
})
}
}
+195
View File
@@ -0,0 +1,195 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { $el } from '../../../scripts/ui.js'
import { get_position_style } from './common.js'
function base64ToBlobFromURL (base64URL, contentType) {
return fetch(base64URL).then(response => response.blob())
}
async function uploadImage (blob, fileType = '.svg', filename) {
// const blob = await (await fetch(src)).blob();
const body = new FormData()
body.append(
'image',
new File([blob], (filename || new Date().getTime()) + fileType)
)
const resp = await api.fetchApi('/upload/image', {
method: 'POST',
body
})
// console.log(resp)
let data = await resp.json()
let { name, subfolder } = data
// let src = api.apiURL(
// `/view?filename=${encodeURIComponent(
// name
// )}&type=input&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
// )
return data
}
// 上传得到url
async function uploadBase64ToFile (base64) {
let bg_blob = await base64ToBlobFromURL(base64)
let url = await uploadImage(bg_blob, '.png')
return url
}
const p5InputNode = {
name: 'Mixlab.Comfy.P5Input',
async getCustomWidgets (app) {
return {
IMAGEBASE64 (node, inputName, inputData, app) {
const widget = {
value: {
images: []
}, // 不能[x,x,x]
type: inputData[0], // the type
name: inputName, // the name, slice
size: [320, 120], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128, 32] // a method to compute the current size of the widget
}
}
node.addCustomWidget(widget)
return widget
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'P5Input') {
// console.log('P5Input')
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const widget = {
type: 'div',
name: 'image_base64',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(
ctx,
widget_width - 24,
44,
node.size[1] * 2.8,
44
)
)
},
serialize: false
}
widget.div = $el('div', {})
widget.div.style = `margin:12px;width:400px;height:480px;background:white`
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
// document.addEventListener('wheel', handleMouseWheel)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
// window.removeEventListener('message', ms)
return onRemoved?.()
}
// 节点的大小控制
this.setSize([480, 560])
app.canvas.draw(true, true)
const onResize = this.onResize
this.onResize = () => {
// 设置最小尺寸
if (
Math.max(this.size[0], 480) != this.size[0] &&
Math.max(this.size[1], 560) != this.size[1]
) {
this.setSize([
Math.max(this.size[0], 480),
Math.max(this.size[1], 560)
])
}
return onResize?.apply(this, arguments)
}
this.serialize_widgets = true //需要保存参数
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
// console.log('##onExecuted', this, message._info)
// app.graph.getNodeById(8).widgets[1].div.querySelector('iframe').contentWindow.postMessage('Hello from parent', '*');
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'P5Input') {
}
},
nodeCreated (node, app) {
//数据延迟??
setTimeout(() => {
let widget = node.widgets?.filter(w => w.name == 'image_base64')[0]
let framesWidget = node.widgets?.filter(w => w.name == 'frames')[0]
if (node.type === 'P5Input' && widget) {
console.log('#nodeCreated P5Input')
if (framesWidget && !framesWidget.value)
framesWidget.value = { images: [] }
framesWidget.value._seed = Math.random()
let nodeId = node.id
//延迟才能获得this.id
widget.div.innerHTML = `<iframe src="mixlab/app/p5_export/p5.html?id=${nodeId}"
style="border:0;width:100%;height:100%;"
></iframe>`
// 监听来自iframe的消息
const ms = async event => {
const data = event.data
console.log('#P5 Input #', data)
if (
data.from === 'p5.widget' &&
data.status === 'save' &&
data.frames &&
data.frames.length >= 0 &&
data.nodeId == nodeId &&
data.id != framesWidget.value.id
) {
const frames = data.frames
//workflow会存储到local,会卡死
framesWidget.value.images = []
for (const f of frames) {
let file = await uploadBase64ToFile(f)
framesWidget.value.images.push(file)
}
// framesWidget.value.base64 = frames
// framesWidget.value._seed = Math.random()
node.title = 'P5 Input #' + frames.length
framesWidget.value.id = data.id
}
}
window.addEventListener('message', ms)
}
}, 1000)
}
}
app.registerExtension(p5InputNode)

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