Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c5392aa237 | ||
|
|
16a2b55fa1 | ||
|
|
b423b09ff3 | ||
|
|
e66add88cb | ||
|
|
32b22c39cb | ||
|
|
259baac177 | ||
|
|
67ef8c13a8 | ||
|
|
b2bb1876de | ||
|
|
cda4e626e7 | ||
|
|
d835aff0cb | ||
|
|
21e1967c5e | ||
|
|
c9b5baf4d9 | ||
|
|
67c974c96e | ||
|
|
b46ccb03c9 | ||
|
|
0ecf98e08b | ||
|
|
f024034724 | ||
|
|
327a21f009 | ||
|
|
cfc51532b8 | ||
|
|
00988f92e4 | ||
|
|
868c6085a8 | ||
|
|
a47a56bda0 | ||
|
|
3667b42b2f | ||
|
|
7d142d7d62 | ||
|
|
24863e2ed3 | ||
|
|
fe8b526bbb | ||
|
|
6298be393a | ||
|
|
3a7853f9cc | ||
|
|
4a9413c83d | ||
|
|
21b04d62ae | ||
|
|
96929b6d7c | ||
|
|
07712d80a5 | ||
|
|
10c9eff16f | ||
|
|
edd7af986d | ||
|
|
1dc31927e3 | ||
|
|
36ef7d25ef | ||
|
|
b766b8b65d | ||
|
|
6579ff20b4 | ||
|
|
2fbee59c3e | ||
|
|
d3aaa19148 | ||
|
|
e32a3675fc | ||
|
|
b72e7dda08 | ||
|
|
0f77f28a95 | ||
|
|
289f83675b | ||
|
|
36633b4c72 | ||
|
|
4f45457811 | ||
|
|
c39890cd64 | ||
|
|
90f1e49263 | ||
|
|
9a1cf205db | ||
|
|
45eacb6a50 | ||
|
|
6cb2b57463 | ||
|
|
59f654fa39 | ||
|
|
be8ccc1dc4 | ||
|
|
228e5d9183 | ||
|
|
8afe6d0383 | ||
|
|
5f7190b08f | ||
|
|
a70a9b4bb1 | ||
|
|
b796e66890 | ||
|
|
f1a663779a | ||
|
|
b0aa972326 | ||
|
|
ef927a7ed1 | ||
|
|
aa8fc59051 | ||
|
|
ce62204392 | ||
|
|
837f28142d | ||
|
|
60c79c991d | ||
|
|
078aaeb679 | ||
|
|
d9edbd535e | ||
|
|
b4a61b21c3 | ||
|
|
bdc4193ffe | ||
|
|
74fdd6e396 | ||
|
|
b2479ebff2 | ||
|
|
ce2162c764 | ||
|
|
16ffd63c80 | ||
|
|
8faf68348d | ||
|
|
02dbc72856 | ||
|
|
da4dcf92dc | ||
|
|
49b750abcc | ||
|
|
4bb4122628 | ||
|
|
cee54f336e | ||
|
|
e95b3813cc | ||
|
|
6815cfb05e | ||
|
|
b6acbbce35 | ||
|
|
399e74877d | ||
|
|
61083e91a6 | ||
|
|
67b4ec3178 | ||
|
|
0fcb725a7a | ||
|
|
0dbdcdfdc7 | ||
|
|
e426d77353 | ||
|
|
bd15e29f17 | ||
|
|
b323d29567 | ||
|
|
0e54af3356 | ||
|
|
97f12f3bed | ||
|
|
5612047b97 | ||
|
|
2d147a3ae1 | ||
|
|
1a93c0f8e8 | ||
|
|
0a2b64881a | ||
|
|
42e7fe4d93 | ||
|
|
7ada28258c | ||
|
|
bd312afd00 | ||
|
|
d94a8af35b | ||
|
|
078fd10147 | ||
|
|
824e25d77c | ||
|
|
899b887e47 | ||
|
|
e58981d8a3 | ||
|
|
fc41d977a5 | ||
|
|
ab6210e667 | ||
|
|
f41805f053 | ||
|
|
baa809fcd6 | ||
|
|
a38d15e495 | ||
|
|
e97641372a | ||
|
|
9aecc2cb08 | ||
|
|
697667945e | ||
|
|
d908024577 | ||
|
|
a5a656d958 | ||
|
|
ddc3cf05dd | ||
|
|
66ad4b0abd | ||
|
|
a66023adc6 | ||
|
|
7277844128 | ||
|
|
6ef82b1d56 | ||
|
|
8ded4829f3 | ||
|
|
c4b6acb916 | ||
|
|
9beb81c303 | ||
|
|
f8dd4c6efa | ||
|
|
6ce5aa6a3a | ||
|
|
d2efa8a90a | ||
|
|
c6374063e9 | ||
|
|
0846013378 | ||
|
|
c141ba405f | ||
|
|
0320f13a9f | ||
|
|
8adc34be4d | ||
|
|
cb6810d3c1 | ||
|
|
d384f64abf | ||
|
|
ef7035f8ee | ||
|
|
bfcadde5c3 | ||
|
|
496ff41782 | ||
|
|
46f0be5484 | ||
|
|
f3db0131c1 | ||
|
|
83a8d47f51 | ||
|
|
fee0222910 | ||
|
|
1ed7b5511f | ||
|
|
b8f7c31537 | ||
|
|
164791c257 | ||
|
|
8f5e599928 | ||
|
|
7a7aaeb84d | ||
|
|
e2136ab2fc | ||
|
|
c75cb21946 | ||
|
|
bf95218c91 | ||
|
|
8cb4507a5f | ||
|
|
555890d1ba | ||
|
|
e4f54e83b6 | ||
|
|
692c4a709e | ||
|
|
cbd1961459 | ||
|
|
2e31a33ebf | ||
|
|
d16c6137d2 | ||
|
|
0416ab79ec | ||
|
|
fc9a1c62b9 | ||
|
|
5d4567b134 | ||
|
|
ae4a17d271 | ||
|
|
d110a08889 | ||
|
|
e0157293cb | ||
|
|
0d985b3b65 | ||
|
|
a65ade9fda | ||
|
|
874d6c8cb1 | ||
|
|
f70ba2afa3 | ||
|
|
e9f821e578 | ||
|
|
8e488d4b1d | ||
|
|
77201a457d | ||
|
|
076e3b1178 | ||
|
|
6b13fa64dc | ||
|
|
846671a890 | ||
|
|
05b3088b75 | ||
|
|
fe57286959 | ||
|
|
03645bbb33 | ||
|
|
93dba9a399 | ||
|
|
5627ea8073 | ||
|
|
7ba679c9ce | ||
|
|
c7a450e6ce | ||
|
|
beda5156bf | ||
|
|
76a9da7163 | ||
|
|
edd0303f59 | ||
|
|
be6f47a333 | ||
|
|
4cd6a072ca | ||
|
|
743a82efe9 | ||
|
|
9589f28ef7 | ||
|
|
35492c5671 | ||
|
|
db1e695bf3 | ||
|
|
ecc4aec43b | ||
|
|
fc063c2205 | ||
|
|
4d60ce138a | ||
|
|
2afd24f6e4 | ||
|
|
437acd023a | ||
|
|
b00523ae14 | ||
|
|
4405a74993 | ||
|
|
cb16090868 | ||
|
|
396e510dce | ||
|
|
3b9790b969 | ||
|
|
a35d07a7ac | ||
|
|
6d004c61fc | ||
|
|
ffdd06da1b | ||
|
|
f03f34cacb | ||
|
|
0c86ea849e | ||
|
|
0efa4c38c0 | ||
|
|
6092ab7793 | ||
|
|
929def87eb | ||
|
|
be074ccff7 | ||
|
|
3445199393 | ||
|
|
216c7e152e | ||
|
|
cc8bc10690 | ||
|
|
69b4218d60 | ||
|
|
1dd18dc4f8 | ||
|
|
4ccbd999d9 | ||
|
|
fa8d404964 | ||
|
|
30086957c9 | ||
|
|
0e57c620c9 | ||
|
|
3ce1c59a2d | ||
|
|
3337e20b9e | ||
|
|
e816b3626e | ||
|
|
3e0cb0f17a | ||
|
|
41bc606217 | ||
|
|
5a5f4ca49a | ||
|
|
c3a8437cd1 | ||
|
|
8d8a1a392d | ||
|
|
5f93fb5e55 | ||
|
|
d05050d7d8 | ||
|
|
8e9744100d | ||
|
|
1e4e7e287d | ||
|
|
e8f0c73f08 | ||
|
|
e923e28f8d | ||
|
|
5cc75bfa7c | ||
|
|
d6701769b8 | ||
|
|
0ddc67bdab | ||
|
|
38b62b7a68 | ||
|
|
7e726000c7 | ||
|
|
d8dfb292ec | ||
|
|
826975241d | ||
|
|
743637ceaf | ||
|
|
e350c7e31e | ||
|
|
66b1e0ab9f | ||
|
|
7b0374d110 | ||
|
|
e86ef8cbb0 | ||
|
|
8c901c54bc | ||
|
|
408d85691e | ||
|
|
c66cd6901b | ||
|
|
aeadbc4f6d | ||
|
|
224136890e | ||
|
|
3669a1e86d | ||
|
|
d588b5b327 | ||
|
|
b705679098 | ||
|
|
f71a0b0da5 | ||
|
|
ebc2c76b6b | ||
|
|
2e3fff278e | ||
|
|
1f4bc5e089 | ||
|
|
52c38b10dd | ||
|
|
7047aa5456 | ||
|
|
33fe4019f7 | ||
|
|
80b9d97690 | ||
|
|
3c3c92723f | ||
|
|
037bd87006 | ||
|
|
f688310d28 | ||
|
|
c4d65e7a45 | ||
|
|
6f208b710d | ||
|
|
b599faaf85 | ||
|
|
6d991d20dc | ||
|
|
c87e0296f6 | ||
|
|
16cdb4c5b4 | ||
|
|
7631b8924d | ||
|
|
785d307ff3 | ||
|
|
8c713ff35e | ||
|
|
7d80493bef | ||
|
|
bff2760c3d | ||
|
|
a0f8848367 | ||
|
|
5b1cbcd8d5 | ||
|
|
05857a92d5 | ||
|
|
6bdc811286 | ||
|
|
469d50a5b8 | ||
|
|
ef86904bfb | ||
|
|
d4181ea67c | ||
|
|
1c6d17309f | ||
|
|
db293ec41d | ||
|
|
db8d468f29 | ||
|
|
d7d7af7265 | ||
|
|
bd763cadc1 | ||
|
|
22799fc549 | ||
|
|
0f231d1271 | ||
|
|
8c0c911020 | ||
|
|
6cb9df700b | ||
|
|
4fdda537b9 | ||
|
|
cb6f32465a | ||
|
|
c235e36cb4 | ||
|
|
aaca440a94 | ||
|
|
1f57950a29 | ||
|
|
3d8855ec72 | ||
|
|
ac9231d9f3 | ||
|
|
37c0a56d89 | ||
|
|
010915dac4 | ||
|
|
83a8b3b970 | ||
|
|
f130202aa1 | ||
|
|
c8f6800bcd | ||
|
|
078f9f5dd4 | ||
|
|
c0de178c7d | ||
|
|
1c767b538d | ||
|
|
51aab44b5d | ||
|
|
b8a0d4a67b | ||
|
|
cd0dcfbb8c | ||
|
|
fcc9e30eae | ||
|
|
5f66218a43 | ||
|
|
61ef4f9a0f | ||
|
|
f0a8734b42 | ||
|
|
4fe95ef4ec | ||
|
|
2a5148845b | ||
|
|
8fa562caaf | ||
|
|
cab5620cd5 | ||
|
|
be38d36677 | ||
|
|
69236fca89 | ||
|
|
4a4f376bfd | ||
|
|
fd9718fe24 | ||
|
|
26a6e11212 | ||
|
|
de1a669f6e | ||
|
|
e482c9e5c4 | ||
|
|
4e96a77a41 | ||
|
|
3346290e5c | ||
|
|
dcac593efe | ||
|
|
7248d0de02 | ||
|
|
8ad3ce632c | ||
|
|
9398b02562 | ||
|
|
164e4da99d | ||
|
|
0025ea6119 | ||
|
|
eef53a5165 | ||
|
|
9ea066d948 | ||
|
|
4bd900c4a1 | ||
|
|
736cd2bebd | ||
|
|
cd658c2a60 | ||
|
|
9730658f21 | ||
|
|
3fa107acb1 | ||
|
|
7da22179a0 | ||
|
|
fc7b71ee78 | ||
|
|
b79573bf1f | ||
|
|
a01db6f7c0 | ||
|
|
117d58c58e | ||
|
|
ff6626ed89 | ||
|
|
10c798a440 | ||
|
|
9097d87819 | ||
|
|
d901f503d1 | ||
|
|
65f2b6ce6f | ||
|
|
d949fe8bf1 | ||
|
|
15cfb48550 | ||
|
|
447dc6d4c3 | ||
|
|
c92e43b920 | ||
|
|
be6c32b0e0 | ||
|
|
3d2062e810 | ||
|
|
dd816e95cd | ||
|
|
6d1b51890d | ||
|
|
11f03ec99a | ||
|
|
bd192f43e7 | ||
|
|
36e4b11983 | ||
|
|
97397ba8c2 | ||
|
|
052eee4111 | ||
|
|
e319496044 | ||
|
|
b83b63c362 | ||
|
|
4d6b1675bb | ||
|
|
13110fab39 | ||
|
|
74ea509848 | ||
|
|
192bff9d2c | ||
|
|
44ed8812dc | ||
|
|
200696ba21 | ||
|
|
51cf3b0c04 | ||
|
|
6a4831c83b | ||
|
|
c5e7ed95a3 | ||
|
|
42a97fa4d9 | ||
|
|
45240d0012 | ||
|
|
8ed085febd | ||
|
|
37803ea61b | ||
|
|
acd416952c | ||
|
|
6ec46cbc44 | ||
|
|
a9d971e476 | ||
|
|
9fe064675d | ||
|
|
c84fa467d0 | ||
|
|
3d7a55f6d3 | ||
|
|
d49baa1540 | ||
|
|
5c686af842 | ||
|
|
22425b5bc6 | ||
|
|
b6d9b338d2 | ||
|
|
a191a13751 | ||
|
|
3eccdbcc9b | ||
|
|
50063903f9 | ||
|
|
95a1b70533 | ||
|
|
c3679ac90b | ||
|
|
29e48eb6a2 | ||
|
|
71d02e9651 | ||
|
|
f5193f3eec | ||
|
|
480c4d6919 | ||
|
|
1c5e030540 | ||
|
|
27e83a5908 | ||
|
|
d3cbf8fa8d | ||
|
|
74b1f8129b | ||
|
|
56ed513cfd | ||
|
|
5f412371c4 | ||
|
|
c6b0b67585 | ||
|
|
16d18e681a | ||
|
|
1fe99f33b2 | ||
|
|
a2ece25ac0 | ||
|
|
4865f4d148 | ||
|
|
41e88824cf | ||
|
|
0961ab138e | ||
|
|
fe8271a12f | ||
|
|
f3866ede89 | ||
|
|
d938adf3cc | ||
|
|
9908cff64b | ||
|
|
1b9b0bb4e6 | ||
|
|
03acd9bea5 | ||
|
|
4c42949023 | ||
|
|
2e4d9836e5 | ||
|
|
43c6b58354 | ||
|
|
df637e8196 | ||
|
|
8e78f9786c | ||
|
|
f4130f06ed | ||
|
|
b766e714a4 | ||
|
|
1100a90be3 | ||
|
|
33aaf80c82 | ||
|
|
b5861dbc24 | ||
|
|
93731416fc | ||
|
|
a305e736ca | ||
|
|
a046ebbafb | ||
|
|
af65e96723 | ||
|
|
c9eb0ab5f0 | ||
|
|
9802e841a8 | ||
|
|
b896df8d54 | ||
|
|
720b8c237b | ||
|
|
371f9f813f | ||
|
|
33e229c41c | ||
|
|
2c33c0d801 | ||
|
|
d64fee5954 | ||
|
|
0217678c8c | ||
|
|
2959a9c31f | ||
|
|
1928a18992 | ||
|
|
a168171009 | ||
|
|
6f767f9700 | ||
|
|
e5459f63fd | ||
|
|
d0ab85a8c4 | ||
|
|
b7b86fe8c4 | ||
|
|
c3bcf6907a | ||
|
|
330fed867b | ||
|
|
45bbc31dc1 | ||
|
|
9dc45239c4 | ||
|
|
b6cfb30908 | ||
|
|
8efd94cc76 | ||
|
|
5d864e7ea2 | ||
|
|
7a1c91a2d5 | ||
|
|
1339c8a1f4 | ||
|
|
27025579d6 | ||
|
|
473d70f818 | ||
|
|
907a5d8d7e | ||
|
|
28c349fd7e | ||
|
|
df9545015f | ||
|
|
07b0b94e46 | ||
|
|
92f85b6d9c | ||
|
|
6919eadb21 | ||
|
|
be76280fd2 | ||
|
|
6b673bdd44 | ||
|
|
c6591db45e | ||
|
|
5813ebaa1c | ||
|
|
743e752b0f | ||
|
|
707cd28cb7 | ||
|
|
ecdb687c25 | ||
|
|
ba68d2dd45 | ||
|
|
d8259de52f | ||
|
|
07cec3566b | ||
|
|
dd8c531889 |
@@ -0,0 +1,25 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'shadowcz007' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
@@ -3,4 +3,6 @@ https/
|
||||
nodes/config.json
|
||||
workflow/my_workflow.json
|
||||
workflow/my_workflow_app.json
|
||||
app/*
|
||||
workflow/prompt_result.json
|
||||
app/*
|
||||
workflow/prompt_result.json
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2024 shadow
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -1,28 +1,94 @@
|
||||
> 适配了最新版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
|
||||
|
||||

|
||||
|
||||
##### `最新`:
|
||||
|
||||
- 新增[fal.ai](https://fal.ai/dashboard)的视频生成:Kling、RunwayGen3、LumaDreamMachine,[工作流下载](./workflow/video-all-in-one-test-workflow.json)
|
||||
|
||||
- 新增 SimulateDevDesignDiscussions,需要安装[swarm](https://github.com/openai/swarm)和[Comfyui-ChatTTS](https://github.com/shadowcz007/Comfyui-ChatTTS),[工作流下载](./workflow/swarm制作的播客节点workflow.json)
|
||||
|
||||
- 新增 SenseVoice
|
||||
|
||||
- [新增JS-SDK,方便直接在前端项目中使用comfyui](https://github.com/shadowcz007/comfyui-js-sdk)
|
||||
|
||||
- 新增API调用图像生成节点 TextToImage Siliconflow,可以直接调用Siliconflow提供的flux生成图像
|
||||
|
||||
- [增加 Her 的DEMO页面,和数字人对话](https://github.com/shadowcz007/ComfyUI-Backend-MixlabNodes/blob/main/workflow/her_demo_workflow.json)
|
||||
|
||||
- 右键菜单支持 text-to-text,方便对 prompt 词补全,支持云LLM或者是本地LLM。
|
||||
|
||||
- 增加 MiniCPM-V 2.6 int4
|
||||
|
||||
This is the int4 quantized version of MiniCPM-V 2.6.
|
||||
Running with int4 version would use lower GPU memory (about 7GB).
|
||||
|
||||
- 移动端适配、修改 app 模式的 Mask 编辑器
|
||||
|
||||
- 增加 p5.js 作为输入节点
|
||||
[workflow](./workflow/p5workflow.json)
|
||||
[workflow2](./workflow/p5-video-workflow.json)
|
||||
|
||||
- App 模式增加 batch prompt,批量提示词,可以把动态提示词批量组成后运行
|
||||
|
||||

|
||||
|
||||
- 增加 API Key Input 节点,用于管理 LLM 的 Key,同时优化 LLM 相关节点,为后续 agent 模式做准备
|
||||
|
||||
- 增加 SiliconflowLLM,可以使用由 Siliconflow 提供的免费 LLM
|
||||
|
||||
<!-- - ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。模型下载后,放置到 `models/llamafile/` -->
|
||||
|
||||
<!--
|
||||
强烈推荐:
|
||||
[Phi-3-mini-4k-instruct-function-calling-GGUF](https://huggingface.co/nold/Phi-3-mini-4k-instruct-function-calling-GGUF)
|
||||
|
||||
[Phi-3-mini-4k-instruct-GGUF](https://huggingface.co/lmstudio-community/Phi-3-mini-4k-instruct-GGUF/tree/main),备选:[llama3_if_ai_sdpromptmkr_q2k](https://hf-mirror.com/impactframes/llama3_if_ai_sdpromptmkr_q2k/tree/main)
|
||||
|
||||
- 右键菜单支持 image-to-text,使用多模态模型,多模态使用 [llava-phi-3-mini-gguf](https://huggingface.co/xtuner/llava-phi-3-mini-gguf/tree/main),注意需要把llava-phi-3-mini-mmproj-f16.gguf也下载
|
||||
|
||||

|
||||
 -->
|
||||
|
||||
#### `相关插件推荐`
|
||||
|
||||
[comfyui-liveportrait](https://github.com/shadowcz007/comfyui-liveportrait)
|
||||
|
||||
[Comfyui-ChatTTS](https://github.com/shadowcz007/Comfyui-ChatTTS)
|
||||
|
||||
[comfyui-sound-lab](https://github.com/shadowcz007/comfyui-sound-lab)
|
||||
|
||||
[comfyui-Image-reward](https://github.com/shadowcz007/comfyui-Image-reward)
|
||||
|
||||
[comfyui-ultralytics-yolo](https://github.com/shadowcz007/comfyui-ultralytics-yolo)
|
||||
|
||||
[comfyui-moondream](https://github.com/shadowcz007/comfyui-moondream)
|
||||
|
||||
[comfyui-CLIPSeg](https://github.com/shadowcz007/comfyui-CLIPSeg)
|
||||
<!-- [comfyui-CLIPSeg](https://github.com/shadowcz007/comfyui-CLIPSeg) -->
|
||||
|
||||
## 🚀🚗🚚🏃 Workflow-to-APP
|
||||
|
||||
## 🚀🚗🚚🏃 Workflow-to-APP
|
||||
- 新增AppInfo节点,可以通过简单的配置,把workflow转变为一个Web APP。
|
||||
- 支持多个web app 切换
|
||||
- 发布为app的workflow,可以在右键里再次编辑了
|
||||
- web app可以设置分类,在comfyui右键菜单可以编辑更新web app
|
||||
- 新增 AppInfo 节点,可以通过简单的配置,把 workflow 转变为一个 Web APP。
|
||||
- 支持多个 web app 切换
|
||||
- 发布为 app 的 workflow,可以在右键里再次编辑了
|
||||
- web app 可以设置分类,在 comfyui 右键菜单可以编辑更新 web app
|
||||
- 支持动态提示
|
||||
- 支持把输出显示到 comfyui 背景(TouchDesigner 风格)
|
||||
- 如果转为 web app 打开是空白的,注意检查下插件目录的名字需要是:comfyui-mixlab-nodes(如果是 zip 包下载会多了个-main 的后缀,需要去掉)
|
||||
|
||||

|
||||
|
||||
- Support multiple web app switching.
|
||||
- Add the AppInfo node, which allows you to transform the workflow into a web app by simple configuration.
|
||||
- The workflow, which is now released as an app, can also be edited again by right-clicking.
|
||||
- The web app can be configured with categories, and the web app can be edited and updated in the right-click menu of ComfyUI.
|
||||
|
||||
|
||||

|
||||
|
||||

|
||||
@@ -30,59 +96,96 @@
|
||||

|
||||
|
||||
Example:
|
||||
|
||||
- workflow
|
||||

|
||||
[text-to-image](./workflow/Text-to-Image-app.json)
|
||||

|
||||
[text-to-image](./workflow/Text-to-Image-app.json)
|
||||
|
||||
APP-JSON:
|
||||
|
||||
- [text-to-image](./example/Text-to-Image_3.json)
|
||||
- [image-to-image](./example/Image-to-Image_2.json)
|
||||
- text-to-text
|
||||
|
||||
> 暂时支持 9 种节点作为界面上的输入节点:Load Image、VHS_LoadVideo、CLIPTextEncode、PromptSlide、TextInput_、Color、FloatSlider、IntNumber、CheckpointLoaderSimple、LoraLoader
|
||||
> 暂时支持 9 种节点作为界面上的输入节点:Load Image、VHS*LoadVideo、CLIPTextEncode、PromptSlide、TextInput*、Color、FloatSlider、IntNumber、CheckpointLoaderSimple、LoraLoader
|
||||
|
||||
> 输出节点:PreviewImage 、SaveImage、ShowTextForGPT、VHS_VideoCombine、PromptImage
|
||||
|
||||
> seed统一输入控件,支持:SamplerCustom、KSampler
|
||||
> seed 统一输入控件,支持:SamplerCustom、KSampler
|
||||
|
||||
> 配套[ps插件](https://github.com/shadowcz007/comfyui-ps-plugin)
|
||||
> 配套[ps 插件](https://github.com/shadowcz007/comfyui-ps-plugin)
|
||||
|
||||
> 如果遇到上传图片不成功,请检查下:局域网或者是云服务,请使用https,端口8189这个服务( 感谢 @Damien 反馈问题)
|
||||
> 如果遇到上传图片不成功,请检查下:局域网或者是云服务,请使用 https,端口 8189 这个服务( 感谢 @Damien 反馈问题)
|
||||
|
||||
> If you encounter difficulties in uploading images, please check the following: for local network or cloud services, please use HTTPS and the service on port 8189. (Thanks to @Damien for reporting the issue.)
|
||||
|
||||
## 🏃🚗🚚🚀 Real-time Design
|
||||
|
||||
|
||||
## 🏃🚗🚚🚀 Real-time Design
|
||||
> ScreenShareNode & FloatingVideoNode. Now comfyui supports capturing screen pixel streams from any software and can be used for LCM-Lora integration. Let's get started with implementation and design! 💻🌐
|
||||
|
||||

|
||||
|
||||
https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43e-410a-ab3a-1952b7b4e7da
|
||||
|
||||
|
||||
<!-- [ScreenShareNode](./workflow/2-screeshare.json) -->
|
||||
|
||||
[ScreenShareNode & FloatingVideoNode](./workflow/3-FloatVideo-workflow.json)
|
||||
|
||||
!! Please use the address with HTTPS (https://127.0.0.1).
|
||||
|
||||
|
||||
### SpeechRecognition & SpeechSynthesis
|
||||
|
||||

|
||||
|
||||
[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.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
|
||||
|
||||
[workflow-5](./workflow/5-gpt-workflow.json)
|
||||
[LLM_base_workflow](./workflow/LLM_base_workflow.json)
|
||||
|
||||
- SiliconflowLLM
|
||||
- ChatGPTOpenAI
|
||||
|
||||
<!-- 最新:ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
|
||||
|
||||
Model download,move to :`models/llamafile/`
|
||||
|
||||
强烈推荐:[Phi-3-mini-4k-instruct-GGUF](https://huggingface.co/lmstudio-community/Phi-3-mini-4k-instruct-GGUF/tree/main)
|
||||
|
||||
备选:[llama3_if_ai_sdpromptmkr_q2k](https://hf-mirror.com/impactframes/llama3_if_ai_sdpromptmkr_q2k/tree/main)
|
||||
|
||||
> 如果碰到安装失败,可以尝试手动安装
|
||||
|
||||
```
|
||||
../../../python_embeded/python.exe -s -m pip install llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
|
||||
|
||||
../../../python_embeded/python.exe -s -m pip install llama-cpp-python[server]
|
||||
|
||||
```
|
||||
|
||||
> [Mac](https://llama-cpp-python.readthedocs.io/en/latest/install/macos/)
|
||||
|
||||
```
|
||||
pip uninstall llama-cpp-python -y
|
||||
CMAKE_ARGS="-DLLAMA_METAL=on" pip install -U llama-cpp-python --no-cache-dir
|
||||
pip install 'llama-cpp-python[server]'
|
||||
```
|
||||
|
||||
```
|
||||
pip install llama-cpp-python \
|
||||
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/metal
|
||||
``` -->
|
||||
|
||||
## Prompt
|
||||
|
||||
> PromptSlide
|
||||

|
||||
> 
|
||||
|
||||
<!--  -->
|
||||
|
||||
@@ -96,104 +199,153 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
|
||||
|
||||
> PromptImage & PromptSimplification,Assist in simplifying prompt words, comparing images and prompt word nodes.
|
||||
|
||||
> ChinesePrompt && PromptGenerate,中文prompt节点,直接用中文书写你的prompt
|
||||
> ChinesePrompt && PromptGenerate,中文 prompt 节点,直接用中文书写你的 prompt
|
||||
|
||||

|
||||
|
||||
|
||||
### Layers
|
||||
|
||||
> A new layer class node has been added, allowing you to separate the image into layers. After merging the images, you can input the controlnet for further processing.
|
||||
|
||||
> The composite images node overlays a foreground image onto a background image at specified positions and scales, with optional blending modes and masking capabilities. position : 'overall',"center_center","left_bottom","center_bottom","right_bottom","left_top","center_top","right_top"
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
|
||||
### 3D
|
||||
|
||||

|
||||

|
||||
[workflow](./assets/Image-to-3D_1.json)
|
||||
|
||||

|
||||
[workflow](./workflow/3D-workflow.json)
|
||||
|
||||
### Image
|
||||
|
||||
### LoadImagesFromLocal
|
||||
> Monitor changes to images in a local folder, and trigger real-time execution of workflows, supporting common image formats, especially PSD format, in conjunction with Photoshop.
|
||||
#### LoadImagesToBatch
|
||||
|
||||
> Upload multiple images for batch input into the IP adapter.
|
||||
|
||||
#### LoadImagesFromLocal
|
||||
|
||||
> Monitor changes to images in a local folder, and trigger real-time execution of workflows, supporting common image formats, especially PSD format, in conjunction with Photoshop.
|
||||
|
||||

|
||||
|
||||
[workflow-4](./workflow/4-loadfromlocal-watcher-workflow.json)
|
||||
|
||||
### LoadImagesFromURL
|
||||
#### LoadImagesFromURL
|
||||
|
||||
> Conveniently load images from a fixed address on the internet to ensure that default images in the workflow can be executed.
|
||||
|
||||
#### TextImage
|
||||
|
||||
> [下载字体](https://drxie.github.io/OSFCC/)放到 `custom_nodes/comfyui-mixlab-nodes/assets/fonts`
|
||||
|
||||
#### MiniCPM-VQA Simple
|
||||
|
||||
This is the int4 quantized version of MiniCPM-V 2.6.
|
||||
Running with int4 version would use lower GPU memory (about 7GB).
|
||||
|
||||
[模型](https://huggingface.co/openbmb/MiniCPM-V-2_6-int4)
|
||||
|
||||

|
||||
|
||||
### Style
|
||||
|
||||
> Apply VisualStyle Prompting , Modified from [ComfyUI_VisualStylePrompting](https://github.com/ExponentialML/ComfyUI_VisualStylePrompting)
|
||||
|
||||

|
||||
|
||||
> StyleAligned , Modified from [style_aligned_comfy](https://github.com/brianfitzgerald/style_aligned_comfy)
|
||||
|
||||
### Utils
|
||||
|
||||
## Utils
|
||||
> The Color node provides a color picker for easy color selection, the Font node offers built-in font selection for use with TextImage to generate text images, and the DynamicDelayByText node allows delayed execution based on the length of the input text.
|
||||
|
||||
- [添加了DynamicDelayByText功能,可以根据输入文本的长度进行延迟执行。](./workflow/audio-chatgpt-workflow.json)
|
||||
- [添加了 DynamicDelayByText 功能,可以根据输入文本的长度进行延迟执行。](./workflow/audio-chatgpt-workflow.json)
|
||||
|
||||
- [Added DynamicDelayByText, enabling delayed execution based on input text length.](./workflow/audio-chatgpt-workflow.json)
|
||||
|
||||
- [使用CkptNames 对比不同的模型效果](./workflow/ckpts-image-workflow.json)
|
||||
- [使用 CkptNames 对比不同的模型效果](./workflow/ckpts-image-workflow.json)
|
||||
|
||||
- [CkptNames compare the effects of different models.](./workflow/ckpts-image-workflow.json)
|
||||
|
||||
### Other Nodes
|
||||
|
||||
|
||||
## Other Nodes
|
||||
- 增加 Edit Mask,方便在生成的时候手动绘制 mask [workflow](./workflow/edit-mask-workflow.json)
|
||||
|
||||

|
||||

|
||||
|
||||
[workflow-1](./workflow/1-workflow.json)
|
||||
|
||||
|
||||
|
||||
> TransparentImage
|
||||
|
||||

|
||||
|
||||
|
||||
> FeatheredMask、SmoothMask
|
||||
|
||||
Add edges to an image.
|
||||
|
||||

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

|
||||
|
||||

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

|
||||
|
||||
|
||||
### Models
|
||||
|
||||
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:models/rembg
|
||||
- [Download TripoSR](https://huggingface.co/stabilityai/TripoSR/blob/main/model.ckpt) and place it in `models/triposr`
|
||||
|
||||
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : models/lama
|
||||
- [Download facebook/dino-vitb16](https://huggingface.co/facebook/dino-vitb16/tree/main) and place it in `models/triposr/facebook/dino-vitb16`
|
||||
|
||||
[Download Salesforce/blip-image-captioning-base](https://huggingface.co/Salesforce/blip-image-captioning-base), move to : models/clip_interrogator/Salesforce/blip-image-captioning-base
|
||||
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:`models/rembg`
|
||||
|
||||
[Download succinctly/text2image-prompt-generator](https://huggingface.co/succinctly/text2image-prompt-generator/tree/main),move to:prompt_generator/text2image-prompt-generator
|
||||
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : `models/lama`
|
||||
|
||||
[Download Helsinki-NLP/opus-mt-zh-en](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main),move to:prompt_generator/opus-mt-zh-en
|
||||
[Download Salesforce/blip-image-captioning-base](https://huggingface.co/Salesforce/blip-image-captioning-base), move to :`models/clip_interrogator/Salesforce/blip-image-captioning-base`
|
||||
|
||||
[Download succinctly/text2image-prompt-generator](https://huggingface.co/succinctly/text2image-prompt-generator/tree/main),move to:`models/prompt_generator/text2image-prompt-generator`
|
||||
|
||||
[Download Helsinki-NLP/opus-mt-zh-en](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main),move to:`models/prompt_generator/opus-mt-zh-en`
|
||||
|
||||
## Installation
|
||||
|
||||
@@ -209,40 +361,35 @@ git clone https://github.com/shadowcz007/comfyui-mixlab-nodes.git
|
||||
Install the requirements:
|
||||
|
||||
run directly:
|
||||
|
||||
```
|
||||
cd ComfyUI/custom_nodes/comfyui-mixlab-nodes
|
||||
install.bat
|
||||
```
|
||||
|
||||
or install the requirements using:
|
||||
|
||||
```
|
||||
../../../python_embeded/python.exe -s -m pip install -r requirements.txt
|
||||
```
|
||||
|
||||
If you are using a venv, make sure you have it activated before installation and use:
|
||||
|
||||
```
|
||||
pip3 install -r requirements.txt
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
#### Chinese community
|
||||
访问 [www.mixcomfy.com](https://www.mixcomfy.com),获得更多内测功能,关注微信公众号:Mixlab无界社区
|
||||
|
||||
访问 [www.mixcomfy.com](https://www.mixcomfy.com),获得更多内测功能,关注微信公众号:Mixlab 无界社区
|
||||
|
||||
####
|
||||
|
||||
|
||||
####
|
||||
File / LoadImagesFromPath SaveImageToLocal LoadImagesFromURL
|
||||
|
||||
|
||||
|
||||
|
||||
#### discussions:
|
||||
[discussions](https://github.com/shadowcz007/comfyui-mixlab-nodes/discussions)
|
||||
|
||||
[discussions](https://github.com/shadowcz007/comfyui-mixlab-nodes/discussions)
|
||||
|
||||
<picture>
|
||||
<source
|
||||
@@ -262,4 +409,3 @@ File / LoadImagesFromPath SaveImageToLocal LoadImagesFromURL
|
||||
src="https://api.star-history.com/svg?repos=shadowcz007/comfyui-mixlab-nodes&type=Date"
|
||||
/>
|
||||
</picture>
|
||||
|
||||
|
||||
|
After Width: | Height: | Size: 537 KiB |
|
After Width: | Height: | Size: 340 KiB |
|
After Width: | Height: | Size: 29 KiB |
|
After Width: | Height: | Size: 366 KiB |
|
After Width: | Height: | Size: 135 KiB |
|
After Width: | Height: | Size: 210 KiB |
|
After Width: | Height: | Size: 2.2 MiB |
|
After Width: | Height: | Size: 1.1 MiB |
|
Before Width: | Height: | Size: 35 KiB After Width: | Height: | Size: 94 KiB |
|
After Width: | Height: | Size: 75 KiB |
|
After Width: | Height: | Size: 63 KiB |
|
After Width: | Height: | Size: 965 KiB |
@@ -0,0 +1,10 @@
|
||||
[
|
||||
{
|
||||
"keyword":"Dog",
|
||||
"imgurl":"http://127.0.0.1:8188/view?filename=1709966910233.png&type=input&subfolder=&rand=0.2734446552394221"
|
||||
},
|
||||
{
|
||||
"keyword":"x",
|
||||
"imgurl":"http://127.0.0.1:8188/view?filename=image%20(33).png&type=input&subfolder=pasted&rand=0.6984318219852814"
|
||||
}
|
||||
]
|
||||
@@ -10,6 +10,12 @@ if exist "%python_exec%" (
|
||||
for /f "delims=" %%i in (%requirements_txt%) do (
|
||||
%python_exec% -s -m pip install "%%i" -i https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
)
|
||||
|
||||
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
|
||||
|
||||
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
|
||||
|
||||
|
||||
) else (
|
||||
echo Installing with system Python
|
||||
for /f "delims=" %%i in (%requirements_txt%) do (
|
||||
|
||||
@@ -1,4 +1,100 @@
|
||||
|
||||
import os
|
||||
import folder_paths
|
||||
import torchaudio
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
|
||||
any_type = AnyType("*")
|
||||
|
||||
|
||||
def analyze_audio_data(audio_data):
|
||||
total_duration = 0
|
||||
total_gap_duration = 0
|
||||
emotion_counts = {}
|
||||
audio_types = set()
|
||||
languages = set()
|
||||
|
||||
for i, entry in enumerate(audio_data):
|
||||
# Calculate the duration of each audio segment
|
||||
start_time = entry['start_time']
|
||||
end_time = entry['end_time']
|
||||
duration = end_time - start_time
|
||||
total_duration += duration
|
||||
|
||||
# Count the emotions
|
||||
if "emotion" in entry:
|
||||
emotion = entry['emotion']
|
||||
if emotion in emotion_counts:
|
||||
emotion_counts[emotion] += 1
|
||||
else:
|
||||
emotion_counts[emotion] = 1
|
||||
|
||||
# Collect the audio types
|
||||
if "audio_type" in entry:
|
||||
audio_types.add(entry['audio_type'])
|
||||
|
||||
if "language" in entry:
|
||||
languages.add(entry['language'])
|
||||
|
||||
# Calculate gap duration if not the last entry
|
||||
if i < len(audio_data) - 1:
|
||||
next_start_time = audio_data[i + 1]['start_time']
|
||||
gap_duration = next_start_time - end_time
|
||||
if gap_duration > 0:
|
||||
total_gap_duration += gap_duration
|
||||
|
||||
# Get the most frequent emotion
|
||||
if len(emotion_counts.keys())>0:
|
||||
most_frequent_emotion = max(emotion_counts, key=emotion_counts.get)
|
||||
else:
|
||||
most_frequent_emotion=None
|
||||
|
||||
# Convert audio_types set to list for better readability
|
||||
audio_types = list(audio_types)
|
||||
|
||||
languages=list(languages)
|
||||
|
||||
# Print the results
|
||||
print(f"Total Effective Duration: {total_duration:.2f} seconds")
|
||||
print(f"Total Gap Duration: {total_gap_duration:.2f} seconds")
|
||||
print(f"Emotion Changes: {emotion_counts}")
|
||||
print(f"Most Frequent Emotion: {most_frequent_emotion}")
|
||||
print(f"Audio Types: {audio_types}")
|
||||
|
||||
|
||||
return {
|
||||
"total_duration": total_duration,
|
||||
"total_gap_duration": total_gap_duration,
|
||||
"emotion_changes": emotion_counts,
|
||||
"most_frequent_emotion": most_frequent_emotion,
|
||||
"audio_types": audio_types,
|
||||
"languages":languages
|
||||
}
|
||||
|
||||
|
||||
# 分析音频数据
|
||||
class AnalyzeAudioNone:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"json":(any_type,),},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ("result",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
def run(self,json):
|
||||
result=analyze_audio_data(json)
|
||||
return (result,)
|
||||
|
||||
|
||||
|
||||
@@ -55,46 +151,65 @@ class SpeechSynthesis:
|
||||
return {"ui": {"text": text}, "result": (text,)}
|
||||
|
||||
|
||||
#
|
||||
class GamePal:
|
||||
|
||||
class AudioPlayNode:
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_temp_directory()
|
||||
self.type = "temp"
|
||||
self.prefix_append = ""
|
||||
self.compress_level = 4
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_text": ("STRING",{"multiline": True,"default": ""}),
|
||||
},
|
||||
"optional": {
|
||||
|
||||
"input_num": ("INT",{
|
||||
"default":100,
|
||||
"min": -1, #Minimum value
|
||||
"max": 0xffffffffffffffff, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "slider" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"python_code": ("STRING",{"multiline": True,"default": "result= 1 if 'Mixlab' in input_text else 0"}),
|
||||
}
|
||||
}
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
RETURN_TYPES = ("INT",)
|
||||
return {"required": {
|
||||
"audio": ("AUDIO",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
|
||||
FUNCTION = "run"
|
||||
OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
def run(self, input_text,input_num,python_code):
|
||||
exec(python_code)
|
||||
res=None
|
||||
try:
|
||||
# 可能会引发异常的代码
|
||||
res=result
|
||||
except:
|
||||
# 处理异常的代码
|
||||
print('')
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = ()
|
||||
|
||||
print(res)
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def run(self,audio):
|
||||
|
||||
# print(session_history)
|
||||
return {"ui": {"text": [input_text],"num":[input_num]}, "result": (res,)}
|
||||
# 判断是否是 Tensor 类型
|
||||
is_tensor = not isinstance(audio, dict)
|
||||
# print('#判断是否是 Tensor 类型',is_tensor,audio)
|
||||
if not is_tensor and 'waveform' in audio and 'sample_rate' in audio:
|
||||
# {'waveform': tensor([], size=(1, 1, 0)), 'sample_rate': 44100}
|
||||
is_tensor=True
|
||||
|
||||
if is_tensor and (not 'audio_path' in audio):
|
||||
filename_prefix=""
|
||||
# 保存
|
||||
filename_prefix += self.prefix_append
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
|
||||
results = list()
|
||||
|
||||
filename_with_batch_num = filename.replace("%batch_num%", str(1))
|
||||
file = f"{filename_with_batch_num}_{counter:05}_.wav"
|
||||
|
||||
torchaudio.save(os.path.join(full_output_folder, file), audio['waveform'].squeeze(0), audio["sample_rate"])
|
||||
results.append({
|
||||
"filename": file,
|
||||
"subfolder": subfolder,
|
||||
"type": self.type
|
||||
})
|
||||
|
||||
else:
|
||||
results=[{
|
||||
"filename": audio['filename'],
|
||||
"subfolder":audio['subfolder'],
|
||||
"type": audio['type'],
|
||||
"audio_path":audio['audio_path']
|
||||
}]
|
||||
|
||||
|
||||
# print(audio)
|
||||
return {"ui": {"audio":results}}
|
||||
@@ -1,10 +1,88 @@
|
||||
import openai
|
||||
from swarm import Swarm, Agent
|
||||
|
||||
import time
|
||||
import urllib.error
|
||||
import re,json,os,string,random
|
||||
import folder_paths
|
||||
import hashlib
|
||||
from zhipuai import ZhipuAI
|
||||
import codecs,sys
|
||||
import importlib.util
|
||||
import subprocess
|
||||
import requests
|
||||
from PIL import Image
|
||||
from io import BytesIO
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
python = sys.executable
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
# 从文本中提取json
|
||||
def extract_json_strings(text):
|
||||
json_strings = []
|
||||
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:
|
||||
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):
|
||||
hash_object = hashlib.sha1(string.encode())
|
||||
unique_hash = hash_object.hexdigest()
|
||||
@@ -42,28 +120,125 @@ 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')==True:
|
||||
from zhipuai import ZhipuAI
|
||||
except:
|
||||
print("#install zhipuai error")
|
||||
|
||||
client = ZhipuAI(
|
||||
api_key=key, # 填写您的 APIKey
|
||||
)
|
||||
return client
|
||||
|
||||
|
||||
# 优先使用phi
|
||||
def phi_sort(lst):
|
||||
return sorted(lst, key=lambda x: x.lower().count('phi'), reverse=True)
|
||||
|
||||
def chat(client, model_name,messages ):
|
||||
def get_llama_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('llamafile')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "llamafile")
|
||||
|
||||
# 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
|
||||
|
||||
# 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 llama_cpp_client(file_name):
|
||||
# try:
|
||||
# if is_installed('llama_cpp')==False:
|
||||
# import subprocess
|
||||
|
||||
# # 安装
|
||||
# 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)
|
||||
|
||||
# #检查命令执行结果
|
||||
# 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)
|
||||
|
||||
# else:
|
||||
# print("#install 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)
|
||||
|
||||
# llm = Llama(model_path=mp, chat_format="chatml",n_gpu_layers=-1,n_ctx=512)
|
||||
|
||||
# return llm
|
||||
|
||||
|
||||
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
|
||||
try:
|
||||
response = client.chat.completions.create(
|
||||
model=model_name,
|
||||
messages=messages
|
||||
)
|
||||
if hasattr(client, "chat"):
|
||||
response = client.chat.completions.create(
|
||||
model=model_name,
|
||||
messages=messages,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature
|
||||
)
|
||||
else:
|
||||
# 是llama的
|
||||
response = client.create_chat_completion_openai_v1(
|
||||
messages=messages,
|
||||
# response_format={
|
||||
# "type": "json_object",
|
||||
# },
|
||||
# temperature=0.7,
|
||||
)
|
||||
|
||||
break
|
||||
except openai.AuthenticationError as ex:
|
||||
raise ex
|
||||
@@ -72,7 +247,8 @@ def chat(client, model_name,messages ):
|
||||
raise ex
|
||||
time.sleep(3)
|
||||
continue
|
||||
|
||||
|
||||
# print(response.keys())
|
||||
finish_reason = response.choices[0].finish_reason
|
||||
if finish_reason != "stop":
|
||||
raise RuntimeError("API finished with unexpected reason: " + finish_reason)
|
||||
@@ -86,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()
|
||||
@@ -95,25 +301,60 @@ class ChatGPTNode:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
|
||||
model_list=[
|
||||
"gpt-3.5-turbo",
|
||||
"gpt-3.5-turbo-16k",
|
||||
"gpt-4o",
|
||||
"gpt-4o-2024-05-13",
|
||||
"gpt-4",
|
||||
"gpt-4-0314",
|
||||
"gpt-4-0613",
|
||||
"gpt-3.5-turbo-0301",
|
||||
"gpt-3.5-turbo-0613",
|
||||
"gpt-3.5-turbo-16k-0613",
|
||||
"qwen-turbo",
|
||||
"qwen-plus",
|
||||
"qwen-long",
|
||||
"qwen-max",
|
||||
"qwen-max-longcontext",
|
||||
"glm-4",
|
||||
"glm-3-turbo",
|
||||
"moonshot-v1-8k",
|
||||
"moonshot-v1-32k",
|
||||
"moonshot-v1-128k",
|
||||
"deepseek-chat",
|
||||
"Qwen/Qwen2-7B-Instruct",
|
||||
"THUDM/glm-4-9b-chat",
|
||||
"01-ai/Yi-1.5-9B-Chat-16K",
|
||||
"meta-llama/Meta-Llama-3.1-8B-Instruct"
|
||||
]
|
||||
|
||||
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": (["gpt-3.5-turbo","gpt-35-turbo","gpt-3.5-turbo-16k", "gpt-3.5-turbo-16k-0613", "gpt-4-0613","gpt-4-1106-preview","glm-4"],
|
||||
{"default": "gpt-3.5-turbo"}),
|
||||
|
||||
"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",)
|
||||
@@ -125,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()
|
||||
@@ -143,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)
|
||||
@@ -152,9 +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)
|
||||
else :
|
||||
client = openai_client(api_key,api_url) # 使用 ChatGPT 的接口
|
||||
print('using ChatGPT interface')
|
||||
# print('using ChatGPT interface',api_key,api_url)
|
||||
|
||||
# 把用户的提示添加到会话历史中
|
||||
# 调用API时传递整个会话历史
|
||||
@@ -170,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}]
|
||||
@@ -190,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
|
||||
@@ -209,7 +639,7 @@ class ShowTextForGPT:
|
||||
OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/GPT"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
def run(self, text,output_dir=[""]):
|
||||
|
||||
@@ -293,7 +723,7 @@ class CharacterInText:
|
||||
# OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/GPT"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
def run(self, text,character,start_index):
|
||||
# print(text,character,start_index)
|
||||
@@ -306,8 +736,8 @@ class TextSplitByDelimiter:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True,"dynamicPrompts": False}),
|
||||
"delimiter":(["newline","comma"],),
|
||||
"text": ("STRING", {"multiline": True,"dynamicPrompts": False}),
|
||||
"delimiter":("STRING", {"multiline": False,"default":",","dynamicPrompts": False}),
|
||||
"start_index": ("INT", {
|
||||
"default": 0,
|
||||
"min": 0, #Minimum value
|
||||
@@ -338,15 +768,320 @@ class TextSplitByDelimiter:
|
||||
# OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/GPT"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
def run(self, text,delimiter,start_index,skip_every,max_count):
|
||||
arr=[]
|
||||
if delimiter=='newline':
|
||||
arr = [line for line in text.split('\n') if line.strip()]
|
||||
elif delimiter=='comma':
|
||||
arr = [line for line in text.split(',') if line.strip()]
|
||||
|
||||
|
||||
if delimiter=="":
|
||||
arr=[text.strip()]
|
||||
else:
|
||||
delimiter=codecs.decode(delimiter, 'unicode_escape')
|
||||
arr= [line for line in text.split(delimiter) if line.strip()]
|
||||
|
||||
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),)
|
||||
|
||||
|
||||
@@ -70,14 +70,20 @@ def load_caption_model(model_path,config,t='blip-base'):
|
||||
return (caption_model,caption_processor)
|
||||
|
||||
|
||||
def get_clip_interrogator_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('clip_interrogator')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "clip_interrogator")
|
||||
|
||||
caption_model_path=os.path.join(folder_paths.models_dir, "clip_interrogator/Salesforce/blip-image-captioning-base")
|
||||
|
||||
cache_path=get_clip_interrogator_path()
|
||||
|
||||
caption_model_path=os.path.join(cache_path, "Salesforce","blip-image-captioning-base")
|
||||
if not os.path.exists(caption_model_path):
|
||||
print(f"## clip_interrogator_model not found: {caption_model_path}, pls download from https://huggingface.co/Salesforce/blip-image-captioning-base")
|
||||
caption_model_path='Salesforce/blip-image-captioning-base'
|
||||
|
||||
cache_path=os.path.join(folder_paths.models_dir, "clip_interrogator")
|
||||
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
|
||||
@@ -0,0 +1,332 @@
|
||||
# 修改自 https://github.com/gokayfem/ComfyUI-fal-API/blob/main/nodes/video_node.py
|
||||
# image-to-video all in one
|
||||
|
||||
import os,sys
|
||||
import torch
|
||||
from PIL import Image
|
||||
import tempfile
|
||||
import numpy as np
|
||||
import requests
|
||||
import cv2
|
||||
import subprocess
|
||||
import importlib.util
|
||||
python = sys.executable
|
||||
|
||||
def is_installed(package, package_overwrite=None,auto_install=True):
|
||||
is_has=False
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
is_has=spec is not None
|
||||
except ModuleNotFoundError:
|
||||
pass
|
||||
|
||||
package = package_overwrite or package
|
||||
|
||||
if spec is None:
|
||||
if auto_install==True:
|
||||
print(f"Installing {package}...")
|
||||
# 清华源 -i https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
command = f'"{python}" -m pip install {package}'
|
||||
|
||||
result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=True, env=os.environ)
|
||||
|
||||
is_has=True
|
||||
|
||||
if result.returncode != 0:
|
||||
print(f"Couldn't install\nCommand: {command}\nError code: {result.returncode}")
|
||||
is_has=False
|
||||
else:
|
||||
print(package+'## OK')
|
||||
|
||||
return is_has
|
||||
|
||||
|
||||
try:
|
||||
if is_installed('fal_client','fal-client')==True:
|
||||
from fal_client import submit, upload_file
|
||||
except:
|
||||
print("#install fal-client error")
|
||||
|
||||
|
||||
def upload_image(image):
|
||||
try:
|
||||
# Convert the image tensor to a numpy array
|
||||
if isinstance(image, torch.Tensor):
|
||||
image_np = image.cpu().numpy()
|
||||
else:
|
||||
image_np = np.array(image)
|
||||
|
||||
# Ensure the image is in the correct format (H, W, C)
|
||||
if image_np.ndim == 4:
|
||||
image_np = image_np.squeeze(0) # Remove batch dimension if present
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.stack([image_np] * 3, axis=-1) # Convert grayscale to RGB
|
||||
elif image_np.shape[0] == 3:
|
||||
image_np = np.transpose(image_np, (1, 2, 0)) # Change from (C, H, W) to (H, W, C)
|
||||
|
||||
# Normalize the image data to 0-255 range
|
||||
if image_np.dtype == np.float32 or image_np.dtype == np.float64:
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
|
||||
# Convert to PIL Image
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
# Save the image to a temporary file
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as temp_file:
|
||||
pil_image.save(temp_file, format="PNG")
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
# Upload the temporary file
|
||||
image_url = upload_file(temp_file_path)
|
||||
return image_url
|
||||
except Exception as e:
|
||||
print(f"Error uploading image: {str(e)}")
|
||||
return None
|
||||
finally:
|
||||
# Clean up the temporary file
|
||||
if 'temp_file_path' in locals():
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
|
||||
class VideoGenKlingNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"duration": (["5", "10"], {"default": "5"}),
|
||||
"aspect_ratio": (["16:9", "9:16", "1:1"], {"default": "16:9"}),
|
||||
"mode": (["standard", "pro"], {"default": "standard"}),
|
||||
"fal_key":("STRING", {"forceInput": True,}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def generate_video(self, prompt, duration, aspect_ratio,mode,fal_key, image=None):
|
||||
arguments = {
|
||||
"prompt": prompt,
|
||||
"duration": duration,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
}
|
||||
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
|
||||
api_url="fal-ai/kling-video/v1/"+mode
|
||||
|
||||
try:
|
||||
if image is not None:
|
||||
image_url = upload_image(image)
|
||||
if image_url:
|
||||
arguments["image_url"] = image_url
|
||||
handler = submit(api_url+"/image-to-video", arguments=arguments)
|
||||
else:
|
||||
return ("Error: Unable to upload image.",)
|
||||
else:
|
||||
handler = submit(api_url+"/text-to-video", arguments=arguments)
|
||||
|
||||
result = handler.get()
|
||||
video_url = result["video"]["url"]
|
||||
return (video_url,)
|
||||
except Exception as e:
|
||||
print(f"Error generating video: {str(e)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
|
||||
class VideoGenRunwayGen3Node:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"image": ("IMAGE",),
|
||||
"duration": (["5", "10"], {"default": "5"}),
|
||||
"aspect_ratio": (["16:9", "9:16"], {"default": "16:9"}),
|
||||
"fal_key":("STRING", {"forceInput": True,}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def generate_video(self, prompt, image, duration,aspect_ratio,fal_key):
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
try:
|
||||
image_url = upload_image(image)
|
||||
if not image_url:
|
||||
return ("Error: Unable to upload image.",)
|
||||
|
||||
arguments = {
|
||||
"prompt": prompt,
|
||||
"image_url": image_url,
|
||||
"duration": duration,
|
||||
"ratio":aspect_ratio
|
||||
}
|
||||
|
||||
handler = submit("fal-ai/runway-gen3/turbo/image-to-video", arguments=arguments)
|
||||
result = handler.get()
|
||||
video_url = result["video"]["url"]
|
||||
return (video_url,)
|
||||
except Exception as e:
|
||||
print(f"Error generating video: {str(e)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
class VideoGenLumaDreamMachineNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"aspect_ratio": (["16:9", "9:16", "4:3", "3:4", "21:9", "9:21"], {"default": "16:9"}),
|
||||
"fal_key":("STRING", {"forceInput": True,}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
"loop": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def generate_video(self, prompt, aspect_ratio,fal_key, image=None, loop=True):
|
||||
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
|
||||
arguments = {
|
||||
"prompt": prompt,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"loop": loop,
|
||||
}
|
||||
|
||||
try:
|
||||
if image is not None:
|
||||
image_url = upload_image(image)
|
||||
if not image_url:
|
||||
return ("Error: Unable to upload image.",)
|
||||
arguments["image_url"] = image_url
|
||||
endpoint = "fal-ai/luma-dream-machine/image-to-video"
|
||||
else:
|
||||
endpoint = "fal-ai/luma-dream-machine"
|
||||
|
||||
handler = submit(endpoint, arguments=arguments)
|
||||
result = handler.get()
|
||||
video_url = result["video"]["url"]
|
||||
return (video_url,)
|
||||
except Exception as e:
|
||||
print(f"Error generating video: {str(e)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
class LoadVideoFromURL:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"url": ("STRING", {"default": "https://example.com/video.mp4"}),
|
||||
"force_rate": ("INT", {"default": 0, "min": 0, "max": 60, "step": 1}),
|
||||
"force_size": (["Disabled", "Custom Height", "Custom Width", "Custom", "256x?", "?x256", "256x256", "512x?", "?x512", "512x512"],),
|
||||
"custom_width": ("INT", {"default": 512, "min": 0, "max": 8192, "step": 8}),
|
||||
"custom_height": ("INT", {"default": 512, "min": 0, "max": 8192, "step": 8}),
|
||||
"frame_load_cap": ("INT", {"default": 0, "min": 0, "max": 1000000, "step": 1}),
|
||||
"skip_first_frames": ("INT", {"default": 0, "min": 0, "max": 1000000, "step": 1}),
|
||||
"select_every_nth": ("INT", {"default": 1, "min": 1, "max": 1000000, "step": 1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "INT", "VHS_VIDEOINFO")
|
||||
RETURN_NAMES = ("frames", "frame_count", "video_info")
|
||||
FUNCTION = "load_video_from_url"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def load_video_from_url(self, url, force_rate, force_size, custom_width, custom_height, frame_load_cap, skip_first_frames, select_every_nth):
|
||||
# Download the video to a temporary file
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=".mp4") as temp_file:
|
||||
response = requests.get(url, stream=True)
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
temp_file.write(chunk)
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
# Load the video using OpenCV
|
||||
cap = cv2.VideoCapture(temp_file_path)
|
||||
|
||||
# Get video properties
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
duration = total_frames / fps
|
||||
|
||||
# Calculate target size
|
||||
if force_size != "Disabled":
|
||||
if force_size == "Custom Width":
|
||||
new_height = int(height * (custom_width / width))
|
||||
new_width = custom_width
|
||||
elif force_size == "Custom Height":
|
||||
new_width = int(width * (custom_height / height))
|
||||
new_height = custom_height
|
||||
elif force_size == "Custom":
|
||||
new_width, new_height = custom_width, custom_height
|
||||
else:
|
||||
target_width, target_height = map(int, force_size.replace("?", "0").split("x"))
|
||||
if target_width == 0:
|
||||
new_width = int(width * (target_height / height))
|
||||
new_height = target_height
|
||||
else:
|
||||
new_height = int(height * (target_width / width))
|
||||
new_width = target_width
|
||||
else:
|
||||
new_width, new_height = width, height
|
||||
|
||||
frames = []
|
||||
frame_count = 0
|
||||
|
||||
for i in range(total_frames):
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
|
||||
if i < skip_first_frames:
|
||||
continue
|
||||
|
||||
if (i - skip_first_frames) % select_every_nth != 0:
|
||||
continue
|
||||
|
||||
if force_size != "Disabled":
|
||||
frame = cv2.resize(frame, (new_width, new_height))
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frame = torch.from_numpy(frame).float() / 255.0
|
||||
frames.append(frame)
|
||||
|
||||
frame_count += 1
|
||||
|
||||
if frame_load_cap > 0 and frame_count >= frame_load_cap:
|
||||
break
|
||||
|
||||
cap.release()
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
frames = torch.stack(frames)
|
||||
|
||||
video_info = {
|
||||
"source_fps": fps,
|
||||
"source_frame_count": total_frames,
|
||||
"source_duration": duration,
|
||||
"source_width": width,
|
||||
"source_height": height,
|
||||
"loaded_fps": fps if force_rate == 0 else force_rate,
|
||||
"loaded_frame_count": frame_count,
|
||||
"loaded_duration": frame_count / (fps if force_rate == 0 else force_rate),
|
||||
"loaded_width": new_width,
|
||||
"loaded_height": new_height,
|
||||
}
|
||||
|
||||
return (frames, frame_count, video_info)
|
||||
|
||||
@@ -0,0 +1,250 @@
|
||||
# 修改自 https://github.com/AnyaCoder/ComfyUI-fish-speech/
|
||||
|
||||
import torch,os
|
||||
from pathlib import Path
|
||||
from .fish_speech.llama_utils import load_model as load_llama_model
|
||||
from .fish_speech.vqgan_utils import load_model as load_vqgan_model
|
||||
from .fish_speech.vqgan_utils import audio2prompt, semantic2audio
|
||||
from .fish_speech.llama_utils import prompt2semantic
|
||||
|
||||
import folder_paths
|
||||
|
||||
def get_checkpoints_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('fish_speech')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "fish_speech")
|
||||
|
||||
current_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
configs_dir=os.path.join(current_directory,"fish_speech","configs")
|
||||
|
||||
CKPTS_FOLDER = Path(get_checkpoints_path())
|
||||
|
||||
CONFIGS_FOLDER = Path(configs_dir)
|
||||
|
||||
|
||||
class LoadVQGAN:
|
||||
def __init__(self):
|
||||
self.vqgan = None
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"config": ([str(c.relative_to(CONFIGS_FOLDER)) for c in CONFIGS_FOLDER.glob("*vq*.yaml")], {"default": "firefly_gan_vq.yaml"}),
|
||||
"model": ([str(p.relative_to(CKPTS_FOLDER)) for p in CKPTS_FOLDER.glob("*vq*.pth")], ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, model):
|
||||
return ""
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(s, model):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("VQGAN", )
|
||||
RETURN_NAMES = ("vqgan", )
|
||||
|
||||
FUNCTION = "load_vqgan"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def load_vqgan(self, config, model, device):
|
||||
config = config.rsplit(".", 1)[0]
|
||||
model = str(CKPTS_FOLDER / model)
|
||||
if self.vqgan is None:
|
||||
self.vqgan = load_vqgan_model(config,model, device=device)
|
||||
return (self.vqgan, )
|
||||
|
||||
|
||||
|
||||
class AudioToPrompt:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"vqgan": ("VQGAN", ),
|
||||
"audio": ("AUDIO", ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("AUDIO", "NUMPY")
|
||||
RETURN_NAMES = ("restored_audio", "prompt_tokens")
|
||||
|
||||
FUNCTION = "encode"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def encode(self, vqgan, audio, device):
|
||||
return audio2prompt(vqgan, audio, device)
|
||||
|
||||
|
||||
|
||||
class Prompt2Semantic:
|
||||
|
||||
def __init__(self):
|
||||
self.llama = None
|
||||
self.decode_func = None
|
||||
pass
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
"prompt_text": ("STRING", {"multiline": True}),
|
||||
"prompt_tokens": ("NUMPY", ),
|
||||
"max_new_tokens": ("INT", {
|
||||
"default": 1024,
|
||||
"min": 0,
|
||||
"max": 2048,
|
||||
"step": 8,
|
||||
"display": "number",
|
||||
}),
|
||||
"top_p": ("FLOAT", {
|
||||
"default": 0.7,
|
||||
"min": 0.6,
|
||||
"max": 0.9,
|
||||
"step": 0.01,
|
||||
"display": "number",
|
||||
}),
|
||||
"repetition_penalty": ("FLOAT", {
|
||||
"default": 1.2,
|
||||
"min": 1.0,
|
||||
"max": 1.5,
|
||||
"step": 0.01,
|
||||
"display": "number",
|
||||
}),
|
||||
"temperature": ("FLOAT", {
|
||||
"default": 0.7,
|
||||
"min": 0.6,
|
||||
"max": 0.9,
|
||||
"step": 0.01,
|
||||
"display": "number",
|
||||
}),
|
||||
|
||||
"seed": ("INT", {
|
||||
"default": 42,
|
||||
"min": 0,
|
||||
"max": 4294967295,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
}),
|
||||
"iterative_prompt": (["yes", "no"], {"default": "yes"}),
|
||||
"chunk_length": ("INT", {
|
||||
"default": 100,
|
||||
"min": 0,
|
||||
"max": 500,
|
||||
"step": 8,
|
||||
"display": "number",
|
||||
}),
|
||||
|
||||
"compile": (["yes", "no"], {"default": "no"}),
|
||||
"precision": (["bf16", "half"], {"default": "bf16"}),
|
||||
|
||||
# "decode_func": ("DECODE_FUNC", ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NUMPY", )
|
||||
RETURN_NAMES = ("codes", )
|
||||
|
||||
FUNCTION = "decode"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def decode(
|
||||
self,
|
||||
|
||||
text: str,
|
||||
prompt_text: str,
|
||||
prompt_tokens,
|
||||
max_new_tokens: int,
|
||||
top_p: float,
|
||||
repetition_penalty: float,
|
||||
temperature: float,
|
||||
|
||||
seed: int,
|
||||
iterative_prompt: str,
|
||||
chunk_length: int,
|
||||
|
||||
compile: str,
|
||||
precision,
|
||||
device: str,
|
||||
):
|
||||
|
||||
model = get_checkpoints_path()
|
||||
precision = torch.bfloat16 if precision == "bf16" else torch.half
|
||||
compile=True if compile == "yes" else False
|
||||
if self.llama is None or self.decode_func is None:
|
||||
self.llama, self.decode_func = load_llama_model(model, device, precision, compile)
|
||||
|
||||
|
||||
return prompt2semantic(
|
||||
self.llama,
|
||||
self.decode_func,
|
||||
text,
|
||||
[prompt_text,],
|
||||
[prompt_tokens,],
|
||||
max_new_tokens,
|
||||
top_p,
|
||||
repetition_penalty,
|
||||
temperature,
|
||||
device,
|
||||
compile=True if compile == "yes" else False,
|
||||
seed=seed,
|
||||
iterative_prompt=True if iterative_prompt == "yes" else False,
|
||||
chunk_length=chunk_length,
|
||||
)
|
||||
|
||||
|
||||
|
||||
class Semantic2Audio:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"vqgan": ("VQGAN", ),
|
||||
"codes": ("NUMPY", ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO", )
|
||||
RETURN_NAMES = ("generated_audio", )
|
||||
|
||||
FUNCTION = "generate"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def generate(self, vqgan, codes, device):
|
||||
return semantic2audio(vqgan, codes, device)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -42,8 +42,13 @@ else:
|
||||
_available=True
|
||||
|
||||
|
||||
|
||||
llma_model_path=os.path.join(folder_paths.models_dir, "lama/big-lama.pt")
|
||||
def get_lama_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('lama')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "lama")
|
||||
|
||||
llma_model_path=os.path.join(get_lama_path(), "big-lama.pt")
|
||||
if not os.path.exists(llma_model_path):
|
||||
os.environ['LAMA_MODEL']=''
|
||||
print(f"## lama torchscript model not found: {llma_model_path},pls download from https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt")
|
||||
@@ -80,8 +85,6 @@ class LaMaInpainting:
|
||||
"image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
},
|
||||
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
|
||||
@@ -2,16 +2,14 @@
|
||||
import scipy.ndimage
|
||||
import torch
|
||||
|
||||
from nodes import MAX_RESOLUTION
|
||||
|
||||
import numpy as np
|
||||
# from PIL import Image, ImageDraw
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
from comfy.cli_args import args
|
||||
import cv2
|
||||
|
||||
|
||||
import cv2,os
|
||||
from nodes import MAX_RESOLUTION, SaveImage, common_ksampler
|
||||
import folder_paths,random
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
@@ -22,6 +20,19 @@ def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def add_masks(mask1, mask2):
|
||||
mask1 = mask1.cpu()
|
||||
mask2 = mask2.cpu()
|
||||
cv2_mask1 = np.array(mask1) * 255
|
||||
cv2_mask2 = np.array(mask2) * 255
|
||||
|
||||
if cv2_mask1.shape == cv2_mask2.shape:
|
||||
cv2_mask = cv2.add(cv2_mask1, cv2_mask2)
|
||||
return torch.clamp(torch.from_numpy(cv2_mask) / 255.0, min=0, max=1)
|
||||
else:
|
||||
return mask1
|
||||
|
||||
|
||||
def grow(mask, expand, tapered_corners):
|
||||
c = 0 if tapered_corners else 1
|
||||
kernel = np.array([[c, 1, c],
|
||||
@@ -58,6 +69,35 @@ def combine(destination, source, x, y):
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class PreviewMask_(SaveImage):
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_temp_directory()
|
||||
self.type = "temp"
|
||||
self.prefix_append =''.join(random.choice("abcdehijklmnopqrstupvxyzfg") for x in range(5))
|
||||
self.compress_level = 4
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mask": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
# 运行的函数
|
||||
def run(self, mask ):
|
||||
img=tensor2pil(mask)
|
||||
img=img.convert('RGB')
|
||||
img=pil2tensor(img)
|
||||
return self.save_images(img, 'temp_', None, None)
|
||||
|
||||
|
||||
class OutlineMask:
|
||||
|
||||
@classmethod
|
||||
@@ -87,6 +127,69 @@ class OutlineMask:
|
||||
return (m3,)
|
||||
|
||||
|
||||
class MaskListReplace:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"masks": ("MASK",),
|
||||
"mask_replace": ("MASK",),
|
||||
"start_index":("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"end_index":("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"invert": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self, masks,mask_replace,start_index,end_index,invert):
|
||||
mask_replace=mask_replace[0]
|
||||
start_index=start_index[0]
|
||||
end_index=end_index[0]
|
||||
invert=invert[0]
|
||||
|
||||
new_masks=[]
|
||||
for i in range(len(masks)):
|
||||
if i>=start_index and i<=end_index:
|
||||
if invert:
|
||||
new_masks.append(masks[i])
|
||||
else:
|
||||
new_masks.append(mask_replace)
|
||||
else:
|
||||
if invert:
|
||||
new_masks.append(mask_replace)
|
||||
else:
|
||||
new_masks.append(masks[i])
|
||||
|
||||
return (new_masks,)
|
||||
|
||||
|
||||
class MaskListMerge:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"masks": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self, masks):
|
||||
mask=masks[0]
|
||||
if isinstance(masks, list):
|
||||
for m in masks:
|
||||
# print(m.shape)
|
||||
mask = add_masks(mask, m)
|
||||
return (mask,)
|
||||
|
||||
|
||||
class FeatheredMask:
|
||||
|
||||
@@ -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,)
|
||||
@@ -0,0 +1,104 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image,ImageSequence,ImageOps
|
||||
import base64
|
||||
import io
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
def load_image_to_tensor( image):
|
||||
image_path = folder_paths.get_annotated_filepath(image)
|
||||
|
||||
img = node_helpers.pillow(Image.open, image_path)
|
||||
|
||||
output_images = []
|
||||
output_masks = []
|
||||
w, h = None, None
|
||||
|
||||
excluded_formats = ['MPO']
|
||||
|
||||
for i in ImageSequence.Iterator(img):
|
||||
i = node_helpers.pillow(ImageOps.exif_transpose, i)
|
||||
|
||||
if i.mode == 'I':
|
||||
i = i.point(lambda i: i * (1 / 255))
|
||||
image = i.convert("RGB")
|
||||
|
||||
if len(output_images) == 0:
|
||||
w = image.size[0]
|
||||
h = image.size[1]
|
||||
|
||||
if image.size[0] != w or image.size[1] != h:
|
||||
continue
|
||||
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[None,]
|
||||
if 'A' in i.getbands():
|
||||
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
|
||||
mask = 1. - torch.from_numpy(mask)
|
||||
else:
|
||||
mask = torch.zeros((64,64), dtype=torch.float32, device="cpu")
|
||||
output_images.append(image)
|
||||
output_masks.append(mask.unsqueeze(0))
|
||||
|
||||
if len(output_images) > 1 and img.format not in excluded_formats:
|
||||
output_image = torch.cat(output_images, dim=0)
|
||||
output_mask = torch.cat(output_masks, dim=0)
|
||||
else:
|
||||
output_image = output_images[0]
|
||||
output_mask = output_masks[0]
|
||||
|
||||
return (output_image, output_mask)
|
||||
|
||||
|
||||
|
||||
class P5Input:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"frames":("IMAGEBASE64",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("frames",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self, frames):
|
||||
ims=[]
|
||||
for im in frames['images']:
|
||||
# print(im)
|
||||
if 'type' in im and (not f"[{im['type']}]" in im['name']):
|
||||
im['name']=im['name']+" "+f"[{im['type']}]"
|
||||
|
||||
output_image, output_mask = load_image_to_tensor(im['name'])
|
||||
ims.append(output_image)
|
||||
|
||||
if len(ims)==0:
|
||||
image1 = Image.new('RGB', (512, 512), color='black')
|
||||
return (pil2tensor(image1),)
|
||||
image1 = ims[0]
|
||||
for image2 in ims[1:]:
|
||||
if image1.shape[1:] != image2.shape[1:]:
|
||||
image2 = comfy.utils.common_upscale(image2.movedim(-1, 1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1, -1)
|
||||
image1 = torch.cat((image1, image2), dim=0)
|
||||
|
||||
# 用于节点提示:p5节点提示有多少帧
|
||||
return {"ui": {"_info": [len(frames['images'])]}, "result": (image1,)}
|
||||
@@ -18,9 +18,16 @@ 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):
|
||||
|
||||
file_list = []
|
||||
for root, dirs, files in os.walk(directory):
|
||||
for file in files:
|
||||
@@ -180,7 +187,8 @@ class PromptImage:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("json_str",)
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@@ -188,19 +196,26 @@ class PromptImage:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
CATEGORY = "♾️Mixlab/Output"
|
||||
|
||||
# 运行的函数
|
||||
def run(self,prompts,images,save_to_image):
|
||||
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]
|
||||
@@ -208,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
|
||||
}),) }
|
||||
|
||||
|
||||
|
||||
@@ -517,9 +544,10 @@ class RandomPrompt:
|
||||
class EmbeddingPrompt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"embedding":(get_files_with_extension(embeddings_path,'.pt'),),
|
||||
"embedding":(folder_paths.get_filename_list("embeddings"),),
|
||||
"weight": ("FLOAT", {"default": 1, "min": -2, "max": 2,"step":0.01 ,"display": "slider"}),
|
||||
},
|
||||
|
||||
@@ -544,7 +572,112 @@ class EmbeddingPrompt:
|
||||
# return (new_prompt)
|
||||
return (prompt,)
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
# RETURN_TYPES = (any_type,)
|
||||
|
||||
# conditioning :提示,正向or负向
|
||||
# clip:clip模型
|
||||
# gligen_textbox_model:gligen模型
|
||||
# grids:矩形框的集合
|
||||
# labels:每个矩形框对应的标签的集合
|
||||
# index:选取第几个矩形框作为gligen的box
|
||||
|
||||
class GLIGENTextBoxApply_Advanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"conditioning": ("CONDITIONING", ),
|
||||
"clip": ("CLIP", ),
|
||||
"gligen_textbox_model": ("GLIGEN", ),
|
||||
"grids": ("_GRID",),
|
||||
"labels": ("STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "",
|
||||
"forceInput": True
|
||||
}),
|
||||
"index": ("INT", {"default": -1, "min": -1, "max": 300, "step": 1}),
|
||||
"max_size": ("INT", {"default": 8, "min": 1, "max": 300, "step": 1}),
|
||||
"random_shuffle":(["on","off"],),
|
||||
},
|
||||
"optional":{
|
||||
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff,"step": 1}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("CONDITIONING","STRING",)
|
||||
RETURN_NAMES = ("CONDITIONING","label",)
|
||||
|
||||
FUNCTION = "run"
|
||||
# INPUT_IS_LIST = True
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
def run(self, conditioning, clip, gligen_textbox_model, grids, labels, index,max_size,random_shuffle,seed=0):
|
||||
# print('grids',grids)
|
||||
# conditioning=conditioning[0]
|
||||
# clip=clip[0]
|
||||
# gligen_textbox_model=gligen_textbox_model[0]
|
||||
# index=index[0]
|
||||
# max_size=max_size[0]
|
||||
# random_shuffle=random_shuffle[0]
|
||||
|
||||
texts=labels
|
||||
|
||||
if index>-1:
|
||||
texts=[labels[index]]
|
||||
grids=[grids[index]]
|
||||
|
||||
if random_shuffle=='on':
|
||||
sss=[[texts[i],grids[i]] for i in range(len(texts))]
|
||||
random.shuffle(sss)
|
||||
texts=[s[0] for s in sss]
|
||||
grids=[s[1] for s in sss]
|
||||
|
||||
if len(texts) > max_size:
|
||||
texts = texts[:max_size]
|
||||
|
||||
c = []
|
||||
|
||||
for t in conditioning:
|
||||
n = [t[0], t[1].copy()]
|
||||
|
||||
|
||||
# 多个
|
||||
position_params=[]
|
||||
for i in range(len(texts)):
|
||||
text=texts[i]
|
||||
grid=grids[i]
|
||||
x,y,width,height=grid
|
||||
# print(text)
|
||||
cond, cond_pooled = clip.encode_from_tokens(clip.tokenize(text), return_pooled=True)
|
||||
position_params =position_params+ [(cond_pooled, height // 8, width // 8, y // 8, x // 8)]
|
||||
|
||||
# 前一个
|
||||
prev = []
|
||||
if "gligen" in n[1]:
|
||||
prev = n[1]['gligen'][2]
|
||||
|
||||
n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params)
|
||||
# print('gligen',n)
|
||||
c.append(n)
|
||||
|
||||
# 下面这个写法有bug
|
||||
# for i in range(len(texts)):
|
||||
# text=texts[i]
|
||||
# grid=grids[i]
|
||||
# x,y,width,height=grid
|
||||
|
||||
# cond, cond_pooled = clip.encode_from_tokens(clip.tokenize(text), return_pooled=True)
|
||||
# for t in conditioning:
|
||||
# n = [t[0], t[1].copy()]
|
||||
# position_params = [(cond_pooled, height // 8, width // 8, y // 8, x // 8)]
|
||||
# prev = []
|
||||
# if "gligen" in n[1]:
|
||||
# prev = n[1]['gligen'][2]
|
||||
|
||||
# n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params)
|
||||
# c.append(n)
|
||||
|
||||
|
||||
return (c,texts, )
|
||||
|
||||
|
||||
class JoinWithDelimiter:
|
||||
@classmethod
|
||||
@@ -559,7 +692,7 @@ class JoinWithDelimiter:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
INPUT_IS_LIST = True # 当true的时候,输入时list,当false的时候,如果输入是list,则会自动包一层for循环调用
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
@@ -467,15 +467,37 @@ class BriaRMBG(nn.Module):
|
||||
|
||||
|
||||
|
||||
def get_U2NET_model_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('rembg')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "rembg")
|
||||
|
||||
|
||||
|
||||
|
||||
U2NET_HOME=os.path.join(folder_paths.models_dir, "rembg")
|
||||
U2NET_HOME=get_U2NET_model_path()
|
||||
os.environ["U2NET_HOME"] = U2NET_HOME
|
||||
|
||||
global _available
|
||||
_available=False
|
||||
|
||||
|
||||
def get_rembg_models(path):
|
||||
"""从目录中获取文件并提取文件名
|
||||
Args:
|
||||
path: 目录路径
|
||||
Returns:
|
||||
文件名列表
|
||||
"""
|
||||
filenames = []
|
||||
for root, _, files in os.walk(path):
|
||||
for filename in files:
|
||||
# 过滤隐藏文件
|
||||
if not filename.startswith('.'):
|
||||
name, ext = os.path.splitext(os.path.basename(filename))
|
||||
filenames.append(name)
|
||||
return filenames
|
||||
|
||||
|
||||
def is_installed(package):
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
@@ -509,8 +531,8 @@ except:
|
||||
_available=False
|
||||
|
||||
|
||||
def briarmbg_run(images=[]):
|
||||
mroot=os.path.join(folder_paths.models_dir, "rembg")
|
||||
def run_briarmbg(images=[]):
|
||||
mroot=U2NET_HOME
|
||||
m=os.path.join(mroot,'briarmbg.pth')
|
||||
if os.path.exists(m)==False:
|
||||
# 下载
|
||||
@@ -573,14 +595,15 @@ def briarmbg_run(images=[]):
|
||||
return (masks,rgba_images,rgb_images)
|
||||
|
||||
|
||||
def run_bg(model_name= "unet",images=[]):
|
||||
def run_rembg(model_name= "unet",images=[],callback=None):
|
||||
# model_name = "unet" # "isnet-general-use"
|
||||
# print('#run_rembg',model_name)
|
||||
rembg_session = new_session(model_name)
|
||||
masks=[]
|
||||
rgba_images=[]
|
||||
rgb_images=[]
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(images) )
|
||||
pbar=callback
|
||||
for img in images:
|
||||
# use the post_process_mask argument to post process the mask to get better results.
|
||||
mask = remove(img, session=rembg_session,only_mask=True,post_process_mask=True)
|
||||
@@ -620,8 +643,9 @@ def run_bg(model_name= "unet",images=[]):
|
||||
rgb_image = Image.new("RGB", image_rgba.size, (0, 0, 0))
|
||||
rgb_image.paste(image_rgba, mask=image_rgba.split()[3])
|
||||
rgb_images.append(rgb_image)
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
return (masks,rgba_images,rgb_images)
|
||||
|
||||
|
||||
@@ -643,17 +667,7 @@ class RembgNode_:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image": ("IMAGE",),
|
||||
"model_name": ([
|
||||
"briarmbg",
|
||||
"u2net",
|
||||
"u2netp",
|
||||
"u2net_human_seg",
|
||||
"u2net_cloth_seg",
|
||||
"silueta",
|
||||
"isnet-general-use",
|
||||
"isnet-anime",
|
||||
|
||||
],),
|
||||
"model_name": (get_rembg_models(U2NET_HOME),),
|
||||
|
||||
},
|
||||
}
|
||||
@@ -681,9 +695,9 @@ class RembgNode_:
|
||||
images.append(im)
|
||||
|
||||
if model_name=='briarmbg':
|
||||
masks,rgba_images,rgb_images=briarmbg_run(images)
|
||||
masks,rgba_images,rgb_images=run_briarmbg(images)
|
||||
else:
|
||||
masks,rgba_images,rgb_images=run_bg(model_name,images)
|
||||
masks,rgba_images,rgb_images=run_rembg(model_name,images, comfy.utils.ProgressBar(len(images) ))
|
||||
|
||||
masks=[pil2tensor(m) for m in masks]
|
||||
|
||||
|
||||
@@ -90,10 +90,10 @@ class ScreenShareNode:
|
||||
} }
|
||||
|
||||
RETURN_TYPES = ('IMAGE','STRING','FLOAT',"INT")
|
||||
RETURN_NAMES = ("IMAGE","PROMPT","FLOAT","INT")
|
||||
RETURN_NAMES = ("current frame (image)","prompt","denoise (float)","seed (int)")
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
CATEGORY = "♾️Mixlab/Screen"
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (False,False,False,False)
|
||||
@@ -109,7 +109,7 @@ class FloatingVideo:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return { "required":{
|
||||
"images": ("IMAGE",)
|
||||
"image": ("IMAGE",)
|
||||
}, }
|
||||
|
||||
# RETURN_TYPES = ('IMAGE','MASK')
|
||||
@@ -118,22 +118,22 @@ class FloatingVideo:
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
CATEGORY = "♾️Mixlab/Screen"
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
# OUTPUT_IS_LIST = (False,False,)
|
||||
|
||||
# 运行的函数
|
||||
def run(self,images):
|
||||
def run(self,image):
|
||||
|
||||
results = list()
|
||||
|
||||
for image in images:
|
||||
image=tensor2pil(image)
|
||||
for im in image:
|
||||
im=tensor2pil(im)
|
||||
# image_base64 = base64.b64encode(image.tobytes())
|
||||
|
||||
buffered = BytesIO()
|
||||
image.save(buffered, format="JPEG")
|
||||
im.save(buffered, format="JPEG")
|
||||
image_base64 = base64.b64encode(buffered.getvalue()).decode("utf-8")
|
||||
|
||||
results.append(image_base64)
|
||||
|
||||
@@ -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,)
|
||||
|
||||
@@ -0,0 +1,503 @@
|
||||
import comfy
|
||||
import torch
|
||||
|
||||
from dataclasses import dataclass
|
||||
import torch.nn as nn
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
import comfy.ops
|
||||
from typing import Union
|
||||
import comfy.sample
|
||||
import latent_preview
|
||||
import comfy.utils
|
||||
|
||||
T = torch.Tensor
|
||||
|
||||
|
||||
from .VisualStylePrompting.attention_functions import VisualStyleProcessor
|
||||
|
||||
class ApplyVisualStylePrompting:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"reference_image": ("IMAGE",),
|
||||
"reference_image_text": ("STRING", {"multiline": True}),
|
||||
"model": ("MODEL",),
|
||||
"clip": ("CLIP", ),
|
||||
"vae": ("VAE", ),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"enabled": ("BOOLEAN", {"default": True}),
|
||||
"denoise": ("FLOAT", {"default": 1., "min": 0., "max": 1., "step": 1e-2}),
|
||||
"batch_size": ("INT", {"default": 1, "min": 1, "max": 4096,"step":2})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CONDITIONING","CONDITIONING", "LATENT")
|
||||
RETURN_NAMES = ("model", "positive", "negative", "latents")
|
||||
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
def run(
|
||||
self,
|
||||
reference_image,
|
||||
reference_image_text,
|
||||
model: comfy.model_patcher.ModelPatcher,
|
||||
clip,
|
||||
vae,
|
||||
positive,
|
||||
negative,
|
||||
enabled,
|
||||
denoise,
|
||||
batch_size=1
|
||||
):
|
||||
|
||||
tokens = clip.tokenize(reference_image_text)
|
||||
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
|
||||
reference_image_prompt=[[cond, {"pooled_output": pooled}]]
|
||||
|
||||
reference_image = reference_image.repeat(((batch_size+1)//2, 1,1,1))
|
||||
|
||||
self.model = model
|
||||
reference_latent = vae.encode(reference_image[:,:,:,:3])
|
||||
|
||||
for n, m in model.model.diffusion_model.named_modules():
|
||||
if m.__class__.__name__ == "CrossAttention":
|
||||
processor = VisualStyleProcessor(m, enabled=enabled)
|
||||
setattr(m, 'forward', processor.visual_style_forward)
|
||||
|
||||
conditioning_prompt = reference_image_prompt + positive
|
||||
negative_prompt = negative * 2
|
||||
|
||||
latents = torch.zeros_like(reference_latent)
|
||||
latents = torch.cat([latents] * 2)
|
||||
|
||||
if denoise < 1.0:
|
||||
latents[::1] = reference_latent[:1]
|
||||
else:
|
||||
latents[::2] = reference_latent
|
||||
|
||||
denoise_mask = torch.ones_like(latents)[:, :1, ...] * denoise
|
||||
|
||||
denoise_mask[0] = 0.
|
||||
|
||||
return (model, conditioning_prompt, negative_prompt, {"samples": latents, "noise_mask": denoise_mask})
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
def default(val, d):
|
||||
if exists(val):
|
||||
return val
|
||||
return d
|
||||
|
||||
|
||||
class StyleAlignedArgs:
|
||||
def __init__(self, share_attn: str) -> None:
|
||||
self.adain_keys = "k" in share_attn
|
||||
self.adain_values = "v" in share_attn
|
||||
self.adain_queries = "q" in share_attn
|
||||
|
||||
share_attention: bool = True
|
||||
adain_queries: bool = True
|
||||
adain_keys: bool = True
|
||||
adain_values: bool = True
|
||||
|
||||
|
||||
def expand_first(
|
||||
feat: T,
|
||||
scale=1.0,
|
||||
) -> T:
|
||||
"""
|
||||
Expand the first element so it has the same shape as the rest of the batch.
|
||||
"""
|
||||
b = feat.shape[0]
|
||||
feat_style = torch.stack((feat[0], feat[b // 2])).unsqueeze(1)
|
||||
if scale == 1:
|
||||
feat_style = feat_style.expand(2, b // 2, *feat.shape[1:])
|
||||
else:
|
||||
feat_style = feat_style.repeat(1, b // 2, 1, 1, 1)
|
||||
feat_style = torch.cat([feat_style[:, :1], scale * feat_style[:, 1:]], dim=1)
|
||||
return feat_style.reshape(*feat.shape)
|
||||
|
||||
|
||||
def concat_first(feat: T, dim=2, scale=1.0) -> T:
|
||||
"""
|
||||
concat the the feature and the style feature expanded above
|
||||
"""
|
||||
feat_style = expand_first(feat, scale=scale)
|
||||
return torch.cat((feat, feat_style), dim=dim)
|
||||
|
||||
|
||||
def calc_mean_std(feat, eps: float = 1e-5) -> "tuple[T, T]":
|
||||
feat_std = (feat.var(dim=-2, keepdims=True) + eps).sqrt()
|
||||
feat_mean = feat.mean(dim=-2, keepdims=True)
|
||||
return feat_mean, feat_std
|
||||
|
||||
def adain(feat: T) -> T:
|
||||
feat_mean, feat_std = calc_mean_std(feat)
|
||||
feat_style_mean = expand_first(feat_mean)
|
||||
feat_style_std = expand_first(feat_std)
|
||||
feat = (feat - feat_mean) / feat_std
|
||||
feat = feat * feat_style_std + feat_style_mean
|
||||
return feat
|
||||
|
||||
class SharedAttentionProcessor:
|
||||
def __init__(self, args: StyleAlignedArgs, scale: float):
|
||||
self.args = args
|
||||
self.scale = scale
|
||||
|
||||
def __call__(self, q, k, v, extra_options):
|
||||
if self.args.adain_queries:
|
||||
q = adain(q)
|
||||
if self.args.adain_keys:
|
||||
k = adain(k)
|
||||
if self.args.adain_values:
|
||||
v = adain(v)
|
||||
if self.args.share_attention:
|
||||
k = concat_first(k, -2, scale=self.scale)
|
||||
v = concat_first(v, -2)
|
||||
|
||||
return q, k, v
|
||||
|
||||
|
||||
def get_norm_layers(
|
||||
layer: nn.Module,
|
||||
norm_layers_: "dict[str, list[Union[nn.GroupNorm, nn.LayerNorm]]]",
|
||||
share_layer_norm: bool,
|
||||
share_group_norm: bool,
|
||||
):
|
||||
if isinstance(layer, nn.LayerNorm) and share_layer_norm:
|
||||
norm_layers_["layer"].append(layer)
|
||||
if isinstance(layer, nn.GroupNorm) and share_group_norm:
|
||||
norm_layers_["group"].append(layer)
|
||||
else:
|
||||
for child_layer in layer.children():
|
||||
get_norm_layers(
|
||||
child_layer, norm_layers_, share_layer_norm, share_group_norm
|
||||
)
|
||||
|
||||
|
||||
def register_norm_forward(
|
||||
norm_layer: Union[nn.GroupNorm, nn.LayerNorm],
|
||||
) -> Union[nn.GroupNorm, nn.LayerNorm]:
|
||||
if not hasattr(norm_layer, "orig_forward"):
|
||||
setattr(norm_layer, "orig_forward", norm_layer.forward)
|
||||
orig_forward = norm_layer.orig_forward
|
||||
|
||||
def forward_(hidden_states: T) -> T:
|
||||
n = hidden_states.shape[-2]
|
||||
hidden_states = concat_first(hidden_states, dim=-2)
|
||||
hidden_states = orig_forward(hidden_states) # type: ignore
|
||||
return hidden_states[..., :n, :]
|
||||
|
||||
norm_layer.forward = forward_ # type: ignore
|
||||
return norm_layer
|
||||
|
||||
|
||||
def register_shared_norm(
|
||||
model: ModelPatcher,
|
||||
share_group_norm: bool = True,
|
||||
share_layer_norm: bool = True,
|
||||
):
|
||||
norm_layers = {"group": [], "layer": []}
|
||||
get_norm_layers(model.model, norm_layers, share_layer_norm, share_group_norm)
|
||||
print(
|
||||
f"Patching {len(norm_layers['group'])} group norms, {len(norm_layers['layer'])} layer norms."
|
||||
)
|
||||
return [register_norm_forward(layer) for layer in norm_layers["group"]] + [
|
||||
register_norm_forward(layer) for layer in norm_layers["layer"]
|
||||
]
|
||||
|
||||
|
||||
SHARE_NORM_OPTIONS = ["both", "group", "layer", "disabled"]
|
||||
SHARE_ATTN_OPTIONS = ["q+k", "q+k+v", "disabled"]
|
||||
|
||||
class StyleAlignedSampleReferenceLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{
|
||||
"reference_image": ("IMAGE",),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"model": ("MODEL",),
|
||||
"vae": ("VAE", ),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"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, ),
|
||||
"denoise": ("FLOAT", {"default": 1, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STEP_LATENTS","LATENT")
|
||||
RETURN_NAMES = ("ref_latents", "noised_output")
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
# CATEGORY = "style_aligned"
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
|
||||
def run(self, reference_image, positive, negative, model, vae, seed, steps, cfg,scheduler,denoise):
|
||||
|
||||
# TODO noise_mask?
|
||||
def vae_encode_crop_pixels(pixels):
|
||||
x = (pixels.shape[1] // 8) * 8
|
||||
y = (pixels.shape[2] // 8) * 8
|
||||
if pixels.shape[1] != x or pixels.shape[2] != y:
|
||||
x_offset = (pixels.shape[1] % 8) // 2
|
||||
y_offset = (pixels.shape[2] % 8) // 2
|
||||
pixels = pixels[:, x_offset:x + x_offset, y_offset:y + y_offset, :]
|
||||
return pixels
|
||||
|
||||
pixels=vae_encode_crop_pixels(reference_image)
|
||||
t = vae.encode(pixels[:,:,:,:3])
|
||||
latent_image = {"samples":t}
|
||||
|
||||
noise_seed=seed
|
||||
|
||||
sampler_name="ddim"
|
||||
|
||||
sampler = comfy.samplers.sampler_object(sampler_name)
|
||||
|
||||
total_steps = steps
|
||||
if denoise < 1.0:
|
||||
total_steps = int(steps/denoise)
|
||||
|
||||
comfy.model_management.load_models_gpu([model])
|
||||
sigmas = comfy.samplers.calculate_sigmas_scheduler(model.model, scheduler, total_steps).cpu()
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
|
||||
sigmas = sigmas.flip(0)
|
||||
if sigmas[0] == 0:
|
||||
sigmas[0] = 0.0001
|
||||
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
|
||||
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
|
||||
ref_latents = []
|
||||
def callback(step: int, x0: T, x: T, steps: int):
|
||||
ref_latents.insert(0, x[0])
|
||||
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
samples = comfy.sample.sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
out_noised = out
|
||||
|
||||
ref_latents = torch.stack(ref_latents)
|
||||
|
||||
return (ref_latents, out_noised)
|
||||
|
||||
class StyleAlignedReferenceSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
|
||||
"ref_latents": ("STEP_LATENTS",),
|
||||
"reference_image_text": ("STRING", {"multiline": True}),
|
||||
"model": ("MODEL",),
|
||||
"clip": ("CLIP", ),
|
||||
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
|
||||
"share_norm": (SHARE_NORM_OPTIONS,),
|
||||
"share_attn": (SHARE_ATTN_OPTIONS,),
|
||||
"scale": ("FLOAT", {"default": 1, "min": 0, "max": 2.0, "step": 0.01}),
|
||||
"batch_size": ("INT", {"default": 2, "min": 1, "max": 8, "step": 1}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"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, ),
|
||||
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT", "LATENT")
|
||||
RETURN_NAMES = ("output", "denoised_output")
|
||||
FUNCTION = "patch"
|
||||
# CATEGORY = "style_aligned"
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
def patch(
|
||||
self,
|
||||
ref_latents,
|
||||
reference_image_text,
|
||||
model,
|
||||
clip,
|
||||
positive,
|
||||
negative,
|
||||
share_norm,
|
||||
share_attn,
|
||||
scale,
|
||||
batch_size,
|
||||
seed,steps,cfg,scheduler,denoise
|
||||
|
||||
) -> "tuple[dict, dict]":
|
||||
|
||||
m = model.clone()
|
||||
|
||||
# ref_latents = vae.encode(reference_image[:,:,:,:3])
|
||||
|
||||
tokens = clip.tokenize(reference_image_text)
|
||||
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
|
||||
ref_positive=[[cond, {"pooled_output": pooled}]]
|
||||
|
||||
noise_seed=seed
|
||||
|
||||
|
||||
total_steps = steps
|
||||
if denoise < 1.0:
|
||||
total_steps = int(steps/denoise)
|
||||
|
||||
# comfy.model_management.load_models_gpu([model])
|
||||
sigmas = comfy.samplers.calculate_sigmas_scheduler(model.model, scheduler, total_steps).cpu()
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
|
||||
sampler_name="ddim"
|
||||
|
||||
sampler = comfy.samplers.sampler_object(sampler_name)
|
||||
|
||||
args = StyleAlignedArgs(share_attn)
|
||||
|
||||
# Concat batch with style latent
|
||||
style_latent_tensor = ref_latents[0].unsqueeze(0)
|
||||
height, width = style_latent_tensor.shape[-2:]
|
||||
latent_t = torch.zeros(
|
||||
[batch_size, 4, height, width], device=ref_latents.device
|
||||
)
|
||||
latent = {"samples": latent_t}
|
||||
noise = comfy.sample.prepare_noise(latent_t, noise_seed)
|
||||
|
||||
latent_t = torch.cat((style_latent_tensor, latent_t), dim=0)
|
||||
ref_noise = torch.zeros_like(noise[0]).unsqueeze(0)
|
||||
noise = torch.cat((ref_noise, noise), dim=0)
|
||||
|
||||
x0_output = {}
|
||||
preview_callback = latent_preview.prepare_callback(m, sigmas.shape[-1] - 1, x0_output)
|
||||
|
||||
# Replace first latent with the corresponding reference latent after each step
|
||||
def callback(step: int, x0: T, x: T, steps: int):
|
||||
preview_callback(step, x0, x, steps)
|
||||
if (step + 1 < steps):
|
||||
# 当ref_latents的step不够时
|
||||
if step+1>len(ref_latents)-1:
|
||||
step=len(ref_latents)-2
|
||||
|
||||
x[0] = ref_latents[step+1]
|
||||
x0[0] = ref_latents[step+1]
|
||||
|
||||
# Register shared norms
|
||||
share_group_norm = share_norm in ["group", "both"]
|
||||
share_layer_norm = share_norm in ["layer", "both"]
|
||||
register_shared_norm(m, share_group_norm, share_layer_norm)
|
||||
|
||||
# Patch cross attn
|
||||
m.set_model_attn1_patch(SharedAttentionProcessor(args, scale))
|
||||
|
||||
# Add reference conditioning to batch
|
||||
batched_condition = []
|
||||
for i,condition in enumerate(positive):
|
||||
additional = condition[1].copy()
|
||||
batch_with_reference = torch.cat([ref_positive[i][0], condition[0].repeat([batch_size] + [1] * len(condition[0].shape[1:]))], dim=0)
|
||||
if 'pooled_output' in additional and 'pooled_output' in ref_positive[i][1]:
|
||||
# combine pooled output
|
||||
pooled_output = torch.cat([ref_positive[i][1]['pooled_output'], additional['pooled_output'].repeat([batch_size]
|
||||
+ [1] * len(additional['pooled_output'].shape[1:]))], dim=0)
|
||||
additional['pooled_output'] = pooled_output
|
||||
if 'control' in additional:
|
||||
if 'control' in ref_positive[i][1]:
|
||||
# combine control conditioning
|
||||
control_hint = torch.cat([ref_positive[i][1]['control'].cond_hint_original, additional['control'].cond_hint_original.repeat([batch_size]
|
||||
+ [1] * len(additional['control'].cond_hint_original.shape[1:]))], dim=0)
|
||||
cloned_controlnet = additional['control'].copy()
|
||||
cloned_controlnet.set_cond_hint(control_hint, strength=additional['control'].strength, timestep_percent_range=additional['control'].timestep_percent_range)
|
||||
additional['control'] = cloned_controlnet
|
||||
else:
|
||||
# add zeros for first in batch
|
||||
control_hint = torch.cat([torch.zeros_like(additional['control'].cond_hint_original), additional['control'].cond_hint_original.repeat([batch_size]
|
||||
+ [1] * len(additional['control'].cond_hint_original.shape[1:]))], dim=0)
|
||||
cloned_controlnet = additional['control'].copy()
|
||||
cloned_controlnet.set_cond_hint(control_hint, strength=additional['control'].strength, timestep_percent_range=additional['control'].timestep_percent_range)
|
||||
additional['control'] = cloned_controlnet
|
||||
batched_condition.append([batch_with_reference, additional])
|
||||
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
samples = comfy.sample.sample_custom(
|
||||
m,
|
||||
noise,
|
||||
cfg,
|
||||
sampler,
|
||||
sigmas,
|
||||
batched_condition,
|
||||
negative,
|
||||
latent_t,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=noise_seed,
|
||||
)
|
||||
|
||||
# remove reference image
|
||||
samples = samples[1:]
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent.copy()
|
||||
x0 = x0_output["x0"][1:]
|
||||
out_denoised["samples"] = m.model.process_latent_out(x0.cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
return (out, out_denoised)
|
||||
|
||||
|
||||
class StyleAlignedBatchAlign:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"share_norm": (SHARE_NORM_OPTIONS,),
|
||||
"share_attn": (SHARE_ATTN_OPTIONS,),
|
||||
"scale": ("FLOAT", {"default": 1, "min": 0, "max": 1.0, "step": 0.1}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
# CATEGORY = "style_aligned"
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
def patch(
|
||||
self,
|
||||
model: ModelPatcher,
|
||||
share_norm: str,
|
||||
share_attn: str,
|
||||
scale: float,
|
||||
):
|
||||
m = model.clone()
|
||||
share_group_norm = share_norm in ["group", "both"]
|
||||
share_layer_norm = share_norm in ["layer", "both"]
|
||||
register_shared_norm(model, share_group_norm, share_layer_norm)
|
||||
args = StyleAlignedArgs(share_attn)
|
||||
m.set_model_attn1_patch(SharedAttentionProcessor(args, scale))
|
||||
return (m,)
|
||||
|
||||
|
||||
@@ -12,23 +12,31 @@ import comfy.utils
|
||||
# import numpy as np
|
||||
import torch
|
||||
import random
|
||||
from lark import Lark, Transformer, v_args
|
||||
|
||||
|
||||
global _available
|
||||
_available=True
|
||||
|
||||
def get_text_generator_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('prompt_generator')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "prompt_generator")
|
||||
|
||||
text_generator_model_path=os.path.join(folder_paths.models_dir, "prompt_generator/text2image-prompt-generator")
|
||||
prompt_generator=get_text_generator_path()
|
||||
|
||||
text_generator_model_path=os.path.join(prompt_generator, "text2image-prompt-generator")
|
||||
if not os.path.exists(text_generator_model_path):
|
||||
print(f"## text_generator_model not found: {text_generator_model_path}, pls download from https://huggingface.co/succinctly/text2image-prompt-generator/tree/main")
|
||||
text_generator_model_path='succinctly/text2image-prompt-generator'
|
||||
|
||||
zh_en_model_path=os.path.join(folder_paths.models_dir, "prompt_generator/opus-mt-zh-en")
|
||||
zh_en_model_path=os.path.join(prompt_generator, "opus-mt-zh-en")
|
||||
if not os.path.exists(zh_en_model_path):
|
||||
print(f"## zh_en_model not found: {zh_en_model_path}, pls download from https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main")
|
||||
zh_en_model_path='Helsinki-NLP/opus-mt-zh-en'
|
||||
|
||||
|
||||
|
||||
def is_installed(package):
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
@@ -62,7 +70,14 @@ except:
|
||||
|
||||
|
||||
|
||||
def translate(zh_en_tokenizer,zh_en_model,text):
|
||||
def translate(text):
|
||||
global text_pipe,zh_en_model,zh_en_tokenizer
|
||||
|
||||
if zh_en_model==None:
|
||||
zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval()
|
||||
zh_en_tokenizer = AutoTokenizer.from_pretrained(zh_en_model_path,padding=True, truncation=True)
|
||||
|
||||
zh_en_model.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
with torch.no_grad():
|
||||
encoded = zh_en_tokenizer([text], return_tensors="pt")
|
||||
encoded.to(zh_en_model.device)
|
||||
@@ -102,18 +117,24 @@ def text_generate(text_pipe,input,seed=None):
|
||||
|
||||
import re
|
||||
|
||||
def correct_prompt_syntax(prompt):
|
||||
def correct_prompt_syntax(prompt=""):
|
||||
|
||||
print("input prompt",prompt)
|
||||
# print("input prompt",prompt)
|
||||
corrected_elements = []
|
||||
# 处理成统一的英文标点
|
||||
prompt = prompt.replace('(', '(').replace(')', ')').replace(',', ',').replace(';', ',').replace('。', '.').replace(':',':')
|
||||
# 删除多余的空格
|
||||
prompt = re.sub(r'\s+', ' ', prompt).strip()
|
||||
prompt = prompt.replace("< ","<").replace(" >",">").replace("( ","(").replace(" )",")").replace("[ ","[").replace(' ]',']')
|
||||
|
||||
# 分词
|
||||
prompt_elements = prompt.split(',')
|
||||
|
||||
def balance_brackets(element, open_bracket, close_bracket):
|
||||
open_brackets_count = element.count(open_bracket)
|
||||
close_brackets_count = element.count(close_bracket)
|
||||
return element + close_bracket * (open_brackets_count - close_brackets_count)
|
||||
|
||||
for element in prompt_elements:
|
||||
element = element.strip()
|
||||
|
||||
@@ -133,21 +154,118 @@ def correct_prompt_syntax(prompt):
|
||||
corrected_elements.append(corrected_element)
|
||||
|
||||
# 重组修正后的prompt
|
||||
corrected_prompt = ', '.join(corrected_elements)
|
||||
print("output prompt",corrected_prompt)
|
||||
return corrected_prompt
|
||||
return ','.join(corrected_elements)
|
||||
|
||||
def balance_brackets(element, open_bracket, close_bracket):
|
||||
open_brackets_count = element.count(open_bracket)
|
||||
close_brackets_count = element.count(close_bracket)
|
||||
return element + close_bracket * (open_brackets_count - close_brackets_count)
|
||||
|
||||
# # 示例使用
|
||||
# test_prompt = "((middle-century castles)), [forsaken: 0.8], (mystery dragons: 1.3, mist forests, sunsets, quiet; (((dummy)), [fisting city: 0.5] background, radiant, soft and flavoured,] promising mountains, ((starry: 1.6), [[crowds], [middle-century castle: urban landscapes of the future: 0.5], [yellow: bright sun: 0.7], overlooking"
|
||||
# corrected_prompt = correct_prompt_syntax(test_prompt)
|
||||
# print(corrected_prompt)
|
||||
|
||||
def detect_language(input_str):
|
||||
# 统计中文和英文字符的数量
|
||||
count_cn = count_en = 0
|
||||
for char in input_str:
|
||||
if '\u4e00' <= char <= '\u9fff':
|
||||
count_cn += 1
|
||||
elif char.isalpha():
|
||||
count_en += 1
|
||||
|
||||
# 根据统计的字符数量判断主要语言
|
||||
if count_cn > count_en:
|
||||
return "cn"
|
||||
elif count_en > count_cn:
|
||||
return "en"
|
||||
else:
|
||||
return "unknow"
|
||||
|
||||
|
||||
|
||||
|
||||
#定义Prompt文法
|
||||
grammar = """
|
||||
start: sentence
|
||||
sentence: phrase ("," phrase)*
|
||||
phrase: emphasis | weight | word | lora | embedding | schedule
|
||||
emphasis: "(" sentence ")" -> emphasis
|
||||
| "[" sentence "]" -> weak_emphasis
|
||||
weight: "(" word ":" NUMBER ")"
|
||||
schedule: "[" word ":" word ":" NUMBER "]"
|
||||
lora: "<" WORD ":" WORD (":" NUMBER)? (":" NUMBER)? ">"
|
||||
embedding: "embedding" ":" WORD (":" NUMBER)? (":" NUMBER)?
|
||||
word: WORD
|
||||
|
||||
NUMBER: /\s*-?\d+(\.\d+)?\s*/
|
||||
WORD: /[^,:\(\)\[\]<>]+/
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@v_args(inline=True) # Decorator to flatten the tree directly into the function arguments
|
||||
class ChinesePromptTranslate(Transformer):
|
||||
|
||||
def sentence(self, *args):
|
||||
return ", ".join(args)
|
||||
|
||||
def phrase(self, *args):
|
||||
return "".join(args)
|
||||
|
||||
def emphasis(self, *args):
|
||||
# Reconstruct the emphasis with translated content
|
||||
return "(" + "".join(args) + ")"
|
||||
|
||||
def weak_emphasis(self, *args):
|
||||
print('weak_emphasis:',args)
|
||||
return "[" + "".join(args) + "]"
|
||||
|
||||
def embedding(self,*args):
|
||||
print('prompt embedding',args[0])
|
||||
if len(args) == 1:
|
||||
# print('prompt embedding',str(args[0]))
|
||||
# 只传递了一个参数,意味着只有embedding名称没有数字
|
||||
embedding_name = str(args[0])
|
||||
return f"embedding:{embedding_name}"
|
||||
elif len(args) > 1:
|
||||
embedding_name,*numbers = args
|
||||
|
||||
if len(numbers)==2:
|
||||
return f"embedding:{embedding_name}:{numbers[0]}:{numbers[1]}"
|
||||
elif len(numbers)==1:
|
||||
return f"embedding:{embedding_name}:{numbers[0]}"
|
||||
else:
|
||||
return f"embedding:{embedding_name}"
|
||||
|
||||
def lora(self,*args):
|
||||
print('lora prompt',*args)
|
||||
if len(args) == 1:
|
||||
return f"<lora:{loar_name}>"
|
||||
elif len(args) > 1:
|
||||
# print('lora', args)
|
||||
_,loar_name,*numbers = args
|
||||
loar_name = str(loar_name).strip()
|
||||
if len(numbers)==2:
|
||||
return f"<lora:{loar_name}:{numbers[0]}:{numbers[1]}>"
|
||||
elif len(numbers)==1:
|
||||
return f"<lora:{loar_name}:{numbers[0]}>"
|
||||
else:
|
||||
return f"<lora:{loar_name}>"
|
||||
|
||||
def weight(self, word,number):
|
||||
translated_word = translate(str(word)).rstrip('.')
|
||||
return f"({translated_word}:{str(number).strip()})"
|
||||
|
||||
def schedule(self,*args):
|
||||
print('prompt schedule',args)
|
||||
data = [str(arg).strip() for arg in args]
|
||||
|
||||
return f"[{':'.join(data)}]"
|
||||
|
||||
def word(self, word):
|
||||
# Translate each word using the dictionary
|
||||
if detect_language(str(word)) == "cn":
|
||||
return translate(str(word)).rstrip('.')
|
||||
else:
|
||||
return str(word).rstrip('.')
|
||||
|
||||
class ChinesePrompt:
|
||||
|
||||
@@ -162,7 +280,7 @@ class ChinesePrompt:
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 0xffffffffffffffff}),
|
||||
|
||||
},
|
||||
|
||||
@@ -185,16 +303,16 @@ class ChinesePrompt:
|
||||
zh_en_tokenizer=None
|
||||
|
||||
def run(self,text,seed,generation):
|
||||
global text_pipe,zh_en_model,zh_en_tokenizer
|
||||
|
||||
|
||||
seed=seed[0]
|
||||
generation=generation[0]
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(text)+1)
|
||||
|
||||
texts = [correct_prompt_syntax(t) for t in text]
|
||||
print('correct_prompt_syntax::',texts)
|
||||
|
||||
global text_pipe,zh_en_model,zh_en_tokenizer
|
||||
if zh_en_model==None:
|
||||
zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval()
|
||||
zh_en_tokenizer = AutoTokenizer.from_pretrained(zh_en_model_path,padding=True, truncation=True)
|
||||
@@ -210,9 +328,18 @@ class ChinesePrompt:
|
||||
|
||||
# print('zh_en_model device',zh_en_model.device,text_pipe.model.device,torch.cuda.current_device() )
|
||||
en_texts=[]
|
||||
|
||||
for t in texts:
|
||||
en_text=translate(zh_en_tokenizer,zh_en_model,t)
|
||||
en_texts.append(en_text)
|
||||
if t:
|
||||
parser = Lark(grammar, start="start", parser="lalr", transformer=ChinesePromptTranslate())
|
||||
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)
|
||||
@@ -232,8 +359,11 @@ class ChinesePrompt:
|
||||
pbar.update(1)
|
||||
|
||||
text_pipe.model.to('cpu')
|
||||
prompt_result = [correct_prompt_syntax(p) for p in prompt_result]
|
||||
|
||||
|
||||
print('prompt_result',prompt_result,)
|
||||
# prompt_result = [','.join(correct_prompt_syntax(p)) for p in prompt_result]
|
||||
if len(prompt_result)==0:
|
||||
prompt_result=[""]
|
||||
return {
|
||||
"ui":{
|
||||
"prompt": prompt_result
|
||||
@@ -241,6 +371,8 @@ class ChinesePrompt:
|
||||
"result": (prompt_result,)}
|
||||
|
||||
|
||||
|
||||
|
||||
class PromptGenerate:
|
||||
|
||||
global _available
|
||||
@@ -254,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}),
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
import sys
|
||||
from os import path
|
||||
sys.path.insert(0, path.dirname(__file__))
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from folder_paths import get_folder_paths, get_full_path, get_save_image_path, get_output_directory,models_dir
|
||||
from comfy.model_management import get_torch_device
|
||||
from .tsr.system import TSR
|
||||
|
||||
import comfy.utils
|
||||
|
||||
|
||||
def get_triposr_model_path():
|
||||
try:
|
||||
return path.join(get_folder_paths('triposr')[0],'model.ckpt')
|
||||
except:
|
||||
return path.join(path.join(models_dir, "triposr"),'model.ckpt')
|
||||
|
||||
triposr_model_path=get_triposr_model_path()
|
||||
|
||||
|
||||
# Tensor to PIL
|
||||
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 fill_background(image):
|
||||
im = np.array(image).astype(np.float32) / 255.0
|
||||
im = im[:, :, :3] * im[:, :, 3:4] + (1 - im[:, :, 3:4]) * 0.5
|
||||
im = Image.fromarray((im * 255.0).astype(np.uint8))
|
||||
return im
|
||||
|
||||
|
||||
class LoadTripoSRModel:
|
||||
def __init__(self):
|
||||
self.initialized_model = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
# "model": (get_filename_list("checkpoints"),),
|
||||
"chunk_size": ("INT", {"default": 8192, "min": 0, "max": 10000})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TRIPOSR_MODEL",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/3D/TripoSR"
|
||||
|
||||
def run(self, chunk_size):
|
||||
device = get_torch_device()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
device = "cpu"
|
||||
|
||||
if not self.initialized_model:
|
||||
# triposr_model_path
|
||||
print("#Loading TripoSR model",triposr_model_path)
|
||||
self.initialized_model = TSR.from_pretrained_custom(
|
||||
weight_path=triposr_model_path,
|
||||
config_path=path.join(path.dirname(__file__), "tsr/config.yaml")
|
||||
)
|
||||
self.initialized_model.renderer.set_chunk_size(chunk_size)
|
||||
self.initialized_model.to(device)
|
||||
|
||||
return (self.initialized_model,)
|
||||
|
||||
|
||||
class TripoSRSampler:
|
||||
def __init__(self):
|
||||
self.initialized_model = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("TRIPOSR_MODEL",),
|
||||
"image": ("IMAGE",),
|
||||
"resolution": ("INT", {"default": 256, "min": 128, "max": 12288}),
|
||||
"threshold": ("FLOAT", {"default": 25.0, "min": 0.0, "step": 0.01}),
|
||||
"device":(["auto","cpu"],),
|
||||
},
|
||||
"optional": {
|
||||
"mask": ("MASK",)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MESH",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/3D/TripoSR"
|
||||
|
||||
def run(self, model, image, resolution, threshold,device='auto', mask=None):
|
||||
|
||||
reference_image=image
|
||||
reference_mask=mask
|
||||
|
||||
device = get_torch_device()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
device = "cpu"
|
||||
|
||||
if device=='cpu':
|
||||
device = "cpu"
|
||||
|
||||
print('#TripoSRSampler device',device)
|
||||
|
||||
to_images=[]
|
||||
|
||||
for i in range(len(reference_image)):
|
||||
|
||||
image = reference_image[i]
|
||||
|
||||
if reference_mask is not None:
|
||||
mask = reference_mask[i].unsqueeze(2)
|
||||
image = torch.cat((image, mask), dim=2).detach().cpu().numpy()
|
||||
image = Image.fromarray(np.clip(255. * image, 0, 255).astype(np.uint8))
|
||||
image = fill_background(image)
|
||||
else:
|
||||
image = tensor2pil(image)
|
||||
|
||||
image = image.convert('RGB')
|
||||
|
||||
to_images.append(image)
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(to_images))
|
||||
def callback(c):
|
||||
pbar.update(1)
|
||||
|
||||
scene_codes = model(to_images, device)
|
||||
meshes = model.extract_mesh(scene_codes, resolution=resolution, threshold=threshold,callback=callback)
|
||||
|
||||
del model
|
||||
return (meshes,)
|
||||
|
||||
|
||||
class SaveTripoSRMesh:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mesh": ("MESH",),
|
||||
# "format":(["glb","obj"],),
|
||||
"filename_prefix":("STRING", {"multiline": False,"default": "TripoSR_"})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/3D/TripoSR"
|
||||
|
||||
def run(self, mesh,filename_prefix):
|
||||
format='glb'
|
||||
saved = list()
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = get_save_image_path(filename_prefix,
|
||||
get_output_directory())
|
||||
|
||||
for (index, single_mesh) in enumerate(mesh):
|
||||
filename_with_batch_num = filename.replace("%batch_num%", str(index))
|
||||
file = f"{filename_with_batch_num}_{counter:05}_.{format}"
|
||||
single_mesh.apply_transform(np.array([[1, 0, 0, 0], [0, 0, 1, 0], [0, -1, 0, 0], [0, 0, 0, 1]]))
|
||||
single_mesh.export(path.join(full_output_folder, file))
|
||||
saved.append({
|
||||
"filename": file,
|
||||
"type": "output",
|
||||
"subfolder": subfolder
|
||||
})
|
||||
|
||||
return {"ui": {"mesh": saved}}
|
||||
|
||||
|
||||
|
||||
@@ -8,6 +8,18 @@ import matplotlib.font_manager as fm
|
||||
import torch
|
||||
import importlib.util
|
||||
|
||||
def create_incrementing_list(min_value, max_value, step, count):
|
||||
l1 = [int(min_value + i * step) for i in range(count) if min_value + i * step <= max_value]
|
||||
l2 = [float(min_value + i * step) for i in range(count) if min_value + i * step <= max_value]
|
||||
return (l1,l2)
|
||||
|
||||
def split_list(lst, chunk_size, transition_size):
|
||||
result = []
|
||||
for i in range(0, len(lst), chunk_size):
|
||||
start = i - transition_size
|
||||
end = i + chunk_size + transition_size
|
||||
result.append(lst[max(start, 0):end])
|
||||
return result
|
||||
|
||||
def recursive_search(directory, excluded_dir_names=None):
|
||||
if not os.path.isdir(directory):
|
||||
@@ -70,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)
|
||||
@@ -121,7 +133,7 @@ def get_font_files(directory):
|
||||
|
||||
return font_files
|
||||
|
||||
r_directory = os.path.join(os.path.dirname(__file__), '../assets/')
|
||||
r_directory = os.path.join(os.path.dirname(__file__), '..','assets','/')
|
||||
|
||||
font_files = get_font_files(r_directory)
|
||||
# print(font_files)
|
||||
@@ -146,7 +158,6 @@ class ColorInput:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
|
||||
"color":("TCOLOR",),
|
||||
},
|
||||
}
|
||||
@@ -156,7 +167,7 @@ class ColorInput:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Color"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,False,False,False,False,)
|
||||
@@ -170,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
|
||||
@@ -185,7 +218,7 @@ class FontInput:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -218,7 +251,7 @@ class TextToNumber:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -273,10 +306,10 @@ class FloatSlider:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
|
||||
RETURN_NAMES = ('FLOAT',)
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -286,9 +319,7 @@ class FloatSlider:
|
||||
number = min_value
|
||||
elif number > max_value:
|
||||
number = max_value
|
||||
scaled_number = (number - min_value) / (max_value - min_value)
|
||||
return (scaled_number,)
|
||||
|
||||
return (number,)
|
||||
|
||||
class IntNumber:
|
||||
@classmethod
|
||||
@@ -329,7 +360,7 @@ class IntNumber:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -389,7 +420,7 @@ class TextInput:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -398,6 +429,61 @@ class TextInput:
|
||||
|
||||
return (text,)
|
||||
|
||||
|
||||
class IncrementingListNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"min_value": ("FLOAT", {
|
||||
"default": 0,
|
||||
"min": -2000, #Minimum value
|
||||
"max": 0xffffffffffffffff,
|
||||
"step": 0.01, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"max_value": ("FLOAT", {
|
||||
"default": 10,
|
||||
"min": -2000, #Minimum value
|
||||
"max": 0xffffffffffffffff,
|
||||
"step": 0.01, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"step": ("FLOAT", {
|
||||
"default": 0,
|
||||
"min": -2000, #Minimum value
|
||||
"max": 0xffffffffffffffff,
|
||||
"step": 0.01, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"count": ("INT", {
|
||||
"default": 1,
|
||||
"min": 1, #Minimum value
|
||||
"max": 0xffffffffffffffff,
|
||||
"step":1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
})
|
||||
},
|
||||
"optional":{
|
||||
"seed":("INT", {"default": -1, "min": -1, "max": 1000000}),
|
||||
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT","FLOAT",)
|
||||
RETURN_NAMES = ('int_list','float_list',)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (True,True,)
|
||||
|
||||
def run(self,min_value,max_value,step,count,seed):
|
||||
print('create_incrementing_list',seed)
|
||||
l1,l2=create_incrementing_list(min_value,max_value,step,count)
|
||||
return (l1,l2,)
|
||||
|
||||
# 接收一个值,然后根据字符串或数值长度计算延迟时间,用户可以自定义延迟"字/s",延迟之后将转化
|
||||
|
||||
import comfy.samplers
|
||||
@@ -502,7 +588,7 @@ class AppInfo:
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"IMAGE": ("IMAGE",),
|
||||
"image": ("IMAGE",),
|
||||
"description":("STRING",{"multiline": True,"default": "","dynamicPrompts": False}),
|
||||
"version":("INT", {
|
||||
"default": 1,
|
||||
@@ -515,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},),
|
||||
}
|
||||
|
||||
}
|
||||
@@ -530,14 +617,21 @@ class AppInfo:
|
||||
INPUT_IS_LIST = True
|
||||
# OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self,name,input_ids,output_ids,IMAGE,description,version,share_prefix,link,category,auto_save):
|
||||
def run(self,name,input_ids,output_ids,image,description,version,share_prefix,link,category,auto_save,idle_animation):
|
||||
name=name[0]
|
||||
|
||||
idle_animation=idle_animation[0]
|
||||
|
||||
im=None
|
||||
if IMAGE:
|
||||
im=IMAGE[0][0]
|
||||
#TODO batch 的方式需要处理
|
||||
im=create_temp_file(im)
|
||||
im=[]
|
||||
if image:
|
||||
images=[image]
|
||||
# batch 的方式需要处理
|
||||
images=flatten_list(images)
|
||||
# img=image[0][0]
|
||||
print('AppInfo_image',len(images))
|
||||
for i in range(len(images)):
|
||||
img=images[i]
|
||||
im.append(create_temp_file(img,i+1)[0])
|
||||
# image [img,] img[batch,w,h,a] 列表里面是batch,
|
||||
|
||||
input_ids=input_ids[0]
|
||||
@@ -550,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),)
|
||||
|
||||
|
||||
|
||||
@@ -577,14 +714,14 @@ class SwitchByIndex:
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,"INT",)
|
||||
RETURN_NAMES = ("C","count",)
|
||||
RETURN_NAMES = ("list", "count",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,False,)
|
||||
OUTPUT_IS_LIST = (True, False,)
|
||||
|
||||
def run(self, A=[],B=[],index=-1,flat='on'):
|
||||
|
||||
@@ -605,10 +742,43 @@ class SwitchByIndex:
|
||||
try:
|
||||
C=[C[index]]
|
||||
except Exception as e:
|
||||
C=[]
|
||||
C=[C[-1]] #最后一个
|
||||
|
||||
return (C,len(C),)
|
||||
return (C, len(C),)
|
||||
|
||||
class ListSplit:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"optional":{
|
||||
"A":(any_type,),
|
||||
},
|
||||
"required": {
|
||||
"chunk_size": ("INT", {"default": 10, "min": 1, "step": 1}),
|
||||
"transition_size": ("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"index": ("INT", {"default": -1, "min": -1, "step": 1}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ("B",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self, A=[],chunk_size=[10],transition_size=[0],index=[-1]):
|
||||
# print(len(A))
|
||||
B=split_list(A,chunk_size[0],transition_size[0])
|
||||
|
||||
if index[0]>-1:
|
||||
B=B[index[0]]
|
||||
|
||||
return (B,)
|
||||
|
||||
|
||||
|
||||
class LimitNumber:
|
||||
@@ -639,7 +809,7 @@ class LimitNumber:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -698,7 +868,7 @@ class TESTNODE_:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"ANY":(any_type,),
|
||||
"ANY":(any_type,),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -706,14 +876,24 @@ class TESTNODE_:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/__TEST"
|
||||
CATEGORY = "♾️Mixlab/Test"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self,ANY):
|
||||
# print(ANY)
|
||||
|
||||
print('#TESTNODE_',len(ANY))
|
||||
|
||||
print(type(ANY))
|
||||
try:
|
||||
print(ANY[0].shape)
|
||||
img= tensor2pil(ANY[0])
|
||||
print(img.size)
|
||||
except:
|
||||
print('')
|
||||
|
||||
# data=ANY
|
||||
list_stats = ListStatistics()
|
||||
|
||||
@@ -751,7 +931,7 @@ class TESTNODE_TOKEN:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/__TEST"
|
||||
CATEGORY = "♾️Mixlab/Test"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = False
|
||||
@@ -788,7 +968,7 @@ class CreateSeedNode:
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, seed):
|
||||
return (seed,)
|
||||
@@ -815,7 +995,7 @@ class CreateCkptNames:
|
||||
# OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, ckpt_names):
|
||||
ckpt_names=ckpt_names.split('\n')
|
||||
@@ -844,7 +1024,7 @@ class CreateLoraNames:
|
||||
# OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, lora_names):
|
||||
lora_names=lora_names.split('\n')
|
||||
@@ -875,7 +1055,7 @@ class CreateSampler_names:
|
||||
# OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, sampler_names):
|
||||
sampler_names=sampler_names.split('\n')
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
from comfy.ldm.modules.attention import default, optimized_attention, optimized_attention_masked
|
||||
from .style_functions import adain, concat_first
|
||||
|
||||
class VisualStyleProcessor(object):
|
||||
def __init__(self,
|
||||
module_self,
|
||||
keys_scale: float = 1.0,
|
||||
enabled: bool = True,
|
||||
adain_queries: bool = True,
|
||||
adain_keys: bool = True,
|
||||
adain_values: bool = False
|
||||
):
|
||||
self.module_self = module_self
|
||||
self.keys_scale = keys_scale
|
||||
self.enabled = enabled
|
||||
self.adain_queries = adain_queries
|
||||
self.adain_keys = adain_keys
|
||||
self.adain_values = adain_values
|
||||
|
||||
def visual_style_forward(self, x, context, value, mask=None):
|
||||
q = self.module_self.to_q(x)
|
||||
context = default(context, x)
|
||||
k = self.module_self.to_k(context)
|
||||
if value is not None:
|
||||
v = self.module_self.to_v(value)
|
||||
del value
|
||||
else:
|
||||
v = self.module_self.to_v(context)
|
||||
|
||||
if self.enabled:
|
||||
if self.adain_queries:
|
||||
q = adain(q)
|
||||
if self.adain_keys:
|
||||
k = adain(k)
|
||||
if self.adain_values:
|
||||
v = adain(v)
|
||||
|
||||
k = concat_first(k, -2, self.keys_scale)
|
||||
v = concat_first(v, -2)
|
||||
|
||||
if mask is None:
|
||||
out = optimized_attention(q, k, v, self.module_self.heads)
|
||||
else:
|
||||
out = optimized_attention_masked(q, k, v, self.module_self.heads, mask)
|
||||
return self.module_self.to_out(out)
|
||||
@@ -0,0 +1,60 @@
|
||||
import torch
|
||||
|
||||
from einops import rearrange
|
||||
from dataclasses import dataclass
|
||||
|
||||
T = torch.Tensor
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StyleAlignedArgs:
|
||||
share_group_norm: bool = True
|
||||
share_layer_norm: bool = True,
|
||||
share_attention: bool = True
|
||||
adain_queries: bool = True
|
||||
adain_keys: bool = True
|
||||
adain_values: bool = False
|
||||
full_attention_share: bool = False
|
||||
keys_scale: float = 1.
|
||||
only_self_level: float = 0.
|
||||
|
||||
def expand_first(feat: T, scale=1., ) -> T:
|
||||
b = feat.shape[0]
|
||||
feat_style = torch.stack((feat[0], feat[b // 2])).unsqueeze(1)
|
||||
if scale == 1:
|
||||
feat_style = feat_style.expand(2, b // 2, *feat.shape[1:])
|
||||
else:
|
||||
feat_style = feat_style.repeat(1, b // 2, 1, 1, 1)
|
||||
feat_style = torch.cat([feat_style[:, :1], scale * feat_style[:, 1:]], dim=1)
|
||||
return feat_style.reshape(*feat.shape)
|
||||
|
||||
|
||||
def concat_first(feat: T, dim=2, scale=1.) -> T:
|
||||
feat_style = expand_first(feat, scale=scale)
|
||||
return torch.cat((feat, feat_style), dim=dim)
|
||||
|
||||
|
||||
def calc_mean_std(feat, eps: float = 1e-5) -> tuple[T, T]:
|
||||
feat_std = (feat.var(dim=-2, keepdims=True) + eps).sqrt()
|
||||
feat_mean = feat.mean(dim=-2, keepdims=True)
|
||||
return feat_mean, feat_std
|
||||
|
||||
|
||||
def adain(feat: T) -> T:
|
||||
feat_mean, feat_std = calc_mean_std(feat)
|
||||
feat_style_mean = expand_first(feat_mean)
|
||||
feat_style_std = expand_first(feat_std)
|
||||
feat = (feat - feat_mean) / feat_std
|
||||
feat = feat * feat_style_std + feat_style_mean
|
||||
return feat
|
||||
|
||||
def swapping_attention(key, value, chunk_size=2):
|
||||
chunk_length = key.size()[0] // chunk_size # [text-condition, null-condition]
|
||||
reference_image_index = [0] * chunk_length # [0 0 0 0 0]
|
||||
key = rearrange(key, "(b f) d c -> b f d c", f=chunk_length)
|
||||
key = key[:, reference_image_index] # ref to all
|
||||
key = rearrange(key, "b f d c -> (b f) d c")
|
||||
value = rearrange(value, "(b f) d c -> b f d c", f=chunk_length)
|
||||
value = value[:, reference_image_index] # ref to all
|
||||
value = rearrange(value, "b f d c -> (b f) d c")
|
||||
|
||||
return key, value
|
||||
@@ -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,)
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
import torch
|
||||
from PIL import Image, ImageOps, ImageSequence, ImageFile
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
|
||||
import numpy as np
|
||||
import os
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
import hashlib
|
||||
from uuid import uuid4
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# tensor 取hash值
|
||||
def tensor_to_hash(tensor):
|
||||
# 将 Tensor 转换为 NumPy 数组
|
||||
np_array = tensor.cpu().numpy()
|
||||
|
||||
# 将 NumPy 数组转换为字节数据
|
||||
byte_data = np_array.tobytes()
|
||||
|
||||
# 计算哈希值
|
||||
hash_value = hashlib.md5(byte_data).hexdigest()
|
||||
|
||||
return hash_value
|
||||
|
||||
|
||||
def create_temp_file(image, uuid):
|
||||
output_dir = folder_paths.get_temp_directory()
|
||||
|
||||
(
|
||||
full_output_folder,
|
||||
filename,
|
||||
counter,
|
||||
subfolder,
|
||||
_,
|
||||
) = folder_paths.get_save_image_path(f'material_{uuid}', output_dir)
|
||||
|
||||
|
||||
image=tensor2pil(image)
|
||||
|
||||
image_file = f"{filename}_{counter:05}.png"
|
||||
|
||||
image_path=os.path.join(full_output_folder, image_file)
|
||||
|
||||
image.save(image_path,compress_level=4)
|
||||
|
||||
return (image_path,[{
|
||||
"filename": image_file,
|
||||
"subfolder": subfolder,
|
||||
"type": "temp"
|
||||
}])
|
||||
|
||||
|
||||
# image - tensor - 文件路径
|
||||
# loadImage的方法( 文件路径 - image-mask )
|
||||
class EditMask:
|
||||
|
||||
def __init__(self):
|
||||
self.image_id = None
|
||||
self.uuid = str(uuid4())
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"image": ("IMAGE",), # 表示一个张量
|
||||
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"image_update": ("IMAGE_FILE",)
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("image", "mask")
|
||||
|
||||
FUNCTION = "edit"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def edit(self, image,image_update=None):
|
||||
|
||||
# 根据image输入来判断是否是新的图片
|
||||
if self.image_id==None:
|
||||
self.image_id=tensor_to_hash(image)
|
||||
image_update=None
|
||||
else:
|
||||
image_id=tensor_to_hash(image)
|
||||
if image_id!=self.image_id:
|
||||
image_update=None
|
||||
self.image_id=image_id
|
||||
|
||||
|
||||
image_path=None
|
||||
# print('#image_update',self.image_id,image_update)
|
||||
if image_update==None:
|
||||
print('--')
|
||||
else:
|
||||
if 'images' in image_update:
|
||||
images=image_update['images']
|
||||
filename=images[0]['filename']
|
||||
subfolder=images[0]['subfolder']
|
||||
type=images[0]['type']
|
||||
name, base_dir=folder_paths.annotated_filepath(filename)
|
||||
if type.endswith("output"):
|
||||
base_dir = folder_paths.get_output_directory()
|
||||
elif type.endswith("input"):
|
||||
base_dir = folder_paths.get_input_directory()
|
||||
elif type.endswith("temp"):
|
||||
base_dir = folder_paths.get_temp_directory()
|
||||
#base_dir = folder_paths.get_input_directory()
|
||||
# print(base_dir,subfolder, name)
|
||||
image_path = os.path.join(base_dir,subfolder, name)
|
||||
|
||||
if image_path==None:
|
||||
image_path,images=create_temp_file(image, self.uuid)
|
||||
|
||||
print('#image_path',os.path.exists(image_path),image_path)
|
||||
# image_path = folder_paths.get_annotated_filepath(image) #文件名
|
||||
|
||||
if not os.path.exists(image_path):
|
||||
image_path,images=create_temp_file(image, self.uuid)
|
||||
|
||||
|
||||
img = node_helpers.pillow(Image.open, image_path)
|
||||
|
||||
output_images = []
|
||||
output_masks = []
|
||||
w, h = None, None
|
||||
|
||||
excluded_formats = ['MPO']
|
||||
|
||||
for i in ImageSequence.Iterator(img):
|
||||
i = node_helpers.pillow(ImageOps.exif_transpose, i)
|
||||
|
||||
if i.mode == 'I':
|
||||
i = i.point(lambda i: i * (1 / 255))
|
||||
image = i.convert("RGB")
|
||||
|
||||
if len(output_images) == 0:
|
||||
w = image.size[0]
|
||||
h = image.size[1]
|
||||
|
||||
if image.size[0] != w or image.size[1] != h:
|
||||
continue
|
||||
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[None,]
|
||||
if 'A' in i.getbands():
|
||||
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
|
||||
mask = 1. - torch.from_numpy(mask)
|
||||
else:
|
||||
# 尺寸不对,需要按照image来
|
||||
mask = torch.zeros((h, w), dtype=torch.float32, device="cpu")
|
||||
|
||||
output_images.append(image)
|
||||
output_masks.append(mask.unsqueeze(0))
|
||||
|
||||
if len(output_images) > 1 and img.format not in excluded_formats:
|
||||
output_image = torch.cat(output_images, dim=0)
|
||||
output_mask = torch.cat(output_masks, dim=0)
|
||||
else:
|
||||
output_image = output_images[0]
|
||||
output_mask = output_masks[0]
|
||||
|
||||
return {"ui":{"images": images},"result": (output_image, output_mask)}
|
||||
|
||||
# return (output_image, output_mask)
|
||||
@@ -0,0 +1,69 @@
|
||||
import itertools
|
||||
import re
|
||||
|
||||
LANGUAGE_UNICODE_RANGE_MAP = {
|
||||
"ZH": [(0x4E00, 0x9FFF)],
|
||||
"JP": [(0x4E00, 0x9FFF), (0x3040, 0x309F), (0x30A0, 0x30FF), (0x31F0, 0x31FF)],
|
||||
"EN": [(0x0000, 0x007F)],
|
||||
}
|
||||
|
||||
SYMBOLS_MAPPING = {
|
||||
":": ",",
|
||||
";": ",",
|
||||
",": ",",
|
||||
"。": ".",
|
||||
"!": "!",
|
||||
"?": "?",
|
||||
"\n": ".",
|
||||
"·": ",",
|
||||
"、": ",",
|
||||
"...": "…",
|
||||
"“": "'",
|
||||
"”": "'",
|
||||
"‘": "'",
|
||||
"’": "'",
|
||||
"(": "'",
|
||||
")": "'",
|
||||
"(": "'",
|
||||
")": "'",
|
||||
"《": "'",
|
||||
"》": "'",
|
||||
"【": "'",
|
||||
"】": "'",
|
||||
"[": "'",
|
||||
"]": "'",
|
||||
"—": "-",
|
||||
"~": "-",
|
||||
"~": "-",
|
||||
"・": "-",
|
||||
"「": "'",
|
||||
"」": "'",
|
||||
";": ",",
|
||||
":": ",",
|
||||
}
|
||||
|
||||
REPLACE_SYMBOL_REGEX = re.compile(
|
||||
"|".join(re.escape(p) for p in SYMBOLS_MAPPING.keys())
|
||||
)
|
||||
ALL_KNOWN_UTF8_RANGE = list(
|
||||
itertools.chain.from_iterable(LANGUAGE_UNICODE_RANGE_MAP.values())
|
||||
)
|
||||
REMOVE_UNKNOWN_SYMBOL_REGEX = re.compile(
|
||||
"[^"
|
||||
+ "".join(
|
||||
f"{re.escape(chr(start))}-{re.escape(chr(end))}"
|
||||
for start, end in ALL_KNOWN_UTF8_RANGE
|
||||
)
|
||||
+ "]"
|
||||
)
|
||||
|
||||
|
||||
def clean_text(text):
|
||||
# Clean the text
|
||||
text = text.strip()
|
||||
|
||||
# Replace all chinese symbols with their english counterparts
|
||||
text = REPLACE_SYMBOL_REGEX.sub(lambda x: SYMBOLS_MAPPING[x.group()], text)
|
||||
text = REMOVE_UNKNOWN_SYMBOL_REGEX.sub("", text)
|
||||
|
||||
return text
|
||||
@@ -0,0 +1,87 @@
|
||||
# Base configuration for training a model
|
||||
paths:
|
||||
run_dir: results/${project}
|
||||
ckpt_dir: ${paths.run_dir}/checkpoints
|
||||
|
||||
hydra:
|
||||
run:
|
||||
dir: ${paths.run_dir}
|
||||
|
||||
# Lightning Trainer
|
||||
trainer:
|
||||
_target_: lightning.pytorch.trainer.Trainer
|
||||
|
||||
default_root_dir: ${paths.run_dir}
|
||||
accelerator: gpu
|
||||
num_nodes: 1
|
||||
devices: auto
|
||||
strategy:
|
||||
_target_: lightning.pytorch.strategies.DDPStrategy
|
||||
process_group_backend: nccl # This should be override when training on windows
|
||||
|
||||
precision: bf16-mixed
|
||||
|
||||
# disable validation by epoch end
|
||||
check_val_every_n_epoch: null
|
||||
val_check_interval: 5000
|
||||
max_steps: 100_000
|
||||
|
||||
# Use torch.backends.cudnn.benchmark to speed up training
|
||||
benchmark: true
|
||||
|
||||
# Callbacks
|
||||
callbacks:
|
||||
model_checkpoint:
|
||||
_target_: lightning.pytorch.callbacks.ModelCheckpoint
|
||||
dirpath: ${paths.ckpt_dir}
|
||||
filename: "step_{step:09d}"
|
||||
save_last: false # additionally always save an exact copy of the last checkpoint to a file last.ckpt
|
||||
save_top_k: 5 # save 5 latest checkpoints
|
||||
monitor: step # use step to monitor checkpoints
|
||||
mode: max # save the latest checkpoint with the highest global_step
|
||||
every_n_epochs: null # don't save checkpoints by epoch end
|
||||
every_n_train_steps: 5000 # save checkpoints every 5000 steps
|
||||
auto_insert_metric_name: false
|
||||
|
||||
model_summary:
|
||||
_target_: lightning.pytorch.callbacks.ModelSummary
|
||||
max_depth: 2 # the maximum depth of layer nesting that the summary will include
|
||||
|
||||
learning_rate_monitor:
|
||||
_target_: lightning.pytorch.callbacks.LearningRateMonitor
|
||||
logging_interval: step
|
||||
log_momentum: false
|
||||
|
||||
grad_norm_monitor:
|
||||
_target_: fish_speech.callbacks.GradNormMonitor
|
||||
norm_type: 2
|
||||
logging_interval: step
|
||||
|
||||
# Logger
|
||||
logger:
|
||||
tensorboard:
|
||||
_target_: lightning.pytorch.loggers.tensorboard.TensorBoardLogger
|
||||
save_dir: "${paths.run_dir}/tensorboard/"
|
||||
name: null
|
||||
log_graph: false
|
||||
default_hp_metric: true
|
||||
prefix: ""
|
||||
|
||||
# wandb:
|
||||
# _target_: lightning.pytorch.loggers.wandb.WandbLogger
|
||||
# # name: "" # name of the run (normally generated by wandb)
|
||||
# save_dir: "${paths.run_dir}"
|
||||
# offline: False
|
||||
# id: null # pass correct id to resume experiment!
|
||||
# anonymous: null # enable anonymous logging
|
||||
# project: "fish-speech"
|
||||
# log_model: False # upload lightning ckpts
|
||||
# prefix: "" # a string to put at the beginning of metric keys
|
||||
# # entity: "" # set to name of your wandb team
|
||||
# group: ""
|
||||
# tags: ["vq", "hq", "finetune"]
|
||||
# job_type: ""
|
||||
|
||||
# Loop
|
||||
train: true
|
||||
test: false
|
||||
@@ -0,0 +1,33 @@
|
||||
_target_: fish_speech.models.vqgan.modules.firefly.FireflyArchitecture
|
||||
spec_transform:
|
||||
_target_: fish_speech.utils.spectrogram.LogMelSpectrogram
|
||||
sample_rate: 44100
|
||||
n_mels: 160
|
||||
n_fft: 2048
|
||||
hop_length: 512
|
||||
win_length: 2048
|
||||
backbone:
|
||||
_target_: fish_speech.models.vqgan.modules.firefly.ConvNeXtEncoder
|
||||
input_channels: 160
|
||||
depths: [3, 3, 9, 3]
|
||||
dims: [128, 256, 384, 512]
|
||||
drop_path_rate: 0.2
|
||||
kernel_size: 7
|
||||
head:
|
||||
_target_: fish_speech.models.vqgan.modules.firefly.HiFiGANGenerator
|
||||
hop_length: 512
|
||||
upsample_rates: [8, 8, 2, 2, 2] # aka. strides
|
||||
upsample_kernel_sizes: [16, 16, 4, 4, 4]
|
||||
resblock_kernel_sizes: [3, 7, 11]
|
||||
resblock_dilation_sizes: [[1, 3, 5], [1, 3, 5], [1, 3, 5]]
|
||||
num_mels: 512
|
||||
upsample_initial_channel: 512
|
||||
pre_conv_kernel_size: 13
|
||||
post_conv_kernel_size: 13
|
||||
quantizer:
|
||||
_target_: fish_speech.models.vqgan.modules.fsq.DownsampleFiniteScalarQuantize
|
||||
input_dim: 512
|
||||
n_groups: 8
|
||||
n_codebooks: 1
|
||||
levels: [8, 5, 5, 5]
|
||||
downsample_factor: [2, 2]
|
||||
@@ -0,0 +1,4 @@
|
||||
_target_: fish_speech.models.text2semantic.lora.LoraConfig
|
||||
r: 8
|
||||
lora_alpha: 16
|
||||
lora_dropout: 0.01
|
||||
@@ -0,0 +1,83 @@
|
||||
defaults:
|
||||
- base
|
||||
- _self_
|
||||
|
||||
project: text2semantic_finetune_dual_ar
|
||||
max_length: 4096
|
||||
pretrained_ckpt_path: checkpoints/fish-speech-1.4
|
||||
|
||||
# Lightning Trainer
|
||||
trainer:
|
||||
accumulate_grad_batches: 1
|
||||
gradient_clip_val: 1.0
|
||||
gradient_clip_algorithm: "norm"
|
||||
max_steps: 1000
|
||||
precision: bf16-true
|
||||
limit_val_batches: 10
|
||||
val_check_interval: 100
|
||||
|
||||
# Dataset Configuration
|
||||
tokenizer:
|
||||
_target_: transformers.AutoTokenizer.from_pretrained
|
||||
pretrained_model_name_or_path: ${pretrained_ckpt_path}
|
||||
|
||||
# Dataset Configuration
|
||||
train_dataset:
|
||||
_target_: fish_speech.datasets.semantic.AutoTextSemanticInstructionDataset
|
||||
proto_files:
|
||||
- data/protos
|
||||
tokenizer: ${tokenizer}
|
||||
causal: true
|
||||
max_length: ${max_length}
|
||||
use_speaker: false
|
||||
interactive_prob: 0.7
|
||||
|
||||
val_dataset:
|
||||
_target_: fish_speech.datasets.semantic.AutoTextSemanticInstructionDataset
|
||||
proto_files:
|
||||
- data/protos
|
||||
tokenizer: ${tokenizer}
|
||||
causal: true
|
||||
max_length: ${max_length}
|
||||
use_speaker: false
|
||||
interactive_prob: 0.7
|
||||
|
||||
data:
|
||||
_target_: fish_speech.datasets.semantic.SemanticDataModule
|
||||
train_dataset: ${train_dataset}
|
||||
val_dataset: ${val_dataset}
|
||||
num_workers: 4
|
||||
batch_size: 8
|
||||
tokenizer: ${tokenizer}
|
||||
max_length: ${max_length}
|
||||
|
||||
# Model Configuration
|
||||
model:
|
||||
_target_: fish_speech.models.text2semantic.lit_module.TextToSemantic
|
||||
model:
|
||||
_target_: fish_speech.models.text2semantic.llama.BaseTransformer.from_pretrained
|
||||
path: ${pretrained_ckpt_path}
|
||||
load_weights: true
|
||||
max_length: ${max_length}
|
||||
lora_config: null
|
||||
|
||||
optimizer:
|
||||
_target_: torch.optim.AdamW
|
||||
_partial_: true
|
||||
lr: 1e-4
|
||||
weight_decay: 0
|
||||
betas: [0.9, 0.95]
|
||||
eps: 1e-5
|
||||
|
||||
lr_scheduler:
|
||||
_target_: torch.optim.lr_scheduler.LambdaLR
|
||||
_partial_: true
|
||||
lr_lambda:
|
||||
_target_: fish_speech.scheduler.get_constant_schedule_with_warmup_lr_lambda
|
||||
_partial_: true
|
||||
num_warmup_steps: 10
|
||||
|
||||
# Callbacks
|
||||
callbacks:
|
||||
model_checkpoint:
|
||||
every_n_train_steps: ${trainer.val_check_interval}
|
||||
@@ -0,0 +1,2 @@
|
||||
SEMANTIC_TOKEN = "<|semantic|>"
|
||||
CODEBOOK_PAD_TOKEN_ID = 0
|
||||
@@ -0,0 +1,53 @@
|
||||
import bisect
|
||||
import random
|
||||
from typing import Iterable
|
||||
|
||||
from torch.utils.data import Dataset, IterableDataset
|
||||
|
||||
|
||||
class ConcatRepeatDataset(Dataset):
|
||||
datasets: list[Dataset]
|
||||
cumulative_sizes: list[int]
|
||||
repeats: list[int]
|
||||
|
||||
@staticmethod
|
||||
def cumsum(sequence, repeats):
|
||||
r, s = [], 0
|
||||
for dataset, repeat in zip(sequence, repeats):
|
||||
l = len(dataset) * repeat
|
||||
r.append(l + s)
|
||||
s += l
|
||||
return r
|
||||
|
||||
def __init__(self, datasets: Iterable[Dataset], repeats: list[int]):
|
||||
super().__init__()
|
||||
|
||||
self.datasets = list(datasets)
|
||||
self.repeats = repeats
|
||||
|
||||
assert len(self.datasets) > 0, "datasets should not be an empty iterable"
|
||||
assert len(self.datasets) == len(
|
||||
repeats
|
||||
), "datasets and repeats should have the same length"
|
||||
|
||||
for d in self.datasets:
|
||||
assert not isinstance(
|
||||
d, IterableDataset
|
||||
), "ConcatRepeatDataset does not support IterableDataset"
|
||||
|
||||
self.cumulative_sizes = self.cumsum(self.datasets, self.repeats)
|
||||
|
||||
def __len__(self):
|
||||
return self.cumulative_sizes[-1]
|
||||
|
||||
def __getitem__(self, idx):
|
||||
dataset_idx = bisect.bisect_right(self.cumulative_sizes, idx)
|
||||
|
||||
if dataset_idx == 0:
|
||||
sample_idx = idx
|
||||
else:
|
||||
sample_idx = idx - self.cumulative_sizes[dataset_idx - 1]
|
||||
|
||||
dataset = self.datasets[dataset_idx]
|
||||
|
||||
return dataset[sample_idx % len(dataset)]
|
||||
@@ -0,0 +1,24 @@
|
||||
syntax = "proto3";
|
||||
|
||||
package text_data;
|
||||
|
||||
message Semantics {
|
||||
repeated uint32 values = 1;
|
||||
}
|
||||
|
||||
message Sentence {
|
||||
repeated string texts = 1;
|
||||
repeated Semantics semantics = 3;
|
||||
}
|
||||
|
||||
message TextData {
|
||||
string source = 1;
|
||||
string name = 2;
|
||||
repeated Sentence sentences = 4;
|
||||
}
|
||||
|
||||
message SampledData {
|
||||
string source = 1;
|
||||
string name = 2;
|
||||
repeated Sentence samples = 3;
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Generated by the protocol buffer compiler. DO NOT EDIT!
|
||||
# source: text-data.proto
|
||||
# Protobuf Python Version: 4.25.1
|
||||
"""Generated protocol buffer code."""
|
||||
from google.protobuf import descriptor as _descriptor
|
||||
from google.protobuf import descriptor_pool as _descriptor_pool
|
||||
from google.protobuf import symbol_database as _symbol_database
|
||||
from google.protobuf.internal import builder as _builder
|
||||
|
||||
# @@protoc_insertion_point(imports)
|
||||
|
||||
_sym_db = _symbol_database.Default()
|
||||
|
||||
|
||||
DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(
|
||||
b'\n\x0ftext-data.proto\x12\ttext_data"\x1b\n\tSemantics\x12\x0e\n\x06values\x18\x01 \x03(\r"B\n\x08Sentence\x12\r\n\x05texts\x18\x01 \x03(\t\x12\'\n\tsemantics\x18\x03 \x03(\x0b\x32\x14.text_data.Semantics"P\n\x08TextData\x12\x0e\n\x06source\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12&\n\tsentences\x18\x04 \x03(\x0b\x32\x13.text_data.Sentence"Q\n\x0bSampledData\x12\x0e\n\x06source\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12$\n\x07samples\x18\x03 \x03(\x0b\x32\x13.text_data.Sentenceb\x06proto3'
|
||||
)
|
||||
|
||||
_globals = globals()
|
||||
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
|
||||
_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, "text_data_pb2", _globals)
|
||||
if _descriptor._USE_C_DESCRIPTORS == False:
|
||||
DESCRIPTOR._options = None
|
||||
_globals["_SEMANTICS"]._serialized_start = 30
|
||||
_globals["_SEMANTICS"]._serialized_end = 57
|
||||
_globals["_SENTENCE"]._serialized_start = 59
|
||||
_globals["_SENTENCE"]._serialized_end = 125
|
||||
_globals["_TEXTDATA"]._serialized_start = 127
|
||||
_globals["_TEXTDATA"]._serialized_end = 207
|
||||
_globals["_SAMPLEDDATA"]._serialized_start = 209
|
||||
_globals["_SAMPLEDDATA"]._serialized_end = 290
|
||||
# @@protoc_insertion_point(module_scope)
|
||||
@@ -0,0 +1,36 @@
|
||||
import struct
|
||||
|
||||
from .text_data_pb2 import TextData
|
||||
|
||||
|
||||
def read_pb_stream(f):
|
||||
while True:
|
||||
buf = f.read(4)
|
||||
if len(buf) == 0:
|
||||
break
|
||||
size = struct.unpack("I", buf)[0]
|
||||
buf = f.read(size)
|
||||
text_data = TextData()
|
||||
text_data.ParseFromString(buf)
|
||||
yield text_data
|
||||
|
||||
|
||||
def write_pb_stream(f, text_data):
|
||||
buf = text_data.SerializeToString()
|
||||
f.write(struct.pack("I", len(buf)))
|
||||
f.write(buf)
|
||||
|
||||
|
||||
def pack_pb_stream(text_data):
|
||||
buf = text_data.SerializeToString()
|
||||
return struct.pack("I", len(buf)) + buf
|
||||
|
||||
|
||||
def split_pb_stream(f):
|
||||
while True:
|
||||
head = f.read(4)
|
||||
if len(head) == 0:
|
||||
break
|
||||
size = struct.unpack("I", head)[0]
|
||||
buf = f.read(size)
|
||||
yield head + buf
|
||||
@@ -0,0 +1,496 @@
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain
|
||||
from pathlib import Path
|
||||
from random import Random
|
||||
from typing import Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from datasets.download.streaming_download_manager import xopen
|
||||
from huggingface_hub import HfApi
|
||||
from lightning import LightningDataModule
|
||||
from torch.distributed import get_rank, get_world_size, is_initialized
|
||||
from torch.utils.data import DataLoader, IterableDataset, get_worker_info
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fish_speech.conversation import CODEBOOK_PAD_TOKEN_ID
|
||||
from fish_speech.datasets.protos.text_data_pb2 import SampledData
|
||||
from fish_speech.datasets.protos.text_data_stream import read_pb_stream
|
||||
from fish_speech.text.clean import clean_text
|
||||
from fish_speech.utils import RankedLogger
|
||||
from fish_speech.utils.braceexpand import braceexpand
|
||||
|
||||
log = RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
def split_by_rank_worker(files):
|
||||
# We need to know the total number of devices
|
||||
# to split the data properly
|
||||
|
||||
total_devices = 1
|
||||
if is_initialized():
|
||||
total_devices = get_world_size()
|
||||
|
||||
worker_info = get_worker_info()
|
||||
if worker_info is not None:
|
||||
total_devices *= worker_info.num_workers
|
||||
|
||||
if len(files) < total_devices:
|
||||
# Repeat the files N times to match the number of devices
|
||||
files = files * (total_devices // len(files) + 1)
|
||||
|
||||
# DDP
|
||||
if is_initialized():
|
||||
files = files[get_rank() :: get_world_size()]
|
||||
|
||||
# Split by worker
|
||||
if worker_info is not None:
|
||||
files = files[worker_info.id :: worker_info.num_workers]
|
||||
|
||||
return files
|
||||
|
||||
|
||||
class AutoTextSemanticInstructionDataset(IterableDataset):
|
||||
"""
|
||||
Auto Augment Dataset by Speaker
|
||||
|
||||
1. Random concatenate multiple sentences from the same speaker to form a longer sentence
|
||||
2. Automatically normalize the text
|
||||
|
||||
For interactive mode, we use the following format (multiple sequences):
|
||||
<s> [INST] [SPK: speaker] text [/INST] ... [INST] text [/INST] </s>
|
||||
|
||||
For non-interactive mode, we use the following format (one long sequence):
|
||||
<s> [INST] text [/INST] ... </s>
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
proto_files: list[str],
|
||||
seed: int = 42,
|
||||
interactive_prob: float = 0.5,
|
||||
max_length: int = 1024,
|
||||
tokenizer: AutoTokenizer = None,
|
||||
use_speaker: bool | float = True,
|
||||
causal: bool = True,
|
||||
num_codebooks: Optional[int] = None,
|
||||
skip_text_prob: float = 0.0,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
proto_files: proto buf files if using local data
|
||||
seed: random seed
|
||||
interactive_prob: probability to use interactive mode
|
||||
max_length: max length of the text
|
||||
tokenizer: tokenizer
|
||||
use_speaker: include speaker information in the prompt
|
||||
causal: use causal sampling when using local data, disable will lead to random sampling
|
||||
num_codebooks: number of codebooks, if None, it will be automatically detected
|
||||
skip_text_prob: probability to skip the text (audio only), this only applies to interactive mode
|
||||
"""
|
||||
|
||||
super().__init__()
|
||||
|
||||
assert 0 <= interactive_prob <= 1, "interactive_prob must be in [0, 1]"
|
||||
|
||||
self.seed = seed
|
||||
self.max_length = max_length
|
||||
self.tokenizer = tokenizer
|
||||
self.interactive_prob = interactive_prob
|
||||
self.use_speaker = use_speaker
|
||||
self.proto_files = proto_files
|
||||
self.causal = causal
|
||||
self.num_codebooks = num_codebooks
|
||||
self.skip_text_prob = skip_text_prob
|
||||
|
||||
self.semantic_token_id = self.tokenizer.convert_tokens_to_ids("<|semantic|>")
|
||||
self.groups = None
|
||||
|
||||
def init_mock_data_server(self):
|
||||
if self.groups is not None:
|
||||
return
|
||||
|
||||
# Expand the proto files
|
||||
expanded_proto_files = []
|
||||
for filename in self.proto_files:
|
||||
for i in braceexpand(filename):
|
||||
i = Path(i)
|
||||
if i.is_file():
|
||||
expanded_proto_files.append(i)
|
||||
elif i.is_dir():
|
||||
expanded_proto_files.extend(i.rglob("*.proto"))
|
||||
expanded_proto_files.extend(i.rglob("*.protos"))
|
||||
else:
|
||||
raise ValueError(f"{i} is not a file or directory")
|
||||
|
||||
expanded_proto_files = sorted(expanded_proto_files)
|
||||
Random(self.seed).shuffle(expanded_proto_files)
|
||||
|
||||
self.groups = []
|
||||
shard_proto_files = split_by_rank_worker(expanded_proto_files)
|
||||
log.info(
|
||||
f"Reading {len(shard_proto_files)} / {len(expanded_proto_files)} files"
|
||||
)
|
||||
|
||||
count = 0
|
||||
for filename in shard_proto_files:
|
||||
with open(filename, "rb") as f:
|
||||
for text_data in read_pb_stream(f):
|
||||
self.groups.append(text_data)
|
||||
count += 1
|
||||
|
||||
log.info(f"Read total {count} groups of data")
|
||||
|
||||
# Shuffle the lines
|
||||
Random(self.seed).shuffle(self.groups)
|
||||
self.group_weights = [len(i.sentences) for i in self.groups]
|
||||
|
||||
def __iter__(self):
|
||||
while True:
|
||||
yield self.augment()
|
||||
|
||||
def tokenize_sentence(self, sentence: str):
|
||||
sentence = clean_text(sentence)
|
||||
tokens = self.tokenizer.encode(
|
||||
f"{sentence}",
|
||||
max_length=10**6,
|
||||
add_special_tokens=False,
|
||||
truncation=False,
|
||||
)
|
||||
return sentence, len(tokens)
|
||||
|
||||
def sample_data(self):
|
||||
if self.groups is None:
|
||||
self.init_mock_data_server()
|
||||
|
||||
# Shuffle unique lines, estimate that each sample is at least 20 tokens
|
||||
num_samples = self.max_length // 20
|
||||
|
||||
# choice group based on their number of samples
|
||||
group = random.choices(self.groups, weights=self.group_weights, k=1)[0]
|
||||
|
||||
if self.causal:
|
||||
# Sample in order
|
||||
if num_samples >= len(group.sentences):
|
||||
samples = group.sentences
|
||||
else:
|
||||
begin = random.randint(0, len(group.sentences) - num_samples)
|
||||
samples = group.sentences[begin : begin + num_samples]
|
||||
else:
|
||||
samples = random.choices(
|
||||
group.sentences, k=min(num_samples, len(group.sentences))
|
||||
)
|
||||
|
||||
return SampledData(
|
||||
source=group.source,
|
||||
name=group.name,
|
||||
samples=samples,
|
||||
)
|
||||
|
||||
def augment(self):
|
||||
final_text, final_semantic = [], []
|
||||
response = self.sample_data()
|
||||
if len(response.samples) == 0:
|
||||
# Invalid group
|
||||
return None
|
||||
|
||||
samples = list(response.samples)
|
||||
idx = 0
|
||||
use_interactive = random.random() < self.interactive_prob
|
||||
|
||||
if use_interactive is False:
|
||||
# Random sample based on speaker using a truncated normal distribution
|
||||
a = torch.tensor([0], dtype=torch.float32)
|
||||
torch.nn.init.trunc_normal_(
|
||||
a,
|
||||
mean=self.max_length // 2,
|
||||
std=self.max_length // 4,
|
||||
a=10,
|
||||
b=self.max_length,
|
||||
)
|
||||
remaining_tokens = a.long().item() - 4
|
||||
else:
|
||||
remaining_tokens = self.max_length
|
||||
|
||||
# Use speaker
|
||||
if isinstance(self.use_speaker, float):
|
||||
use_speaker = random.random() < self.use_speaker
|
||||
else:
|
||||
use_speaker = self.use_speaker
|
||||
|
||||
all_tokens, all_labels = [], []
|
||||
while remaining_tokens > 0 and len(samples) > 0:
|
||||
sentence = samples.pop(0)
|
||||
|
||||
text = random.choice(sentence.texts)
|
||||
text, length = self.tokenize_sentence(text)
|
||||
remaining_tokens -= length + len(sentence.semantics[0].values)
|
||||
|
||||
if use_interactive is False:
|
||||
final_text.append(text)
|
||||
final_semantic.append(sentence.semantics)
|
||||
else:
|
||||
# For interactive mode, we only apply speaker for the first sentence
|
||||
# [INST] [SPK: speaker] text [/INST] ... [INST] text [/INST]
|
||||
tokens, labels = self.pack_sentences(
|
||||
sentences=[text],
|
||||
semantics=[sentence.semantics],
|
||||
speaker=response.name if use_speaker else None,
|
||||
skip_text=random.random() < self.skip_text_prob,
|
||||
)
|
||||
|
||||
all_tokens.append(tokens)
|
||||
all_labels.append(labels)
|
||||
|
||||
idx += 1
|
||||
|
||||
if use_interactive is False:
|
||||
tokens, labels = self.pack_sentences(
|
||||
final_text,
|
||||
semantics=final_semantic,
|
||||
speaker=response.name if use_speaker else None,
|
||||
)
|
||||
all_tokens.append(tokens)
|
||||
all_labels.append(labels)
|
||||
|
||||
tokens = torch.cat(all_tokens, dim=1)
|
||||
labels = torch.cat(all_labels, dim=1)
|
||||
|
||||
# Verify that the length is correct
|
||||
assert tokens.size(1) == labels.size(1), f"{tokens.size(1)} != {labels.size(1)}"
|
||||
|
||||
data = {"tokens": tokens, "labels": labels}
|
||||
|
||||
return data
|
||||
|
||||
def pack_sentences(
|
||||
self,
|
||||
sentences: list[str],
|
||||
semantics: list,
|
||||
speaker: Optional[str] = None,
|
||||
skip_text: bool = False,
|
||||
):
|
||||
if speaker is None:
|
||||
speaker = "assistant"
|
||||
|
||||
cated_sentences = " ".join(sentences)
|
||||
if skip_text:
|
||||
cated_sentences = "<|skip_text|>"
|
||||
|
||||
final_text = "<|im_start|>user\n" + cated_sentences + "<|im_end|>"
|
||||
final_text = final_text + f"<|im_start|>{speaker}\n"
|
||||
|
||||
encoded = self.tokenizer.encode(
|
||||
final_text,
|
||||
add_special_tokens=False,
|
||||
truncation=False,
|
||||
max_length=10**6,
|
||||
)
|
||||
semantic_length = sum([len(i[0].values) for i in semantics])
|
||||
prompt_length = len(encoded)
|
||||
num_codebooks = (
|
||||
len(semantics[0]) if self.num_codebooks is None else self.num_codebooks
|
||||
)
|
||||
|
||||
# Pack the tokens and semantics (add <s> and </s> to semantic tokens)
|
||||
tokens = (
|
||||
encoded
|
||||
+ [self.semantic_token_id] * semantic_length
|
||||
+ self.tokenizer.convert_tokens_to_ids(["<|im_end|>"])
|
||||
)
|
||||
|
||||
# Codebook bos/padding: 0, eos: 1
|
||||
codes = [[CODEBOOK_PAD_TOKEN_ID] * prompt_length for _ in range(num_codebooks)]
|
||||
for segment in semantics:
|
||||
for book_idx, book in zip(range(num_codebooks), segment):
|
||||
for j in book.values:
|
||||
codes[book_idx].append(int(j) + 1)
|
||||
|
||||
for book in codes:
|
||||
book.extend([CODEBOOK_PAD_TOKEN_ID] * 1)
|
||||
|
||||
tokens = [tokens] + codes
|
||||
|
||||
tokens = torch.tensor(tokens, dtype=torch.long)
|
||||
labels = tokens.clone()
|
||||
|
||||
if skip_text:
|
||||
# If text is not provided, the sentence is used for condition only, all labels are -100
|
||||
torch.fill_(labels, -100)
|
||||
return tokens, labels
|
||||
|
||||
# Mask out the <s> tokens for semantic, predict semantic tokens only
|
||||
# Since we don't mask out the input tokens, the language modeling still works
|
||||
labels[1:, :prompt_length] = -100
|
||||
|
||||
tokens = tokens[:, :-1]
|
||||
labels = labels[:, 1:]
|
||||
|
||||
# Verify the padding is correct, and the last token is eos
|
||||
assert (tokens[1:, :prompt_length] == CODEBOOK_PAD_TOKEN_ID).all()
|
||||
assert (labels[1:, -1:] == CODEBOOK_PAD_TOKEN_ID).all()
|
||||
|
||||
return tokens, labels
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextDataCollator:
|
||||
tokenizer: AutoTokenizer
|
||||
max_length: int = 1024
|
||||
|
||||
def __call__(self, examples):
|
||||
if "negative_tokens" in examples:
|
||||
positive_examples = []
|
||||
negative_examples = []
|
||||
|
||||
for i in examples:
|
||||
positive_examples.append(
|
||||
{
|
||||
"tokens": i["tokens"],
|
||||
"labels": i["labels"],
|
||||
}
|
||||
)
|
||||
negative_examples.append(
|
||||
{
|
||||
"tokens": i["negative_tokens"],
|
||||
"labels": i["negative_labels"],
|
||||
}
|
||||
)
|
||||
|
||||
examples = positive_examples + negative_examples
|
||||
|
||||
return self.batchify(examples)
|
||||
|
||||
def batchify(self, examples, tokens_key="tokens", labels_key="labels"):
|
||||
tokens, attention_masks, labels = [], [], []
|
||||
|
||||
# Calculate the max length
|
||||
max_tokens_length = 0
|
||||
for example in examples:
|
||||
max_tokens_length = max(max_tokens_length, example[tokens_key].size(1))
|
||||
max_tokens_length = min(max_tokens_length, self.max_length)
|
||||
|
||||
for example in examples:
|
||||
_tokens = example[tokens_key][:, :max_tokens_length]
|
||||
_labels = example[labels_key][:, :max_tokens_length]
|
||||
_attention_mask = torch.ones((max_tokens_length,), dtype=torch.bool)
|
||||
tokens_length = _tokens.size(1)
|
||||
_attention_mask[:tokens_length] = False
|
||||
|
||||
assert tokens_length == _labels.size(
|
||||
1
|
||||
), f"{tokens_length} != {_labels.size(1)}"
|
||||
|
||||
if tokens_length < max_tokens_length:
|
||||
_tokens = F.pad(
|
||||
_tokens,
|
||||
(0, max_tokens_length - tokens_length),
|
||||
value=self.tokenizer.eos_token_id,
|
||||
)
|
||||
_tokens[1:, tokens_length:] = CODEBOOK_PAD_TOKEN_ID
|
||||
_labels = F.pad(
|
||||
_labels, (0, max_tokens_length - _labels.size(1)), value=-100
|
||||
)
|
||||
|
||||
tokens.append(_tokens)
|
||||
attention_masks.append(_attention_mask)
|
||||
labels.append(_labels)
|
||||
|
||||
tokens = torch.stack(tokens, dim=0)
|
||||
attention_masks = torch.stack(attention_masks, dim=0)
|
||||
labels = torch.stack(labels, dim=0)
|
||||
|
||||
return {
|
||||
"inputs": tokens,
|
||||
"attention_masks": attention_masks,
|
||||
"labels": labels,
|
||||
}
|
||||
|
||||
|
||||
class InterleaveDataset(IterableDataset):
|
||||
def __init__(
|
||||
self,
|
||||
datasets: list[IterableDataset],
|
||||
probabilities: list[float],
|
||||
seed: int = 42,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.datasets = datasets
|
||||
self.probabilities = probabilities
|
||||
self.seed = seed
|
||||
|
||||
def __iter__(self):
|
||||
rng = np.random.default_rng(self.seed)
|
||||
dataset_iterators = [iter(dataset) for dataset in self.datasets]
|
||||
|
||||
while True:
|
||||
# Random choice one
|
||||
dataset_idx = rng.choice(len(self.datasets), p=self.probabilities)
|
||||
dataset_iterator = dataset_iterators[dataset_idx]
|
||||
|
||||
try:
|
||||
yield next(dataset_iterator)
|
||||
except StopIteration:
|
||||
# Exhausted, create a new iterator
|
||||
dataset_iterators[dataset_idx] = iter(self.datasets[dataset_idx])
|
||||
yield next(dataset_iterators[dataset_idx])
|
||||
|
||||
|
||||
class SemanticDataModule(LightningDataModule):
|
||||
def __init__(
|
||||
self,
|
||||
train_dataset: Union[AutoTextSemanticInstructionDataset, InterleaveDataset],
|
||||
val_dataset: Union[AutoTextSemanticInstructionDataset, InterleaveDataset],
|
||||
batch_size: int = 32,
|
||||
tokenizer: AutoTokenizer = None,
|
||||
max_length: int = 1024,
|
||||
num_workers: int = 4,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.train_dataset = train_dataset
|
||||
self.val_dataset = val_dataset
|
||||
self.batch_size = batch_size
|
||||
self.tokenizer = tokenizer
|
||||
self.max_length = max_length
|
||||
self.num_workers = num_workers
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=self.batch_size,
|
||||
collate_fn=TextDataCollator(self.tokenizer, self.max_length),
|
||||
num_workers=self.num_workers,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
def val_dataloader(self):
|
||||
return DataLoader(
|
||||
self.val_dataset,
|
||||
batch_size=self.batch_size,
|
||||
collate_fn=TextDataCollator(self.tokenizer, self.max_length),
|
||||
num_workers=self.num_workers,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from tqdm import tqdm
|
||||
|
||||
ds = AutoTextSemanticInstructionDataset(
|
||||
["data/protos"],
|
||||
tokenizer=AutoTokenizer.from_pretrained("fishaudio/fish-speech-1"),
|
||||
use_speaker=False,
|
||||
interactive_prob=1.0,
|
||||
skip_text_prob=0.5,
|
||||
)
|
||||
|
||||
for i in ds:
|
||||
print(ds.tokenizer.decode(i["tokens"][0], skip_special_tokens=False))
|
||||
# i["labels"][0][i["labels"][0] == -100] = 0
|
||||
# print(ds.tokenizer.decode(i["labels"][0], skip_special_tokens=False))
|
||||
break
|
||||
@@ -0,0 +1,147 @@
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
from lightning import LightningDataModule
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
|
||||
from fish_speech.utils import RankedLogger
|
||||
|
||||
logger = RankedLogger(__name__, rank_zero_only=False)
|
||||
|
||||
|
||||
class VQGANDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
filelist: str,
|
||||
sample_rate: int = 32000,
|
||||
hop_length: int = 640,
|
||||
slice_frames: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
filelist = Path(filelist)
|
||||
root = filelist.parent
|
||||
|
||||
self.files = [
|
||||
root / line.strip()
|
||||
for line in filelist.read_text(encoding="utf-8").splitlines()
|
||||
if line.strip()
|
||||
]
|
||||
self.sample_rate = sample_rate
|
||||
self.hop_length = hop_length
|
||||
self.slice_frames = slice_frames
|
||||
|
||||
def __len__(self):
|
||||
return len(self.files)
|
||||
|
||||
def get_item(self, idx):
|
||||
file = self.files[idx]
|
||||
|
||||
audio, _ = librosa.load(file, sr=self.sample_rate, mono=True)
|
||||
|
||||
# Slice audio and features
|
||||
if (
|
||||
self.slice_frames is not None
|
||||
and audio.shape[0] > self.slice_frames * self.hop_length
|
||||
):
|
||||
start = np.random.randint(
|
||||
0, audio.shape[0] - self.slice_frames * self.hop_length
|
||||
)
|
||||
audio = audio[start : start + self.slice_frames * self.hop_length]
|
||||
|
||||
if len(audio) == 0:
|
||||
return None
|
||||
|
||||
max_value = np.abs(audio).max()
|
||||
if max_value > 1.0:
|
||||
audio = audio / max_value
|
||||
|
||||
return {
|
||||
"audio": torch.from_numpy(audio),
|
||||
}
|
||||
|
||||
def __getitem__(self, idx):
|
||||
try:
|
||||
return self.get_item(idx)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
logger.error(f"Error loading {self.files[idx]}: {e}")
|
||||
return None
|
||||
|
||||
|
||||
@dataclass
|
||||
class VQGANCollator:
|
||||
def __call__(self, batch):
|
||||
batch = [x for x in batch if x is not None]
|
||||
|
||||
audio_lengths = torch.tensor([len(x["audio"]) for x in batch])
|
||||
audio_maxlen = audio_lengths.max()
|
||||
|
||||
# Rounds up to nearest multiple of 2 (audio_lengths)
|
||||
audios = []
|
||||
for x in batch:
|
||||
audios.append(
|
||||
torch.nn.functional.pad(x["audio"], (0, audio_maxlen - len(x["audio"])))
|
||||
)
|
||||
|
||||
return {
|
||||
"audios": torch.stack(audios),
|
||||
"audio_lengths": audio_lengths,
|
||||
}
|
||||
|
||||
|
||||
class VQGANDataModule(LightningDataModule):
|
||||
def __init__(
|
||||
self,
|
||||
train_dataset: VQGANDataset,
|
||||
val_dataset: VQGANDataset,
|
||||
batch_size: int = 32,
|
||||
num_workers: int = 4,
|
||||
val_batch_size: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.train_dataset = train_dataset
|
||||
self.val_dataset = val_dataset
|
||||
self.batch_size = batch_size
|
||||
self.val_batch_size = val_batch_size or batch_size
|
||||
self.num_workers = num_workers
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=self.batch_size,
|
||||
collate_fn=VQGANCollator(),
|
||||
num_workers=self.num_workers,
|
||||
shuffle=True,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
def val_dataloader(self):
|
||||
return DataLoader(
|
||||
self.val_dataset,
|
||||
batch_size=self.val_batch_size,
|
||||
collate_fn=VQGANCollator(),
|
||||
num_workers=self.num_workers,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = VQGANDataset("data/LibriTTS_R/vq_train_filelist.txt")
|
||||
dataloader = DataLoader(
|
||||
dataset, batch_size=4, shuffle=False, collate_fn=VQGANCollator()
|
||||
)
|
||||
|
||||
for batch in dataloader:
|
||||
print(batch["audios"].shape)
|
||||
print(batch["features"].shape)
|
||||
print(batch["audio_lengths"])
|
||||
print(batch["feature_lengths"])
|
||||
break
|
||||
@@ -0,0 +1,104 @@
|
||||
|
||||
import torch
|
||||
from .models.text2semantic.llama import BaseTransformer, NaiveTransformer, DualARTransformer
|
||||
from .tools.llama.generate import decode_one_token_ar, decode_one_token_naive, generate_long
|
||||
import numpy as np
|
||||
import time
|
||||
from typing import Union
|
||||
from loguru import logger
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
def load_model(checkpoint_path, device, precision, compile=False):
|
||||
model: Union[NaiveTransformer, DualARTransformer] = BaseTransformer.from_pretrained(
|
||||
checkpoint_path, load_weights=True
|
||||
)
|
||||
|
||||
model = model.to(device=device, dtype=precision)
|
||||
logger.info(f"Restored model from checkpoint")
|
||||
|
||||
if isinstance(model, DualARTransformer):
|
||||
decode_one_token = decode_one_token_ar
|
||||
logger.info("Using DualARTransformer")
|
||||
else:
|
||||
decode_one_token = decode_one_token_naive
|
||||
logger.info("Using NaiveTransformer")
|
||||
|
||||
if compile:
|
||||
logger.info("Compiling function...")
|
||||
decode_one_token = torch.compile(
|
||||
decode_one_token, mode="reduce-overhead", fullgraph=True
|
||||
)
|
||||
|
||||
return model.eval(), decode_one_token
|
||||
|
||||
|
||||
def prompt2semantic(
|
||||
model: DualARTransformer,
|
||||
decode_one_token: callable,
|
||||
text: str,
|
||||
prompt_text: Optional[list[str]],
|
||||
prompt_tokens: Optional[list[np.ndarray]],
|
||||
max_new_tokens: int,
|
||||
top_p: float,
|
||||
repetition_penalty: float,
|
||||
temperature: float,
|
||||
device: str,
|
||||
compile: bool,
|
||||
seed: int,
|
||||
iterative_prompt: bool,
|
||||
chunk_length: int,
|
||||
):
|
||||
|
||||
if prompt_text is not None and len(prompt_text) != len(prompt_tokens):
|
||||
raise ValueError(
|
||||
f"Number of prompt text ({len(prompt_text)}) and prompt tokens ({len(prompt_tokens)}) should be the same"
|
||||
)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
if prompt_tokens is not None:
|
||||
prompt_tokens = [torch.from_numpy(pt).to(device) for pt in prompt_tokens]
|
||||
|
||||
torch.manual_seed(seed)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
generator = generate_long(
|
||||
model=model,
|
||||
device=device,
|
||||
decode_one_token=decode_one_token,
|
||||
text=text,
|
||||
num_samples=1,
|
||||
max_new_tokens=max_new_tokens,
|
||||
top_p=top_p,
|
||||
repetition_penalty=repetition_penalty,
|
||||
temperature=temperature,
|
||||
compile=compile,
|
||||
iterative_prompt=iterative_prompt,
|
||||
chunk_length=chunk_length,
|
||||
prompt_text=prompt_text,
|
||||
prompt_tokens=prompt_tokens,
|
||||
)
|
||||
|
||||
idx = 0
|
||||
all_codes = []
|
||||
codes = []
|
||||
|
||||
for response in generator:
|
||||
if response.action == "sample":
|
||||
codes.append(response.codes)
|
||||
logger.info(f"Sampled text: {response.text}")
|
||||
elif response.action == "next":
|
||||
if codes:
|
||||
all_codes.append(torch.cat(codes, dim=1).cpu().numpy())
|
||||
logger.info(f"Saved codes to codes_{idx}.npy")
|
||||
logger.info(f"Next sample")
|
||||
codes = []
|
||||
idx += 1
|
||||
else:
|
||||
logger.error(f"Error: {response}")
|
||||
|
||||
return all_codes
|
||||
@@ -0,0 +1,202 @@
|
||||
from typing import Any, Optional
|
||||
|
||||
import lightning as L
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from lightning.pytorch.utilities.types import OptimizerLRScheduler
|
||||
|
||||
import fish_speech.utils as utils
|
||||
from fish_speech.conversation import CODEBOOK_PAD_TOKEN_ID
|
||||
from fish_speech.models.text2semantic.llama import NaiveTransformer
|
||||
|
||||
log = utils.RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
class TextToSemantic(L.LightningModule):
|
||||
def __init__(
|
||||
self,
|
||||
model: NaiveTransformer,
|
||||
optimizer: Any,
|
||||
lr_scheduler: Any,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.model = model
|
||||
self.optimizer_builder = optimizer
|
||||
self.lr_scheduler_builder = lr_scheduler
|
||||
|
||||
def forward(self, x):
|
||||
return self.model(x)
|
||||
|
||||
def on_save_checkpoint(self, checkpoint):
|
||||
# Save only LoRA parameters
|
||||
state_dict = checkpoint["state_dict"]
|
||||
use_lora = any("lora" in name for name in state_dict.keys())
|
||||
if not use_lora:
|
||||
return
|
||||
|
||||
for name in list(state_dict.keys()):
|
||||
if "lora" not in name:
|
||||
state_dict.pop(name)
|
||||
|
||||
def configure_optimizers(self) -> OptimizerLRScheduler:
|
||||
# Get weight decay parameters
|
||||
weight_decay_parameters, other_parameters = [], []
|
||||
for name, param in self.named_parameters():
|
||||
if ".bias" in name or "norm.weight" in name or ".embeddings." in name:
|
||||
other_parameters.append(param)
|
||||
else:
|
||||
weight_decay_parameters.append(param)
|
||||
|
||||
optimizer = self.optimizer_builder(
|
||||
[
|
||||
{"params": weight_decay_parameters},
|
||||
{"params": other_parameters, "weight_decay": 0.0},
|
||||
]
|
||||
)
|
||||
|
||||
# Print the parameters and their weight decay
|
||||
for i in optimizer.param_groups:
|
||||
log.info(
|
||||
f"Set weight decay: {i['weight_decay']} for {len(i['params'])} parameters"
|
||||
)
|
||||
|
||||
lr_scheduler = self.lr_scheduler_builder(optimizer)
|
||||
|
||||
return {
|
||||
"optimizer": optimizer,
|
||||
"lr_scheduler": {
|
||||
"scheduler": lr_scheduler,
|
||||
"interval": "step",
|
||||
},
|
||||
}
|
||||
|
||||
# Copied from https://github.com/eric-mitchell/direct-preference-optimization/blob/main/trainers.py#L90
|
||||
def get_batch_logps(
|
||||
self,
|
||||
logits: torch.FloatTensor,
|
||||
labels: torch.LongTensor,
|
||||
average_log_prob: bool = False,
|
||||
) -> torch.FloatTensor:
|
||||
"""Compute the log probabilities of the given labels under the given logits.
|
||||
|
||||
Args:
|
||||
logits: Logits of the model (unnormalized). Shape: (batch_size, sequence_length, codebook_size, vocab_size)
|
||||
labels: Labels for which to compute the log probabilities. Label tokens with a value of -100 are ignored. Shape: (batch_size, sequence_length, codebook_size)
|
||||
average_log_prob: If True, return the average log probability per (non-masked) token. Otherwise, return the sum of the log probabilities of the (non-masked) tokens.
|
||||
|
||||
Returns:
|
||||
A tensor of shape (batch_size,) containing the average/sum log probabilities of the given labels under the given logits.
|
||||
"""
|
||||
assert logits.shape[:-1] == labels.shape
|
||||
|
||||
labels = labels.clone()
|
||||
loss_mask = labels != -100
|
||||
|
||||
# dummy token; we'll ignore the losses on these tokens later
|
||||
labels[labels == -100] = 0
|
||||
|
||||
per_token_logps = torch.gather(
|
||||
logits.log_softmax(-1), dim=-1, index=labels.unsqueeze(-1)
|
||||
).squeeze(-1)
|
||||
|
||||
if average_log_prob:
|
||||
return (per_token_logps * loss_mask).sum(-1) / loss_mask.sum(-1)
|
||||
else:
|
||||
return (per_token_logps * loss_mask).sum(-1)
|
||||
|
||||
def _step(self, batch, batch_idx, stage: str):
|
||||
is_train = stage == "train"
|
||||
|
||||
if is_train:
|
||||
# Key part to make lora work
|
||||
# Otherwise the parameters are merged, which lead to incorrect gradients
|
||||
self.model.train()
|
||||
|
||||
# Do positive and negative samples in the same batch to speed up training
|
||||
labels = batch["labels"]
|
||||
outputs = self.model(
|
||||
inp=batch["inputs"],
|
||||
key_padding_mask=batch["attention_masks"],
|
||||
)
|
||||
token_logits = outputs.token_logits
|
||||
codebook_logits = outputs.codebook_logits
|
||||
|
||||
# Generate labels
|
||||
base_loss = F.cross_entropy(
|
||||
token_logits.view(-1, token_logits.size(-1)),
|
||||
labels[:, 0].reshape(-1),
|
||||
ignore_index=-100,
|
||||
)
|
||||
|
||||
codebook_labels = labels[:, 1 : 1 + self.model.config.num_codebooks].mT
|
||||
semantic_loss = F.cross_entropy(
|
||||
codebook_logits.view(-1, codebook_logits.size(-1)),
|
||||
codebook_labels.reshape(-1),
|
||||
ignore_index=-100,
|
||||
)
|
||||
|
||||
loss = base_loss + semantic_loss
|
||||
|
||||
self.log(
|
||||
f"{stage}/loss",
|
||||
loss,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=True,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
self.log(
|
||||
f"{stage}/base_loss",
|
||||
base_loss,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=False,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
self.log(
|
||||
f"{stage}/semantic_loss",
|
||||
semantic_loss,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=False,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
# Top-5 accuracy
|
||||
accuracy = self.get_accuracy(codebook_logits, codebook_labels)
|
||||
self.log(
|
||||
f"{stage}/top_5_accuracy",
|
||||
accuracy,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=True,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
return loss
|
||||
|
||||
def get_accuracy(self, logits, labels):
|
||||
mask = (labels != -100) & (labels != CODEBOOK_PAD_TOKEN_ID)
|
||||
if mask.sum() == 0:
|
||||
return torch.tensor(0.0, device=logits.device)
|
||||
|
||||
_, indices = logits.topk(5, dim=-1)
|
||||
correct = indices.eq(labels.unsqueeze(-1))
|
||||
correct[~mask] = 0
|
||||
correct = correct.sum()
|
||||
accuracy = correct / mask.sum()
|
||||
|
||||
return accuracy
|
||||
|
||||
def training_step(self, batch, batch_idx):
|
||||
return self._step(batch, batch_idx, "train")
|
||||
|
||||
def validation_step(self, batch, batch_idx):
|
||||
return self._step(batch, batch_idx, "val")
|
||||
@@ -0,0 +1,779 @@
|
||||
import json
|
||||
import math
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from loguru import logger
|
||||
from torch import Tensor
|
||||
from torch.nn import functional as F
|
||||
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fish_speech.conversation import SEMANTIC_TOKEN
|
||||
from fish_speech.utils import RankedLogger
|
||||
|
||||
from .lora import LoraConfig, setup_lora
|
||||
|
||||
log = RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
def find_multiple(n: int, k: int) -> int:
|
||||
if n % k == 0:
|
||||
return n
|
||||
return n + k - (n % k)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseModelArgs:
|
||||
model_type: str = "base"
|
||||
|
||||
vocab_size: int = 32000
|
||||
n_layer: int = 32
|
||||
n_head: int = 32
|
||||
dim: int = 4096
|
||||
intermediate_size: int = None
|
||||
n_local_heads: int = -1
|
||||
head_dim: int = 64
|
||||
rope_base: float = 10000
|
||||
norm_eps: float = 1e-5
|
||||
max_seq_len: int = 2048
|
||||
dropout: float = 0.0
|
||||
tie_word_embeddings: bool = True
|
||||
attention_qkv_bias: bool = False
|
||||
|
||||
# Codebook configs
|
||||
codebook_size: int = 160
|
||||
num_codebooks: int = 4
|
||||
|
||||
# Gradient checkpointing
|
||||
use_gradient_checkpointing: bool = True
|
||||
|
||||
# Initialize the model
|
||||
initializer_range: float = 0.02
|
||||
|
||||
def __post_init__(self):
|
||||
if self.n_local_heads == -1:
|
||||
self.n_local_heads = self.n_head
|
||||
if self.intermediate_size is None:
|
||||
hidden_dim = 4 * self.dim
|
||||
n_hidden = int(2 * hidden_dim / 3)
|
||||
self.intermediate_size = find_multiple(n_hidden, 256)
|
||||
self.head_dim = self.dim // self.n_head
|
||||
|
||||
@staticmethod
|
||||
def from_pretrained(path: str):
|
||||
path = Path(path)
|
||||
|
||||
if path.is_dir():
|
||||
path = path / "config.json"
|
||||
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
|
||||
match data["model_type"]:
|
||||
case "naive":
|
||||
cls = NaiveModelArgs
|
||||
case "dual_ar":
|
||||
cls = DualARModelArgs
|
||||
case _:
|
||||
raise ValueError(f"Unknown model type: {data['model_type']}")
|
||||
|
||||
return cls(**data)
|
||||
|
||||
def save(self, path: str):
|
||||
with open(path, "w") as f:
|
||||
json.dump(self.__dict__, f, indent=4, sort_keys=True, ensure_ascii=False)
|
||||
|
||||
|
||||
@dataclass
|
||||
class NaiveModelArgs(BaseModelArgs):
|
||||
model_type: str = "naive"
|
||||
|
||||
|
||||
@dataclass
|
||||
class DualARModelArgs(BaseModelArgs):
|
||||
model_type: str = "dual_ar"
|
||||
n_fast_layer: int = 4
|
||||
|
||||
|
||||
class KVCache(nn.Module):
|
||||
def __init__(
|
||||
self, max_batch_size, max_seq_len, n_heads, head_dim, dtype=torch.bfloat16
|
||||
):
|
||||
super().__init__()
|
||||
cache_shape = (max_batch_size, n_heads, max_seq_len, head_dim)
|
||||
self.register_buffer("k_cache", torch.zeros(cache_shape, dtype=dtype))
|
||||
self.register_buffer("v_cache", torch.zeros(cache_shape, dtype=dtype))
|
||||
|
||||
def update(self, input_pos, k_val, v_val):
|
||||
# input_pos: [S], k_val: [B, H, S, D]
|
||||
assert input_pos.shape[0] == k_val.shape[2]
|
||||
|
||||
k_out = self.k_cache
|
||||
v_out = self.v_cache
|
||||
k_out[:, :, input_pos] = k_val
|
||||
v_out[:, :, input_pos] = v_val
|
||||
|
||||
return k_out, v_out
|
||||
|
||||
|
||||
@dataclass
|
||||
class TransformerForwardResult:
|
||||
token_logits: Tensor
|
||||
codebook_logits: Tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseTransformerForwardResult:
|
||||
logits: Tensor
|
||||
hidden_states: Tensor
|
||||
|
||||
|
||||
class BaseTransformer(nn.Module):
|
||||
def __init__(
|
||||
self, config: BaseModelArgs, tokenizer: AutoTokenizer, init_weights: bool = True
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.semantic_token_id = tokenizer.convert_tokens_to_ids(SEMANTIC_TOKEN)
|
||||
|
||||
# Slow transformer
|
||||
self.embeddings = nn.Embedding(
|
||||
config.vocab_size,
|
||||
config.dim,
|
||||
)
|
||||
self.codebook_embeddings = nn.Embedding(
|
||||
config.codebook_size * config.num_codebooks,
|
||||
config.dim,
|
||||
)
|
||||
self.layers = nn.ModuleList(
|
||||
TransformerBlock(config, use_sdpa=True) for _ in range(config.n_layer)
|
||||
)
|
||||
self.norm = RMSNorm(config.dim, eps=config.norm_eps)
|
||||
|
||||
if self.config.tie_word_embeddings is False:
|
||||
self.output = nn.Linear(
|
||||
config.dim,
|
||||
config.vocab_size,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.register_buffer(
|
||||
"freqs_cis",
|
||||
precompute_freqs_cis(
|
||||
config.max_seq_len,
|
||||
config.dim // config.n_head,
|
||||
config.rope_base,
|
||||
),
|
||||
persistent=False,
|
||||
)
|
||||
self.register_buffer(
|
||||
"causal_mask",
|
||||
torch.tril(
|
||||
torch.ones(
|
||||
config.max_seq_len,
|
||||
config.max_seq_len,
|
||||
dtype=torch.bool,
|
||||
)
|
||||
),
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
# For kv cache
|
||||
self.max_batch_size = -1
|
||||
self.max_seq_len = -1
|
||||
|
||||
if init_weights:
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def setup_caches(
|
||||
self, max_batch_size: int, max_seq_len: int, dtype: torch.dtype = torch.bfloat16
|
||||
):
|
||||
if self.max_seq_len >= max_seq_len and self.max_batch_size >= max_batch_size:
|
||||
return
|
||||
|
||||
head_dim = self.config.dim // self.config.n_head
|
||||
max_seq_len = find_multiple(max_seq_len, 8)
|
||||
self.max_seq_len = max_seq_len
|
||||
self.max_batch_size = max_batch_size
|
||||
|
||||
for b in self.layers:
|
||||
b.attention.kv_cache = KVCache(
|
||||
max_batch_size,
|
||||
max_seq_len,
|
||||
self.config.n_local_heads,
|
||||
head_dim,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
def embed(self, x: Tensor) -> Tensor:
|
||||
vocab_embeds = [self.embeddings(x[:, 0])]
|
||||
for i in range(self.config.num_codebooks):
|
||||
emb = self.codebook_embeddings(x[:, i + 1] + i * self.config.codebook_size)
|
||||
emb[x[:, 0] != self.semantic_token_id] = 0
|
||||
vocab_embeds.append(emb)
|
||||
|
||||
x = torch.stack(vocab_embeds, dim=3)
|
||||
x = x.sum(dim=3)
|
||||
|
||||
return x
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inp: Tensor,
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
) -> BaseTransformerForwardResult:
|
||||
seq_len = inp.size(2)
|
||||
|
||||
# Here we want to merge the embeddings of the codebooks
|
||||
x = self.embed(inp)
|
||||
|
||||
freqs_cis = self.freqs_cis[:seq_len]
|
||||
|
||||
# Not that the causal mask here follows the definition of scaled_dot_product_attention
|
||||
# That is, FALSE means masked out
|
||||
# To maintain consistency, key_padding_mask use TRUE to mask out
|
||||
mask = None
|
||||
if key_padding_mask is not None:
|
||||
mask = self.causal_mask[None, None, :seq_len, :seq_len] # (B, N, Q, K)
|
||||
mask = mask & key_padding_mask[:, None, None, :].logical_not()
|
||||
|
||||
for layer in self.layers:
|
||||
if self.config.use_gradient_checkpointing and self.training:
|
||||
x = checkpoint(layer, x, freqs_cis, mask, use_reentrant=True)
|
||||
else:
|
||||
x = layer(x, freqs_cis, mask)
|
||||
|
||||
# We got slow_out here
|
||||
slow_out = self.norm(x)
|
||||
|
||||
if self.config.tie_word_embeddings:
|
||||
token_logits = F.linear(slow_out, self.embeddings.weight)
|
||||
else:
|
||||
token_logits = self.output(slow_out)
|
||||
|
||||
return BaseTransformerForwardResult(
|
||||
logits=token_logits,
|
||||
hidden_states=x,
|
||||
)
|
||||
|
||||
def forward_generate(
|
||||
self,
|
||||
x: Tensor,
|
||||
input_pos: Optional[Tensor] = None,
|
||||
return_all: bool = False,
|
||||
) -> BaseTransformerForwardResult:
|
||||
# This is used for generation, optimized for torch compile
|
||||
assert (
|
||||
self.max_seq_len != -1 and self.max_batch_size != -1
|
||||
), "Please call setup_caches before forward_generate"
|
||||
|
||||
x = self.embed(x)
|
||||
|
||||
mask = self.causal_mask[
|
||||
None, None, input_pos, : self.max_seq_len
|
||||
] # (B, N, Q, K)
|
||||
freqs_cis = self.freqs_cis[input_pos]
|
||||
|
||||
for layer in self.layers:
|
||||
x = layer(x, freqs_cis, mask, input_pos=input_pos)
|
||||
|
||||
# If prefill, we only calculate the logits of last token
|
||||
if x.size(1) > 1 and not return_all:
|
||||
x = x[:, -1:]
|
||||
|
||||
# We got slow_out here
|
||||
slow_out = self.norm(x)
|
||||
|
||||
if self.config.tie_word_embeddings:
|
||||
token_logits = F.linear(slow_out, self.embeddings.weight)
|
||||
else:
|
||||
token_logits = self.output(slow_out)
|
||||
|
||||
return BaseTransformerForwardResult(
|
||||
logits=token_logits,
|
||||
hidden_states=x,
|
||||
)
|
||||
|
||||
def _init_weights(self, module):
|
||||
std = self.config.initializer_range
|
||||
if isinstance(module, nn.Linear):
|
||||
module.weight.data.normal_(mean=0.0, std=std)
|
||||
if module.bias is not None:
|
||||
module.bias.data.zero_()
|
||||
elif isinstance(module, nn.Embedding):
|
||||
module.weight.data.normal_(mean=0.0, std=std)
|
||||
if module.padding_idx is not None:
|
||||
module.weight.data[module.padding_idx].zero_()
|
||||
|
||||
@staticmethod
|
||||
def from_pretrained(
|
||||
path: str,
|
||||
load_weights: bool = False,
|
||||
max_length: int | None = None,
|
||||
lora_config: LoraConfig | None = None,
|
||||
rope_base: int | None = None,
|
||||
) -> "BaseTransformer":
|
||||
config = BaseModelArgs.from_pretrained(str(path))
|
||||
if max_length is not None:
|
||||
config.max_seq_len = max_length
|
||||
log.info(f"Override max_seq_len to {max_length}")
|
||||
|
||||
if rope_base is not None:
|
||||
config.rope_base = rope_base
|
||||
log.info(f"Override rope_base to {rope_base}")
|
||||
|
||||
match config.model_type:
|
||||
case "naive":
|
||||
model_cls = NaiveTransformer
|
||||
case "dual_ar":
|
||||
model_cls = DualARTransformer
|
||||
case _:
|
||||
raise ValueError(f"Unknown model type: {config.model_type}")
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(str(path))
|
||||
log.info(f"Loading model from {path}, config: {config}")
|
||||
model = model_cls(config, tokenizer=tokenizer)
|
||||
|
||||
if lora_config is not None:
|
||||
setup_lora(model, lora_config)
|
||||
log.info(f"LoRA setup: {lora_config}")
|
||||
|
||||
if load_weights is False:
|
||||
log.info("Randomly initialized model")
|
||||
else:
|
||||
|
||||
if "int8" in str(Path(path)):
|
||||
logger.info("Using int8 weight-only quantization!")
|
||||
from tools.llama.quantize import WeightOnlyInt8QuantHandler
|
||||
|
||||
simple_quantizer = WeightOnlyInt8QuantHandler(model)
|
||||
model = simple_quantizer.convert_for_runtime()
|
||||
|
||||
if "int4" in str(Path(path)):
|
||||
logger.info("Using int4 quantization!")
|
||||
path_comps = path.name.split("-")
|
||||
assert path_comps[-2].startswith("g")
|
||||
groupsize = int(path_comps[-2][1:])
|
||||
from tools.llama.quantize import WeightOnlyInt4QuantHandler
|
||||
|
||||
simple_quantizer = WeightOnlyInt4QuantHandler(model, groupsize)
|
||||
model = simple_quantizer.convert_for_runtime()
|
||||
|
||||
weights = torch.load(
|
||||
Path(path) / "model.pth", map_location="cpu", mmap=True
|
||||
)
|
||||
|
||||
if "state_dict" in weights:
|
||||
logger.warning(
|
||||
"Using a TextToSemantic LightningModule checkpoint, "
|
||||
"please make sure it is a full model, not a LoRA model."
|
||||
)
|
||||
weights = weights["state_dict"]
|
||||
|
||||
if next(iter(weights.keys())).startswith("model."):
|
||||
logger.info(
|
||||
f"Remove prefix 'model.' created by TextToSemantic LightningModule from keys"
|
||||
)
|
||||
new_weights = OrderedDict()
|
||||
for k, v in weights.items():
|
||||
new_weights[k.replace("model.", "")] = v
|
||||
weights = new_weights
|
||||
|
||||
# Verify the name and shape of parameters since strict=False in load_state_dict.
|
||||
for k, v in model.named_parameters():
|
||||
if k not in weights:
|
||||
logger.warning(f"No weight for {k}")
|
||||
elif v.shape != weights[k].shape:
|
||||
logger.warning(
|
||||
f"Shape mismatch for {k}: {v.shape} vs {weights[k].shape}"
|
||||
)
|
||||
|
||||
err = model.load_state_dict(weights, strict=False, assign=True)
|
||||
log.info(f"Loaded weights with error: {err}")
|
||||
|
||||
return model
|
||||
|
||||
def save_pretrained(self, path: str, drop_lora: bool = False):
|
||||
path = Path(path)
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
self.config.save(path / "config.json")
|
||||
state_dict = self.state_dict()
|
||||
|
||||
if drop_lora:
|
||||
for key in list(state_dict.keys()):
|
||||
if "lora" not in key:
|
||||
continue
|
||||
|
||||
state_dict.pop(key)
|
||||
log.info(f"Drop LoRA parameter: {key}")
|
||||
|
||||
torch.save(state_dict, path / "model.pth")
|
||||
self.tokenizer.save_pretrained(path)
|
||||
|
||||
|
||||
class NaiveTransformer(BaseTransformer):
|
||||
def __init__(self, config: NaiveModelArgs, tokenizer: AutoTokenizer) -> None:
|
||||
super().__init__(config, init_weights=False, tokenizer=tokenizer)
|
||||
|
||||
self.codebook_norm = RMSNorm(config.dim, eps=config.norm_eps)
|
||||
self.codebook_output = nn.Linear(
|
||||
config.dim,
|
||||
config.codebook_size * config.num_codebooks,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def decode(self, result: BaseTransformerForwardResult) -> TransformerForwardResult:
|
||||
token_logits = result.logits
|
||||
x = result.hidden_states
|
||||
|
||||
# Codebook
|
||||
codebook_logits = self.codebook_output(self.codebook_norm(x))
|
||||
codebook_logits = rearrange(
|
||||
codebook_logits, "b n (c d) -> b n c d", c=self.config.num_codebooks
|
||||
)
|
||||
|
||||
return TransformerForwardResult(
|
||||
token_logits=token_logits,
|
||||
codebook_logits=codebook_logits,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inp: Tensor,
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
) -> TransformerForwardResult:
|
||||
result = super().forward(
|
||||
inp=inp,
|
||||
key_padding_mask=key_padding_mask,
|
||||
)
|
||||
return self.decode(result)
|
||||
|
||||
def forward_generate(
|
||||
self, x: Tensor, input_pos: Optional[Tensor] = None
|
||||
) -> TransformerForwardResult:
|
||||
result = super().forward_generate(x, input_pos)
|
||||
return self.decode(result)
|
||||
|
||||
|
||||
class DualARTransformer(BaseTransformer):
|
||||
def __init__(self, config: NaiveModelArgs, tokenizer: AutoTokenizer) -> None:
|
||||
super().__init__(config, init_weights=False, tokenizer=tokenizer)
|
||||
|
||||
# Fast transformer
|
||||
self.fast_embeddings = nn.Embedding(config.codebook_size, config.dim)
|
||||
|
||||
# The equivalent bs is so large that sdpa doesn't work
|
||||
self.fast_layers = nn.ModuleList(
|
||||
TransformerBlock(config, use_sdpa=False) for _ in range(config.n_fast_layer)
|
||||
)
|
||||
self.fast_norm = RMSNorm(config.dim, eps=config.norm_eps)
|
||||
self.fast_output = nn.Linear(
|
||||
config.dim,
|
||||
config.codebook_size,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def setup_caches(
|
||||
self, max_batch_size: int, max_seq_len: int, dtype: torch.dtype = torch.bfloat16
|
||||
):
|
||||
super().setup_caches(max_batch_size, max_seq_len, dtype)
|
||||
|
||||
head_dim = self.config.dim // self.config.n_head
|
||||
|
||||
# Fast transformer
|
||||
# The max seq len here is the number of codebooks
|
||||
for b in self.fast_layers:
|
||||
b.attention.kv_cache = KVCache(
|
||||
max_batch_size,
|
||||
self.config.num_codebooks,
|
||||
self.config.n_local_heads,
|
||||
head_dim,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inp: Tensor,
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
) -> TransformerForwardResult:
|
||||
parent_result = super().forward(inp, key_padding_mask)
|
||||
token_logits = parent_result.logits
|
||||
x = parent_result.hidden_states
|
||||
|
||||
# Fast transformer
|
||||
fast_seq_len = self.config.num_codebooks
|
||||
fast_mask = self.causal_mask[
|
||||
None, None, :fast_seq_len, :fast_seq_len
|
||||
] # (B, N, Q, K)
|
||||
fast_freqs_cis = self.freqs_cis[:fast_seq_len]
|
||||
|
||||
# Drop the last token and rotate left
|
||||
codebooks = inp[:, 1:-1, 1:]
|
||||
codebooks = F.pad(codebooks, (0, 1), value=0)
|
||||
codebook_embeddings = self.fast_embeddings(codebooks)
|
||||
x = torch.cat([x[:, None], codebook_embeddings], dim=1)
|
||||
b, s = x.size(0), x.size(2)
|
||||
x = rearrange(x, "b n s d -> (b s) n d") # flatten the batch and seq_len
|
||||
|
||||
# Remove padded part
|
||||
codebooks = rearrange(codebooks, "b n s -> (b s) n")
|
||||
codebook_mask = (codebooks == 0).all(dim=-1)
|
||||
|
||||
if torch.all(codebook_mask):
|
||||
# If all codebooks are padded, we keep first 8 to make sure the model runs
|
||||
codebook_mask[:8] = False
|
||||
|
||||
x_bs, x_len = x.size(0), x.size(1)
|
||||
x = x[~codebook_mask]
|
||||
|
||||
for layer in self.fast_layers:
|
||||
if self.config.use_gradient_checkpointing and self.training:
|
||||
x = checkpoint(layer, x, fast_freqs_cis, fast_mask, use_reentrant=True)
|
||||
else:
|
||||
x = layer(x, fast_freqs_cis, fast_mask)
|
||||
|
||||
# unflatten the batch and num_codebooks
|
||||
fast_out = self.fast_norm(x)
|
||||
codebook_logits = self.fast_output(fast_out)
|
||||
|
||||
# Re-pad the codebook_logits
|
||||
buffer = torch.zeros(
|
||||
x_bs,
|
||||
x_len,
|
||||
codebook_logits.size(-1),
|
||||
device=codebook_logits.device,
|
||||
dtype=codebook_logits.dtype,
|
||||
)
|
||||
buffer[~codebook_mask] = codebook_logits
|
||||
codebook_logits = buffer
|
||||
|
||||
assert codebook_logits.shape[1] == self.config.num_codebooks
|
||||
codebook_logits = rearrange(
|
||||
codebook_logits,
|
||||
"(b s) n d -> b s n d",
|
||||
b=b,
|
||||
s=s,
|
||||
n=self.config.num_codebooks,
|
||||
)
|
||||
|
||||
return TransformerForwardResult(
|
||||
token_logits=token_logits,
|
||||
codebook_logits=codebook_logits,
|
||||
)
|
||||
|
||||
def forward_generate_fast(
|
||||
self, x: Tensor, input_pos: Optional[Tensor] = None
|
||||
) -> Tensor:
|
||||
# Fast transformer
|
||||
x = x.view(1, 1, -1)
|
||||
|
||||
fast_mask = self.causal_mask[
|
||||
None, None, input_pos, : self.config.num_codebooks
|
||||
] # (B, N, Q, K)
|
||||
fast_freqs_cis = self.freqs_cis[input_pos]
|
||||
|
||||
for layer in self.fast_layers:
|
||||
x = layer(x, fast_freqs_cis, fast_mask, input_pos=input_pos)
|
||||
|
||||
# unflatten the batch and num_codebooks
|
||||
fast_out = self.fast_norm(x) # only take the last token
|
||||
codebook_logits = self.fast_output(fast_out)
|
||||
|
||||
return codebook_logits
|
||||
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
def __init__(self, config: BaseModelArgs, use_sdpa: bool = True) -> None:
|
||||
super().__init__()
|
||||
self.attention = Attention(config, use_sdpa=use_sdpa)
|
||||
self.feed_forward = FeedForward(config)
|
||||
self.ffn_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.attention_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
|
||||
def forward(
|
||||
self, x: Tensor, freqs_cis: Tensor, mask: Tensor, input_pos: Tensor = None
|
||||
) -> Tensor:
|
||||
h = x + self.attention(self.attention_norm(x), freqs_cis, mask, input_pos)
|
||||
out = h + self.feed_forward(self.ffn_norm(h))
|
||||
return out
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(self, config: BaseModelArgs, use_sdpa: bool = True):
|
||||
super().__init__()
|
||||
assert config.dim % config.n_head == 0
|
||||
|
||||
total_head_dim = (config.n_head + 2 * config.n_local_heads) * config.head_dim
|
||||
# key, query, value projections for all heads, but in a batch
|
||||
self.wqkv = nn.Linear(
|
||||
config.dim, total_head_dim, bias=config.attention_qkv_bias
|
||||
)
|
||||
self.wo = nn.Linear(config.dim, config.dim, bias=False)
|
||||
self.kv_cache = None
|
||||
|
||||
self.dropout = config.dropout
|
||||
self.n_head = config.n_head
|
||||
self.head_dim = config.head_dim
|
||||
self.n_local_heads = config.n_local_heads
|
||||
self.dim = config.dim
|
||||
self.use_sdpa = use_sdpa
|
||||
self._register_load_state_dict_pre_hook(self.load_hook)
|
||||
|
||||
def load_hook(self, state_dict, prefix, *args):
|
||||
if prefix + "wq.weight" in state_dict:
|
||||
wq = state_dict.pop(prefix + "wq.weight")
|
||||
wk = state_dict.pop(prefix + "wk.weight")
|
||||
wv = state_dict.pop(prefix + "wv.weight")
|
||||
state_dict[prefix + "wqkv.weight"] = torch.cat([wq, wk, wv])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
freqs_cis: Tensor,
|
||||
mask: Tensor,
|
||||
input_pos: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
bsz, seqlen, _ = x.shape
|
||||
|
||||
kv_size = self.n_local_heads * self.head_dim
|
||||
q, k, v = self.wqkv(x).split([self.dim, kv_size, kv_size], dim=-1)
|
||||
|
||||
q = q.view(bsz, seqlen, self.n_head, self.head_dim)
|
||||
k = k.view(bsz, seqlen, self.n_local_heads, self.head_dim)
|
||||
v = v.view(bsz, seqlen, self.n_local_heads, self.head_dim)
|
||||
|
||||
q = apply_rotary_emb(q, freqs_cis)
|
||||
k = apply_rotary_emb(k, freqs_cis)
|
||||
|
||||
q, k, v = map(lambda x: x.transpose(1, 2), (q, k, v))
|
||||
|
||||
if self.kv_cache is not None:
|
||||
k, v = self.kv_cache.update(input_pos, k, v)
|
||||
|
||||
k = k.repeat_interleave(self.n_head // self.n_local_heads, dim=1)
|
||||
v = v.repeat_interleave(self.n_head // self.n_local_heads, dim=1)
|
||||
|
||||
if self.use_sdpa:
|
||||
if mask is None:
|
||||
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
|
||||
y = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=self.dropout if self.training else 0.0,
|
||||
is_causal=True,
|
||||
# No third party attn_mask here to use flash_attention
|
||||
)
|
||||
else:
|
||||
y = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=mask,
|
||||
dropout_p=self.dropout if self.training else 0.0,
|
||||
)
|
||||
else:
|
||||
y = self.eq_scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=mask,
|
||||
dropout_p=self.dropout if self.training else 0.0,
|
||||
)
|
||||
|
||||
y = y.transpose(1, 2).contiguous().view(bsz, seqlen, self.dim)
|
||||
|
||||
return self.wo(y)
|
||||
|
||||
def eq_scaled_dot_product_attention(
|
||||
self,
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=None,
|
||||
dropout_p=0.0,
|
||||
) -> torch.Tensor:
|
||||
# This is a standard scaled dot product attention
|
||||
# It's low efficient, but it doesn't raise cuda error
|
||||
|
||||
L, S = query.size(-2), key.size(-2)
|
||||
scale_factor = 1 / math.sqrt(query.size(-1))
|
||||
attn_bias = torch.zeros(1, 1, L, S, dtype=query.dtype, device=query.device)
|
||||
|
||||
if attn_mask is not None:
|
||||
if attn_mask.dtype == torch.bool:
|
||||
attn_bias.masked_fill_(attn_mask.logical_not(), float("-inf"))
|
||||
else:
|
||||
attn_bias += attn_mask
|
||||
|
||||
attn_weight = query @ key.transpose(-2, -1) * scale_factor
|
||||
attn_weight += attn_bias
|
||||
attn_weight = torch.softmax(attn_weight, dim=-1)
|
||||
attn_weight = torch.dropout(attn_weight, dropout_p, train=True)
|
||||
|
||||
return attn_weight @ value
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, config: BaseModelArgs) -> None:
|
||||
super().__init__()
|
||||
self.w1 = nn.Linear(config.dim, config.intermediate_size, bias=False)
|
||||
self.w3 = nn.Linear(config.dim, config.intermediate_size, bias=False)
|
||||
self.w2 = nn.Linear(config.intermediate_size, config.dim, bias=False)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim: int, eps: float = 1e-5):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(torch.mean(x * x, dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
output = self._norm(x.float()).type_as(x)
|
||||
return output * self.weight
|
||||
|
||||
|
||||
def precompute_freqs_cis(seq_len: int, n_elem: int, base: int = 10000) -> Tensor:
|
||||
freqs = 1.0 / (
|
||||
base ** (torch.arange(0, n_elem, 2)[: (n_elem // 2)].float() / n_elem)
|
||||
)
|
||||
t = torch.arange(seq_len, device=freqs.device)
|
||||
freqs = torch.outer(t, freqs)
|
||||
freqs_cis = torch.polar(torch.ones_like(freqs), freqs)
|
||||
cache = torch.stack([freqs_cis.real, freqs_cis.imag], dim=-1)
|
||||
return cache.to(dtype=torch.bfloat16)
|
||||
|
||||
|
||||
def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
|
||||
xshaped = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||
freqs_cis = freqs_cis.view(1, xshaped.size(1), 1, xshaped.size(3), 2)
|
||||
x_out2 = torch.stack(
|
||||
[
|
||||
xshaped[..., 0] * freqs_cis[..., 0] - xshaped[..., 1] * freqs_cis[..., 1],
|
||||
xshaped[..., 1] * freqs_cis[..., 0] + xshaped[..., 0] * freqs_cis[..., 1],
|
||||
],
|
||||
-1,
|
||||
)
|
||||
|
||||
x_out2 = x_out2.flatten(3)
|
||||
return x_out2.type_as(x)
|
||||
@@ -0,0 +1,92 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
import loralib as lora
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoraConfig:
|
||||
r: int
|
||||
lora_alpha: float
|
||||
lora_dropout: float = 0.0
|
||||
|
||||
|
||||
def setup_lora(model, lora_config):
|
||||
# Replace the embedding layer with a LoRA layer
|
||||
model.embeddings = lora.Embedding(
|
||||
num_embeddings=model.embeddings.num_embeddings,
|
||||
embedding_dim=model.embeddings.embedding_dim,
|
||||
padding_idx=model.embeddings.padding_idx,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
)
|
||||
|
||||
model.codebook_embeddings = lora.Embedding(
|
||||
num_embeddings=model.codebook_embeddings.num_embeddings,
|
||||
embedding_dim=model.codebook_embeddings.embedding_dim,
|
||||
padding_idx=model.codebook_embeddings.padding_idx,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
)
|
||||
|
||||
# Replace output layer with a LoRA layer
|
||||
linears = [(model, "output")]
|
||||
|
||||
# Replace all linear layers with LoRA layers
|
||||
for layer in model.layers:
|
||||
linears.extend([(layer.attention, "wqkv"), (layer.attention, "wo")])
|
||||
linears.extend(
|
||||
[
|
||||
(layer.feed_forward, "w1"),
|
||||
(layer.feed_forward, "w2"),
|
||||
(layer.feed_forward, "w3"),
|
||||
]
|
||||
)
|
||||
|
||||
if hasattr(model, "fast_layers"):
|
||||
model.fast_embeddings = lora.Embedding(
|
||||
num_embeddings=model.fast_embeddings.num_embeddings,
|
||||
embedding_dim=model.fast_embeddings.embedding_dim,
|
||||
padding_idx=model.fast_embeddings.padding_idx,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
)
|
||||
|
||||
# Dual-AR model
|
||||
linears.append((model, "fast_output"))
|
||||
|
||||
for layer in model.fast_layers:
|
||||
linears.extend([(layer.attention, "wqkv"), (layer.attention, "wo")])
|
||||
linears.extend(
|
||||
[
|
||||
(layer.feed_forward, "w1"),
|
||||
(layer.feed_forward, "w2"),
|
||||
(layer.feed_forward, "w3"),
|
||||
]
|
||||
)
|
||||
|
||||
for module, layer in linears:
|
||||
updated_linear = lora.Linear(
|
||||
in_features=getattr(module, layer).in_features,
|
||||
out_features=getattr(module, layer).out_features,
|
||||
bias=getattr(module, layer).bias,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
lora_dropout=lora_config.lora_dropout,
|
||||
)
|
||||
setattr(module, layer, updated_linear)
|
||||
|
||||
# Mark only the LoRA layers as trainable
|
||||
lora.mark_only_lora_as_trainable(model, bias="none")
|
||||
|
||||
|
||||
def get_merged_state_dict(model):
|
||||
# This line will merge the state dict of the model and the LoRA parameters
|
||||
model.eval()
|
||||
|
||||
# Then we need to remove the LoRA parameters from the state dict
|
||||
state_dict = model.state_dict()
|
||||
for name in list(state_dict.keys()):
|
||||
if "lora" in name:
|
||||
state_dict.pop(name)
|
||||
|
||||
return state_dict
|
||||
@@ -0,0 +1,596 @@
|
||||
import math
|
||||
from functools import partial
|
||||
from math import prod
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from torch.nn.utils.parametrizations import weight_norm
|
||||
from torch.nn.utils.parametrize import remove_parametrizations
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
|
||||
def sequence_mask(length, max_length=None):
|
||||
if max_length is None:
|
||||
max_length = length.max()
|
||||
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
|
||||
return x.unsqueeze(0) < length.unsqueeze(1)
|
||||
|
||||
|
||||
def init_weights(m, mean=0.0, std=0.01):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv1D") != -1:
|
||||
m.weight.data.normal_(mean, std)
|
||||
|
||||
|
||||
def get_padding(kernel_size, dilation=1):
|
||||
return (kernel_size * dilation - dilation) // 2
|
||||
|
||||
|
||||
def unpad1d(x: torch.Tensor, paddings: tuple[int, int]):
|
||||
"""Remove padding from x, handling properly zero padding. Only for 1d!"""
|
||||
padding_left, padding_right = paddings
|
||||
assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
|
||||
assert (padding_left + padding_right) <= x.shape[-1]
|
||||
end = x.shape[-1] - padding_right
|
||||
return x[..., padding_left:end]
|
||||
|
||||
|
||||
def get_extra_padding_for_conv1d(
|
||||
x: torch.Tensor, kernel_size: int, stride: int, padding_total: int = 0
|
||||
) -> int:
|
||||
"""See `pad_for_conv1d`."""
|
||||
length = x.shape[-1]
|
||||
n_frames = (length - kernel_size + padding_total) / stride + 1
|
||||
ideal_length = (math.ceil(n_frames) - 1) * stride + (kernel_size - padding_total)
|
||||
return ideal_length - length
|
||||
|
||||
|
||||
def pad1d(
|
||||
x: torch.Tensor,
|
||||
paddings: tuple[int, int],
|
||||
mode: str = "zeros",
|
||||
value: float = 0.0,
|
||||
):
|
||||
"""Tiny wrapper around F.pad, just to allow for reflect padding on small input.
|
||||
If this is the case, we insert extra 0 padding to the right
|
||||
before the reflection happen.
|
||||
"""
|
||||
length = x.shape[-1]
|
||||
padding_left, padding_right = paddings
|
||||
assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
|
||||
if mode == "reflect":
|
||||
max_pad = max(padding_left, padding_right)
|
||||
extra_pad = 0
|
||||
if length <= max_pad:
|
||||
extra_pad = max_pad - length + 1
|
||||
x = F.pad(x, (0, extra_pad))
|
||||
padded = F.pad(x, paddings, mode, value)
|
||||
end = padded.shape[-1] - extra_pad
|
||||
return padded[..., :end]
|
||||
else:
|
||||
return F.pad(x, paddings, mode, value)
|
||||
|
||||
|
||||
class FishConvNet(nn.Module):
|
||||
def __init__(
|
||||
self, in_channels, out_channels, kernel_size, dilation=1, stride=1, groups=1
|
||||
):
|
||||
super(FishConvNet, self).__init__()
|
||||
self.conv = nn.Conv1d(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=stride,
|
||||
dilation=dilation,
|
||||
groups=groups,
|
||||
)
|
||||
self.stride = stride
|
||||
self.kernel_size = (kernel_size - 1) * dilation + 1
|
||||
self.dilation = dilation
|
||||
|
||||
def forward(self, x):
|
||||
pad = self.kernel_size - self.stride
|
||||
extra_padding = get_extra_padding_for_conv1d(
|
||||
x, self.kernel_size, self.stride, pad
|
||||
)
|
||||
x = pad1d(x, (pad, extra_padding), mode="constant", value=0)
|
||||
return self.conv(x).contiguous()
|
||||
|
||||
def weight_norm(self, name="weight", dim=0):
|
||||
self.conv = weight_norm(self.conv, name=name, dim=dim)
|
||||
return self
|
||||
|
||||
def remove_weight_norm(self):
|
||||
self.conv = remove_parametrizations(self.conv)
|
||||
return self
|
||||
|
||||
|
||||
class FishTransConvNet(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel_size, dilation=1, stride=1):
|
||||
super(FishTransConvNet, self).__init__()
|
||||
self.conv = nn.ConvTranspose1d(
|
||||
in_channels, out_channels, kernel_size, stride=stride, dilation=dilation
|
||||
)
|
||||
self.stride = stride
|
||||
self.kernel_size = kernel_size
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
pad = self.kernel_size - self.stride
|
||||
padding_right = math.ceil(pad)
|
||||
padding_left = pad - padding_right
|
||||
x = unpad1d(x, (padding_left, padding_right))
|
||||
return x.contiguous()
|
||||
|
||||
def weight_norm(self, name="weight", dim=0):
|
||||
self.conv = weight_norm(self.conv, name=name, dim=dim)
|
||||
return self
|
||||
|
||||
def remove_weight_norm(self):
|
||||
self.conv = remove_parametrizations(self.conv)
|
||||
return self
|
||||
|
||||
|
||||
class ResBlock1(torch.nn.Module):
|
||||
def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)):
|
||||
super().__init__()
|
||||
|
||||
self.convs1 = nn.ModuleList(
|
||||
[
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[0]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[1]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[2]
|
||||
).weight_norm(),
|
||||
]
|
||||
)
|
||||
self.convs1.apply(init_weights)
|
||||
|
||||
self.convs2 = nn.ModuleList(
|
||||
[
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[0]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[1]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[2]
|
||||
).weight_norm(),
|
||||
]
|
||||
)
|
||||
self.convs2.apply(init_weights)
|
||||
|
||||
def forward(self, x):
|
||||
for c1, c2 in zip(self.convs1, self.convs2):
|
||||
xt = F.silu(x)
|
||||
xt = c1(xt)
|
||||
xt = F.silu(xt)
|
||||
xt = c2(xt)
|
||||
x = xt + x
|
||||
return x
|
||||
|
||||
def remove_parametrizations(self):
|
||||
for conv in self.convs1:
|
||||
remove_parametrizations(conv, tensor_name="weight")
|
||||
for conv in self.convs2:
|
||||
remove_parametrizations(conv, tensor_name="weight")
|
||||
|
||||
|
||||
class ParallelBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
channels: int,
|
||||
kernel_sizes: tuple[int] = (3, 7, 11),
|
||||
dilation_sizes: tuple[tuple[int]] = ((1, 3, 5), (1, 3, 5), (1, 3, 5)),
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
assert len(kernel_sizes) == len(dilation_sizes)
|
||||
|
||||
self.blocks = nn.ModuleList()
|
||||
for k, d in zip(kernel_sizes, dilation_sizes):
|
||||
self.blocks.append(ResBlock1(channels, k, d))
|
||||
|
||||
def forward(self, x):
|
||||
return torch.stack([block(x) for block in self.blocks], dim=0).mean(dim=0)
|
||||
|
||||
def remove_parametrizations(self):
|
||||
for block in self.blocks:
|
||||
block.remove_parametrizations()
|
||||
|
||||
|
||||
class HiFiGANGenerator(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
hop_length: int = 512,
|
||||
upsample_rates: tuple[int] = (8, 8, 2, 2, 2),
|
||||
upsample_kernel_sizes: tuple[int] = (16, 16, 8, 2, 2),
|
||||
resblock_kernel_sizes: tuple[int] = (3, 7, 11),
|
||||
resblock_dilation_sizes: tuple[tuple[int]] = ((1, 3, 5), (1, 3, 5), (1, 3, 5)),
|
||||
num_mels: int = 128,
|
||||
upsample_initial_channel: int = 512,
|
||||
pre_conv_kernel_size: int = 7,
|
||||
post_conv_kernel_size: int = 7,
|
||||
post_activation: Callable = partial(nn.SiLU, inplace=True),
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
assert (
|
||||
prod(upsample_rates) == hop_length
|
||||
), f"hop_length must be {prod(upsample_rates)}"
|
||||
|
||||
self.conv_pre = FishConvNet(
|
||||
num_mels,
|
||||
upsample_initial_channel,
|
||||
pre_conv_kernel_size,
|
||||
stride=1,
|
||||
).weight_norm()
|
||||
|
||||
self.num_upsamples = len(upsample_rates)
|
||||
self.num_kernels = len(resblock_kernel_sizes)
|
||||
|
||||
self.noise_convs = nn.ModuleList()
|
||||
self.ups = nn.ModuleList()
|
||||
|
||||
for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
|
||||
self.ups.append(
|
||||
FishTransConvNet(
|
||||
upsample_initial_channel // (2**i),
|
||||
upsample_initial_channel // (2 ** (i + 1)),
|
||||
k,
|
||||
stride=u,
|
||||
).weight_norm()
|
||||
)
|
||||
|
||||
self.resblocks = nn.ModuleList()
|
||||
for i in range(len(self.ups)):
|
||||
ch = upsample_initial_channel // (2 ** (i + 1))
|
||||
self.resblocks.append(
|
||||
ParallelBlock(ch, resblock_kernel_sizes, resblock_dilation_sizes)
|
||||
)
|
||||
|
||||
self.activation_post = post_activation()
|
||||
self.conv_post = FishConvNet(
|
||||
ch, 1, post_conv_kernel_size, stride=1
|
||||
).weight_norm()
|
||||
self.ups.apply(init_weights)
|
||||
self.conv_post.apply(init_weights)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_pre(x)
|
||||
|
||||
for i in range(self.num_upsamples):
|
||||
x = F.silu(x, inplace=True)
|
||||
x = self.ups[i](x)
|
||||
|
||||
if self.training and self.checkpointing:
|
||||
x = checkpoint(
|
||||
self.resblocks[i],
|
||||
x,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
x = self.resblocks[i](x)
|
||||
|
||||
x = self.activation_post(x)
|
||||
x = self.conv_post(x)
|
||||
x = torch.tanh(x)
|
||||
|
||||
return x
|
||||
|
||||
def remove_parametrizations(self):
|
||||
for up in self.ups:
|
||||
remove_parametrizations(up, tensor_name="weight")
|
||||
for block in self.resblocks:
|
||||
block.remove_parametrizations()
|
||||
remove_parametrizations(self.conv_pre, tensor_name="weight")
|
||||
remove_parametrizations(self.conv_post, tensor_name="weight")
|
||||
|
||||
|
||||
# DropPath copied from timm library
|
||||
def drop_path(
|
||||
x, drop_prob: float = 0.0, training: bool = False, scale_by_keep: bool = True
|
||||
):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||
|
||||
This is the same as the DropConnect impl I created for EfficientNet, etc networks, however,
|
||||
the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
|
||||
See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for
|
||||
changing the layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use
|
||||
'survival rate' as the argument.
|
||||
|
||||
""" # noqa: E501
|
||||
|
||||
if drop_prob == 0.0 or not training:
|
||||
return x
|
||||
keep_prob = 1 - drop_prob
|
||||
shape = (x.shape[0],) + (1,) * (
|
||||
x.ndim - 1
|
||||
) # work with diff dim tensors, not just 2D ConvNets
|
||||
random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
|
||||
if keep_prob > 0.0 and scale_by_keep:
|
||||
random_tensor.div_(keep_prob)
|
||||
return x * random_tensor
|
||||
|
||||
|
||||
class DropPath(nn.Module):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).""" # noqa: E501
|
||||
|
||||
def __init__(self, drop_prob: float = 0.0, scale_by_keep: bool = True):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
self.scale_by_keep = scale_by_keep
|
||||
|
||||
def forward(self, x):
|
||||
return drop_path(x, self.drop_prob, self.training, self.scale_by_keep)
|
||||
|
||||
def extra_repr(self):
|
||||
return f"drop_prob={round(self.drop_prob,3):0.3f}"
|
||||
|
||||
|
||||
class LayerNorm(nn.Module):
|
||||
r"""LayerNorm that supports two data formats: channels_last (default) or channels_first.
|
||||
The ordering of the dimensions in the inputs. channels_last corresponds to inputs with
|
||||
shape (batch_size, height, width, channels) while channels_first corresponds to inputs
|
||||
with shape (batch_size, channels, height, width).
|
||||
""" # noqa: E501
|
||||
|
||||
def __init__(self, normalized_shape, eps=1e-6, data_format="channels_last"):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(normalized_shape))
|
||||
self.bias = nn.Parameter(torch.zeros(normalized_shape))
|
||||
self.eps = eps
|
||||
self.data_format = data_format
|
||||
if self.data_format not in ["channels_last", "channels_first"]:
|
||||
raise NotImplementedError
|
||||
self.normalized_shape = (normalized_shape,)
|
||||
|
||||
def forward(self, x):
|
||||
if self.data_format == "channels_last":
|
||||
return F.layer_norm(
|
||||
x, self.normalized_shape, self.weight, self.bias, self.eps
|
||||
)
|
||||
elif self.data_format == "channels_first":
|
||||
u = x.mean(1, keepdim=True)
|
||||
s = (x - u).pow(2).mean(1, keepdim=True)
|
||||
x = (x - u) / torch.sqrt(s + self.eps)
|
||||
x = self.weight[:, None] * x + self.bias[:, None]
|
||||
return x
|
||||
|
||||
|
||||
# ConvNeXt Block copied from https://github.com/fishaudio/fish-diffusion/blob/main/fish_diffusion/modules/convnext.py
|
||||
class ConvNeXtBlock(nn.Module):
|
||||
r"""ConvNeXt Block. There are two equivalent implementations:
|
||||
(1) DwConv -> LayerNorm (channels_first) -> 1x1 Conv -> GELU -> 1x1 Conv; all in (N, C, H, W)
|
||||
(2) DwConv -> Permute to (N, H, W, C); LayerNorm (channels_last) -> Linear -> GELU -> Linear; Permute back
|
||||
We use (2) as we find it slightly faster in PyTorch
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
drop_path (float): Stochastic depth rate. Default: 0.0
|
||||
layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.0.
|
||||
kernel_size (int): Kernel size for depthwise conv. Default: 7.
|
||||
dilation (int): Dilation for depthwise conv. Default: 1.
|
||||
""" # noqa: E501
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
drop_path: float = 0.0,
|
||||
layer_scale_init_value: float = 1e-6,
|
||||
mlp_ratio: float = 4.0,
|
||||
kernel_size: int = 7,
|
||||
dilation: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.dwconv = FishConvNet(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size=kernel_size,
|
||||
# padding=int(dilation * (kernel_size - 1) / 2),
|
||||
groups=dim,
|
||||
) # depthwise conv
|
||||
self.norm = LayerNorm(dim, eps=1e-6)
|
||||
self.pwconv1 = nn.Linear(
|
||||
dim, int(mlp_ratio * dim)
|
||||
) # pointwise/1x1 convs, implemented with linear layers
|
||||
self.act = nn.GELU()
|
||||
self.pwconv2 = nn.Linear(int(mlp_ratio * dim), dim)
|
||||
self.gamma = (
|
||||
nn.Parameter(layer_scale_init_value * torch.ones((dim)), requires_grad=True)
|
||||
if layer_scale_init_value > 0
|
||||
else None
|
||||
)
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
|
||||
|
||||
def forward(self, x, apply_residual: bool = True):
|
||||
input = x
|
||||
|
||||
x = self.dwconv(x)
|
||||
x = x.permute(0, 2, 1) # (N, C, L) -> (N, L, C)
|
||||
x = self.norm(x)
|
||||
x = self.pwconv1(x)
|
||||
x = self.act(x)
|
||||
x = self.pwconv2(x)
|
||||
|
||||
if self.gamma is not None:
|
||||
x = self.gamma * x
|
||||
|
||||
x = x.permute(0, 2, 1) # (N, L, C) -> (N, C, L)
|
||||
x = self.drop_path(x)
|
||||
|
||||
if apply_residual:
|
||||
x = input + x
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class ConvNeXtEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_channels: int = 3,
|
||||
depths: list[int] = [3, 3, 9, 3],
|
||||
dims: list[int] = [96, 192, 384, 768],
|
||||
drop_path_rate: float = 0.0,
|
||||
layer_scale_init_value: float = 1e-6,
|
||||
kernel_size: int = 7,
|
||||
):
|
||||
super().__init__()
|
||||
assert len(depths) == len(dims)
|
||||
|
||||
self.downsample_layers = nn.ModuleList()
|
||||
stem = nn.Sequential(
|
||||
FishConvNet(
|
||||
input_channels,
|
||||
dims[0],
|
||||
kernel_size=7,
|
||||
# padding=3,
|
||||
# padding_mode="replicate",
|
||||
# padding_mode="zeros",
|
||||
),
|
||||
LayerNorm(dims[0], eps=1e-6, data_format="channels_first"),
|
||||
)
|
||||
self.downsample_layers.append(stem)
|
||||
|
||||
for i in range(len(depths) - 1):
|
||||
mid_layer = nn.Sequential(
|
||||
LayerNorm(dims[i], eps=1e-6, data_format="channels_first"),
|
||||
nn.Conv1d(dims[i], dims[i + 1], kernel_size=1),
|
||||
)
|
||||
self.downsample_layers.append(mid_layer)
|
||||
|
||||
self.stages = nn.ModuleList()
|
||||
dp_rates = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]
|
||||
|
||||
cur = 0
|
||||
for i in range(len(depths)):
|
||||
stage = nn.Sequential(
|
||||
*[
|
||||
ConvNeXtBlock(
|
||||
dim=dims[i],
|
||||
drop_path=dp_rates[cur + j],
|
||||
layer_scale_init_value=layer_scale_init_value,
|
||||
kernel_size=kernel_size,
|
||||
)
|
||||
for j in range(depths[i])
|
||||
]
|
||||
)
|
||||
self.stages.append(stage)
|
||||
cur += depths[i]
|
||||
|
||||
self.norm = LayerNorm(dims[-1], eps=1e-6, data_format="channels_first")
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, (nn.Conv1d, nn.Linear)):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
for i in range(len(self.downsample_layers)):
|
||||
x = self.downsample_layers[i](x)
|
||||
x = self.stages[i](x)
|
||||
|
||||
return self.norm(x)
|
||||
|
||||
|
||||
class FireflyArchitecture(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
backbone: nn.Module,
|
||||
head: nn.Module,
|
||||
quantizer: nn.Module,
|
||||
spec_transform: nn.Module,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.backbone = backbone
|
||||
self.head = head
|
||||
self.quantizer = quantizer
|
||||
self.spec_transform = spec_transform
|
||||
self.downsample_factor = math.prod(self.quantizer.downsample_factor)
|
||||
|
||||
def forward(self, x: torch.Tensor, template=None, mask=None) -> torch.Tensor:
|
||||
if self.spec_transform is not None:
|
||||
x = self.spec_transform(x)
|
||||
|
||||
x = self.backbone(x)
|
||||
if mask is not None:
|
||||
x = x * mask
|
||||
|
||||
if self.quantizer is not None:
|
||||
vq_result = self.quantizer(x)
|
||||
x = vq_result.z
|
||||
|
||||
if mask is not None:
|
||||
x = x * mask
|
||||
|
||||
x = self.head(x, template=template)
|
||||
|
||||
if x.ndim == 2:
|
||||
x = x[:, None, :]
|
||||
|
||||
if self.vq is not None:
|
||||
return x, vq_result
|
||||
|
||||
return x
|
||||
|
||||
def encode(self, audios, audio_lengths):
|
||||
audios = audios.float()
|
||||
|
||||
mels = self.spec_transform(audios)
|
||||
mel_lengths = audio_lengths // self.spec_transform.hop_length
|
||||
mel_masks = sequence_mask(mel_lengths, mels.shape[2])
|
||||
mel_masks_float_conv = mel_masks[:, None, :].float()
|
||||
mels = mels * mel_masks_float_conv
|
||||
|
||||
# Encode
|
||||
encoded_features = self.backbone(mels) * mel_masks_float_conv
|
||||
feature_lengths = mel_lengths // self.downsample_factor
|
||||
|
||||
return self.quantizer.encode(encoded_features), feature_lengths
|
||||
|
||||
def decode(self, indices, feature_lengths) -> torch.Tensor:
|
||||
mel_masks = sequence_mask(
|
||||
feature_lengths * self.downsample_factor,
|
||||
indices.shape[2] * self.downsample_factor,
|
||||
)
|
||||
mel_masks_float_conv = mel_masks[:, None, :].float()
|
||||
audio_lengths = (
|
||||
feature_lengths * self.downsample_factor * self.spec_transform.hop_length
|
||||
)
|
||||
|
||||
audio_masks = sequence_mask(
|
||||
audio_lengths,
|
||||
indices.shape[2] * self.downsample_factor * self.spec_transform.hop_length,
|
||||
)
|
||||
audio_masks_float_conv = audio_masks[:, None, :].float()
|
||||
|
||||
z = self.quantizer.decode(indices) * mel_masks_float_conv
|
||||
x = self.head(z) * audio_masks_float_conv
|
||||
|
||||
return x, audio_lengths
|
||||
|
||||
def remove_parametrizations(self):
|
||||
if hasattr(self.backbone, "remove_parametrizations"):
|
||||
self.backbone.remove_parametrizations()
|
||||
|
||||
if hasattr(self.head, "remove_parametrizations"):
|
||||
self.head.remove_parametrizations()
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
@@ -0,0 +1,116 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from vector_quantize_pytorch import GroupedResidualFSQ
|
||||
|
||||
from .firefly import ConvNeXtBlock, FishConvNet, FishTransConvNet
|
||||
|
||||
|
||||
@dataclass
|
||||
class FSQResult:
|
||||
z: torch.Tensor
|
||||
codes: torch.Tensor
|
||||
latents: torch.Tensor
|
||||
|
||||
|
||||
class DownsampleFiniteScalarQuantize(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int = 512,
|
||||
n_codebooks: int = 9,
|
||||
n_groups: int = 1,
|
||||
levels: tuple[int] = (8, 5, 5, 5), # Approximate 2**10
|
||||
downsample_factor: tuple[int] = (2, 2),
|
||||
downsample_dims: tuple[int] | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if downsample_dims is None:
|
||||
downsample_dims = [input_dim for _ in range(len(downsample_factor))]
|
||||
|
||||
all_dims = (input_dim,) + tuple(downsample_dims)
|
||||
|
||||
self.residual_fsq = GroupedResidualFSQ(
|
||||
dim=all_dims[-1],
|
||||
levels=levels,
|
||||
num_quantizers=n_codebooks,
|
||||
groups=n_groups,
|
||||
)
|
||||
|
||||
self.downsample_factor = downsample_factor
|
||||
self.downsample_dims = downsample_dims
|
||||
|
||||
self.downsample = nn.Sequential(
|
||||
*[
|
||||
nn.Sequential(
|
||||
FishConvNet(
|
||||
all_dims[idx],
|
||||
all_dims[idx + 1],
|
||||
kernel_size=factor,
|
||||
stride=factor,
|
||||
),
|
||||
ConvNeXtBlock(dim=all_dims[idx + 1]),
|
||||
)
|
||||
for idx, factor in enumerate(downsample_factor)
|
||||
]
|
||||
)
|
||||
|
||||
self.upsample = nn.Sequential(
|
||||
*[
|
||||
nn.Sequential(
|
||||
FishTransConvNet(
|
||||
all_dims[idx + 1],
|
||||
all_dims[idx],
|
||||
kernel_size=factor,
|
||||
stride=factor,
|
||||
),
|
||||
ConvNeXtBlock(dim=all_dims[idx]),
|
||||
)
|
||||
for idx, factor in reversed(list(enumerate(downsample_factor)))
|
||||
]
|
||||
)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, (nn.Conv1d, nn.Linear)):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, z) -> FSQResult:
|
||||
original_shape = z.shape
|
||||
z = self.downsample(z)
|
||||
quantized, indices = self.residual_fsq(z.mT)
|
||||
result = FSQResult(
|
||||
z=quantized.mT,
|
||||
codes=indices.mT,
|
||||
latents=z,
|
||||
)
|
||||
result.z = self.upsample(result.z)
|
||||
|
||||
# Pad or crop z to match original shape
|
||||
diff = original_shape[-1] - result.z.shape[-1]
|
||||
left = diff // 2
|
||||
right = diff - left
|
||||
|
||||
if diff > 0:
|
||||
result.z = F.pad(result.z, (left, right))
|
||||
elif diff < 0:
|
||||
result.z = result.z[..., left:-right]
|
||||
|
||||
return result
|
||||
|
||||
def encode(self, z):
|
||||
z = self.downsample(z)
|
||||
_, indices = self.residual_fsq(z.mT)
|
||||
indices = rearrange(indices, "g b l r -> b (g r) l")
|
||||
return indices
|
||||
|
||||
def decode(self, indices: torch.Tensor):
|
||||
indices = rearrange(indices, "b (g r) l -> g b l r", g=self.residual_fsq.groups)
|
||||
z_q = self.residual_fsq.get_output_from_indices(indices)
|
||||
z_q = self.upsample(z_q.mT)
|
||||
return z_q
|
||||
@@ -0,0 +1,94 @@
|
||||
import matplotlib
|
||||
import torch
|
||||
from matplotlib import pyplot as plt
|
||||
|
||||
matplotlib.use("Agg")
|
||||
|
||||
|
||||
def convert_pad_shape(pad_shape):
|
||||
l = pad_shape[::-1]
|
||||
pad_shape = [item for sublist in l for item in sublist]
|
||||
return pad_shape
|
||||
|
||||
|
||||
def sequence_mask(length, max_length=None):
|
||||
if max_length is None:
|
||||
max_length = length.max()
|
||||
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
|
||||
return x.unsqueeze(0) < length.unsqueeze(1)
|
||||
|
||||
|
||||
def init_weights(m, mean=0.0, std=0.01):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv") != -1:
|
||||
m.weight.data.normal_(mean, std)
|
||||
|
||||
|
||||
def get_padding(kernel_size, dilation=1):
|
||||
return int((kernel_size * dilation - dilation) / 2)
|
||||
|
||||
|
||||
def plot_mel(data, titles=None):
|
||||
fig, axes = plt.subplots(len(data), 1, squeeze=False)
|
||||
|
||||
if titles is None:
|
||||
titles = [None for i in range(len(data))]
|
||||
|
||||
plt.tight_layout()
|
||||
|
||||
for i in range(len(data)):
|
||||
mel = data[i]
|
||||
|
||||
if isinstance(mel, torch.Tensor):
|
||||
mel = mel.float().detach().cpu().numpy()
|
||||
|
||||
axes[i][0].imshow(mel, origin="lower")
|
||||
axes[i][0].set_aspect(2.5, adjustable="box")
|
||||
axes[i][0].set_ylim(0, mel.shape[0])
|
||||
axes[i][0].set_title(titles[i], fontsize="medium")
|
||||
axes[i][0].tick_params(labelsize="x-small", left=False, labelleft=False)
|
||||
axes[i][0].set_anchor("W")
|
||||
|
||||
return fig
|
||||
|
||||
|
||||
def slice_segments(x, ids_str, segment_size=4):
|
||||
ret = torch.zeros_like(x[:, :, :segment_size])
|
||||
for i in range(x.size(0)):
|
||||
idx_str = ids_str[i]
|
||||
idx_end = idx_str + segment_size
|
||||
ret[i] = x[i, :, idx_str:idx_end]
|
||||
|
||||
return ret
|
||||
|
||||
|
||||
def rand_slice_segments(x, x_lengths=None, segment_size=4):
|
||||
b, d, t = x.size()
|
||||
if x_lengths is None:
|
||||
x_lengths = t
|
||||
ids_str_max = torch.clamp(x_lengths - segment_size + 1, min=0)
|
||||
ids_str = (torch.rand([b], device=x.device) * ids_str_max).to(dtype=torch.long)
|
||||
ret = slice_segments(x, ids_str, segment_size)
|
||||
return ret, ids_str
|
||||
|
||||
|
||||
@torch.jit.script
|
||||
def fused_add_tanh_sigmoid_multiply(in_act, n_channels):
|
||||
n_channels_int = n_channels[0]
|
||||
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
||||
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
|
||||
acts = t_act * s_act
|
||||
|
||||
return acts
|
||||
|
||||
|
||||
def avg_with_mask(x, mask):
|
||||
assert mask.dtype == torch.float, "Mask should be float"
|
||||
|
||||
if mask.ndim == 2:
|
||||
mask = mask.unsqueeze(1)
|
||||
|
||||
if mask.shape[1] == 1:
|
||||
mask = mask.expand_as(x)
|
||||
|
||||
return (x * mask).sum() / mask.sum()
|
||||
@@ -0,0 +1,130 @@
|
||||
import re
|
||||
import string
|
||||
|
||||
from .clean import clean_text
|
||||
|
||||
|
||||
def utf_8_len(text):
|
||||
return len(text.encode("utf-8"))
|
||||
|
||||
|
||||
def break_text(texts, length, splits: set):
|
||||
for text in texts:
|
||||
if utf_8_len(text) <= length:
|
||||
yield text
|
||||
continue
|
||||
|
||||
curr = ""
|
||||
for char in text:
|
||||
curr += char
|
||||
|
||||
if char in splits:
|
||||
yield curr
|
||||
curr = ""
|
||||
|
||||
if curr:
|
||||
yield curr
|
||||
|
||||
|
||||
def break_text_by_length(texts, length):
|
||||
for text in texts:
|
||||
if utf_8_len(text) <= length:
|
||||
yield text
|
||||
continue
|
||||
|
||||
curr = ""
|
||||
for char in text:
|
||||
curr += char
|
||||
|
||||
if utf_8_len(curr) >= length:
|
||||
yield curr
|
||||
curr = ""
|
||||
|
||||
if curr:
|
||||
yield curr
|
||||
|
||||
|
||||
def add_cleaned(curr, segments):
|
||||
curr = curr.strip()
|
||||
if curr and not all(c.isspace() or c in string.punctuation for c in curr):
|
||||
segments.append(curr)
|
||||
|
||||
|
||||
def protect_float(text):
|
||||
# Turns 3.14 into <3_f_14> to prevent splitting
|
||||
return re.sub(r"(\d+)\.(\d+)", r"<\1_f_\2>", text)
|
||||
|
||||
|
||||
def unprotect_float(text):
|
||||
# Turns <3_f_14> into 3.14
|
||||
return re.sub(r"<(\d+)_f_(\d+)>", r"\1.\2", text)
|
||||
|
||||
|
||||
def split_text(text, length):
|
||||
text = clean_text(text)
|
||||
|
||||
# Break the text into pieces with following rules:
|
||||
# 1. Split the text at ".", "!", "?" if text is NOT a float
|
||||
# 2. If the text is longer than length, split at ","
|
||||
# 3. If the text is still longer than length, split at " "
|
||||
# 4. If the text is still longer than length, split at any character to length
|
||||
|
||||
texts = [text]
|
||||
texts = map(protect_float, texts)
|
||||
texts = break_text(texts, length, {".", "!", "?"})
|
||||
texts = map(unprotect_float, texts)
|
||||
texts = break_text(texts, length, {","})
|
||||
texts = break_text(texts, length, {" "})
|
||||
texts = list(break_text_by_length(texts, length))
|
||||
|
||||
# Then, merge the texts into segments with length <= length
|
||||
segments = []
|
||||
curr = ""
|
||||
|
||||
for text in texts:
|
||||
if utf_8_len(curr) + utf_8_len(text) <= length:
|
||||
curr += text
|
||||
else:
|
||||
add_cleaned(curr, segments)
|
||||
curr = text
|
||||
|
||||
if curr:
|
||||
add_cleaned(curr, segments)
|
||||
|
||||
return segments
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Test the split_text function
|
||||
|
||||
text = "This is a test sentence. This is another test sentence. And a third one."
|
||||
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence.",
|
||||
"This is another test sentence. And a third one.",
|
||||
]
|
||||
assert split_text("a,aaaaaa3.14", 10) == ["a,", "aaaaaa3.14"]
|
||||
assert split_text(" ", 10) == []
|
||||
assert split_text("a", 10) == ["a"]
|
||||
|
||||
text = "This is a test sentence with only commas, and no dots, and no exclamation marks, and no question marks, and no newlines."
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence with only commas,",
|
||||
"and no dots, and no exclamation marks,",
|
||||
"and no question marks, and no newlines.",
|
||||
]
|
||||
|
||||
text = "This is a test sentence This is a test sentence This is a test sentence. This is a test sentence, This is a test sentence, This is a test sentence."
|
||||
# First half split at " ", second half split at ","
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence This is a test sentence",
|
||||
"This is a test sentence. This is a test sentence,",
|
||||
"This is a test sentence, This is a test sentence.",
|
||||
]
|
||||
|
||||
text = "这是一段很长的中文文本,而且没有句号,也没有感叹号,也没有问号,也没有换行符。"
|
||||
assert split_text(text, 50) == [
|
||||
"这是一段很长的中文文本,",
|
||||
"而且没有句号,也没有感叹号,",
|
||||
"也没有问号,也没有换行符.",
|
||||
]
|
||||
@@ -0,0 +1,4 @@
|
||||
from .clean import clean_text
|
||||
from .spliter import split_text
|
||||
|
||||
__all__ = ["clean_text", "split_text"]
|
||||
@@ -0,0 +1,114 @@
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
|
||||
# C extensions
|
||||
*.so
|
||||
|
||||
# Distribution / packaging
|
||||
.Python
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
MANIFEST
|
||||
|
||||
# PyInstaller
|
||||
# Usually these files are written by a python script from a template
|
||||
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||
*.manifest
|
||||
*.spec
|
||||
|
||||
# Installer logs
|
||||
pip-log.txt
|
||||
pip-delete-this-directory.txt
|
||||
|
||||
# Unit test / coverage reports
|
||||
htmlcov/
|
||||
.tox/
|
||||
.coverage
|
||||
.coverage.*
|
||||
.cache
|
||||
nosetests.xml
|
||||
coverage.xml
|
||||
*.cover
|
||||
.hypothesis/
|
||||
.pytest_cache/
|
||||
|
||||
# Translations
|
||||
*.mo
|
||||
*.pot
|
||||
|
||||
# Django stuff:
|
||||
*.log
|
||||
local_settings.py
|
||||
db.sqlite3
|
||||
|
||||
# Flask stuff:
|
||||
instance/
|
||||
.webassets-cache
|
||||
|
||||
# Scrapy stuff:
|
||||
.scrapy
|
||||
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
|
||||
# PyBuilder
|
||||
target/
|
||||
|
||||
# Jupyter Notebook
|
||||
.ipynb_checkpoints
|
||||
|
||||
# pyenv
|
||||
.python-version
|
||||
|
||||
# celery beat schedule file
|
||||
celerybeat-schedule
|
||||
|
||||
# SageMath parsed files
|
||||
*.sage.py
|
||||
|
||||
# Environments
|
||||
.env
|
||||
.venv
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
env.bak/
|
||||
venv.bak/
|
||||
|
||||
# Spyder project settings
|
||||
.spyderproject
|
||||
.spyproject
|
||||
|
||||
# Rope project settings
|
||||
.ropeproject
|
||||
|
||||
# mkdocs documentation
|
||||
/site
|
||||
|
||||
# mypy
|
||||
.mypy_cache/
|
||||
|
||||
# JetBrains PyCharm
|
||||
.idea
|
||||
|
||||
# Customize
|
||||
references
|
||||
url.txt
|
||||
|
||||
# Git
|
||||
.git
|
||||
@@ -0,0 +1,36 @@
|
||||
# This account is no longer in use, see [Atomicoo](https://github.com/atomicoo) for my latest works.
|
||||
|
||||
# Chn Text Norm
|
||||
|
||||
this is a repository for chinese text normalization (no longer maintained).
|
||||
|
||||
## Quick Start ##
|
||||
|
||||
### Git Clone Repo ###
|
||||
|
||||
git clone this repo to the root directory of your project which need to use it.
|
||||
|
||||
cd /path/to/proj
|
||||
git clone https://github.com/Joee1995/chn-text-norm.git
|
||||
|
||||
after that, your doc tree should be:
|
||||
```
|
||||
proj # root of your project
|
||||
|--- chn_text_norm # this chn-text-norm tool
|
||||
|--- text.py
|
||||
|--- ...
|
||||
|--- text_normalize.py # your text normalization code
|
||||
|--- ...
|
||||
```
|
||||
|
||||
### How to Use ? ###
|
||||
|
||||
# text_normalize.py
|
||||
from chn_text_norm.text import *
|
||||
|
||||
raw_text = 'your raw text'
|
||||
text = Text(raw_text=raw_text).normalize()
|
||||
|
||||
### How to add quantums ###
|
||||
|
||||
打开test.py,然后你就知道怎么做了。
|
||||
@@ -0,0 +1,172 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""基本类
|
||||
中文字符类
|
||||
中文数字/数位类
|
||||
中文数字类
|
||||
中文数位类
|
||||
中文数字系统类
|
||||
中文数学符号类
|
||||
*中文其他符号类
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-02"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_constant import NUMBERING_TYPES
|
||||
|
||||
|
||||
class ChineseChar(object):
|
||||
"""
|
||||
中文字符
|
||||
每个字符对应简体和繁体,
|
||||
e.g. 简体 = '负', 繁体 = '負'
|
||||
转换时可转换为简体或繁体
|
||||
"""
|
||||
|
||||
def __init__(self, simplified, traditional):
|
||||
self.simplified = simplified
|
||||
self.traditional = traditional
|
||||
self.__repr__ = self.__str__
|
||||
|
||||
def __str__(self):
|
||||
return self.simplified or self.traditional or None
|
||||
|
||||
def __repr__(self):
|
||||
return self.__str__()
|
||||
|
||||
|
||||
class ChineseNumberUnit(ChineseChar):
|
||||
"""
|
||||
中文数字/数位字符
|
||||
每个字符除繁简体外还有一个额外的大写字符
|
||||
e.g. '陆' 和 '陸'
|
||||
"""
|
||||
|
||||
def __init__(self, power, simplified, traditional, big_s, big_t):
|
||||
super(ChineseNumberUnit, self).__init__(simplified, traditional)
|
||||
self.power = power
|
||||
self.big_s = big_s
|
||||
self.big_t = big_t
|
||||
|
||||
def __str__(self):
|
||||
return "10^{}".format(self.power)
|
||||
|
||||
@classmethod
|
||||
def create(cls, index, value, numbering_type=NUMBERING_TYPES[1], small_unit=False):
|
||||
|
||||
if small_unit:
|
||||
return ChineseNumberUnit(
|
||||
power=index + 1,
|
||||
simplified=value[0],
|
||||
traditional=value[1],
|
||||
big_s=value[1],
|
||||
big_t=value[1],
|
||||
)
|
||||
elif numbering_type == NUMBERING_TYPES[0]:
|
||||
return ChineseNumberUnit(
|
||||
power=index + 8,
|
||||
simplified=value[0],
|
||||
traditional=value[1],
|
||||
big_s=value[0],
|
||||
big_t=value[1],
|
||||
)
|
||||
elif numbering_type == NUMBERING_TYPES[1]:
|
||||
return ChineseNumberUnit(
|
||||
power=(index + 2) * 4,
|
||||
simplified=value[0],
|
||||
traditional=value[1],
|
||||
big_s=value[0],
|
||||
big_t=value[1],
|
||||
)
|
||||
elif numbering_type == NUMBERING_TYPES[2]:
|
||||
return ChineseNumberUnit(
|
||||
power=pow(2, index + 3),
|
||||
simplified=value[0],
|
||||
traditional=value[1],
|
||||
big_s=value[0],
|
||||
big_t=value[1],
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Counting type should be in {0} ({1} provided).".format(
|
||||
NUMBERING_TYPES, numbering_type
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class ChineseNumberDigit(ChineseChar):
|
||||
"""
|
||||
中文数字字符
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, value, simplified, traditional, big_s, big_t, alt_s=None, alt_t=None
|
||||
):
|
||||
super(ChineseNumberDigit, self).__init__(simplified, traditional)
|
||||
self.value = value
|
||||
self.big_s = big_s
|
||||
self.big_t = big_t
|
||||
self.alt_s = alt_s
|
||||
self.alt_t = alt_t
|
||||
|
||||
def __str__(self):
|
||||
return str(self.value)
|
||||
|
||||
@classmethod
|
||||
def create(cls, i, v):
|
||||
return ChineseNumberDigit(i, v[0], v[1], v[2], v[3])
|
||||
|
||||
|
||||
class ChineseMath(ChineseChar):
|
||||
"""
|
||||
中文数位字符
|
||||
"""
|
||||
|
||||
def __init__(self, simplified, traditional, symbol, expression=None):
|
||||
super(ChineseMath, self).__init__(simplified, traditional)
|
||||
self.symbol = symbol
|
||||
self.expression = expression
|
||||
self.big_s = simplified
|
||||
self.big_t = traditional
|
||||
|
||||
|
||||
CC, CNU, CND, CM = ChineseChar, ChineseNumberUnit, ChineseNumberDigit, ChineseMath
|
||||
|
||||
|
||||
class NumberSystem(object):
|
||||
"""
|
||||
中文数字系统
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class MathSymbol(object):
|
||||
"""
|
||||
用于中文数字系统的数学符号 (繁/简体), e.g.
|
||||
positive = ['正', '正']
|
||||
negative = ['负', '負']
|
||||
point = ['点', '點']
|
||||
"""
|
||||
|
||||
def __init__(self, positive, negative, point):
|
||||
self.positive = positive
|
||||
self.negative = negative
|
||||
self.point = point
|
||||
|
||||
def __iter__(self):
|
||||
for v in self.__dict__.values():
|
||||
yield v
|
||||
|
||||
|
||||
# class OtherSymbol(object):
|
||||
# """
|
||||
# 其他符号
|
||||
# """
|
||||
#
|
||||
# def __init__(self, sil):
|
||||
# self.sil = sil
|
||||
#
|
||||
# def __iter__(self):
|
||||
# for v in self.__dict__.values():
|
||||
# yield v
|
||||
@@ -0,0 +1,30 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""基本常量
|
||||
中文数字/数位/符号字符常量
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-02"
|
||||
|
||||
CHINESE_DIGIS = "零一二三四五六七八九"
|
||||
BIG_CHINESE_DIGIS_SIMPLIFIED = "零壹贰叁肆伍陆柒捌玖"
|
||||
BIG_CHINESE_DIGIS_TRADITIONAL = "零壹貳參肆伍陸柒捌玖"
|
||||
SMALLER_BIG_CHINESE_UNITS_SIMPLIFIED = "十百千万"
|
||||
SMALLER_BIG_CHINESE_UNITS_TRADITIONAL = "拾佰仟萬"
|
||||
LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED = "亿兆京垓秭穰沟涧正载"
|
||||
LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL = "億兆京垓秭穰溝澗正載"
|
||||
SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED = "十百千万"
|
||||
SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL = "拾佰仟萬"
|
||||
|
||||
ZERO_ALT = "〇"
|
||||
ONE_ALT = "幺"
|
||||
TWO_ALTS = ["两", "兩"]
|
||||
|
||||
POSITIVE = ["正", "正"]
|
||||
NEGATIVE = ["负", "負"]
|
||||
POINT = ["点", "點"]
|
||||
# PLUS = [u'加', u'加']
|
||||
# SIL = [u'杠', u'槓']
|
||||
|
||||
# 中文数字系统类型
|
||||
NUMBERING_TYPES = ["low", "mid", "high"]
|
||||
@@ -0,0 +1,342 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""基本方法
|
||||
创建中文数字系统 方法
|
||||
中文字符串 <=> 数字串 方法
|
||||
数字串 <=> 中文字符串 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-02"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_class import *
|
||||
from fish_speech.text.chn_text_norm.basic_constant import *
|
||||
|
||||
|
||||
def create_system(numbering_type=NUMBERING_TYPES[1]):
|
||||
"""
|
||||
根据数字系统类型返回创建相应的数字系统,默认为 mid
|
||||
NUMBERING_TYPES = ['low', 'mid', 'high']: 中文数字系统类型
|
||||
low: '兆' = '亿' * '十' = $10^{9}$, '京' = '兆' * '十', etc.
|
||||
mid: '兆' = '亿' * '万' = $10^{12}$, '京' = '兆' * '万', etc.
|
||||
high: '兆' = '亿' * '亿' = $10^{16}$, '京' = '兆' * '兆', etc.
|
||||
返回对应的数字系统
|
||||
"""
|
||||
|
||||
# chinese number units of '亿' and larger
|
||||
all_larger_units = zip(
|
||||
LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED,
|
||||
LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL,
|
||||
)
|
||||
larger_units = [
|
||||
CNU.create(i, v, numbering_type, False) for i, v in enumerate(all_larger_units)
|
||||
]
|
||||
# chinese number units of '十, 百, 千, 万'
|
||||
all_smaller_units = zip(
|
||||
SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED,
|
||||
SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL,
|
||||
)
|
||||
smaller_units = [
|
||||
CNU.create(i, v, small_unit=True) for i, v in enumerate(all_smaller_units)
|
||||
]
|
||||
# digis
|
||||
chinese_digis = zip(
|
||||
CHINESE_DIGIS,
|
||||
CHINESE_DIGIS,
|
||||
BIG_CHINESE_DIGIS_SIMPLIFIED,
|
||||
BIG_CHINESE_DIGIS_TRADITIONAL,
|
||||
)
|
||||
digits = [CND.create(i, v) for i, v in enumerate(chinese_digis)]
|
||||
digits[0].alt_s, digits[0].alt_t = ZERO_ALT, ZERO_ALT
|
||||
digits[1].alt_s, digits[1].alt_t = ONE_ALT, ONE_ALT
|
||||
digits[2].alt_s, digits[2].alt_t = TWO_ALTS[0], TWO_ALTS[1]
|
||||
|
||||
# symbols
|
||||
positive_cn = CM(POSITIVE[0], POSITIVE[1], "+", lambda x: x)
|
||||
negative_cn = CM(NEGATIVE[0], NEGATIVE[1], "-", lambda x: -x)
|
||||
point_cn = CM(POINT[0], POINT[1], ".", lambda x, y: float(str(x) + "." + str(y)))
|
||||
# sil_cn = CM(SIL[0], SIL[1], '-', lambda x, y: float(str(x) + '-' + str(y)))
|
||||
system = NumberSystem()
|
||||
system.units = smaller_units + larger_units
|
||||
system.digits = digits
|
||||
system.math = MathSymbol(positive_cn, negative_cn, point_cn)
|
||||
# system.symbols = OtherSymbol(sil_cn)
|
||||
return system
|
||||
|
||||
|
||||
def chn2num(chinese_string, numbering_type=NUMBERING_TYPES[1]):
|
||||
|
||||
def get_symbol(char, system):
|
||||
for u in system.units:
|
||||
if char in [u.traditional, u.simplified, u.big_s, u.big_t]:
|
||||
return u
|
||||
for d in system.digits:
|
||||
if char in [
|
||||
d.traditional,
|
||||
d.simplified,
|
||||
d.big_s,
|
||||
d.big_t,
|
||||
d.alt_s,
|
||||
d.alt_t,
|
||||
]:
|
||||
return d
|
||||
for m in system.math:
|
||||
if char in [m.traditional, m.simplified]:
|
||||
return m
|
||||
|
||||
def string2symbols(chinese_string, system):
|
||||
int_string, dec_string = chinese_string, ""
|
||||
for p in [system.math.point.simplified, system.math.point.traditional]:
|
||||
if p in chinese_string:
|
||||
int_string, dec_string = chinese_string.split(p)
|
||||
break
|
||||
return [get_symbol(c, system) for c in int_string], [
|
||||
get_symbol(c, system) for c in dec_string
|
||||
]
|
||||
|
||||
def correct_symbols(integer_symbols, system):
|
||||
"""
|
||||
一百八 to 一百八十
|
||||
一亿一千三百万 to 一亿 一千万 三百万
|
||||
"""
|
||||
|
||||
if integer_symbols and isinstance(integer_symbols[0], CNU):
|
||||
if integer_symbols[0].power == 1:
|
||||
integer_symbols = [system.digits[1]] + integer_symbols
|
||||
|
||||
if len(integer_symbols) > 1:
|
||||
if isinstance(integer_symbols[-1], CND) and isinstance(
|
||||
integer_symbols[-2], CNU
|
||||
):
|
||||
integer_symbols.append(
|
||||
CNU(integer_symbols[-2].power - 1, None, None, None, None)
|
||||
)
|
||||
|
||||
result = []
|
||||
unit_count = 0
|
||||
for s in integer_symbols:
|
||||
if isinstance(s, CND):
|
||||
result.append(s)
|
||||
unit_count = 0
|
||||
elif isinstance(s, CNU):
|
||||
current_unit = CNU(s.power, None, None, None, None)
|
||||
unit_count += 1
|
||||
|
||||
if unit_count == 1:
|
||||
result.append(current_unit)
|
||||
elif unit_count > 1:
|
||||
for i in range(len(result)):
|
||||
if (
|
||||
isinstance(result[-i - 1], CNU)
|
||||
and result[-i - 1].power < current_unit.power
|
||||
):
|
||||
result[-i - 1] = CNU(
|
||||
result[-i - 1].power + current_unit.power,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
return result
|
||||
|
||||
def compute_value(integer_symbols):
|
||||
"""
|
||||
Compute the value.
|
||||
When current unit is larger than previous unit, current unit * all previous units will be used as all previous units.
|
||||
e.g. '两千万' = 2000 * 10000 not 2000 + 10000
|
||||
"""
|
||||
value = [0]
|
||||
last_power = 0
|
||||
for s in integer_symbols:
|
||||
if isinstance(s, CND):
|
||||
value[-1] = s.value
|
||||
elif isinstance(s, CNU):
|
||||
value[-1] *= pow(10, s.power)
|
||||
if s.power > last_power:
|
||||
value[:-1] = list(map(lambda v: v * pow(10, s.power), value[:-1]))
|
||||
last_power = s.power
|
||||
value.append(0)
|
||||
return sum(value)
|
||||
|
||||
system = create_system(numbering_type)
|
||||
int_part, dec_part = string2symbols(chinese_string, system)
|
||||
int_part = correct_symbols(int_part, system)
|
||||
int_str = str(compute_value(int_part))
|
||||
dec_str = "".join([str(d.value) for d in dec_part])
|
||||
if dec_part:
|
||||
return "{0}.{1}".format(int_str, dec_str)
|
||||
else:
|
||||
return int_str
|
||||
|
||||
|
||||
def num2chn(
|
||||
number_string,
|
||||
numbering_type=NUMBERING_TYPES[1],
|
||||
big=False,
|
||||
traditional=False,
|
||||
alt_zero=False,
|
||||
alt_one=False,
|
||||
alt_two=True,
|
||||
use_zeros=True,
|
||||
use_units=True,
|
||||
):
|
||||
|
||||
def get_value(value_string, use_zeros=True):
|
||||
|
||||
striped_string = value_string.lstrip("0")
|
||||
|
||||
# record nothing if all zeros
|
||||
if not striped_string:
|
||||
return []
|
||||
|
||||
# record one digits
|
||||
elif len(striped_string) == 1:
|
||||
if use_zeros and len(value_string) != len(striped_string):
|
||||
return [system.digits[0], system.digits[int(striped_string)]]
|
||||
else:
|
||||
return [system.digits[int(striped_string)]]
|
||||
|
||||
# recursively record multiple digits
|
||||
else:
|
||||
result_unit = next(
|
||||
u for u in reversed(system.units) if u.power < len(striped_string)
|
||||
)
|
||||
result_string = value_string[: -result_unit.power]
|
||||
return (
|
||||
get_value(result_string)
|
||||
+ [result_unit]
|
||||
+ get_value(striped_string[-result_unit.power :])
|
||||
)
|
||||
|
||||
system = create_system(numbering_type)
|
||||
|
||||
int_dec = number_string.split(".")
|
||||
if len(int_dec) == 1:
|
||||
int_string = int_dec[0]
|
||||
dec_string = ""
|
||||
elif len(int_dec) == 2:
|
||||
int_string = int_dec[0]
|
||||
dec_string = int_dec[1]
|
||||
else:
|
||||
raise ValueError(
|
||||
"invalid input num string with more than one dot: {}".format(number_string)
|
||||
)
|
||||
|
||||
if use_units and len(int_string) > 1:
|
||||
result_symbols = get_value(int_string)
|
||||
else:
|
||||
result_symbols = [system.digits[int(c)] for c in int_string]
|
||||
dec_symbols = [system.digits[int(c)] for c in dec_string]
|
||||
if dec_string:
|
||||
result_symbols += [system.math.point] + dec_symbols
|
||||
|
||||
if alt_two:
|
||||
liang = CND(
|
||||
2,
|
||||
system.digits[2].alt_s,
|
||||
system.digits[2].alt_t,
|
||||
system.digits[2].big_s,
|
||||
system.digits[2].big_t,
|
||||
)
|
||||
for i, v in enumerate(result_symbols):
|
||||
if isinstance(v, CND) and v.value == 2:
|
||||
next_symbol = (
|
||||
result_symbols[i + 1] if i < len(result_symbols) - 1 else None
|
||||
)
|
||||
previous_symbol = result_symbols[i - 1] if i > 0 else None
|
||||
if isinstance(next_symbol, CNU) and isinstance(
|
||||
previous_symbol, (CNU, type(None))
|
||||
):
|
||||
if next_symbol.power != 1 and (
|
||||
(previous_symbol is None) or (previous_symbol.power != 1)
|
||||
):
|
||||
result_symbols[i] = liang
|
||||
|
||||
# if big is True, '两' will not be used and `alt_two` has no impact on output
|
||||
if big:
|
||||
attr_name = "big_"
|
||||
if traditional:
|
||||
attr_name += "t"
|
||||
else:
|
||||
attr_name += "s"
|
||||
else:
|
||||
if traditional:
|
||||
attr_name = "traditional"
|
||||
else:
|
||||
attr_name = "simplified"
|
||||
|
||||
result = "".join([getattr(s, attr_name) for s in result_symbols])
|
||||
|
||||
# if not use_zeros:
|
||||
# result = result.strip(getattr(system.digits[0], attr_name))
|
||||
|
||||
if alt_zero:
|
||||
result = result.replace(
|
||||
getattr(system.digits[0], attr_name), system.digits[0].alt_s
|
||||
)
|
||||
|
||||
if alt_one:
|
||||
result = result.replace(
|
||||
getattr(system.digits[1], attr_name), system.digits[1].alt_s
|
||||
)
|
||||
|
||||
for i, p in enumerate(POINT):
|
||||
if result.startswith(p):
|
||||
return CHINESE_DIGIS[0] + result
|
||||
|
||||
# ^10, 11, .., 19
|
||||
if (
|
||||
len(result) >= 2
|
||||
and result[1]
|
||||
in [
|
||||
SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED[0],
|
||||
SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL[0],
|
||||
]
|
||||
and result[0]
|
||||
in [
|
||||
CHINESE_DIGIS[1],
|
||||
BIG_CHINESE_DIGIS_SIMPLIFIED[1],
|
||||
BIG_CHINESE_DIGIS_TRADITIONAL[1],
|
||||
]
|
||||
):
|
||||
result = result[1:]
|
||||
|
||||
return result
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
all_chinese_number_string = (
|
||||
CHINESE_DIGIS
|
||||
+ BIG_CHINESE_DIGIS_SIMPLIFIED
|
||||
+ BIG_CHINESE_DIGIS_TRADITIONAL
|
||||
+ LARGER_CHINESE_NUMERING_UNITS_SIMPLIFIED
|
||||
+ LARGER_CHINESE_NUMERING_UNITS_TRADITIONAL
|
||||
+ SMALLER_CHINESE_NUMERING_UNITS_SIMPLIFIED
|
||||
+ SMALLER_CHINESE_NUMERING_UNITS_TRADITIONAL
|
||||
+ ZERO_ALT
|
||||
+ ONE_ALT
|
||||
+ "".join(TWO_ALTS + POSITIVE + NEGATIVE + POINT)
|
||||
)
|
||||
|
||||
print("num:", chn2num("一万零四百零三点八零五"))
|
||||
print("num:", chn2num("一亿六点三"))
|
||||
print("num:", chn2num("一亿零六点三"))
|
||||
print("num:", chn2num("两千零一亿六点三"))
|
||||
# print('num:', chn2num('一零零八六'))
|
||||
print("txt:", num2chn("10260.03", alt_zero=True))
|
||||
print("txt:", num2chn("20037.090", numbering_type="low", traditional=True))
|
||||
print("txt:", num2chn("100860001.77", numbering_type="high", big=True))
|
||||
print(
|
||||
"txt:",
|
||||
num2chn(
|
||||
"059523810880",
|
||||
alt_one=True,
|
||||
alt_two=False,
|
||||
use_lzeros=True,
|
||||
use_rzeros=True,
|
||||
use_units=False,
|
||||
),
|
||||
)
|
||||
|
||||
print(all_chinese_number_string)
|
||||
@@ -0,0 +1,32 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""CARDINAL类 (包含小数DECIMAL类)
|
||||
纯数 <=> 中文字符串 方法
|
||||
中文字符串 <=> 纯数 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class Cardinal:
|
||||
"""
|
||||
CARDINAL类
|
||||
"""
|
||||
|
||||
def __init__(self, cardinal=None, chntext=None):
|
||||
self.cardinal = cardinal
|
||||
self.chntext = chntext
|
||||
|
||||
def chntext2cardinal(self):
|
||||
return chn2num(self.chntext)
|
||||
|
||||
def cardinal2chntext(self):
|
||||
return num2chn(self.cardinal)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Cardinal(cardinal="21357.230").cardinal2chntext())
|
||||
@@ -0,0 +1,75 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""DATE类
|
||||
日期 <=> 中文字符串 方法
|
||||
中文字符串 <=> 日期 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-07"
|
||||
|
||||
from fish_speech.text.chn_text_norm.cardinal import Cardinal
|
||||
from fish_speech.text.chn_text_norm.digit import Digit
|
||||
|
||||
|
||||
class Date:
|
||||
"""
|
||||
DATE类
|
||||
"""
|
||||
|
||||
def __init__(self, date=None, chntext=None):
|
||||
self.date = date
|
||||
self.chntext = chntext
|
||||
|
||||
# def chntext2date(self):
|
||||
# chntext = self.chntext
|
||||
# try:
|
||||
# year, other = chntext.strip().split('年', maxsplit=1)
|
||||
# year = Digit(chntext=year).digit2chntext() + '年'
|
||||
# except ValueError:
|
||||
# other = chntext
|
||||
# year = ''
|
||||
# if other:
|
||||
# try:
|
||||
# month, day = other.strip().split('月', maxsplit=1)
|
||||
# month = Cardinal(chntext=month).chntext2cardinal() + '月'
|
||||
# except ValueError:
|
||||
# day = chntext
|
||||
# month = ''
|
||||
# if day:
|
||||
# day = Cardinal(chntext=day[:-1]).chntext2cardinal() + day[-1]
|
||||
# else:
|
||||
# month = ''
|
||||
# day = ''
|
||||
# date = year + month + day
|
||||
# self.date = date
|
||||
# return self.date
|
||||
|
||||
def date2chntext(self):
|
||||
date = self.date
|
||||
try:
|
||||
year, other = date.strip().split("年", maxsplit=1)
|
||||
year = Digit(digit=year).digit2chntext() + "年"
|
||||
except ValueError:
|
||||
other = date
|
||||
year = ""
|
||||
if other:
|
||||
try:
|
||||
month, day = other.strip().split("月", maxsplit=1)
|
||||
month = Cardinal(cardinal=month).cardinal2chntext() + "月"
|
||||
except ValueError:
|
||||
day = date
|
||||
month = ""
|
||||
if day:
|
||||
day = Cardinal(cardinal=day[:-1]).cardinal2chntext() + day[-1]
|
||||
else:
|
||||
month = ""
|
||||
day = ""
|
||||
chntext = year + month + day
|
||||
self.chntext = chntext
|
||||
return self.chntext
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试
|
||||
print(Date(date="09年3月16日").date2chntext())
|
||||
@@ -0,0 +1,32 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""DIGIT类
|
||||
数字串 <=> 中文字符串 方法
|
||||
中文字符串 <=> 数字串 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class Digit:
|
||||
"""
|
||||
DIGIT类
|
||||
"""
|
||||
|
||||
def __init__(self, digit=None, chntext=None):
|
||||
self.digit = digit
|
||||
self.chntext = chntext
|
||||
|
||||
# def chntext2digit(self):
|
||||
# return chn2num(self.chntext)
|
||||
|
||||
def digit2chntext(self):
|
||||
return num2chn(self.digit, alt_two=False, use_units=False)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Digit(digit="2016").digit2chntext())
|
||||
@@ -0,0 +1,35 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""FRACTION类
|
||||
分数 <=> 中文字符串 方法
|
||||
中文字符串 <=> 分数 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class Fraction:
|
||||
"""
|
||||
FRACTION类
|
||||
"""
|
||||
|
||||
def __init__(self, fraction=None, chntext=None):
|
||||
self.fraction = fraction
|
||||
self.chntext = chntext
|
||||
|
||||
def chntext2fraction(self):
|
||||
denominator, numerator = self.chntext.split("分之")
|
||||
return chn2num(numerator) + "/" + chn2num(denominator)
|
||||
|
||||
def fraction2chntext(self):
|
||||
numerator, denominator = self.fraction.split("/")
|
||||
return num2chn(denominator) + "分之" + num2chn(numerator)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Fraction(fraction="2135/7230").fraction2chntext())
|
||||
print(Fraction(chntext="五百八十一分之三百六十九").chntext2fraction())
|
||||
@@ -0,0 +1,43 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""MONEY类
|
||||
金钱 <=> 中文字符串 方法
|
||||
中文字符串 <=> 金钱 方法
|
||||
"""
|
||||
import re
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-08"
|
||||
|
||||
from fish_speech.text.chn_text_norm.cardinal import Cardinal
|
||||
|
||||
|
||||
class Money:
|
||||
"""
|
||||
MONEY类
|
||||
"""
|
||||
|
||||
def __init__(self, money=None, chntext=None):
|
||||
self.money = money
|
||||
self.chntext = chntext
|
||||
|
||||
# def chntext2money(self):
|
||||
# return self.money
|
||||
|
||||
def money2chntext(self):
|
||||
money = self.money
|
||||
pattern = re.compile(r"(\d+(\.\d+)?)")
|
||||
matchers = pattern.findall(money)
|
||||
if matchers:
|
||||
for matcher in matchers:
|
||||
money = money.replace(
|
||||
matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext()
|
||||
)
|
||||
self.chntext = money
|
||||
return self.chntext
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试
|
||||
print(Money(money="21.5万元").money2chntext())
|
||||
print(Money(money="230块5毛").money2chntext())
|
||||
@@ -0,0 +1,33 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""PERCENTAGE类
|
||||
百分数 <=> 中文字符串 方法
|
||||
中文字符串 <=> 百分数 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-06"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class Percentage:
|
||||
"""
|
||||
PERCENTAGE类
|
||||
"""
|
||||
|
||||
def __init__(self, percentage=None, chntext=None):
|
||||
self.percentage = percentage
|
||||
self.chntext = chntext
|
||||
|
||||
def chntext2percentage(self):
|
||||
return chn2num(self.chntext.strip().strip("百分之")) + "%"
|
||||
|
||||
def percentage2chntext(self):
|
||||
return "百分之" + num2chn(self.percentage.strip().strip("%"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Percentage(chntext="百分之五十六点零三").chntext2percentage())
|
||||
print(Percentage(percentage="65.3%").percentage2chntext())
|
||||
@@ -0,0 +1,51 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""TELEPHONE类
|
||||
电话号码 <=> 中文字符串 方法
|
||||
中文字符串 <=> 电话号码 方法
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
from fish_speech.text.chn_text_norm.basic_util import *
|
||||
|
||||
|
||||
class TelePhone:
|
||||
"""
|
||||
TELEPHONE类
|
||||
"""
|
||||
|
||||
def __init__(self, telephone=None, raw_chntext=None, chntext=None):
|
||||
self.telephone = telephone
|
||||
self.raw_chntext = raw_chntext
|
||||
self.chntext = chntext
|
||||
|
||||
# def chntext2telephone(self):
|
||||
# sil_parts = self.raw_chntext.split('<SIL>')
|
||||
# self.telephone = '-'.join([
|
||||
# str(chn2num(p)) for p in sil_parts
|
||||
# ])
|
||||
# return self.telephone
|
||||
|
||||
def telephone2chntext(self, fixed=False):
|
||||
|
||||
if fixed:
|
||||
sil_parts = self.telephone.split("-")
|
||||
self.raw_chntext = "<SIL>".join(
|
||||
[num2chn(part, alt_two=False, use_units=False) for part in sil_parts]
|
||||
)
|
||||
self.chntext = self.raw_chntext.replace("<SIL>", "")
|
||||
else:
|
||||
sp_parts = self.telephone.strip("+").split()
|
||||
self.raw_chntext = "<SP>".join(
|
||||
[num2chn(part, alt_two=False, use_units=False) for part in sp_parts]
|
||||
)
|
||||
self.chntext = self.raw_chntext.replace("<SP>", "")
|
||||
return self.chntext
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(TelePhone(telephone="0595-23980880").telephone2chntext())
|
||||
# print(TelePhone(raw_chntext='零五九五杠二三八六五零九八').chntext2telephone())
|
||||
@@ -0,0 +1,177 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
TEXT类
|
||||
"""
|
||||
|
||||
__author__ = "Zhiyang Zhou <zyzhou@stu.xmu.edu.cn>"
|
||||
__data__ = "2019-05-03"
|
||||
|
||||
import re
|
||||
|
||||
from fish_speech.text.chn_text_norm.cardinal import Cardinal
|
||||
from fish_speech.text.chn_text_norm.date import Date
|
||||
from fish_speech.text.chn_text_norm.digit import Digit
|
||||
from fish_speech.text.chn_text_norm.fraction import Fraction
|
||||
from fish_speech.text.chn_text_norm.money import Money
|
||||
from fish_speech.text.chn_text_norm.percentage import Percentage
|
||||
from fish_speech.text.chn_text_norm.telephone import TelePhone
|
||||
|
||||
CURRENCY_NAMES = (
|
||||
"(人民币|美元|日元|英镑|欧元|马克|法郎|加拿大元|澳元|港币|先令|芬兰马克|爱尔兰镑|"
|
||||
"里拉|荷兰盾|埃斯库多|比塞塔|印尼盾|林吉特|新西兰元|比索|卢布|新加坡元|韩元|泰铢)"
|
||||
)
|
||||
CURRENCY_UNITS = "((亿|千万|百万|万|千|百)|(亿|千万|百万|万|千|百|)元|(亿|千万|百万|万|千|百|)块|角|毛|分)"
|
||||
COM_QUANTIFIERS = (
|
||||
"(匹|张|座|回|场|尾|条|个|首|阙|阵|网|炮|顶|丘|棵|只|支|袭|辆|挑|担|颗|壳|窠|曲|墙|群|腔|"
|
||||
"砣|座|客|贯|扎|捆|刀|令|打|手|罗|坡|山|岭|江|溪|钟|队|单|双|对|出|口|头|脚|板|跳|枝|件|贴|"
|
||||
"针|线|管|名|位|身|堂|课|本|页|家|户|层|丝|毫|厘|分|钱|两|斤|担|铢|石|钧|锱|忽|(千|毫|微)克|"
|
||||
"毫|厘|分|寸|尺|丈|里|寻|常|铺|程|(千|分|厘|毫|微)米|撮|勺|合|升|斗|石|盘|碗|碟|叠|桶|笼|盆|"
|
||||
"盒|杯|钟|斛|锅|簋|篮|盘|桶|罐|瓶|壶|卮|盏|箩|箱|煲|啖|袋|钵|年|月|日|季|刻|时|周|天|秒|分|旬|"
|
||||
"纪|岁|世|更|夜|春|夏|秋|冬|代|伏|辈|丸|泡|粒|颗|幢|堆|条|根|支|道|面|片|张|颗|块|人|抽)"
|
||||
)
|
||||
|
||||
|
||||
class Text:
|
||||
"""
|
||||
Text类
|
||||
"""
|
||||
|
||||
def __init__(self, raw_text, norm_text=None):
|
||||
self.raw_text = "^" + raw_text + "$"
|
||||
self.norm_text = norm_text
|
||||
|
||||
def _particular(self):
|
||||
text = self.norm_text
|
||||
pattern = re.compile(r"(([a-zA-Z]+)二([a-zA-Z]+))")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('particular')
|
||||
for matcher in matchers:
|
||||
text = text.replace(matcher[0], matcher[1] + "2" + matcher[2], 1)
|
||||
self.norm_text = text
|
||||
return self.norm_text
|
||||
|
||||
def normalize(self):
|
||||
text = self.raw_text
|
||||
|
||||
# 规范化日期
|
||||
pattern = re.compile(
|
||||
r"\D+((([089]\d|(19|20)\d{2})年)?(\d{1,2}月(\d{1,2}[日号])?)?)"
|
||||
)
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('date')
|
||||
for matcher in matchers:
|
||||
text = text.replace(matcher[0], Date(date=matcher[0]).date2chntext(), 1)
|
||||
|
||||
# 规范化金钱
|
||||
pattern = re.compile(
|
||||
r"\D+((\d+(\.\d+)?)[多余几]?"
|
||||
+ CURRENCY_UNITS
|
||||
+ "(\d"
|
||||
+ CURRENCY_UNITS
|
||||
+ "?)?)"
|
||||
)
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('money')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0], Money(money=matcher[0]).money2chntext(), 1
|
||||
)
|
||||
|
||||
# 规范化固话/手机号码
|
||||
# 手机
|
||||
# http://www.jihaoba.com/news/show/13680
|
||||
# 移动:139、138、137、136、135、134、159、158、157、150、151、152、188、187、182、183、184、178、198
|
||||
# 联通:130、131、132、156、155、186、185、176
|
||||
# 电信:133、153、189、180、181、177
|
||||
pattern = re.compile(r"\D((\+?86 ?)?1([38]\d|5[0-35-9]|7[678]|9[89])\d{8})\D")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('telephone')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0], TelePhone(telephone=matcher[0]).telephone2chntext(), 1
|
||||
)
|
||||
# 固话
|
||||
pattern = re.compile(r"\D((0(10|2[1-3]|[3-9]\d{2})-?)?[1-9]\d{6,7})\D")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('fixed telephone')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0],
|
||||
TelePhone(telephone=matcher[0]).telephone2chntext(fixed=True),
|
||||
1,
|
||||
)
|
||||
|
||||
# 规范化分数
|
||||
pattern = re.compile(r"(\d+/\d+)")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('fraction')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher, Fraction(fraction=matcher).fraction2chntext(), 1
|
||||
)
|
||||
|
||||
# 规范化百分数
|
||||
text = text.replace("%", "%")
|
||||
pattern = re.compile(r"(\d+(\.\d+)?%)")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('percentage')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0],
|
||||
Percentage(percentage=matcher[0]).percentage2chntext(),
|
||||
1,
|
||||
)
|
||||
|
||||
# 规范化纯数+量词
|
||||
pattern = re.compile(r"(\d+(\.\d+)?)[多余几]?" + COM_QUANTIFIERS)
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('cardinal+quantifier')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1
|
||||
)
|
||||
|
||||
# 规范化数字编号
|
||||
pattern = re.compile(r"(\d{4,32})")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('digit')
|
||||
for matcher in matchers:
|
||||
text = text.replace(matcher, Digit(digit=matcher).digit2chntext(), 1)
|
||||
|
||||
# 规范化纯数
|
||||
pattern = re.compile(r"(\d+(\.\d+)?)")
|
||||
matchers = pattern.findall(text)
|
||||
if matchers:
|
||||
# print('cardinal')
|
||||
for matcher in matchers:
|
||||
text = text.replace(
|
||||
matcher[0], Cardinal(cardinal=matcher[0]).cardinal2chntext(), 1
|
||||
)
|
||||
|
||||
self.norm_text = text
|
||||
self._particular()
|
||||
|
||||
return self.norm_text.lstrip("^").rstrip("$")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
# 测试程序
|
||||
print(Text(raw_text="固话:0595-23865596或23880880。").normalize())
|
||||
print(Text(raw_text="手机:+86 19859213959或15659451527。").normalize())
|
||||
print(Text(raw_text="分数:32477/76391。").normalize())
|
||||
print(Text(raw_text="百分数:80.03%。").normalize())
|
||||
print(Text(raw_text="编号:31520181154418。").normalize())
|
||||
print(Text(raw_text="纯数:2983.07克或12345.60米。").normalize())
|
||||
print(Text(raw_text="日期:1999年2月20日或09年3月15号。").normalize())
|
||||
print(Text(raw_text="金钱:12块5,34.5元,20.1万").normalize())
|
||||
print(Text(raw_text="特殊:O2O或B2C。").normalize())
|
||||
@@ -0,0 +1,31 @@
|
||||
import re
|
||||
|
||||
SYMBOLS_MAPPING = {
|
||||
"“": "'",
|
||||
"”": "'",
|
||||
"‘": "'",
|
||||
"’": "'",
|
||||
"【": "",
|
||||
"】": "",
|
||||
"[": "",
|
||||
"]": "",
|
||||
"(": "",
|
||||
")": "",
|
||||
"(": "",
|
||||
")": "",
|
||||
"・": "·",
|
||||
}
|
||||
|
||||
REPLACE_SYMBOL_REGEX = re.compile(
|
||||
"|".join(re.escape(p) for p in SYMBOLS_MAPPING.keys())
|
||||
)
|
||||
|
||||
|
||||
def clean_text(text):
|
||||
# Clean the text
|
||||
text = text.strip()
|
||||
|
||||
# Replace all chinese symbols with their english counterparts
|
||||
text = REPLACE_SYMBOL_REGEX.sub(lambda x: SYMBOLS_MAPPING[x.group()], text)
|
||||
|
||||
return text
|
||||
@@ -0,0 +1,130 @@
|
||||
import re
|
||||
import string
|
||||
|
||||
from fish_speech.text.clean import clean_text
|
||||
|
||||
|
||||
def utf_8_len(text):
|
||||
return len(text.encode("utf-8"))
|
||||
|
||||
|
||||
def break_text(texts, length, splits: set):
|
||||
for text in texts:
|
||||
if utf_8_len(text) <= length:
|
||||
yield text
|
||||
continue
|
||||
|
||||
curr = ""
|
||||
for char in text:
|
||||
curr += char
|
||||
|
||||
if char in splits:
|
||||
yield curr
|
||||
curr = ""
|
||||
|
||||
if curr:
|
||||
yield curr
|
||||
|
||||
|
||||
def break_text_by_length(texts, length):
|
||||
for text in texts:
|
||||
if utf_8_len(text) <= length:
|
||||
yield text
|
||||
continue
|
||||
|
||||
curr = ""
|
||||
for char in text:
|
||||
curr += char
|
||||
|
||||
if utf_8_len(curr) >= length:
|
||||
yield curr
|
||||
curr = ""
|
||||
|
||||
if curr:
|
||||
yield curr
|
||||
|
||||
|
||||
def add_cleaned(curr, segments):
|
||||
curr = curr.strip()
|
||||
if curr and not all(c.isspace() or c in string.punctuation for c in curr):
|
||||
segments.append(curr)
|
||||
|
||||
|
||||
def protect_float(text):
|
||||
# Turns 3.14 into <3_f_14> to prevent splitting
|
||||
return re.sub(r"(\d+)\.(\d+)", r"<\1_f_\2>", text)
|
||||
|
||||
|
||||
def unprotect_float(text):
|
||||
# Turns <3_f_14> into 3.14
|
||||
return re.sub(r"<(\d+)_f_(\d+)>", r"\1.\2", text)
|
||||
|
||||
|
||||
def split_text(text, length):
|
||||
text = clean_text(text)
|
||||
|
||||
# Break the text into pieces with following rules:
|
||||
# 1. Split the text at ".", "!", "?" if text is NOT a float
|
||||
# 2. If the text is longer than length, split at ","
|
||||
# 3. If the text is still longer than length, split at " "
|
||||
# 4. If the text is still longer than length, split at any character to length
|
||||
|
||||
texts = [text]
|
||||
texts = map(protect_float, texts)
|
||||
texts = break_text(texts, length, {".", "!", "?", "。", "!", "?"})
|
||||
texts = map(unprotect_float, texts)
|
||||
texts = break_text(texts, length, {",", ","})
|
||||
texts = break_text(texts, length, {" "})
|
||||
texts = list(break_text_by_length(texts, length))
|
||||
|
||||
# Then, merge the texts into segments with length <= length
|
||||
segments = []
|
||||
curr = ""
|
||||
|
||||
for text in texts:
|
||||
if utf_8_len(curr) + utf_8_len(text) <= length:
|
||||
curr += text
|
||||
else:
|
||||
add_cleaned(curr, segments)
|
||||
curr = text
|
||||
|
||||
if curr:
|
||||
add_cleaned(curr, segments)
|
||||
|
||||
return segments
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Test the split_text function
|
||||
|
||||
text = "This is a test sentence. This is another test sentence. And a third one."
|
||||
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence.",
|
||||
"This is another test sentence. And a third one.",
|
||||
]
|
||||
assert split_text("a,aaaaaa3.14", 10) == ["a,", "aaaaaa3.14"]
|
||||
assert split_text(" ", 10) == []
|
||||
assert split_text("a", 10) == ["a"]
|
||||
|
||||
text = "This is a test sentence with only commas, and no dots, and no exclamation marks, and no question marks, and no newlines."
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence with only commas,",
|
||||
"and no dots, and no exclamation marks,",
|
||||
"and no question marks, and no newlines.",
|
||||
]
|
||||
|
||||
text = "This is a test sentence This is a test sentence This is a test sentence. This is a test sentence, This is a test sentence, This is a test sentence."
|
||||
# First half split at " ", second half split at ","
|
||||
assert split_text(text, 50) == [
|
||||
"This is a test sentence This is a test sentence",
|
||||
"This is a test sentence. This is a test sentence,",
|
||||
"This is a test sentence, This is a test sentence.",
|
||||
]
|
||||
|
||||
text = "这是一段很长的中文文本,而且没有句号,也没有感叹号,也没有问号,也没有换行符。"
|
||||
assert split_text(text, 50) == [
|
||||
"这是一段很长的中文文本,",
|
||||
"而且没有句号,也没有感叹号,",
|
||||
"也没有问号,也没有换行符.",
|
||||
]
|
||||
@@ -0,0 +1,169 @@
|
||||
import itertools
|
||||
import os
|
||||
import re
|
||||
from collections import defaultdict
|
||||
from functools import partial
|
||||
from multiprocessing import Pool
|
||||
from pathlib import Path
|
||||
|
||||
import click
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
|
||||
from fish_speech.datasets.protos.text_data_pb2 import Semantics, Sentence, TextData
|
||||
from fish_speech.datasets.protos.text_data_stream import pack_pb_stream
|
||||
from tools.file import load_filelist
|
||||
|
||||
# To avoid CPU overload
|
||||
os.environ["MKL_NUM_THREADS"] = "1"
|
||||
os.environ["OMP_NUM_THREADS"] = "1"
|
||||
|
||||
|
||||
def task_generator_folder(root: Path, text_extension: str):
|
||||
files = list(tqdm(Path(root).rglob("*.npy"), desc=f"Loading {root}"))
|
||||
files = sorted(files)
|
||||
|
||||
grouped_files = defaultdict(list)
|
||||
for file in tqdm(files, desc=f"Grouping {root}"):
|
||||
p = str(file.parent)
|
||||
speaker = file.parent.name
|
||||
|
||||
try:
|
||||
if isinstance(text_extension, str):
|
||||
texts = [file.with_suffix(text_extension).read_text(encoding="utf-8")]
|
||||
else:
|
||||
texts = [
|
||||
file.with_suffix(ext).read_text(encoding="utf-8")
|
||||
for ext in text_extension
|
||||
]
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to read text {file}: {e}")
|
||||
continue
|
||||
|
||||
grouped_files[p].append((speaker, file, texts))
|
||||
|
||||
logger.info(
|
||||
f"Found {len(grouped_files)} groups in {root}, {list(grouped_files.keys())[:5]}..."
|
||||
)
|
||||
|
||||
for i in grouped_files.values():
|
||||
subset = [(f, t) for _, f, t in i]
|
||||
yield i[0][0], subset, "folder"
|
||||
|
||||
|
||||
def task_generator_filelist(filelist):
|
||||
grouped_files = defaultdict(list)
|
||||
for filename, speaker, _, text in load_filelist(filelist):
|
||||
grouped_files[speaker].append((Path(filename), [text]))
|
||||
|
||||
logger.info(f"Found {len(grouped_files)} groups in {filelist}")
|
||||
for speaker, values in grouped_files.items():
|
||||
yield speaker, values, "filelist"
|
||||
|
||||
|
||||
def run_task(task):
|
||||
name, subset, source = task
|
||||
|
||||
# Parse the files
|
||||
sentences = []
|
||||
for file, texts in subset:
|
||||
np_file = file.with_suffix(".npy")
|
||||
if np_file.exists() is False:
|
||||
logger.warning(f"Can't find {np_file}")
|
||||
continue
|
||||
|
||||
new_texts = []
|
||||
|
||||
for text in texts:
|
||||
# Simple cleaning: replace { xxx } and < xxx > with space
|
||||
text = re.sub(r"\{.*?\}", " ", text)
|
||||
text = re.sub(r"<.*?>", " ", text)
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
new_texts.append(text)
|
||||
|
||||
try:
|
||||
semantics = np.load(np_file)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to parse {file}: {e}")
|
||||
continue
|
||||
|
||||
if isinstance(semantics, np.ndarray):
|
||||
semantics = semantics.tolist()
|
||||
|
||||
sentences.append(
|
||||
Sentence(
|
||||
texts=new_texts,
|
||||
semantics=[Semantics(values=s) for s in semantics],
|
||||
)
|
||||
)
|
||||
|
||||
# Pack the sentences
|
||||
return pack_pb_stream(
|
||||
TextData(
|
||||
source=source,
|
||||
name=name,
|
||||
sentences=sentences,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@click.command()
|
||||
@click.option(
|
||||
"--input",
|
||||
type=click.Path(path_type=Path),
|
||||
required=True,
|
||||
help="A folder containing the dataset or a filelist",
|
||||
multiple=True,
|
||||
)
|
||||
@click.option(
|
||||
"--output", type=click.Path(path_type=Path), default="data/quantized-dataset-ft"
|
||||
)
|
||||
@click.option("--num-workers", type=int, default=16)
|
||||
@click.option("--text-extension", type=str, default=[".txt"], multiple=True)
|
||||
@click.option(
|
||||
"--shard-size", type=int, default=10, help="The maximum size of each shard in mb"
|
||||
)
|
||||
def main(input, output, num_workers, text_extension, shard_size):
|
||||
generator_fns = []
|
||||
|
||||
for f in input:
|
||||
assert f.exists(), f"{f} not found"
|
||||
|
||||
if f.is_dir():
|
||||
generator_fn = task_generator_folder(f, text_extension)
|
||||
else:
|
||||
generator_fn = task_generator_filelist(f)
|
||||
|
||||
generator_fns.append(generator_fn)
|
||||
|
||||
generator_fn = itertools.chain(*generator_fns)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
dataset_fp = None
|
||||
tar_idx = 0
|
||||
written_size = 0
|
||||
|
||||
with Pool(num_workers) as p:
|
||||
for result in tqdm(p.imap_unordered(run_task, generator_fn)):
|
||||
if dataset_fp is None:
|
||||
dataset_fp = open(Path(output) / f"{tar_idx:08d}.protos", "wb")
|
||||
|
||||
dataset_fp.write(result)
|
||||
written_size += len(result)
|
||||
|
||||
if written_size > shard_size * 1024 * 1024:
|
||||
logger.info(f"Finished writing {tar_idx} shards to {output}")
|
||||
dataset_fp.close()
|
||||
dataset_fp = None
|
||||
written_size = 0
|
||||
tar_idx += 1
|
||||
|
||||
if dataset_fp is not None:
|
||||
dataset_fp.close()
|
||||
|
||||
logger.info(f"Finished writing {tar_idx + 1} shards to {output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,171 @@
|
||||
import pyrootutils
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from matplotlib import pyplot as plt
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
# register eval resolver and root
|
||||
pyrootutils.setup_root(__file__, indicator=".project-root", pythonpath=True)
|
||||
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from fish_speech.datasets.semantic import AutoAugTextDataset, TextDataCollator
|
||||
from tools.llama.generate import load_model
|
||||
|
||||
|
||||
def smooth(
|
||||
scalars: list[float], weight: float
|
||||
) -> list[float]: # Weight between 0 and 1
|
||||
last = scalars[0] # First value in the plot (first timestep)
|
||||
smoothed = list()
|
||||
for point in scalars:
|
||||
smoothed_val = last * weight + (1 - weight) * point # Calculate smoothed value
|
||||
smoothed.append(smoothed_val) # Save it
|
||||
last = smoothed_val # Anchor the last smoothed value
|
||||
|
||||
return smoothed
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def analyze_one_model(loader, config, weight, max_length):
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model = load_model(
|
||||
config,
|
||||
weight,
|
||||
device,
|
||||
torch.bfloat16,
|
||||
max_length,
|
||||
compile=False,
|
||||
)[0]
|
||||
|
||||
current_step = 0
|
||||
model.eval()
|
||||
|
||||
semantic_loss_sum = torch.zeros(
|
||||
max_length,
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
counter = torch.zeros(
|
||||
max_length,
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
)
|
||||
|
||||
for batch in loader:
|
||||
batch = {k: v.to(device) for k, v in batch.items()}
|
||||
|
||||
labels = batch["labels"]
|
||||
outputs = model(
|
||||
inp=batch["inputs"],
|
||||
key_padding_mask=batch["attention_masks"],
|
||||
)
|
||||
|
||||
token_logits = outputs.token_logits
|
||||
codebook_logits = outputs.codebook_logits
|
||||
|
||||
# Generate labels
|
||||
base_loss = F.cross_entropy(
|
||||
token_logits.reshape(-1, token_logits.size(-1)),
|
||||
labels[:, 0].reshape(-1),
|
||||
ignore_index=-100,
|
||||
reduction="none",
|
||||
)
|
||||
|
||||
codebook_labels = labels[:, 1 : 1 + model.config.num_codebooks].mT
|
||||
semantic_loss = F.cross_entropy(
|
||||
codebook_logits.reshape(-1, codebook_logits.size(-1)),
|
||||
codebook_labels.reshape(-1),
|
||||
ignore_index=-100,
|
||||
reduction="none",
|
||||
)
|
||||
|
||||
base_loss = base_loss.reshape(labels[:, 0].shape)
|
||||
semantic_loss = semantic_loss.reshape(codebook_labels.shape)
|
||||
|
||||
semantic_loss_frame = semantic_loss.mean(-1)
|
||||
pad_pos = codebook_labels.sum(-1) == -100 * model.config.num_codebooks
|
||||
|
||||
for loss_sample, pad in zip(semantic_loss_frame, pad_pos):
|
||||
semantic_loss_sum[~pad] += loss_sample[~pad]
|
||||
counter[~pad] += 1
|
||||
|
||||
current_step += 1
|
||||
if current_step == 10:
|
||||
break
|
||||
|
||||
semantic_loss = semantic_loss.cpu()
|
||||
counter = counter.cpu()
|
||||
xs, ys = [], []
|
||||
|
||||
for i, (loss, count) in enumerate(zip(semantic_loss_sum, counter)):
|
||||
if count > 0:
|
||||
xs.append(i)
|
||||
ys.append((loss / count).item()) # for better loss visualization
|
||||
|
||||
smoothed_ys = smooth(ys, 0.95)
|
||||
|
||||
# Unload model
|
||||
del model
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return xs, ys, smoothed_ys
|
||||
|
||||
|
||||
def main():
|
||||
tokenizer = AutoTokenizer.from_pretrained("fishaudio/fish-speech-1")
|
||||
max_length = 4096
|
||||
|
||||
ds = AutoAugTextDataset(
|
||||
["data/protos/sft/云天河"],
|
||||
tokenizer=tokenizer,
|
||||
use_speaker=False,
|
||||
interactive_prob=1.0,
|
||||
max_length=max_length,
|
||||
)
|
||||
|
||||
loader = DataLoader(
|
||||
ds,
|
||||
batch_size=8,
|
||||
collate_fn=TextDataCollator(tokenizer, max_length=max_length),
|
||||
num_workers=0,
|
||||
shuffle=False,
|
||||
)
|
||||
|
||||
plt.figure(figsize=(10, 5), dpi=200)
|
||||
|
||||
plt.xlabel("Frame")
|
||||
plt.ylabel("Loss")
|
||||
plt.yscale("log")
|
||||
plt.title("Semantic Loss")
|
||||
plt.grid(which="both", axis="both")
|
||||
plt.xlim(0, max_length)
|
||||
|
||||
tests = [
|
||||
(
|
||||
"pertrain-medium",
|
||||
"dual_ar_2_codebook_medium",
|
||||
"checkpoints/text2semantic-pretrain-medium-2k-v1.pth",
|
||||
),
|
||||
(
|
||||
"sft-medium",
|
||||
"dual_ar_2_codebook_medium",
|
||||
"checkpoints/text2semantic-sft-medium-v1.1-4k.pth",
|
||||
),
|
||||
(
|
||||
"sft-large",
|
||||
"dual_ar_2_codebook_large",
|
||||
"checkpoints/text2semantic-sft-large-v1.1-4k.pth",
|
||||
),
|
||||
]
|
||||
|
||||
for name, config, weight in tests:
|
||||
xs, _, smoothed_ys = analyze_one_model(loader, config, weight, max_length)
|
||||
plt.plot(xs, smoothed_ys, label=name)
|
||||
|
||||
plt.legend()
|
||||
plt.savefig("semantic_loss.png")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,699 @@
|
||||
import os
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Literal, Optional, Tuple, Union
|
||||
|
||||
import click
|
||||
import hydra
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch._dynamo.config
|
||||
import torch._inductor.config
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
|
||||
from fish_speech.conversation import CODEBOOK_PAD_TOKEN_ID
|
||||
from fish_speech.clean import clean_text
|
||||
from fish_speech.spliter import split_text
|
||||
|
||||
|
||||
import comfy.utils
|
||||
|
||||
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
torch._inductor.config.coordinate_descent_tuning = True
|
||||
torch._inductor.config.triton.unique_kernel_names = True
|
||||
|
||||
if hasattr(torch._inductor.config, "fx_graph_cache"):
|
||||
# Experimental feature to reduce compilation times, will be on by default in future
|
||||
torch._inductor.config.fx_graph_cache = True
|
||||
|
||||
|
||||
from ...models.text2semantic.llama import BaseTransformer, DualARTransformer, NaiveTransformer
|
||||
|
||||
|
||||
def multinomial_sample_one_no_sync(
|
||||
probs_sort,
|
||||
): # Does multinomial sampling without a cuda synchronization
|
||||
q = torch.empty_like(probs_sort).exponential_(1)
|
||||
return torch.argmax(probs_sort / q, dim=-1, keepdim=True).to(dtype=torch.int)
|
||||
|
||||
|
||||
def logits_to_probs(
|
||||
logits,
|
||||
previous_tokens: Optional[torch.Tensor] = None,
|
||||
temperature: torch.Tensor = 1.0,
|
||||
top_p: torch.Tensor = 1.0,
|
||||
repetition_penalty: torch.Tensor = 1.0,
|
||||
) -> torch.Tensor:
|
||||
# Apply repetition penalty
|
||||
if previous_tokens is not None:
|
||||
previous_tokens = previous_tokens.long()
|
||||
score = torch.gather(logits, dim=0, index=previous_tokens)
|
||||
score = torch.where(
|
||||
score < 0, score * repetition_penalty, score / repetition_penalty
|
||||
)
|
||||
logits.scatter_(dim=0, index=previous_tokens, src=score)
|
||||
|
||||
# Apply top-p sampling
|
||||
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
|
||||
cum_probs = torch.cumsum(torch.nn.functional.softmax(sorted_logits, dim=-1), dim=-1)
|
||||
sorted_indices_to_remove = cum_probs > top_p
|
||||
sorted_indices_to_remove[0] = False # keep at least one option
|
||||
indices_to_remove = sorted_indices_to_remove.scatter(
|
||||
dim=0, index=sorted_indices, src=sorted_indices_to_remove
|
||||
)
|
||||
logits = logits.masked_fill(indices_to_remove, -float("Inf"))
|
||||
|
||||
logits = logits / max(temperature, 1e-5)
|
||||
|
||||
probs = torch.nn.functional.softmax(logits, dim=-1)
|
||||
return probs
|
||||
|
||||
|
||||
def sample(
|
||||
logits,
|
||||
previous_tokens: Optional[torch.Tensor] = None,
|
||||
**sampling_kwargs,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
probs = logits_to_probs(
|
||||
logits=logits[0, -1], previous_tokens=previous_tokens, **sampling_kwargs
|
||||
)
|
||||
idx_next = multinomial_sample_one_no_sync(probs)
|
||||
return idx_next, probs
|
||||
|
||||
|
||||
def decode_one_token_ar(
|
||||
model: DualARTransformer,
|
||||
x: torch.Tensor,
|
||||
input_pos: torch.Tensor,
|
||||
previous_tokens: torch.Tensor = None,
|
||||
**sampling_kwargs,
|
||||
) -> torch.Tensor:
|
||||
x = model.forward_generate(x, input_pos)
|
||||
codebooks = [
|
||||
sample(
|
||||
x.logits,
|
||||
previous_tokens=(
|
||||
previous_tokens[0] if previous_tokens is not None else None
|
||||
), # Disable repetition penalty for the token codebook
|
||||
**sampling_kwargs,
|
||||
)[0]
|
||||
]
|
||||
x = x.hidden_states
|
||||
|
||||
# Cleanup the cache
|
||||
for layer in model.fast_layers:
|
||||
layer.attention.kv_cache.k_cache.fill_(0)
|
||||
layer.attention.kv_cache.v_cache.fill_(0)
|
||||
|
||||
for codebook_idx in range(model.config.num_codebooks):
|
||||
input_pos = torch.tensor([codebook_idx], device=x.device, dtype=torch.long)
|
||||
logits = model.forward_generate_fast(x, input_pos)
|
||||
a = sample(
|
||||
logits,
|
||||
previous_tokens=(
|
||||
previous_tokens[codebook_idx + 1]
|
||||
if previous_tokens is not None
|
||||
else None
|
||||
),
|
||||
**sampling_kwargs,
|
||||
)[0]
|
||||
x = model.fast_embeddings(a)
|
||||
codebooks.append(a)
|
||||
|
||||
return torch.stack(codebooks, dim=0)
|
||||
|
||||
|
||||
def decode_one_token_naive(
|
||||
model: NaiveTransformer,
|
||||
x: torch.Tensor,
|
||||
input_pos: torch.Tensor,
|
||||
previous_tokens: torch.Tensor = None,
|
||||
**sampling_kwargs,
|
||||
) -> torch.Tensor:
|
||||
x = model.forward_generate(x, input_pos)
|
||||
|
||||
codebooks = [
|
||||
sample(
|
||||
x.token_logits,
|
||||
previous_tokens=None, # Disable repetition penalty for the token codebook
|
||||
**sampling_kwargs,
|
||||
)[0]
|
||||
]
|
||||
|
||||
for i in range(model.config.num_codebooks):
|
||||
codebooks.append(
|
||||
sample(
|
||||
x.codebook_logits[:, :, i],
|
||||
previous_tokens=(
|
||||
previous_tokens[i + 1] if previous_tokens is not None else None
|
||||
),
|
||||
**sampling_kwargs,
|
||||
)[0]
|
||||
)
|
||||
|
||||
return torch.stack(codebooks, dim=0)
|
||||
|
||||
|
||||
def decode_n_tokens(
|
||||
model: NaiveTransformer,
|
||||
cur_token: torch.Tensor,
|
||||
input_pos: torch.Tensor,
|
||||
num_new_tokens: int,
|
||||
im_end_id: int = 4,
|
||||
decode_one_token=decode_one_token_naive,
|
||||
**sampling_kwargs,
|
||||
):
|
||||
previous_tokens = torch.zeros(
|
||||
(model.config.num_codebooks + 1, model.config.max_seq_len),
|
||||
dtype=torch.int,
|
||||
device=cur_token.device,
|
||||
)
|
||||
|
||||
for i in tqdm(range(num_new_tokens)):
|
||||
# We need to get windowed repeat penalty
|
||||
win_size = 16
|
||||
if i < win_size:
|
||||
window = previous_tokens[:, :win_size]
|
||||
else:
|
||||
window = previous_tokens[:, i - win_size : i]
|
||||
|
||||
with torch.backends.cuda.sdp_kernel(
|
||||
enable_flash=False, enable_mem_efficient=False, enable_math=True
|
||||
): # Actually better for Inductor to codegen attention here
|
||||
next_token = decode_one_token(
|
||||
model=model,
|
||||
x=cur_token,
|
||||
input_pos=input_pos,
|
||||
previous_tokens=window,
|
||||
**sampling_kwargs,
|
||||
)
|
||||
|
||||
input_pos += 1
|
||||
cur_token = next_token.view(1, model.config.num_codebooks + 1, -1)
|
||||
previous_tokens[:, i : i + 1] = next_token.view(
|
||||
model.config.num_codebooks + 1, -1
|
||||
)
|
||||
|
||||
if cur_token[0, 0, -1] == im_end_id:
|
||||
break
|
||||
|
||||
return previous_tokens[:, : i + 1]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.inference_mode()
|
||||
def generate(
|
||||
*,
|
||||
model: NaiveTransformer,
|
||||
prompt: torch.Tensor,
|
||||
max_new_tokens: int,
|
||||
im_end_id: int = 4,
|
||||
decode_one_token=decode_one_token_naive,
|
||||
**sampling_kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Takes a conditioning sequence (prompt) as input and continues to generate as many tokens as requested.
|
||||
"""
|
||||
|
||||
# create an empty tensor of the expected final shape and fill in the current tokens
|
||||
T = prompt.size(1)
|
||||
|
||||
if max_new_tokens:
|
||||
if T + max_new_tokens > model.config.max_seq_len:
|
||||
max_new_tokens = model.config.max_seq_len - T
|
||||
logger.info(f"Truncating max_new_tokens to {max_new_tokens}")
|
||||
|
||||
T_new = T + max_new_tokens
|
||||
else:
|
||||
T_new = model.config.max_seq_len
|
||||
max_new_tokens = T_new - T
|
||||
|
||||
device, dtype = prompt.device, prompt.dtype
|
||||
with torch.device(device):
|
||||
model.setup_caches(
|
||||
max_batch_size=1, max_seq_len=T_new, dtype=next(model.parameters()).dtype
|
||||
)
|
||||
|
||||
codebook_dim = 1 + model.config.num_codebooks
|
||||
# create an empty tensor of the expected final shape and fill in the current tokens
|
||||
empty = torch.empty((codebook_dim, T_new), dtype=dtype, device=device)
|
||||
empty[:, :T] = prompt
|
||||
seq = empty
|
||||
input_pos = torch.arange(0, T, device=device)
|
||||
|
||||
# Use non-accelerated version for now, to avoid compilation overhead
|
||||
prefill_decode = (
|
||||
decode_one_token_naive
|
||||
if isinstance(model, NaiveTransformer)
|
||||
else decode_one_token_ar
|
||||
)
|
||||
|
||||
next_token = prefill_decode(
|
||||
model, prompt.view(1, codebook_dim, -1), input_pos, **sampling_kwargs
|
||||
)
|
||||
seq[:, T : T + 1] = next_token
|
||||
|
||||
input_pos = torch.tensor([T], device=device, dtype=torch.int)
|
||||
x = decode_n_tokens(
|
||||
model,
|
||||
next_token.view(1, codebook_dim, -1),
|
||||
input_pos,
|
||||
max_new_tokens - 1,
|
||||
im_end_id=im_end_id,
|
||||
decode_one_token=decode_one_token,
|
||||
**sampling_kwargs,
|
||||
)
|
||||
# x = torch.cat(generated_tokens, dim=1)
|
||||
seq = seq[:, : T + 1 + x.size(1)]
|
||||
seq[:, T + 1 :] = x
|
||||
|
||||
return seq
|
||||
|
||||
|
||||
def encode_tokens(
|
||||
tokenizer,
|
||||
string,
|
||||
device="cuda",
|
||||
prompt_tokens=None,
|
||||
num_codebooks=4,
|
||||
):
|
||||
string = clean_text(string)
|
||||
string = f"<|im_start|>user\n{string}<|im_end|><|im_start|>assistant\n"
|
||||
|
||||
new_tokens = tokenizer.encode(
|
||||
string,
|
||||
add_special_tokens=False,
|
||||
max_length=10**6,
|
||||
truncation=False,
|
||||
)
|
||||
tokens = torch.tensor([new_tokens], dtype=torch.int, device=device)
|
||||
|
||||
# Codebooks
|
||||
zeros = (
|
||||
torch.ones((num_codebooks, tokens.size(1)), dtype=torch.int, device=device)
|
||||
* CODEBOOK_PAD_TOKEN_ID
|
||||
)
|
||||
prompt = torch.cat((tokens, zeros), dim=0)
|
||||
|
||||
if prompt_tokens is None:
|
||||
return prompt
|
||||
|
||||
# Get prompt tokens
|
||||
if prompt_tokens.ndim == 3:
|
||||
assert (
|
||||
prompt_tokens.shape[0] == 1
|
||||
), f"3 dim prompt tokens should have shape (1, num_codebooks, seq_len)"
|
||||
prompt_tokens = prompt_tokens[0]
|
||||
|
||||
assert prompt_tokens.ndim == 2
|
||||
data = prompt_tokens + 1
|
||||
|
||||
if prompt_tokens.shape[0] > num_codebooks:
|
||||
logger.warning(
|
||||
f"Prompt tokens shape {prompt_tokens.shape} is larger than num_codebooks {num_codebooks}, getting first {num_codebooks} codebooks"
|
||||
)
|
||||
data = data[:num_codebooks]
|
||||
|
||||
# Add pad token for each codebook
|
||||
data = torch.cat(
|
||||
(data, torch.zeros((data.size(0), 1), dtype=torch.int, device=device)),
|
||||
dim=1,
|
||||
)
|
||||
|
||||
# Since 1.0, we use <|semantic|>
|
||||
s0_token_id = tokenizer.convert_tokens_to_ids("<|semantic|>")
|
||||
end_token_id = tokenizer.convert_tokens_to_ids("<|im_end|>")
|
||||
main_token_ids = (
|
||||
torch.ones((1, data.size(1)), dtype=torch.int, device=device) * s0_token_id
|
||||
)
|
||||
main_token_ids[0, -1] = end_token_id
|
||||
|
||||
data = torch.cat((main_token_ids, data), dim=0)
|
||||
prompt = torch.cat((prompt, data), dim=1)
|
||||
|
||||
return prompt
|
||||
|
||||
|
||||
def load_model(checkpoint_path, device, precision, compile=False):
|
||||
model: Union[NaiveTransformer, DualARTransformer] = BaseTransformer.from_pretrained(
|
||||
checkpoint_path, load_weights=True
|
||||
)
|
||||
|
||||
model = model.to(device=device, dtype=precision)
|
||||
logger.info(f"Restored model from checkpoint")
|
||||
|
||||
if isinstance(model, DualARTransformer):
|
||||
decode_one_token = decode_one_token_ar
|
||||
logger.info("Using DualARTransformer")
|
||||
else:
|
||||
decode_one_token = decode_one_token_naive
|
||||
logger.info("Using NaiveTransformer")
|
||||
|
||||
if compile:
|
||||
logger.info("Compiling function...")
|
||||
decode_one_token = torch.compile(
|
||||
decode_one_token, mode="reduce-overhead", fullgraph=True
|
||||
)
|
||||
|
||||
return model.eval(), decode_one_token
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerateResponse:
|
||||
action: Literal["sample", "next"]
|
||||
codes: Optional[torch.Tensor] = None
|
||||
text: Optional[str] = None
|
||||
|
||||
|
||||
def generate_long(
|
||||
*,
|
||||
model,
|
||||
device: str | torch.device,
|
||||
decode_one_token: callable,
|
||||
text: str,
|
||||
num_samples: int = 1,
|
||||
max_new_tokens: int = 0,
|
||||
top_p: int = 0.7,
|
||||
repetition_penalty: float = 1.5,
|
||||
temperature: float = 0.7,
|
||||
compile: bool = False,
|
||||
iterative_prompt: bool = True,
|
||||
max_length: int = 2048,
|
||||
chunk_length: int = 150,
|
||||
prompt_text: Optional[str | list[str]] = None,
|
||||
prompt_tokens: Optional[torch.Tensor | list[torch.Tensor]] = None,
|
||||
):
|
||||
assert 0 < top_p <= 1, "top_p must be in (0, 1]"
|
||||
assert 0 < repetition_penalty < 2, "repetition_penalty must be in (0, 2)"
|
||||
assert 0 < temperature < 2, "temperature must be in (0, 2)"
|
||||
|
||||
use_prompt = prompt_text is not None and prompt_tokens is not None
|
||||
if use_prompt and isinstance(prompt_text, str):
|
||||
prompt_text = [prompt_text]
|
||||
prompt_tokens = [prompt_tokens]
|
||||
|
||||
assert use_prompt is False or len(prompt_text) == len(
|
||||
prompt_tokens
|
||||
), "Prompt text and tokens must have the same length"
|
||||
|
||||
model_size = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
tokenizer = model.tokenizer
|
||||
im_end_id = tokenizer.convert_tokens_to_ids("<|im_end|>")
|
||||
|
||||
encoded = []
|
||||
texts = split_text(text, chunk_length) if iterative_prompt else [text]
|
||||
encoded_prompts = []
|
||||
|
||||
if use_prompt:
|
||||
for idx, (t, c) in enumerate(zip(prompt_text, prompt_tokens)):
|
||||
encoded_prompts.append(
|
||||
encode_tokens(
|
||||
tokenizer,
|
||||
string=t,
|
||||
device=device,
|
||||
prompt_tokens=c,
|
||||
num_codebooks=model.config.num_codebooks,
|
||||
)
|
||||
)
|
||||
|
||||
for idx, text in enumerate(texts):
|
||||
encoded.append(
|
||||
encode_tokens(
|
||||
tokenizer,
|
||||
string=text,
|
||||
device=device,
|
||||
num_codebooks=model.config.num_codebooks,
|
||||
)
|
||||
)
|
||||
logger.info(f"Encoded text: {text}")
|
||||
|
||||
# Move temperature, top_p, repetition_penalty to device
|
||||
# This is important so that changing params doesn't trigger recompile
|
||||
temperature = torch.tensor(temperature, device=device, dtype=torch.float)
|
||||
top_p = torch.tensor(top_p, device=device, dtype=torch.float)
|
||||
repetition_penalty = torch.tensor(
|
||||
repetition_penalty, device=device, dtype=torch.float
|
||||
)
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(num_samples*len(encoded))
|
||||
|
||||
for sample_idx in range(num_samples):
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
global_encoded = []
|
||||
seg_idx = 0
|
||||
|
||||
while seg_idx < len(encoded):
|
||||
logger.info(
|
||||
f"Generating sentence {seg_idx + 1}/{len(encoded)} of sample {sample_idx + 1}/{num_samples}"
|
||||
)
|
||||
pbar.update(1)
|
||||
seg = encoded[seg_idx]
|
||||
global_encoded.append(seg)
|
||||
|
||||
lengths = reversed([seg.size(1) for seg in global_encoded])
|
||||
|
||||
# Pick last 2000 tokens
|
||||
count = 0
|
||||
for i, length in enumerate(lengths):
|
||||
count += length
|
||||
if count + length > max_length - 1024 - sum(
|
||||
t.shape[1] for t in encoded_prompts
|
||||
):
|
||||
break
|
||||
|
||||
if i != 0 and i % 2 == 0:
|
||||
i -= 1
|
||||
|
||||
# Rotate the list, always make sure first segment is included to avoid drift
|
||||
if i < len(global_encoded) - 2:
|
||||
partial_encoded = global_encoded[:2] + global_encoded[-i:]
|
||||
else:
|
||||
partial_encoded = global_encoded
|
||||
|
||||
if use_prompt:
|
||||
partial_encoded = encoded_prompts + partial_encoded
|
||||
|
||||
cat_encoded = torch.cat(partial_encoded, dim=1)
|
||||
prompt_length = cat_encoded.size(1)
|
||||
|
||||
t0 = time.perf_counter()
|
||||
y = generate(
|
||||
model=model,
|
||||
prompt=cat_encoded,
|
||||
max_new_tokens=max_new_tokens,
|
||||
im_end_id=im_end_id,
|
||||
decode_one_token=decode_one_token,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
repetition_penalty=repetition_penalty,
|
||||
)
|
||||
|
||||
if sample_idx == 0 and seg_idx == 0 and compile:
|
||||
logger.info(f"Compilation time: {time.perf_counter() - t0:.2f} seconds")
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
t = time.perf_counter() - t0
|
||||
|
||||
tokens_generated = y.size(1) - prompt_length
|
||||
tokens_sec = tokens_generated / t
|
||||
logger.info(
|
||||
f"Generated {tokens_generated} tokens in {t:.02f} seconds, {tokens_sec:.02f} tokens/sec"
|
||||
)
|
||||
logger.info(
|
||||
f"Bandwidth achieved: {model_size * tokens_sec / 1e9:.02f} GB/s"
|
||||
)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
logger.info(
|
||||
f"GPU Memory used: {torch.cuda.max_memory_reserved() / 1e9:.02f} GB"
|
||||
)
|
||||
|
||||
# Put the generated tokens
|
||||
# since there is <im_end> and <eos> tokens, we remove last 2 tokens
|
||||
codes = y[1:, prompt_length:-1].clone()
|
||||
codes = codes - 1
|
||||
assert (codes >= 0).all(), f"Negative code found"
|
||||
|
||||
decoded = y[:, prompt_length:-1].clone()
|
||||
# But for global encoding, we should keep the <im_end> token
|
||||
|
||||
global_encoded.append(decoded)
|
||||
assert (codes >= 0).all(), f"Negative code found: {codes}"
|
||||
yield GenerateResponse(action="sample", codes=codes, text=texts[seg_idx])
|
||||
seg_idx += 1
|
||||
|
||||
# This indicates the end of the current sample
|
||||
yield GenerateResponse(action="next")
|
||||
|
||||
|
||||
@dataclass
|
||||
class WrappedGenerateResponse:
|
||||
status: Literal["success", "error"]
|
||||
response: Optional[GenerateResponse | Exception] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerateRequest:
|
||||
request: dict
|
||||
response_queue: queue.Queue
|
||||
|
||||
|
||||
def launch_thread_safe_queue(
|
||||
checkpoint_path,
|
||||
device,
|
||||
precision,
|
||||
compile: bool = False,
|
||||
):
|
||||
input_queue = queue.Queue()
|
||||
init_event = threading.Event()
|
||||
|
||||
def worker():
|
||||
model, decode_one_token = load_model(
|
||||
checkpoint_path, device, precision, compile=compile
|
||||
)
|
||||
init_event.set()
|
||||
|
||||
while True:
|
||||
item: GenerateRequest | None = input_queue.get()
|
||||
if item is None:
|
||||
break
|
||||
|
||||
kwargs = item.request
|
||||
response_queue = item.response_queue
|
||||
|
||||
try:
|
||||
for chunk in generate_long(
|
||||
model=model, decode_one_token=decode_one_token, **kwargs
|
||||
):
|
||||
response_queue.put(
|
||||
WrappedGenerateResponse(status="success", response=chunk)
|
||||
)
|
||||
except Exception as e:
|
||||
response_queue.put(WrappedGenerateResponse(status="error", response=e))
|
||||
|
||||
threading.Thread(target=worker, daemon=True).start()
|
||||
init_event.wait()
|
||||
|
||||
return input_queue
|
||||
|
||||
|
||||
@click.command()
|
||||
@click.option(
|
||||
"--text",
|
||||
type=str,
|
||||
default="你说的对, 但是原神是一款由米哈游自主研发的开放世界手游.",
|
||||
)
|
||||
@click.option("--prompt-text", type=str, default=None, multiple=True)
|
||||
@click.option(
|
||||
"--prompt-tokens",
|
||||
type=click.Path(path_type=Path, exists=True),
|
||||
default=None,
|
||||
multiple=True,
|
||||
)
|
||||
@click.option("--num-samples", type=int, default=1)
|
||||
@click.option("--max-new-tokens", type=int, default=0)
|
||||
@click.option("--top-p", type=float, default=0.7)
|
||||
@click.option("--repetition-penalty", type=float, default=1.2)
|
||||
@click.option("--temperature", type=float, default=0.7)
|
||||
@click.option(
|
||||
"--checkpoint-path",
|
||||
type=click.Path(path_type=Path, exists=True),
|
||||
default="checkpoints/fish-speech-1.2-sft",
|
||||
)
|
||||
@click.option("--device", type=str, default="cuda")
|
||||
@click.option("--compile/--no-compile", default=False)
|
||||
@click.option("--seed", type=int, default=42)
|
||||
@click.option("--half/--no-half", default=False)
|
||||
@click.option("--iterative-prompt/--no-iterative-prompt", default=True)
|
||||
@click.option("--chunk-length", type=int, default=100)
|
||||
def main(
|
||||
text: str,
|
||||
prompt_text: Optional[list[str]],
|
||||
prompt_tokens: Optional[list[Path]],
|
||||
num_samples: int,
|
||||
max_new_tokens: int,
|
||||
top_p: int,
|
||||
repetition_penalty: float,
|
||||
temperature: float,
|
||||
checkpoint_path: Path,
|
||||
device: str,
|
||||
compile: bool,
|
||||
seed: int,
|
||||
half: bool,
|
||||
iterative_prompt: bool,
|
||||
chunk_length: int,
|
||||
) -> None:
|
||||
|
||||
precision = torch.half if half else torch.bfloat16
|
||||
|
||||
if prompt_text is not None and len(prompt_text) != len(prompt_tokens):
|
||||
raise ValueError(
|
||||
f"Number of prompt text ({len(prompt_text)}) and prompt tokens ({len(prompt_tokens)}) should be the same"
|
||||
)
|
||||
|
||||
logger.info("Loading model ...")
|
||||
t0 = time.time()
|
||||
model, decode_one_token = load_model(
|
||||
checkpoint_path, device, precision, compile=compile
|
||||
)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
logger.info(f"Time to load model: {time.time() - t0:.02f} seconds")
|
||||
|
||||
if prompt_tokens is not None:
|
||||
prompt_tokens = [torch.from_numpy(np.load(p)).to(device) for p in prompt_tokens]
|
||||
|
||||
torch.manual_seed(seed)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
generator = generate_long(
|
||||
model=model,
|
||||
device=device,
|
||||
decode_one_token=decode_one_token,
|
||||
text=text,
|
||||
num_samples=num_samples,
|
||||
max_new_tokens=max_new_tokens,
|
||||
top_p=top_p,
|
||||
repetition_penalty=repetition_penalty,
|
||||
temperature=temperature,
|
||||
compile=compile,
|
||||
iterative_prompt=iterative_prompt,
|
||||
chunk_length=chunk_length,
|
||||
prompt_text=prompt_text,
|
||||
prompt_tokens=prompt_tokens,
|
||||
)
|
||||
|
||||
idx = 0
|
||||
codes = []
|
||||
|
||||
for response in generator:
|
||||
if response.action == "sample":
|
||||
codes.append(response.codes)
|
||||
logger.info(f"Sampled text: {response.text}")
|
||||
elif response.action == "next":
|
||||
if codes:
|
||||
np.save(f"codes_{idx}.npy", torch.cat(codes, dim=1).cpu().numpy())
|
||||
logger.info(f"Saved codes to codes_{idx}.npy")
|
||||
logger.info(f"Next sample")
|
||||
codes = []
|
||||
idx += 1
|
||||
else:
|
||||
logger.error(f"Error: {response}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,95 @@
|
||||
import shutil
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
|
||||
import click
|
||||
import hydra
|
||||
import torch
|
||||
from hydra import compose, initialize
|
||||
from hydra.utils import instantiate
|
||||
from loguru import logger
|
||||
|
||||
from fish_speech.models.text2semantic.llama import BaseTransformer
|
||||
from fish_speech.models.text2semantic.lora import get_merged_state_dict
|
||||
|
||||
|
||||
@click.command()
|
||||
@click.option("--lora-config", type=str, default="r_8_alpha_16")
|
||||
@click.option("--base-weight", type=str, default="checkpoints/fish-speech-1.4")
|
||||
@click.option("--lora-weight", type=str, required=True)
|
||||
@click.option("--output", type=str, required=True)
|
||||
def merge(lora_config, base_weight, lora_weight, output):
|
||||
output = Path(output)
|
||||
logger.info(
|
||||
f"Merging {base_weight} and {lora_weight} into {output} with {lora_config}"
|
||||
)
|
||||
|
||||
with initialize(version_base="1.3", config_path="../../fish_speech/configs/lora"):
|
||||
cfg = compose(config_name=lora_config)
|
||||
|
||||
lora_config = instantiate(cfg)
|
||||
logger.info(f"Loaded lora model with config {lora_config}")
|
||||
|
||||
llama_model = BaseTransformer.from_pretrained(
|
||||
path=base_weight,
|
||||
load_weights=True,
|
||||
lora_config=lora_config,
|
||||
)
|
||||
logger.info(f"Loaded llama model")
|
||||
|
||||
llama_state_dict = llama_model.state_dict()
|
||||
llama_state_dict = {k: v for k, v in llama_state_dict.items() if "lora" not in k}
|
||||
llama_state_dict_copy = deepcopy(llama_state_dict)
|
||||
lora_state_dict = torch.load(lora_weight, map_location="cpu")
|
||||
|
||||
if "state_dict" in llama_state_dict:
|
||||
llama_state_dict = llama_state_dict["state_dict"]
|
||||
|
||||
if "state_dict" in lora_state_dict:
|
||||
lora_state_dict = lora_state_dict["state_dict"]
|
||||
|
||||
# remove prefix model.
|
||||
if any(k.startswith("model.") for k in llama_state_dict.keys()):
|
||||
llama_state_dict = {
|
||||
k.replace("model.", ""): v
|
||||
for k, v in llama_state_dict.items()
|
||||
if k.startswith("model.")
|
||||
}
|
||||
if any(k.startswith("model.") for k in lora_state_dict.keys()):
|
||||
lora_state_dict = {
|
||||
k.replace("model.", ""): v
|
||||
for k, v in lora_state_dict.items()
|
||||
if k.startswith("model.")
|
||||
}
|
||||
|
||||
logger.info(f"Found {len(llama_state_dict)} keys in llama model")
|
||||
logger.info(f"Found {len(lora_state_dict)} keys in lora model")
|
||||
|
||||
merged_state_dict = llama_state_dict | lora_state_dict
|
||||
llama_model.load_state_dict(merged_state_dict, strict=True)
|
||||
logger.info(f"Merged model loaded")
|
||||
|
||||
# Trigger eval mode to merge lora
|
||||
llama_model.eval()
|
||||
llama_model.save_pretrained(output, drop_lora=True)
|
||||
logger.info(f"Saved merged model to {output}, validating")
|
||||
|
||||
new_state_dict = torch.load(output / "model.pth", map_location="cpu")
|
||||
original_keys = set(llama_state_dict_copy.keys())
|
||||
merged_keys = set(new_state_dict.keys())
|
||||
|
||||
assert original_keys == merged_keys, "Keys should be same"
|
||||
|
||||
for key in original_keys:
|
||||
diff_l1 = (new_state_dict[key] - llama_state_dict_copy[key]).abs().sum().item()
|
||||
if diff_l1 != 0:
|
||||
break
|
||||
else:
|
||||
logger.error("Merged model is same as the original model")
|
||||
exit(1)
|
||||
|
||||
logger.info("Merged model is different from the original model, check passed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
merge()
|
||||
@@ -0,0 +1,497 @@
|
||||
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
||||
# All rights reserved.
|
||||
import datetime
|
||||
import shutil
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import click
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fish_speech.models.text2semantic.llama import find_multiple
|
||||
from tools.llama.generate import load_model
|
||||
|
||||
##### Quantization Primitives ######
|
||||
|
||||
|
||||
def dynamically_quantize_per_channel(x, quant_min, quant_max, target_dtype):
|
||||
# assumes symmetric quantization
|
||||
# assumes axis == 0
|
||||
# assumes dense memory format
|
||||
# TODO(future): relax ^ as needed
|
||||
|
||||
# default setup for affine quantization of activations
|
||||
eps = torch.finfo(torch.float32).eps
|
||||
|
||||
# get min and max
|
||||
min_val, max_val = torch.aminmax(x, dim=1)
|
||||
|
||||
# calculate scales and zero_points based on min and max
|
||||
# reference: https://fburl.com/code/srbiybme
|
||||
min_val_neg = torch.min(min_val, torch.zeros_like(min_val))
|
||||
max_val_pos = torch.max(max_val, torch.zeros_like(max_val))
|
||||
device = min_val_neg.device
|
||||
|
||||
# reference: https://fburl.com/code/4wll53rk
|
||||
max_val_pos = torch.max(-min_val_neg, max_val_pos)
|
||||
scales = max_val_pos / (float(quant_max - quant_min) / 2)
|
||||
# ensure scales is the same dtype as the original tensor
|
||||
scales = torch.clamp(scales, min=eps).to(x.dtype)
|
||||
zero_points = torch.zeros(min_val_neg.size(), dtype=torch.int64, device=device)
|
||||
|
||||
# quantize based on qmin/qmax/scales/zp
|
||||
# reference: https://www.internalfb.com/code/fbsource/[8edc275012b1]/fbcode/caffe2/torch/ao/quantization/fx/_decomposed.py?lines=63
|
||||
x_div = x / scales.unsqueeze(-1)
|
||||
x_round = torch.round(x_div)
|
||||
x_zp = x_round + zero_points.unsqueeze(-1)
|
||||
quant = torch.clamp(x_zp, quant_min, quant_max).to(target_dtype)
|
||||
|
||||
return quant, scales, zero_points
|
||||
|
||||
|
||||
def get_group_qparams(w, n_bit=4, groupsize=128):
|
||||
# needed for GPTQ with padding
|
||||
if groupsize > w.shape[-1]:
|
||||
groupsize = w.shape[-1]
|
||||
assert groupsize > 1
|
||||
assert w.shape[-1] % groupsize == 0
|
||||
assert w.dim() == 2
|
||||
|
||||
to_quant = w.reshape(-1, groupsize)
|
||||
assert torch.isnan(to_quant).sum() == 0
|
||||
|
||||
max_val = to_quant.amax(dim=1, keepdim=True)
|
||||
min_val = to_quant.amin(dim=1, keepdim=True)
|
||||
max_int = 2**n_bit - 1
|
||||
scales = (max_val - min_val).clamp(min=1e-6) / max_int
|
||||
zeros = min_val + scales * (2 ** (n_bit - 1))
|
||||
return scales.to(torch.bfloat16).reshape(w.shape[0], -1), zeros.to(
|
||||
torch.bfloat16
|
||||
).reshape(w.shape[0], -1)
|
||||
|
||||
|
||||
def pack_scales_and_zeros(scales, zeros):
|
||||
assert scales.shape == zeros.shape
|
||||
assert scales.dtype == torch.bfloat16
|
||||
assert zeros.dtype == torch.bfloat16
|
||||
return (
|
||||
torch.cat(
|
||||
[
|
||||
scales.reshape(scales.size(0), scales.size(1), 1),
|
||||
zeros.reshape(zeros.size(0), zeros.size(1), 1),
|
||||
],
|
||||
2,
|
||||
)
|
||||
.transpose(0, 1)
|
||||
.contiguous()
|
||||
)
|
||||
|
||||
|
||||
def unpack_scales_and_zeros(scales_and_zeros):
|
||||
assert len(scales_and_zeros.shape) == 3 and scales_and_zeros.shape[2] == 2
|
||||
assert scales_and_zeros.dtype == torch.float
|
||||
return torch.split(scales_and_zeros.transpose(0, 1), 1, 2)
|
||||
|
||||
|
||||
def group_quantize_tensor_from_qparams(w, scales, zeros, n_bit=4, groupsize=128):
|
||||
assert groupsize > 1
|
||||
# needed for GPTQ single column quantize
|
||||
if groupsize > w.shape[-1] and scales.shape[-1] == 1:
|
||||
groupsize = w.shape[-1]
|
||||
|
||||
assert w.shape[-1] % groupsize == 0
|
||||
assert w.dim() == 2
|
||||
|
||||
to_quant = w.reshape(-1, groupsize)
|
||||
assert torch.isnan(to_quant).sum() == 0
|
||||
|
||||
scales = scales.reshape(-1, 1)
|
||||
zeros = zeros.reshape(-1, 1)
|
||||
min_val = zeros - scales * (2 ** (n_bit - 1))
|
||||
max_int = 2**n_bit - 1
|
||||
min_int = 0
|
||||
w_int32 = (
|
||||
to_quant.sub(min_val)
|
||||
.div(scales)
|
||||
.round()
|
||||
.clamp_(min_int, max_int)
|
||||
.to(torch.int32)
|
||||
.reshape_as(w)
|
||||
)
|
||||
|
||||
return w_int32
|
||||
|
||||
|
||||
def group_quantize_tensor(w, n_bit=4, groupsize=128):
|
||||
scales, zeros = get_group_qparams(w, n_bit, groupsize)
|
||||
w_int32 = group_quantize_tensor_from_qparams(w, scales, zeros, n_bit, groupsize)
|
||||
scales_and_zeros = pack_scales_and_zeros(scales, zeros)
|
||||
return w_int32, scales_and_zeros
|
||||
|
||||
|
||||
def group_dequantize_tensor_from_qparams(
|
||||
w_int32, scales, zeros, n_bit=4, groupsize=128
|
||||
):
|
||||
assert groupsize > 1
|
||||
# needed for GPTQ single column dequantize
|
||||
if groupsize > w_int32.shape[-1] and scales.shape[-1] == 1:
|
||||
groupsize = w_int32.shape[-1]
|
||||
assert w_int32.shape[-1] % groupsize == 0
|
||||
assert w_int32.dim() == 2
|
||||
|
||||
w_int32_grouped = w_int32.reshape(-1, groupsize)
|
||||
scales = scales.reshape(-1, 1)
|
||||
zeros = zeros.reshape(-1, 1)
|
||||
|
||||
w_dq = (
|
||||
w_int32_grouped.sub(2 ** (n_bit - 1)).mul(scales).add(zeros).reshape_as(w_int32)
|
||||
)
|
||||
return w_dq
|
||||
|
||||
|
||||
def group_dequantize_tensor(w_int32, scales_and_zeros, n_bit=4, groupsize=128):
|
||||
scales, zeros = unpack_scales_and_zeros(scales_and_zeros)
|
||||
return group_dequantize_tensor_from_qparams(
|
||||
w_int32, scales, zeros, n_bit, groupsize
|
||||
)
|
||||
|
||||
|
||||
class QuantHandler:
|
||||
def __init__(self, mod):
|
||||
self.mod = mod
|
||||
|
||||
def create_quantized_state_dict(self) -> "StateDict":
|
||||
pass
|
||||
|
||||
def convert_for_runtime(self) -> "nn.Module":
|
||||
pass
|
||||
|
||||
|
||||
##### Weight-only int8 per-channel quantized code ######
|
||||
|
||||
|
||||
def replace_linear_weight_only_int8_per_channel(module):
|
||||
for name, child in module.named_children():
|
||||
if isinstance(child, nn.Linear):
|
||||
setattr(
|
||||
module,
|
||||
name,
|
||||
WeightOnlyInt8Linear(child.in_features, child.out_features),
|
||||
)
|
||||
else:
|
||||
replace_linear_weight_only_int8_per_channel(child)
|
||||
|
||||
|
||||
class WeightOnlyInt8QuantHandler:
|
||||
def __init__(self, mod):
|
||||
self.mod = mod
|
||||
|
||||
@torch.no_grad()
|
||||
def create_quantized_state_dict(self):
|
||||
cur_state_dict = self.mod.state_dict()
|
||||
for fqn, mod in self.mod.named_modules():
|
||||
if isinstance(mod, torch.nn.Linear):
|
||||
int8_weight, scales, _ = dynamically_quantize_per_channel(
|
||||
mod.weight.float(), -128, 127, torch.int8
|
||||
)
|
||||
cur_state_dict[f"{fqn}.weight"] = int8_weight
|
||||
cur_state_dict[f"{fqn}.scales"] = scales.to(mod.weight.dtype)
|
||||
|
||||
return cur_state_dict
|
||||
|
||||
def convert_for_runtime(self):
|
||||
replace_linear_weight_only_int8_per_channel(self.mod)
|
||||
return self.mod
|
||||
|
||||
|
||||
class WeightOnlyInt8Linear(torch.nn.Module):
|
||||
__constants__ = ["in_features", "out_features"]
|
||||
in_features: int
|
||||
out_features: int
|
||||
weight: torch.Tensor
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
bias: bool = True,
|
||||
device=None,
|
||||
dtype=None,
|
||||
) -> None:
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
self.register_buffer(
|
||||
"weight", torch.empty((out_features, in_features), dtype=torch.int8)
|
||||
)
|
||||
self.register_buffer("scales", torch.ones(out_features, dtype=torch.bfloat16))
|
||||
|
||||
def forward(self, input: torch.Tensor) -> torch.Tensor:
|
||||
return F.linear(input, self.weight.to(dtype=input.dtype)) * self.scales
|
||||
|
||||
|
||||
##### weight only int4 per channel groupwise quantized code ######
|
||||
|
||||
|
||||
def prepare_int4_weight_and_scales_and_zeros(weight_bf16, groupsize, inner_k_tiles):
|
||||
weight_int32, scales_and_zeros = group_quantize_tensor(
|
||||
weight_bf16, n_bit=4, groupsize=groupsize
|
||||
)
|
||||
weight_int4pack = torch.ops.aten._convert_weight_to_int4pack(
|
||||
weight_int32, inner_k_tiles
|
||||
)
|
||||
return weight_int4pack, scales_and_zeros
|
||||
|
||||
|
||||
def linear_forward_int4(x, weight_int4pack, scales_and_zeros, out_features, groupsize):
|
||||
origin_x_size = x.size()
|
||||
x = x.reshape(-1, origin_x_size[-1])
|
||||
c = torch.ops.aten._weight_int4pack_mm(
|
||||
x, weight_int4pack, groupsize, scales_and_zeros
|
||||
)
|
||||
new_shape = origin_x_size[:-1] + (out_features,)
|
||||
c = c.reshape(new_shape)
|
||||
return c
|
||||
|
||||
|
||||
def _check_linear_int4_k(k, groupsize=1, inner_k_tiles=1):
|
||||
return k % groupsize == 0 and k % (inner_k_tiles * 16) == 0
|
||||
|
||||
|
||||
def replace_linear_int4(module, groupsize, inner_k_tiles, padding):
|
||||
for name, child in module.named_children():
|
||||
if isinstance(child, nn.Linear):
|
||||
if _check_linear_int4_k(child.in_features, groupsize, inner_k_tiles):
|
||||
setattr(
|
||||
module,
|
||||
name,
|
||||
WeightOnlyInt4Linear(
|
||||
child.in_features,
|
||||
child.out_features,
|
||||
bias=False,
|
||||
groupsize=groupsize,
|
||||
inner_k_tiles=inner_k_tiles,
|
||||
padding=False,
|
||||
),
|
||||
)
|
||||
elif padding:
|
||||
setattr(
|
||||
module,
|
||||
name,
|
||||
WeightOnlyInt4Linear(
|
||||
child.in_features,
|
||||
child.out_features,
|
||||
bias=False,
|
||||
groupsize=groupsize,
|
||||
inner_k_tiles=inner_k_tiles,
|
||||
padding=True,
|
||||
),
|
||||
)
|
||||
else:
|
||||
replace_linear_int4(child, groupsize, inner_k_tiles, padding)
|
||||
|
||||
|
||||
class WeightOnlyInt4QuantHandler:
|
||||
def __init__(self, mod, groupsize=128, inner_k_tiles=8, padding=True):
|
||||
self.mod = mod
|
||||
self.groupsize = groupsize
|
||||
self.inner_k_tiles = inner_k_tiles
|
||||
self.padding = padding
|
||||
assert groupsize in [32, 64, 128, 256]
|
||||
assert inner_k_tiles in [2, 4, 8]
|
||||
|
||||
@torch.no_grad()
|
||||
def create_quantized_state_dict(self):
|
||||
cur_state_dict = self.mod.state_dict()
|
||||
for fqn, mod in self.mod.named_modules():
|
||||
if isinstance(mod, torch.nn.Linear):
|
||||
assert not mod.bias
|
||||
out_features = mod.out_features
|
||||
in_features = mod.in_features
|
||||
assert out_features % 8 == 0, "require out_features % 8 == 0"
|
||||
print(f"linear: {fqn}, in={in_features}, out={out_features}")
|
||||
|
||||
weight = mod.weight.data
|
||||
if not _check_linear_int4_k(
|
||||
in_features, self.groupsize, self.inner_k_tiles
|
||||
):
|
||||
if self.padding:
|
||||
import torch.nn.functional as F
|
||||
|
||||
print(
|
||||
f"warning: {fqn} is padded to satisfy in_features % 1024 == 0"
|
||||
)
|
||||
padded_in_features = find_multiple(in_features, 1024)
|
||||
weight = F.pad(
|
||||
weight, pad=(0, padded_in_features - in_features)
|
||||
)
|
||||
else:
|
||||
print(
|
||||
f"warning: {fqn} is skipped, int4 requires that in_features is 32, 64, or is divisible by 1024, "
|
||||
+ "and that groupsize and inner_k_tiles*16 evenly divide into it"
|
||||
)
|
||||
continue
|
||||
(
|
||||
weight_int4pack,
|
||||
scales_and_zeros,
|
||||
) = prepare_int4_weight_and_scales_and_zeros(
|
||||
weight.to(torch.bfloat16).to("cuda"),
|
||||
self.groupsize,
|
||||
self.inner_k_tiles,
|
||||
)
|
||||
cur_state_dict[f"{fqn}.weight"] = weight_int4pack.to("cpu")
|
||||
cur_state_dict[f"{fqn}.scales_and_zeros"] = scales_and_zeros.to("cpu")
|
||||
|
||||
return cur_state_dict
|
||||
|
||||
def convert_for_runtime(self):
|
||||
replace_linear_int4(self.mod, self.groupsize, self.inner_k_tiles, self.padding)
|
||||
return self.mod
|
||||
|
||||
|
||||
class WeightOnlyInt4Linear(torch.nn.Module):
|
||||
__constants__ = ["in_features", "out_features"]
|
||||
in_features: int
|
||||
out_features: int
|
||||
weight: torch.Tensor
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
bias=True,
|
||||
device=None,
|
||||
dtype=None,
|
||||
groupsize: int = 128,
|
||||
inner_k_tiles: int = 8,
|
||||
padding: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.padding = padding
|
||||
if padding:
|
||||
self.origin_in_features = in_features
|
||||
in_features = find_multiple(in_features, 1024)
|
||||
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
assert not bias, "require bias=False"
|
||||
self.groupsize = groupsize
|
||||
self.inner_k_tiles = inner_k_tiles
|
||||
|
||||
assert out_features % 8 == 0, "require out_features % 8 == 0"
|
||||
assert (
|
||||
in_features % (inner_k_tiles * 16) == 0
|
||||
), "require in_features % (innerKTiles * 16) == 0"
|
||||
self.register_buffer(
|
||||
"weight",
|
||||
torch.empty(
|
||||
(
|
||||
out_features // 8,
|
||||
in_features // (inner_k_tiles * 16),
|
||||
32,
|
||||
inner_k_tiles // 2,
|
||||
),
|
||||
dtype=torch.int32,
|
||||
),
|
||||
)
|
||||
self.register_buffer(
|
||||
"scales_and_zeros",
|
||||
torch.empty(
|
||||
(in_features // groupsize, out_features, 2), dtype=torch.bfloat16
|
||||
),
|
||||
)
|
||||
|
||||
def forward(self, input: torch.Tensor) -> torch.Tensor:
|
||||
input = input.to(torch.bfloat16)
|
||||
if self.padding:
|
||||
import torch.nn.functional as F
|
||||
|
||||
input = F.pad(input, pad=(0, self.in_features - self.origin_in_features))
|
||||
return linear_forward_int4(
|
||||
input, self.weight, self.scales_and_zeros, self.out_features, self.groupsize
|
||||
)
|
||||
|
||||
|
||||
def generate_folder_name():
|
||||
now = datetime.datetime.now()
|
||||
folder_name = now.strftime("%Y%m%d_%H%M%S")
|
||||
return folder_name
|
||||
|
||||
|
||||
@click.command()
|
||||
@click.option(
|
||||
"--checkpoint-path",
|
||||
type=click.Path(path_type=Path, exists=True),
|
||||
default="checkpoints/fish-speech-1.4",
|
||||
)
|
||||
@click.option(
|
||||
"--mode", type=str, default="int8", help="type of quantization to perform"
|
||||
)
|
||||
@click.option(
|
||||
"--groupsize", type=int, default=128, help="Group size for int4 quantization."
|
||||
)
|
||||
@click.option("--timestamp", type=str, default="None", help="When to do quantization")
|
||||
def quantize(checkpoint_path: Path, mode: str, groupsize: int, timestamp: str) -> None:
|
||||
|
||||
device = "cpu"
|
||||
precision = torch.bfloat16
|
||||
|
||||
print("Loading model ...")
|
||||
t0 = time.time()
|
||||
|
||||
model, _ = load_model(
|
||||
checkpoint_path=checkpoint_path,
|
||||
device=device,
|
||||
precision=precision,
|
||||
compile=False,
|
||||
)
|
||||
vq_model = "firefly-gan-vq-fsq-8x1024-21hz-generator.pth"
|
||||
now = timestamp if timestamp != "None" else generate_folder_name()
|
||||
|
||||
if mode == "int8":
|
||||
print(
|
||||
"Quantizing model weights for int8 weight-only symmetric per-channel quantization"
|
||||
)
|
||||
quant_handler = WeightOnlyInt8QuantHandler(model)
|
||||
quantized_state_dict = quant_handler.create_quantized_state_dict()
|
||||
|
||||
dir_name = checkpoint_path
|
||||
dst_name = Path(f"checkpoints/fs-1.2-int8-{now}")
|
||||
shutil.copytree(str(dir_name.resolve()), str(dst_name.resolve()))
|
||||
if (dst_name / vq_model).exists():
|
||||
(dst_name / vq_model).unlink()
|
||||
quantize_path = dst_name / "model.pth"
|
||||
|
||||
elif mode == "int4":
|
||||
print(
|
||||
"Quantizing model weights for int4 weight-only affine per-channel groupwise quantization"
|
||||
)
|
||||
quant_handler = WeightOnlyInt4QuantHandler(model, groupsize)
|
||||
quantized_state_dict = quant_handler.create_quantized_state_dict()
|
||||
|
||||
dir_name = checkpoint_path
|
||||
dst_name = Path(f"checkpoints/fs-1.2-int4-g{groupsize}-{now}")
|
||||
shutil.copytree(str(dir_name.resolve()), str(dst_name.resolve()))
|
||||
if (dst_name / vq_model).exists():
|
||||
(dst_name / vq_model).unlink()
|
||||
quantize_path = dst_name / "model.pth"
|
||||
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid quantization mode {mode} needs to be one of [int8, int4, int4-gpptq]"
|
||||
)
|
||||
|
||||
print(f"Writing quantized weights to {quantize_path}")
|
||||
quantize_path.unlink(missing_ok=True) # remove existing file if one already there
|
||||
torch.save(quantized_state_dict, quantize_path)
|
||||
print(f"Quantization complete took {time.time() - t0:.02f} seconds")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
quantize()
|
||||
@@ -0,0 +1,57 @@
|
||||
from tokenizers import Tokenizer, decoders, models, pre_tokenizers, processors, trainers
|
||||
from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast
|
||||
|
||||
# Initialize a tokenizer
|
||||
tokenizer = Tokenizer(models.BPE())
|
||||
|
||||
# Customize pre-tokenization and decoding
|
||||
tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
|
||||
tokenizer.decoder = decoders.ByteLevel()
|
||||
tokenizer.post_processor = processors.ByteLevel(trim_offsets=False)
|
||||
|
||||
# Don't train the tokenizer
|
||||
trainer = trainers.BpeTrainer(
|
||||
vocab_size=0,
|
||||
min_frequency=2,
|
||||
initial_alphabet=pre_tokenizers.ByteLevel.alphabet(),
|
||||
special_tokens=[
|
||||
"<|begin_of_sequence|>",
|
||||
"<|end_of_sequence|>",
|
||||
"<|im_start|>",
|
||||
"<|im_sep|>", # system, user, assistant, etc.
|
||||
"<|im_end|>",
|
||||
"<|semantic|>", # audio features
|
||||
"<|pad|>",
|
||||
],
|
||||
)
|
||||
|
||||
# <|im_start|>user<|im_sep|>...<|im_end|>
|
||||
# <|im_start|>assistant<|im_sep|><|semantic|><|semantic|><|semantic|><|semantic|><|semantic|><|im_end|>
|
||||
tokenizer.train_from_iterator([], trainer=trainer)
|
||||
|
||||
print(len(tokenizer.get_vocab()))
|
||||
x = tokenizer.encode(
|
||||
"Hello, how are you? dfgnviadfjoiviouajeiodfjv 你好世界 🈶<|semantic|>"
|
||||
).ids
|
||||
print(x, len(x))
|
||||
print(tokenizer.decode(x, skip_special_tokens=True))
|
||||
|
||||
|
||||
tokenizer = PreTrainedTokenizerFast(
|
||||
tokenizer_object=tokenizer,
|
||||
pad_token="<|pad|>",
|
||||
bos_token="<|begin_of_sequence|>",
|
||||
eos_token="<|end_of_sequence|>",
|
||||
)
|
||||
|
||||
# Try tokenizing a new sequence
|
||||
sequence = "All around, too, lay vast quantities of the costliest merchandise, and treasures were heaped in every cranny of the rocks, but all these things only added to the desolation of the scene. 测试中文, 你好世界 🈶<|semantic|>"
|
||||
encoded = tokenizer(sequence).input_ids
|
||||
|
||||
print("Test encoding....")
|
||||
print(f"\tSentence: {sequence}")
|
||||
print(f"\tEncoded: {encoded}")
|
||||
print(f"\tDecoded: {tokenizer.batch_decode(encoded)}")
|
||||
print(f"\tDecoded: {tokenizer.decode(encoded)}")
|
||||
|
||||
tokenizer.push_to_hub("fishaudio/fish-speech-1", private=True)
|
||||
@@ -0,0 +1,23 @@
|
||||
from .braceexpand import braceexpand
|
||||
from .context import autocast_exclude_mps
|
||||
from .file import get_latest_checkpoint
|
||||
from .instantiators import instantiate_callbacks, instantiate_loggers
|
||||
from .logger import RankedLogger
|
||||
# from .logging_utils import log_hyperparameters
|
||||
from .rich_utils import enforce_tags, print_config_tree
|
||||
from .utils import extras, get_metric_value, task_wrapper
|
||||
|
||||
__all__ = [
|
||||
"enforce_tags",
|
||||
"extras",
|
||||
"get_metric_value",
|
||||
"RankedLogger",
|
||||
"instantiate_callbacks",
|
||||
"instantiate_loggers",
|
||||
# "log_hyperparameters",
|
||||
"print_config_tree",
|
||||
"task_wrapper",
|
||||
"braceexpand",
|
||||
"get_latest_checkpoint",
|
||||
"autocast_exclude_mps",
|
||||
]
|
||||
@@ -0,0 +1,217 @@
|
||||
"""
|
||||
Bash-style brace expansion
|
||||
Copied from: https://github.com/trendels/braceexpand/blob/main/src/braceexpand/__init__.py
|
||||
License: MIT
|
||||
"""
|
||||
|
||||
import re
|
||||
import string
|
||||
from itertools import chain, product
|
||||
from typing import Iterable, Iterator, Optional
|
||||
|
||||
__all__ = ["braceexpand", "alphabet", "UnbalancedBracesError"]
|
||||
|
||||
|
||||
class UnbalancedBracesError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
alphabet = string.ascii_uppercase + string.ascii_lowercase
|
||||
|
||||
int_range_re = re.compile(r"^(-?\d+)\.\.(-?\d+)(?:\.\.-?(\d+))?$")
|
||||
char_range_re = re.compile(r"^([A-Za-z])\.\.([A-Za-z])(?:\.\.-?(\d+))?$")
|
||||
escape_re = re.compile(r"\\(.)")
|
||||
|
||||
|
||||
def braceexpand(pattern: str, escape: bool = True) -> Iterator[str]:
|
||||
"""braceexpand(pattern) -> iterator over generated strings
|
||||
|
||||
Returns an iterator over the strings resulting from brace expansion
|
||||
of pattern. This function implements Brace Expansion as described in
|
||||
bash(1), with the following limitations:
|
||||
|
||||
* A pattern containing unbalanced braces will raise an
|
||||
UnbalancedBracesError exception. In bash, unbalanced braces will either
|
||||
be partly expanded or ignored.
|
||||
|
||||
* A mixed-case character range like '{Z..a}' or '{a..Z}' will not
|
||||
include the characters '[]^_`' between 'Z' and 'a'.
|
||||
|
||||
When escape is True (the default), characters in pattern can be
|
||||
prefixed with a backslash to cause them not to be interpreted as
|
||||
special characters for brace expansion (such as '{', '}', ',').
|
||||
To pass through a a literal backslash, double it ('\\\\').
|
||||
|
||||
When escape is False, backslashes in pattern have no special
|
||||
meaning and will be preserved in the output.
|
||||
|
||||
Examples:
|
||||
|
||||
>>> from braceexpand import braceexpand
|
||||
|
||||
# Integer range
|
||||
>>> list(braceexpand('item{1..3}'))
|
||||
['item1', 'item2', 'item3']
|
||||
|
||||
# Character range
|
||||
>>> list(braceexpand('{a..c}'))
|
||||
['a', 'b', 'c']
|
||||
|
||||
# Sequence
|
||||
>>> list(braceexpand('index.html{,.backup}'))
|
||||
['index.html', 'index.html.backup']
|
||||
|
||||
# Nested patterns
|
||||
>>> list(braceexpand('python{2.{5..7},3.{2,3}}'))
|
||||
['python2.5', 'python2.6', 'python2.7', 'python3.2', 'python3.3']
|
||||
|
||||
# Prefixing an integer with zero causes all numbers to be padded to
|
||||
# the same width.
|
||||
>>> list(braceexpand('{07..10}'))
|
||||
['07', '08', '09', '10']
|
||||
|
||||
# An optional increment can be specified for ranges.
|
||||
>>> list(braceexpand('{a..g..2}'))
|
||||
['a', 'c', 'e', 'g']
|
||||
|
||||
# Ranges can go in both directions.
|
||||
>>> list(braceexpand('{4..1}'))
|
||||
['4', '3', '2', '1']
|
||||
|
||||
# Numbers can be negative
|
||||
>>> list(braceexpand('{2..-1}'))
|
||||
['2', '1', '0', '-1']
|
||||
|
||||
# Unbalanced braces raise an exception.
|
||||
>>> list(braceexpand('{1{2,3}'))
|
||||
Traceback (most recent call last):
|
||||
...
|
||||
UnbalancedBracesError: Unbalanced braces: '{1{2,3}'
|
||||
|
||||
# By default, the backslash is the escape character.
|
||||
>>> list(braceexpand(r'{1\\{2,3}'))
|
||||
['1{2', '3']
|
||||
|
||||
# Setting 'escape' to False disables backslash escaping.
|
||||
>>> list(braceexpand(r'\\{1,2}', escape=False))
|
||||
['\\\\1', '\\\\2']
|
||||
|
||||
"""
|
||||
return (
|
||||
escape_re.sub(r"\1", s) if escape else s for s in parse_pattern(pattern, escape)
|
||||
)
|
||||
|
||||
|
||||
def parse_pattern(pattern: str, escape: bool) -> Iterator[str]:
|
||||
start = 0
|
||||
pos = 0
|
||||
bracketdepth = 0
|
||||
items: list[Iterable[str]] = []
|
||||
|
||||
# print 'pattern:', pattern
|
||||
while pos < len(pattern):
|
||||
if escape and pattern[pos] == "\\":
|
||||
pos += 2
|
||||
continue
|
||||
elif pattern[pos] == "{":
|
||||
if bracketdepth == 0 and pos > start:
|
||||
# print 'literal:', pattern[start:pos]
|
||||
items.append([pattern[start:pos]])
|
||||
start = pos
|
||||
bracketdepth += 1
|
||||
elif pattern[pos] == "}":
|
||||
bracketdepth -= 1
|
||||
if bracketdepth == 0:
|
||||
# print 'expression:', pattern[start+1:pos]
|
||||
expr = pattern[start + 1 : pos]
|
||||
item = parse_expression(expr, escape)
|
||||
if item is None: # not a range or sequence
|
||||
items.extend([["{"], parse_pattern(expr, escape), ["}"]])
|
||||
else:
|
||||
items.append(item)
|
||||
start = pos + 1 # skip the closing brace
|
||||
pos += 1
|
||||
|
||||
if bracketdepth != 0: # unbalanced braces
|
||||
raise UnbalancedBracesError("Unbalanced braces: '%s'" % pattern)
|
||||
|
||||
if start < pos:
|
||||
items.append([pattern[start:]])
|
||||
|
||||
return ("".join(item) for item in product(*items))
|
||||
|
||||
|
||||
def parse_expression(expr: str, escape: bool) -> Optional[Iterable[str]]:
|
||||
int_range_match = int_range_re.match(expr)
|
||||
if int_range_match:
|
||||
return make_int_range(*int_range_match.groups())
|
||||
|
||||
char_range_match = char_range_re.match(expr)
|
||||
if char_range_match:
|
||||
return make_char_range(*char_range_match.groups())
|
||||
|
||||
return parse_sequence(expr, escape)
|
||||
|
||||
|
||||
def parse_sequence(seq: str, escape: bool) -> Optional[Iterator[str]]:
|
||||
# sequence -> chain(*sequence_items)
|
||||
start = 0
|
||||
pos = 0
|
||||
bracketdepth = 0
|
||||
items: list[Iterable[str]] = []
|
||||
|
||||
# print 'sequence:', seq
|
||||
while pos < len(seq):
|
||||
if escape and seq[pos] == "\\":
|
||||
pos += 2
|
||||
continue
|
||||
elif seq[pos] == "{":
|
||||
bracketdepth += 1
|
||||
elif seq[pos] == "}":
|
||||
bracketdepth -= 1
|
||||
elif seq[pos] == "," and bracketdepth == 0:
|
||||
items.append(parse_pattern(seq[start:pos], escape))
|
||||
start = pos + 1 # skip the comma
|
||||
pos += 1
|
||||
|
||||
if bracketdepth != 0:
|
||||
raise UnbalancedBracesError
|
||||
if not items:
|
||||
return None
|
||||
|
||||
# part after the last comma (may be the empty string)
|
||||
items.append(parse_pattern(seq[start:], escape))
|
||||
return chain(*items)
|
||||
|
||||
|
||||
def make_int_range(left: str, right: str, incr: Optional[str] = None) -> Iterator[str]:
|
||||
if any([s.startswith(("0", "-0")) for s in (left, right) if s not in ("0", "-0")]):
|
||||
padding = max(len(left), len(right))
|
||||
else:
|
||||
padding = 0
|
||||
step = (int(incr) or 1) if incr else 1
|
||||
start = int(left)
|
||||
end = int(right)
|
||||
r = range(start, end + 1, step) if start < end else range(start, end - 1, -step)
|
||||
fmt = "%0{}d".format(padding)
|
||||
return (fmt % i for i in r)
|
||||
|
||||
|
||||
def make_char_range(left: str, right: str, incr: Optional[str] = None) -> str:
|
||||
step = (int(incr) or 1) if incr else 1
|
||||
start = alphabet.index(left)
|
||||
end = alphabet.index(right)
|
||||
if start < end:
|
||||
return alphabet[start : end + 1 : step]
|
||||
else:
|
||||
end = end or -len(alphabet)
|
||||
return alphabet[start : end - 1 : -step]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import doctest
|
||||
import sys
|
||||
|
||||
failed, _ = doctest.testmod(optionflags=doctest.IGNORE_EXCEPTION_DETAIL)
|
||||
if failed:
|
||||
sys.exit(1)
|
||||