Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3bac87ee52 | ||
|
|
9a01701019 | ||
|
|
693954ee23 | ||
|
|
bf4ba91e7a | ||
|
|
0828353253 | ||
|
|
1997c7ad8f | ||
|
|
b1e62440e4 | ||
|
|
d549a5eb6a | ||
|
|
77bfb08d76 | ||
|
|
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 | ||
|
|
686ebcfd8b | ||
|
|
51aaba39cf | ||
|
|
d7c6632499 | ||
|
|
1b0ea06876 | ||
|
|
8ebe88629b | ||
|
|
e226992703 | ||
|
|
11a8394d69 | ||
|
|
9f54a1b91a | ||
|
|
35d11061e9 | ||
|
|
65d8b490ca | ||
|
|
0fef12c3b1 | ||
|
|
8516bff224 | ||
|
|
1bcc501352 | ||
|
|
f492b17fbe | ||
|
|
48ae90f80e | ||
|
|
9321ccbc48 | ||
|
|
402cd01e1a | ||
|
|
7b2d0e29c6 | ||
|
|
52d38c401a | ||
|
|
a34dd61076 | ||
|
|
a53a3e772a | ||
|
|
8f24c294a7 | ||
|
|
acc3f76654 | ||
|
|
8c977fb442 | ||
|
|
aa20a2de67 | ||
|
|
5fcb154d89 | ||
|
|
0980129f4e | ||
|
|
1258746886 | ||
|
|
ed128b0ad6 | ||
|
|
f0db08acd6 | ||
|
|
7c655e3080 | ||
|
|
1b9871c3df | ||
|
|
7568aaf243 | ||
|
|
a76be8450d | ||
|
|
5564ee1246 | ||
|
|
a6e9251521 | ||
|
|
0bee093916 | ||
|
|
13a9878823 | ||
|
|
acd35d50f8 | ||
|
|
5c0d99e72d | ||
|
|
e37af93be3 | ||
|
|
244c1700e1 | ||
|
|
a3a15473ba | ||
|
|
d734b5077c | ||
|
|
5430072b19 | ||
|
|
465aebaed4 | ||
|
|
ee0b16c2ea | ||
|
|
21d5eacb41 | ||
|
|
c06688eb0b | ||
|
|
b74bbcd279 | ||
|
|
64d8d9b05d | ||
|
|
6ab60f281b | ||
|
|
e35be3b2fa | ||
|
|
0a0c27ac96 | ||
|
|
fb249e84eb | ||
|
|
76ad86fcae | ||
|
|
037614d227 | ||
|
|
4a50e445fd | ||
|
|
6f3c1c4393 | ||
|
|
c6a9b4b592 | ||
|
|
e915ac4eca | ||
|
|
a857793f63 | ||
|
|
31914f7510 | ||
|
|
0be859f0ee | ||
|
|
14b9c3697b | ||
|
|
96b66a57bb | ||
|
|
f13701c489 | ||
|
|
31515b810e | ||
|
|
29e84e08a4 | ||
|
|
0ac9ad9757 | ||
|
|
9a432e0608 | ||
|
|
c83ba5fe7f | ||
|
|
1b55c743ea | ||
|
|
a93579376c | ||
|
|
eba49f3c68 | ||
|
|
3a3da49c69 | ||
|
|
3d68e48219 | ||
|
|
4351fa6a0e | ||
|
|
77222d2808 | ||
|
|
36db7e5a9a | ||
|
|
3572368f16 | ||
|
|
0755dc1462 | ||
|
|
94f81b7102 | ||
|
|
08f8fe3d7e | ||
|
|
a6cd383d67 | ||
|
|
10face6ab0 | ||
|
|
a6259ff600 | ||
|
|
d7d9e6cbfe | ||
|
|
8263609470 | ||
|
|
fc2367de76 | ||
|
|
9a4f2ebc70 | ||
|
|
812879610a | ||
|
|
c5e521ccc1 | ||
|
|
bc1c8fa351 | ||
|
|
7271fcf9c1 | ||
|
|
ace3b7707b | ||
|
|
9eb65cc4ee | ||
|
|
c5c2bc779c | ||
|
|
ae1751d9c0 | ||
|
|
d17583ef7d | ||
|
|
a363713ae0 | ||
|
|
1566165bd4 | ||
|
|
fe065fa318 | ||
|
|
202d5cf071 | ||
|
|
b785a9dc5b | ||
|
|
73bc658b2f | ||
|
|
0b94216138 | ||
|
|
86eec2b4cc | ||
|
|
c99b531d28 | ||
|
|
74a4338cb5 | ||
|
|
d691c52e49 | ||
|
|
bd3e9e4b3c | ||
|
|
3c3ca5fb9c | ||
|
|
e4ff4fce1c | ||
|
|
2918d4b07d | ||
|
|
c40e49be46 | ||
|
|
4ccda20975 | ||
|
|
09957617d3 | ||
|
|
c703aa7058 | ||
|
|
64d366d323 | ||
|
|
333e0a2faa | ||
|
|
a677d95bc8 | ||
|
|
aafd87e84b | ||
|
|
efa3bae54b | ||
|
|
8e4362689d | ||
|
|
b86634284e | ||
|
|
b3293ddccd | ||
|
|
4b40831b83 | ||
|
|
e426b0521c | ||
|
|
1777bf6e06 | ||
|
|
dafc892f0f | ||
|
|
fea0cfd5dd | ||
|
|
a7e158db6d | ||
|
|
455ac4abd3 | ||
|
|
778dfa2cf5 | ||
|
|
329f2e6f81 | ||
|
|
63ad6d97d7 | ||
|
|
8db56db7cf | ||
|
|
bd542f1e0b | ||
|
|
a162e53dea | ||
|
|
a8a4c848ed | ||
|
|
bddd38996a | ||
|
|
9f084eae94 | ||
|
|
c54c635161 | ||
|
|
fa8b42e05e | ||
|
|
9c3c323884 | ||
|
|
c8a46439be | ||
|
|
b7ec701259 | ||
|
|
a2cd0e0a38 | ||
|
|
7bccc0e236 | ||
|
|
a3a649a79f | ||
|
|
2a14d30552 | ||
|
|
313bef0609 | ||
|
|
960a80aeca | ||
|
|
20e6d50a98 | ||
|
|
507c3417d6 | ||
|
|
d1adb8d4ed | ||
|
|
915ff12747 | ||
|
|
fbade79137 | ||
|
|
a240a677d0 | ||
|
|
ab64cf31f6 | ||
|
|
67f2e32dae | ||
|
|
0dc40fe052 | ||
|
|
b7225be552 | ||
|
|
629e00ec94 | ||
|
|
6d033c9314 | ||
|
|
4f8926ed00 | ||
|
|
1daa1a4603 | ||
|
|
062773d929 | ||
|
|
57decadaef | ||
|
|
ec804ab7c9 | ||
|
|
3ee7533098 | ||
|
|
0d383ccc1f | ||
|
|
e2f2257c34 | ||
|
|
d62b9fc4c6 | ||
|
|
be7ad0c7fb | ||
|
|
156864cc8b | ||
|
|
7240e496cc | ||
|
|
8acdf4018d | ||
|
|
1ea7c3e203 | ||
|
|
e54aeb6125 | ||
|
|
d59f51fbcf | ||
|
|
a667eb6982 | ||
|
|
8d72732247 | ||
|
|
f0f3b30a62 | ||
|
|
d99fe24542 | ||
|
|
ea4c7381bd | ||
|
|
e900d20641 | ||
|
|
8a46647d8c | ||
|
|
968178bf57 | ||
|
|
406a255db0 | ||
|
|
574557810e | ||
|
|
998a02c3a4 | ||
|
|
c6f964c921 | ||
|
|
efb0e147c5 | ||
|
|
9cf7356f98 | ||
|
|
af05c43174 | ||
|
|
380c68ff2b | ||
|
|
068b00b99f | ||
|
|
cd6a42ab64 | ||
|
|
f115abec92 | ||
|
|
f0e23cf878 | ||
|
|
a94f11d809 | ||
|
|
38972bea5f | ||
|
|
a761ff552a | ||
|
|
dbeb84ea9a | ||
|
|
0a4938f39a | ||
|
|
b2182c716d | ||
|
|
8253be73f6 | ||
|
|
20318e296e | ||
|
|
cea1b69286 | ||
|
|
ab8aa69389 | ||
|
|
681491f1d0 | ||
|
|
c2fb815074 | ||
|
|
fffa14dc44 | ||
|
|
df37166d42 | ||
|
|
25fa3a8f6a | ||
|
|
b11507c5e6 | ||
|
|
f06d02489f | ||
|
|
c203af2f71 | ||
|
|
c715155a70 | ||
|
|
8163133294 | ||
|
|
765be5dab4 | ||
|
|
7c1523389d | ||
|
|
7a2b1ba166 | ||
|
|
45b4dcfcd0 | ||
|
|
3b2e535566 | ||
|
|
db556d13a3 | ||
|
|
a987063c68 | ||
|
|
4ce30ef899 | ||
|
|
d988282d98 | ||
|
|
695fdf7ceb | ||
|
|
0befe164cc | ||
|
|
d506c68a80 | ||
|
|
c59c429b75 | ||
|
|
c726b6e4a2 | ||
|
|
96075ad4e1 | ||
|
|
7a8dc07a8a | ||
|
|
1fdac0bc09 | ||
|
|
804b942a36 | ||
|
|
6335d4378b | ||
|
|
cfc2189616 |
@@ -0,0 +1,21 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
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,14 +1,52 @@
|
||||
> 适配了最新版comfyui的py3.11 ,torch 2.1.2+cu121
|
||||

|
||||
|
||||
## 🚀🚗🚚🏃 Workflow-to-APP
|
||||
- 新增AppInfo节点,可以通过简单的配置,把workflow转变为一个Web APP。
|
||||
- 支持多个web app 切换
|
||||
- 发布为app的workflow,可以在右键里再次编辑了
|
||||
> 适配了最新版 comfyui 的 py3.11 ,torch 2.1.2+cu121
|
||||
> [Mixlab nodes discord](https://discord.gg/cXs9vZSqeK)
|
||||
|
||||
|
||||
##### `最新`:
|
||||
|
||||
ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。模型下载后,放置到 `models/llamafile/`
|
||||
|
||||
- 右键菜单支持 text-to-text,方便对 prompt 词补全
|
||||
|
||||
强烈推荐:
|
||||
[Phi-3-mini-4k-instruct-function-calling-GGUF](https://huggingface.co/nold/Phi-3-mini-4k-instruct-function-calling-GGUF)
|
||||
|
||||
[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-sd-prompt-mixlab](https://github.com/shadowcz007/comfyui-sd-prompt-mixlab) -->
|
||||
|
||||
[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) -->
|
||||
|
||||
## 🚀🚗🚚🏃 Workflow-to-APP
|
||||
|
||||
- 新增 AppInfo 节点,可以通过简单的配置,把 workflow 转变为一个 Web APP。
|
||||
- 支持多个 web app 切换
|
||||
- 发布为 app 的 workflow,可以在右键里再次编辑了
|
||||
- web app 可以设置分类,在 comfyui 右键菜单可以编辑更新 web app
|
||||
- 支持动态提示
|
||||
|
||||

|
||||
|
||||
- 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.
|
||||
|
||||

|
||||
|
||||
@@ -17,115 +55,191 @@
|
||||

|
||||
|
||||
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
|
||||
|
||||
> 暂时支持6种节点作为界面上的输入节点:Load Image、CLIPTextEncode、TextInput_、FloatSlider、IntNumber、CheckpointLoaderSimple、LoraLoader
|
||||
> 暂时支持 9 种节点作为界面上的输入节点:Load Image、VHS*LoadVideo、CLIPTextEncode、PromptSlide、TextInput*、Color、FloatSlider、IntNumber、CheckpointLoaderSimple、LoraLoader
|
||||
|
||||
> 输出节点:PreviewImage 、SaveImage、ShowTextForGPT、VHS_VideoCombine
|
||||
> 输出节点:PreviewImage 、SaveImage、ShowTextForGPT、VHS_VideoCombine、PromptImage
|
||||
|
||||
> seed 统一输入控件,支持:SamplerCustom、KSampler
|
||||
|
||||
> 配套[ps 插件](https://github.com/shadowcz007/comfyui-ps-plugin)
|
||||
|
||||
> 如果遇到上传图片不成功,请检查下:局域网或者是云服务,请使用 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)
|
||||
|
||||
### GPT
|
||||
> Support for calling multiple GPTs.ChatGPT、ChatGLM3 , 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(llama.cpp)、 ChatGPT、ChatGLM3 、ChatGLM4 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1 . Azure OpenAI:https://xxxx.openai.azure.com
|
||||
|
||||

|
||||
|
||||
[workflow-5](./workflow/5-gpt-workflow.json)
|
||||
|
||||
### 3D
|
||||

|
||||
[workflow](./workflow/3D-workflow.json)
|
||||
最新: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
|
||||
> 
|
||||
|
||||
<!--  -->
|
||||
|
||||
> randomPrompt
|
||||
|
||||

|
||||
|
||||
> ClipInterrogator
|
||||
|
||||
[add clip-interrogator](https://github.com/pharmapsychotic/clip-interrogator)
|
||||
|
||||
> PromptImage & PromptSimplification,Assist in simplifying prompt words, comparing images and prompt word nodes.
|
||||
|
||||
> 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.
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
### 3D
|
||||
|
||||
### 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](./assets/Image-to-3D_1.json)
|
||||
|
||||

|
||||
[workflow](./workflow/3D-workflow.json)
|
||||
|
||||
### Image
|
||||
|
||||
#### 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.
|
||||
|
||||
### 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)
|
||||
|
||||
## Other Nodes
|
||||
- [使用 CkptNames 对比不同的模型效果](./workflow/ckpts-image-workflow.json)
|
||||
|
||||
- [CkptNames compare the effects of different models.](./workflow/ckpts-image-workflow.json)
|
||||
|
||||
### Other Nodes
|
||||
|
||||

|
||||

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

|
||||
|
||||
> TransparentImage
|
||||
|
||||

|
||||
|
||||
|
||||
> Consistency Decoder
|
||||
|
||||
[openai Consistency Decoder]( https://github.com/openai/consistencydecoder)
|
||||
|
||||

|
||||
After downloading the OpenAI VAE model, place it in the "model/vae" directory for use.
|
||||
https://openaipublic.azureedge.net/diff-vae/c9cebd3132dd9c42936d803e33424145a748843c8f716c0814838bdc8a2fe7cb/decoder.pt
|
||||
|
||||
|
||||
> FeatheredMask、SmoothMask
|
||||
|
||||
Add edges to an image.
|
||||
|
||||

|
||||
|
||||
|
||||
> LaMaInpainting
|
||||
|
||||
from [simple-lama-inpainting](https://github.com/enesmsahin/simple-lama-inpainting)
|
||||
|
||||
> rembgNode
|
||||
|
||||
### Improvement
|
||||
"briarmbg","u2net","u2netp","u2net_human_seg","u2net_cloth_seg","silueta","isnet-general-use","isnet-anime"
|
||||
|
||||
**_ briarmbg _** model was developed by BRlA Al and can be used as an open-source model for non-commercial purposes
|
||||
|
||||
### Improvement
|
||||
|
||||
- Add "help" option to the context menu for each node.
|
||||
- Add "Nodes Map" option to the global context menu.
|
||||
@@ -136,26 +250,21 @@ An improvement has been made to directly redirect to GitHub to search for missin
|
||||
|
||||

|
||||
|
||||
|
||||
### Update
|
||||
v0.8.0 🚀🚗🚚🏃 LaMaInpainting
|
||||
- 新增 LaMaInpainting
|
||||
- 优化color节点的输出
|
||||
- 修复高清显示屏上定位节点不准的情况
|
||||
|
||||
- Add LaMaInpainting
|
||||
- Optimize the output of the color node
|
||||
- Fix the issue of inaccurate positioning node on high-definition display screens
|
||||
|
||||
|
||||
### Models
|
||||
[Download CLIPSeg](https://huggingface.co/CIDAS/clipseg-rd64-refined/tree/main), move to : models/clipseg
|
||||
|
||||
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : models/lama
|
||||
- [Download TripoSR](https://huggingface.co/stabilityai/TripoSR/blob/main/model.ckpt) and place it in `models/triposr`
|
||||
|
||||
<!-- ### Workflow
|
||||
[Workflow](./workflow.md) -->
|
||||
- [Download facebook/dino-vitb16](https://huggingface.co/facebook/dino-vitb16/tree/main) and place it in `models/triposr/facebook/dino-vitb16`
|
||||
|
||||
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:`models/rembg`
|
||||
|
||||
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : `models/lama`
|
||||
|
||||
[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
|
||||
|
||||
@@ -171,36 +280,36 @@ 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 无界社区
|
||||
|
||||
#### Thanks:
|
||||
[ComfyUI-CLIPSeg](https://github.com/biegert/ComfyUI-CLIPSeg/tree/main)
|
||||
####
|
||||
|
||||
File / LoadImagesFromPath SaveImageToLocal LoadImagesFromURL
|
||||
|
||||
#### discussions:
|
||||
|
||||
[discussions](https://github.com/shadowcz007/comfyui-mixlab-nodes/discussions)
|
||||
|
||||
### TODO:
|
||||
- 音频播放节点:带可视化、支持多音轨、可配置音轨音量
|
||||
- vector https://github.com/GeorgLegato/stable-diffusion-webui-vectorstudio
|
||||
|
||||
|
||||
<picture>
|
||||
<source
|
||||
media="(prefers-color-scheme: dark)"
|
||||
@@ -219,4 +328,3 @@ pip3 install -r requirements.txt
|
||||
src="https://api.star-history.com/svg?repos=shadowcz007/comfyui-mixlab-nodes&type=Date"
|
||||
/>
|
||||
</picture>
|
||||
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
#
|
||||
import os
|
||||
import subprocess
|
||||
import importlib.util
|
||||
@@ -6,10 +5,35 @@ import sys,json
|
||||
import urllib
|
||||
import hashlib
|
||||
import datetime
|
||||
|
||||
|
||||
import folder_paths
|
||||
import logging
|
||||
import base64,io,re
|
||||
from PIL import Image
|
||||
from comfy.cli_args import args
|
||||
python = sys.executable
|
||||
|
||||
# print("sys.path", sys.path)
|
||||
|
||||
#修复 sys.stdout.isatty() object has no attribute 'isatty'
|
||||
try:
|
||||
sys.stdout.isatty()
|
||||
except:
|
||||
# print('#fix sys.stdout.isatty')
|
||||
sys.stdout.isatty = lambda: False
|
||||
|
||||
llama_port=None
|
||||
llama_model=""
|
||||
llama_chat_format=""
|
||||
|
||||
try:
|
||||
from .nodes.ChatGPT import get_llama_models,get_llama_model_path,llama_cpp_client
|
||||
llama_cpp_client("")
|
||||
|
||||
except:
|
||||
print("##nodes.ChatGPT ImportError")
|
||||
|
||||
|
||||
from .nodes.RembgNode import get_rembg_models,U2NET_HOME,run_briarmbg,run_rembg
|
||||
|
||||
from server import PromptServer
|
||||
|
||||
@@ -42,7 +66,7 @@ def is_installed(package, package_overwrite=None):
|
||||
print(f"Couldn't install\nCommand: {command}\nError code: {result.returncode}")
|
||||
else:
|
||||
print(package+'## OK')
|
||||
|
||||
|
||||
try:
|
||||
import OpenSSL
|
||||
except ImportError:
|
||||
@@ -79,6 +103,26 @@ install_openai()
|
||||
current_path = os.path.abspath(os.path.dirname(__file__))
|
||||
|
||||
|
||||
def remove_base64_prefix(base64_str):
|
||||
"""
|
||||
去除 base64 字符串中的 data:image/*;base64, 前缀
|
||||
|
||||
Args:
|
||||
base64_str: base64 编码的字符串
|
||||
|
||||
Returns:
|
||||
去除前缀后的 base64 字符串
|
||||
"""
|
||||
|
||||
# 使用正则表达式匹配常见的前缀
|
||||
pattern = r'^data:image\/(.*);base64,(.+)$'
|
||||
match = re.match(pattern, base64_str)
|
||||
if match:
|
||||
# 如果匹配到常见的前缀,则去除前缀并返回
|
||||
return match.group(2)
|
||||
else:
|
||||
# 如果不匹配到常见的前缀,则直接返回
|
||||
return base64_str
|
||||
|
||||
def calculate_md5(string):
|
||||
encoded_string = string.encode()
|
||||
@@ -133,8 +177,38 @@ def create_for_https():
|
||||
return (crt,key)
|
||||
|
||||
|
||||
|
||||
# workflow 目录下的所有json
|
||||
def read_workflow_json_files_all(folder_path):
|
||||
print('#read_workflow_json_files_all',folder_path)
|
||||
json_files = []
|
||||
for root, dirs, files in os.walk(folder_path):
|
||||
for file in files:
|
||||
if file.endswith('.json'):
|
||||
json_files.append(os.path.join(root, file))
|
||||
|
||||
data = []
|
||||
for file_path in json_files:
|
||||
try:
|
||||
with open(file_path) as json_file:
|
||||
json_data = json.load(json_file)
|
||||
creation_time = datetime.datetime.fromtimestamp(os.path.getctime(file_path))
|
||||
numeric_timestamp = creation_time.timestamp()
|
||||
file_info = {
|
||||
'filename': os.path.basename(file_path),
|
||||
'category': os.path.dirname(file_path),
|
||||
'data': json_data,
|
||||
'date': numeric_timestamp
|
||||
}
|
||||
data.append(file_info)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
sorted_data = sorted(data, key=lambda x: x['date'], reverse=True)
|
||||
return sorted_data
|
||||
|
||||
# workflow
|
||||
def read_workflow_json_files(folder_path):
|
||||
def read_workflow_json_files(folder_path ):
|
||||
json_files = []
|
||||
for filename in os.listdir(folder_path):
|
||||
if filename.endswith('.json'):
|
||||
@@ -162,41 +236,68 @@ def read_workflow_json_files(folder_path):
|
||||
def get_workflows():
|
||||
# print("#####path::", current_path)
|
||||
workflow_path=os.path.join(current_path, "workflow")
|
||||
print('workflow_path: ',workflow_path)
|
||||
# print('workflow_path: ',workflow_path)
|
||||
if not os.path.exists(workflow_path):
|
||||
# 使用mkdir()方法创建新目录
|
||||
os.mkdir(workflow_path)
|
||||
workflows=read_workflow_json_files(workflow_path)
|
||||
return workflows
|
||||
|
||||
def get_my_workflow_for_app(filename="my_workflow_app.json"):
|
||||
def get_my_workflow_for_app(filename="my_workflow_app.json",category="",is_all=False):
|
||||
app_path=os.path.join(current_path, "app")
|
||||
if not os.path.exists(app_path):
|
||||
os.mkdir(app_path)
|
||||
|
||||
category_path=os.path.join(app_path,category)
|
||||
if not os.path.exists(category_path):
|
||||
os.mkdir(category_path)
|
||||
|
||||
apps=[]
|
||||
if filename==None:
|
||||
data=read_workflow_json_files(app_path)
|
||||
|
||||
#TODO 支持目录内遍历
|
||||
if is_all:
|
||||
data=read_workflow_json_files_all(category_path)
|
||||
else:
|
||||
data=read_workflow_json_files(category_path)
|
||||
|
||||
i=0
|
||||
for item in data:
|
||||
# print(item)
|
||||
try:
|
||||
x=item["data"]
|
||||
if i==0:
|
||||
# 管理员模式,读取全部数据
|
||||
if i==0 or is_all:
|
||||
apps.append({
|
||||
"filename":item["filename"],
|
||||
# "category":item['category'],
|
||||
"data":x,
|
||||
"date":item["date"]
|
||||
"date":item["date"],
|
||||
})
|
||||
else:
|
||||
category=''
|
||||
input=None
|
||||
output=None
|
||||
if 'category' in x['app']:
|
||||
category=x['app']['category']
|
||||
if 'input' in x['app']:
|
||||
input=x['app']['input']
|
||||
if 'output' in x['app']:
|
||||
output=x['app']['output']
|
||||
apps.append({
|
||||
"filename":item["filename"],
|
||||
"category":category,
|
||||
"data":{
|
||||
"app":{
|
||||
"category":category,
|
||||
"description":x['app']['description'],
|
||||
"filename":(x['app']['filename'] if 'filename' in x['app'] else "") ,
|
||||
"icon":(x['app']['icon'] if 'icon' in x['app'] else None),
|
||||
"name":x['app']['name'],
|
||||
"version":x['app']['version'],
|
||||
"input":input,
|
||||
"output":output,
|
||||
"id":x['app']['id']
|
||||
}
|
||||
},
|
||||
"date":item["date"]
|
||||
@@ -205,8 +306,8 @@ def get_my_workflow_for_app(filename="my_workflow_app.json"):
|
||||
except Exception as e:
|
||||
print("发生异常:", str(e))
|
||||
else:
|
||||
app_workflow_path=os.path.join(app_path, filename)
|
||||
# print('app_workflow_path: ',app_workflow_path)
|
||||
app_workflow_path=os.path.join(category_path, filename)
|
||||
print('app_workflow_path: ',app_workflow_path)
|
||||
try:
|
||||
with open(app_workflow_path) as json_file:
|
||||
apps = [{
|
||||
@@ -216,22 +317,37 @@ def get_my_workflow_for_app(filename="my_workflow_app.json"):
|
||||
except Exception as e:
|
||||
print("发生异常:", str(e))
|
||||
|
||||
if len(apps)==1:
|
||||
data=read_workflow_json_files(app_path)
|
||||
# 这个代码不需要
|
||||
# if len(apps)==1 and category!='' and category!=None:
|
||||
data=read_workflow_json_files(category_path)
|
||||
|
||||
for item in data:
|
||||
x=item["data"]
|
||||
# print(apps[0]['filename'] ,item["filename"])
|
||||
if apps[0]['filename']!=item["filename"]:
|
||||
category=''
|
||||
input=None
|
||||
output=None
|
||||
if 'category' in x['app']:
|
||||
category=x['app']['category']
|
||||
if 'input' in x['app']:
|
||||
input=x['app']['input']
|
||||
if 'output' in x['app']:
|
||||
output=x['app']['output']
|
||||
apps.append({
|
||||
"filename":item["filename"],
|
||||
# "category":category,
|
||||
"data":{
|
||||
"app":{
|
||||
"category":category,
|
||||
"description":x['app']['description'],
|
||||
"filename":(x['app']['filename'] if 'filename' in x['app'] else "") ,
|
||||
"icon":(x['app']['icon'] if 'icon' in x['app'] else None),
|
||||
"name":x['app']['name'],
|
||||
"version":x['app']['version'],
|
||||
"input":input,
|
||||
"output":output,
|
||||
"id":x['app']['id']
|
||||
}
|
||||
},
|
||||
"date":item["date"]
|
||||
@@ -239,17 +355,47 @@ def get_my_workflow_for_app(filename="my_workflow_app.json"):
|
||||
|
||||
return apps
|
||||
|
||||
# 历史记录
|
||||
def save_prompt_result(id,data):
|
||||
prompt_result_path=os.path.join(current_path, "workflow/prompt_result.json")
|
||||
prompt_result={}
|
||||
if os.path.exists(prompt_result_path):
|
||||
with open(prompt_result_path) as json_file:
|
||||
prompt_result = json.load(json_file)
|
||||
|
||||
prompt_result[id]=data
|
||||
|
||||
with open(prompt_result_path, 'w') as file:
|
||||
json.dump(prompt_result, file)
|
||||
return prompt_result_path
|
||||
|
||||
def get_prompt_result():
|
||||
prompt_result_path=os.path.join(current_path, "workflow/prompt_result.json")
|
||||
prompt_result={}
|
||||
if os.path.exists(prompt_result_path):
|
||||
with open(prompt_result_path) as json_file:
|
||||
prompt_result = json.load(json_file)
|
||||
res=list(prompt_result.values())
|
||||
# print(res)
|
||||
return res
|
||||
|
||||
|
||||
def save_workflow_json(data):
|
||||
workflow_path=os.path.join(current_path, "workflow/my_workflow.json")
|
||||
with open(workflow_path, 'w') as file:
|
||||
json.dump(data, file)
|
||||
return workflow_path
|
||||
|
||||
def save_workflow_for_app(data,filename="my_workflow_app.json"):
|
||||
def save_workflow_for_app(data,filename="my_workflow_app.json",category=""):
|
||||
app_path=os.path.join(current_path, "app")
|
||||
if not os.path.exists(app_path):
|
||||
os.mkdir(app_path)
|
||||
app_workflow_path=os.path.join(app_path, filename)
|
||||
|
||||
category_path=os.path.join(app_path,category)
|
||||
if not os.path.exists(category_path):
|
||||
os.mkdir(category_path)
|
||||
|
||||
app_workflow_path=os.path.join(category_path, filename)
|
||||
|
||||
try:
|
||||
output_str = json.dumps(data['output'])
|
||||
@@ -294,32 +440,107 @@ async def new_request(self, method, url, *args, **kwargs):
|
||||
|
||||
# 应用 Monkey Patch
|
||||
aiohttp.ClientSession._request = new_request
|
||||
import socket
|
||||
|
||||
async def check_port_available(address, port):
|
||||
#检查端口是否可用
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
|
||||
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
try:
|
||||
sock.bind((address, port))
|
||||
return True
|
||||
except socket.error:
|
||||
return False
|
||||
|
||||
# https
|
||||
async def new_start(self, address, port, verbose=True, call_on_start=None):
|
||||
|
||||
|
||||
try:
|
||||
runner = web.AppRunner(self.app, access_log=None)
|
||||
await runner.setup()
|
||||
site = web.TCPSite(runner, address, port)
|
||||
await site.start()
|
||||
|
||||
import ssl
|
||||
crt,key=create_for_https()
|
||||
ssl_context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH)
|
||||
ssl_context.load_cert_chain(crt,key)
|
||||
site2 = web.TCPSite(runner, address, port+1,ssl_context=ssl_context)
|
||||
await site2.start()
|
||||
# if not await check_port_available(address, port):
|
||||
# raise RuntimeError(f"Port {port} is already in use.")
|
||||
|
||||
http_success = False
|
||||
http_port=port
|
||||
for i in range(11): # 尝试最多11次
|
||||
if await check_port_available(address, port + i):
|
||||
http_port = port + i
|
||||
site = web.TCPSite(runner, address, http_port)
|
||||
await site.start()
|
||||
http_success = True
|
||||
break
|
||||
|
||||
if not http_success:
|
||||
raise RuntimeError(f"Ports {port} to {port + 10} are all in use.")
|
||||
|
||||
|
||||
# site = web.TCPSite(runner, address, port)
|
||||
# await site.start()
|
||||
|
||||
ssl_context = None
|
||||
scheme = "http"
|
||||
try:
|
||||
# 跟着本体修改
|
||||
if args.tls_keyfile and args.tls_certfile:
|
||||
scheme = "https"
|
||||
ssl_context = ssl.SSLContext(protocol=ssl.PROTOCOL_TLS_SERVER, verify_mode=ssl.CERT_NONE)
|
||||
ssl_context.load_cert_chain(certfile=args.tls_certfile,
|
||||
keyfile=args.tls_keyfile)
|
||||
else:
|
||||
# 如果没传,则自动创建
|
||||
import ssl
|
||||
crt, key = create_for_https()
|
||||
ssl_context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH)
|
||||
ssl_context.load_cert_chain(crt, key)
|
||||
except:
|
||||
import ssl
|
||||
crt, key = create_for_https()
|
||||
ssl_context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH)
|
||||
ssl_context.load_cert_chain(crt, key)
|
||||
|
||||
|
||||
success = False
|
||||
for i in range(11): # 尝试最多11次
|
||||
if await check_port_available(address, http_port + 1 + i):
|
||||
https_port = http_port + 1 + i
|
||||
site2 = web.TCPSite(runner, address, https_port, ssl_context=ssl_context)
|
||||
await site2.start()
|
||||
success = True
|
||||
break
|
||||
|
||||
if not success:
|
||||
raise RuntimeError(f"Ports {http_port + 1} to {http_port + 10} are all in use.")
|
||||
|
||||
if address == '':
|
||||
address = '0.0.0.0'
|
||||
address = '127.0.0.1'
|
||||
if address=='0.0.0.0':
|
||||
address = '127.0.0.1'
|
||||
|
||||
if verbose:
|
||||
# print('\033[91mMixlab Nodes: \033[93mLoaded\033[0m')
|
||||
print("\033[93mStarting server\n")
|
||||
print("\033[93mTo see the GUI go to: http://{}:{}".format(address, port))
|
||||
print("\033[93mTo see the GUI go to: https://{}:{}\033[0m".format(address, port+1))
|
||||
|
||||
logging.info("\n")
|
||||
logging.info("\n\nStarting server")
|
||||
|
||||
# print("\033[93mStarting server\n")
|
||||
logging.info("\033[93mTo see the GUI go to: http://{}:{}".format(address, http_port))
|
||||
logging.info("\033[93mTo see the GUI go to: https://{}:{}\033[0m".format(address, https_port))
|
||||
|
||||
# print("\033[93mTo see the GUI go to: http://{}:{}".format(address, http_port))
|
||||
# print("\033[93mTo see the GUI go to: https://{}:{}\033[0m".format(address, https_port))
|
||||
|
||||
if call_on_start is not None:
|
||||
call_on_start(address, port)
|
||||
try:
|
||||
if scheme=='https':
|
||||
call_on_start(scheme,address, https_port)
|
||||
else:
|
||||
call_on_start(scheme,address, http_port)
|
||||
except:
|
||||
call_on_start(address,http_port)
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error starting the server: {e}")
|
||||
|
||||
# import webbrowser
|
||||
# if os.name == 'nt' and address == '0.0.0.0':
|
||||
@@ -327,11 +548,10 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
|
||||
# webbrowser.open(f"https://{address}")
|
||||
# webbrowser.open(f"http://{address}:{port}")
|
||||
|
||||
|
||||
PromptServer.start=new_start
|
||||
|
||||
# 创建路由表
|
||||
routes = web.RouteTableDef()
|
||||
routes = PromptServer.instance.routes
|
||||
|
||||
@routes.post('/mixlab')
|
||||
async def mixlab_hander(request):
|
||||
@@ -356,7 +576,7 @@ async def mixlab_app_handler(request):
|
||||
return web.Response(text=html_data, content_type='text/html')
|
||||
else:
|
||||
return web.Response(text="HTML file not found", status=404)
|
||||
|
||||
|
||||
|
||||
@routes.post('/mixlab/workflow')
|
||||
async def mixlab_workflow_hander(request):
|
||||
@@ -371,17 +591,26 @@ async def mixlab_workflow_hander(request):
|
||||
'file_path':file_path
|
||||
}
|
||||
elif data['task']=='save_app':
|
||||
file_path=save_workflow_for_app(data['data'],data['filename'])
|
||||
category=""
|
||||
if "category" in data:
|
||||
category=data['category']
|
||||
file_path=save_workflow_for_app(data['data'],data['filename'],category)
|
||||
result={
|
||||
'status':'success',
|
||||
'file_path':file_path
|
||||
}
|
||||
elif data['task']=='my_app':
|
||||
filename=None
|
||||
category=""
|
||||
admin=False
|
||||
if 'filename' in data:
|
||||
filename=data['filename']
|
||||
if 'category' in data:
|
||||
category=data['category']
|
||||
if 'admin' in data:
|
||||
admin=data['admin']
|
||||
result={
|
||||
'data':get_my_workflow_for_app(filename),
|
||||
'data':get_my_workflow_for_app(filename,category,admin),
|
||||
'status':'success',
|
||||
}
|
||||
elif data['task']=='list':
|
||||
@@ -408,62 +637,338 @@ async def nodes_map_hander(request):
|
||||
|
||||
return web.json_response(result)
|
||||
|
||||
# 把插件自定义的路由添加到comfyui server里
|
||||
def new_add_routes(self):
|
||||
import nodes
|
||||
self.app.add_routes(routes)
|
||||
self.app.add_routes(self.routes)
|
||||
for name, dir in nodes.EXTENSION_WEB_DIRS.items():
|
||||
self.app.add_routes([
|
||||
web.static('/extensions/' + urllib.parse.quote(name), dir, follow_symlinks=True),
|
||||
])
|
||||
self.app.add_routes([
|
||||
web.static('/', self.web_root, follow_symlinks=True),
|
||||
])
|
||||
|
||||
PromptServer.add_routes=new_add_routes
|
||||
@routes.post("/mixlab/folder_paths")
|
||||
async def get_checkpoints(request):
|
||||
data = await request.json()
|
||||
t="checkpoints"
|
||||
names=[]
|
||||
try:
|
||||
t=data['type']
|
||||
names = folder_paths.get_filename_list(t)
|
||||
except Exception as e:
|
||||
print('/mixlab/folder_paths',False,e)
|
||||
|
||||
try:
|
||||
if data['type']=='llamafile':
|
||||
names=get_llama_models()
|
||||
except:
|
||||
print("llamafile none")
|
||||
|
||||
try:
|
||||
if data['type']=='rembg':
|
||||
names=get_rembg_models(U2NET_HOME)
|
||||
except:
|
||||
print("rembg none")
|
||||
|
||||
return web.json_response({"names":names,"types":list(folder_paths.folder_names_and_paths.keys())})
|
||||
|
||||
|
||||
@routes.post('/mixlab/rembg')
|
||||
async def rembg_hander(request):
|
||||
data = await request.json()
|
||||
model=data['model']
|
||||
result={}
|
||||
|
||||
data_base64=remove_base64_prefix(data['base64'])
|
||||
image_data = base64.b64decode(data_base64)
|
||||
|
||||
# 创建一个BytesIO对象
|
||||
image_stream = io.BytesIO(image_data)
|
||||
|
||||
# 使用PIL Image模块读取图像
|
||||
image = Image.open(image_stream)
|
||||
|
||||
if model=='briarmbg':
|
||||
_,rgba_images,_=run_briarmbg([image])
|
||||
else:
|
||||
_,rgba_images,_=run_rembg(model,[image])
|
||||
|
||||
with io.BytesIO() as buf:
|
||||
rgba_images[0].save(buf, format='PNG')
|
||||
img_bytes = buf.getvalue()
|
||||
img_base64 = base64.b64encode(img_bytes).decode('utf-8')
|
||||
|
||||
try:
|
||||
result={
|
||||
'data':img_base64,
|
||||
'model':model,
|
||||
'status':'success',
|
||||
}
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
return web.json_response(result)
|
||||
|
||||
|
||||
@routes.post("/mixlab/prompt_result")
|
||||
async def post_prompt_result(request):
|
||||
data = await request.json()
|
||||
res=None
|
||||
# print(data)
|
||||
try:
|
||||
action=data['action']
|
||||
if action=='save':
|
||||
result=data['data']
|
||||
res=save_prompt_result(result['prompt_id'],result)
|
||||
elif action=='all':
|
||||
res=get_prompt_result()
|
||||
except Exception as e:
|
||||
print('/mixlab/prompt_result',False,e)
|
||||
|
||||
return web.json_response({"result":res})
|
||||
|
||||
|
||||
|
||||
# 扩展api接口
|
||||
# from server import PromptServer
|
||||
# from aiohttp import web
|
||||
def start_local_live_thread(data):
|
||||
import asyncio
|
||||
from VoiceStreamAI.server import Server
|
||||
from VoiceStreamAI.asr.asr_factory import ASRFactory
|
||||
from VoiceStreamAI.vad.vad_factory import VADFactory
|
||||
|
||||
model="large-v3"
|
||||
if "model" in data:
|
||||
model=data['model']
|
||||
|
||||
vad_pipeline = VADFactory.create_vad_pipeline("pyannote")
|
||||
#device
|
||||
asr_pipeline = ASRFactory.create_asr_pipeline("faster_whisper", **{"model_size":model})
|
||||
|
||||
port=8765
|
||||
if 'port' in data:
|
||||
port=data['port']
|
||||
|
||||
llm_port=9000
|
||||
if 'llm_port' in data:
|
||||
llm_port=data['llm_port']
|
||||
|
||||
server = Server(vad_pipeline,
|
||||
asr_pipeline,
|
||||
host="127.0.0.1",
|
||||
port=port,
|
||||
sampling_rate=16000,
|
||||
samples_width=2,
|
||||
llm_port=llm_port
|
||||
)
|
||||
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
loop.run_until_complete(server.start())
|
||||
loop.run_forever()
|
||||
|
||||
|
||||
async def start_local_llm(data):
|
||||
global llama_port,llama_model,llama_chat_format
|
||||
if llama_port and llama_model and llama_chat_format:
|
||||
return {"port":llama_port,"model":llama_model,"chat_format":llama_chat_format}
|
||||
|
||||
import threading
|
||||
import uvicorn
|
||||
from llama_cpp.server.app import create_app
|
||||
from llama_cpp.server.settings import (
|
||||
Settings,
|
||||
ServerSettings,
|
||||
ModelSettings,
|
||||
ConfigFileSettings,
|
||||
)
|
||||
|
||||
if not "model" in data and "model_path" in data:
|
||||
data['model']= os.path.basename(data["model_path"])
|
||||
model=data["model_path"]
|
||||
|
||||
elif "model" in data:
|
||||
model=get_llama_model_path(data['model'])
|
||||
|
||||
n_gpu_layers=-1
|
||||
|
||||
if "n_gpu_layers" in data:
|
||||
n_gpu_layers=data['n_gpu_layers']
|
||||
|
||||
|
||||
chat_format="chatml"
|
||||
if "model" in data and "function-calling" in data['model']:
|
||||
chat_format="functionary-v2"
|
||||
|
||||
model_alias=os.path.basename(model)
|
||||
|
||||
# 多模态
|
||||
clip_model_path=None
|
||||
|
||||
prefix = "llava-phi-3-mini"
|
||||
file_name = prefix+"-mmproj-"
|
||||
if model_alias.startswith(prefix):
|
||||
for file in os.listdir(os.path.dirname(model)):
|
||||
if file.startswith(file_name):
|
||||
clip_model_path=os.path.join(os.path.dirname(model),file)
|
||||
chat_format='llava-1-5'
|
||||
print('#clip_model_path',chat_format,clip_model_path)
|
||||
|
||||
|
||||
address="127.0.0.1"
|
||||
port=9090
|
||||
success = False
|
||||
for i in range(11): # 尝试最多11次
|
||||
if await check_port_available(address, port + i):
|
||||
port = port + i
|
||||
success = True
|
||||
break
|
||||
|
||||
if success == False:
|
||||
return {"port":None,"model":""}
|
||||
|
||||
|
||||
server_settings=ServerSettings(host=address,port=port)
|
||||
|
||||
name, ext = os.path.splitext(os.path.basename(model))
|
||||
print('#model',name)
|
||||
app = create_app(
|
||||
server_settings=server_settings,
|
||||
model_settings=[
|
||||
ModelSettings(
|
||||
model=model,
|
||||
model_alias=name,
|
||||
n_gpu_layers=n_gpu_layers,
|
||||
n_ctx=4098,
|
||||
chat_format=chat_format,
|
||||
embedding=False,
|
||||
clip_model_path=clip_model_path
|
||||
)])
|
||||
|
||||
def run_uvicorn():
|
||||
uvicorn.run(
|
||||
app,
|
||||
host=os.getenv("HOST", server_settings.host),
|
||||
port=int(os.getenv("PORT", server_settings.port)),
|
||||
ssl_keyfile=server_settings.ssl_keyfile,
|
||||
ssl_certfile=server_settings.ssl_certfile,
|
||||
)
|
||||
|
||||
# 创建一个子线程
|
||||
thread = threading.Thread(target=run_uvicorn)
|
||||
|
||||
# 启动子线程
|
||||
thread.start()
|
||||
|
||||
llama_port=port
|
||||
llama_model=data['model']
|
||||
llama_chat_format=chat_format
|
||||
|
||||
return {"port":llama_port,"model":llama_model,"chat_format":llama_chat_format}
|
||||
|
||||
# llam服务的开启
|
||||
@routes.post('/mixlab/start_llama')
|
||||
async def my_hander_method(request):
|
||||
data =await request.json()
|
||||
# print(data)
|
||||
if llama_port and llama_model and llama_chat_format:
|
||||
return web.json_response({"port":llama_port,"model":llama_model,"chat_format":llama_chat_format} )
|
||||
try:
|
||||
result=await start_local_llm(data)
|
||||
except:
|
||||
result= {"port":None,"model":"","llama_cpp_error":True}
|
||||
print('start_local_llm error')
|
||||
|
||||
return web.json_response(result)
|
||||
|
||||
|
||||
@routes.post('/mixlab/start_live')
|
||||
async def mixlab_live_start_handler(request):
|
||||
import threading
|
||||
llm=await start_local_llm({
|
||||
"model":"Phi-3-mini-4k-instruct-Q5_K_S.gguf",
|
||||
"n_gpu_layers":2
|
||||
})
|
||||
|
||||
# {"port":llama_port,"model":llama_model,"chat_format":llama_chat_format}
|
||||
|
||||
os.environ['HF_ENDPOINT'] = 'https://hf-mirror.com' #hf_hub_download 里的下载地址修改
|
||||
os.environ['PYANNOTE_AUTH_TOKEN'] = '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'
|
||||
|
||||
# Create and start the thread
|
||||
data = {
|
||||
"llm_port":llm['port'],
|
||||
"port":8725,
|
||||
"model":"large-v3"
|
||||
} # Replace with your actual data if needed
|
||||
thread = threading.Thread(target=start_local_live_thread, args=(data,))
|
||||
thread.start()
|
||||
|
||||
return web.json_response(data)
|
||||
|
||||
|
||||
@routes.get('/mixlab/live')
|
||||
async def mixlab_live_handler(request):
|
||||
html_file = os.path.join(current_path, "web/live.html")
|
||||
if os.path.exists(html_file):
|
||||
with open(html_file, 'r', encoding='utf-8', errors='ignore') as f:
|
||||
html_data = f.read()
|
||||
return web.Response(text=html_data, content_type='text/html')
|
||||
else:
|
||||
return web.Response(text="HTML file not found", status=404)
|
||||
|
||||
# 重启服务
|
||||
@routes.post('/mixlab/re_start')
|
||||
def re_start(request):
|
||||
try:
|
||||
sys.stdout.close_log()
|
||||
except Exception as e:
|
||||
pass
|
||||
return os.execv(sys.executable, [sys.executable] + sys.argv)
|
||||
|
||||
# @routes.post('/ws_image')
|
||||
# async def my_hander_method(request):
|
||||
# post = await request.post()
|
||||
# x = post.get("something")
|
||||
# return web.json_response({})
|
||||
|
||||
|
||||
# 导入节点
|
||||
from .nodes.PromptNode import RandomPrompt,PromptSlide
|
||||
from .nodes.ImageNode import NoiseImage,TransparentImage,LoadImagesFromPath,LoadImagesFromURL,ResizeImage,TextImage,SvgImage,Image3D,ShowLayer,NewLayer,MergeLayers,AreaToMask,SmoothMask,FeatheredMask,SplitLongMask,ImageCropByAlpha,EnhanceImage,FaceToMask
|
||||
from .nodes.Vae import VAELoader,VAEDecode
|
||||
from .nodes.PromptNode import GLIGENTextBoxApply_Advanced,EmbeddingPrompt,RandomPrompt,PromptSlide,PromptSimplification,PromptImage,JoinWithDelimiter
|
||||
from .nodes.ImageNode import ComparingTwoFrames,LoadImages_,CompositeImages,GridDisplayAndSave,GridInput,ImagesPrompt,SaveImageAndMetadata,SaveImageToLocal,SplitImage,GridOutput,GetImageSize_,MirroredImage,ImageColorTransfer,NoiseImage,TransparentImage,GradientImage,LoadImagesFromPath,LoadImagesFromURL,ResizeImage,TextImage,SvgImage,Image3D,ShowLayer,NewLayer,MergeLayers,CenterImage,AreaToMask,SmoothMask,SplitLongMask,ImageCropByAlpha,EnhanceImage,FaceToMask
|
||||
# from .nodes.Vae import VAELoader,VAEDecode
|
||||
from .nodes.ScreenShareNode import ScreenShareNode,FloatingVideo
|
||||
from .nodes.Clipseg import CLIPSeg,CombineMasks
|
||||
from .nodes.ChatGPT import ChatGPTNode,ShowTextForGPT,CharacterInText
|
||||
|
||||
from .nodes.ChatGPT import ChatGPTNode,ShowTextForGPT,CharacterInText,TextSplitByDelimiter
|
||||
from .nodes.Audio import GamePal,SpeechRecognition,SpeechSynthesis
|
||||
from .nodes.Utils import AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,GetImageSize_,MultiplicationNode
|
||||
from .nodes.Lama import LaMaInpainting
|
||||
from .nodes.Utils import IncrementingListNode,ListSplit,CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,TESTNODE_TOKEN,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
|
||||
from .nodes.Mask import PreviewMask_,MaskListReplace,MaskListMerge,OutlineMask,FeatheredMask
|
||||
|
||||
from .nodes.Style import ApplyVisualStylePrompting,StyleAlignedReferenceSampler,StyleAlignedBatchAlign,StyleAlignedSampleReferenceLatents
|
||||
|
||||
from .nodes.Video import VideoCombine_Adv,LoadVideoAndSegment,ImageListReplace,VAEEncodeForInpaint_Frames
|
||||
|
||||
from .nodes.TripoSR import LoadTripoSRModel,TripoSRSampler,SaveTripoSRMesh
|
||||
|
||||
|
||||
# 要导出的所有节点及其名称的字典
|
||||
# 注意:名称应全局唯一
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"AppInfo":AppInfo,
|
||||
"TESTNODE_":TESTNODE_,
|
||||
"TESTNODE_TOKEN":TESTNODE_TOKEN,
|
||||
"RandomPrompt":RandomPrompt,
|
||||
# "LoraPrompt":LoraPrompt,
|
||||
"EmbeddingPrompt":EmbeddingPrompt,
|
||||
"PromptSlide":PromptSlide,
|
||||
"GLIGENTextBoxApply_Advanced":GLIGENTextBoxApply_Advanced,
|
||||
"PromptSimplification":PromptSimplification,
|
||||
"PromptImage":PromptImage,
|
||||
"MirroredImage":MirroredImage,
|
||||
"NoiseImage":NoiseImage,
|
||||
"GradientImage":GradientImage,
|
||||
"TransparentImage":TransparentImage,
|
||||
"ResizeImageMixlab":ResizeImage,
|
||||
"LoadImagesFromPath":LoadImagesFromPath,
|
||||
"LoadImagesFromURL":LoadImagesFromURL,
|
||||
"LoadImagesToBatch":LoadImages_,
|
||||
"TextImage":TextImage,
|
||||
"EnhanceImage":EnhanceImage,
|
||||
"SvgImage":SvgImage,
|
||||
"3DImage":Image3D,
|
||||
"ImageColorTransfer":ImageColorTransfer,
|
||||
"ShowLayer":ShowLayer,
|
||||
"NewLayer":NewLayer,
|
||||
"CompositeImages_":CompositeImages,
|
||||
"SplitImage":SplitImage,
|
||||
"CenterImage":CenterImage,
|
||||
"GridOutput":GridOutput,
|
||||
"GridDisplayAndSave":GridDisplayAndSave,
|
||||
"GridInput":GridInput,
|
||||
"MergeLayers":MergeLayers,
|
||||
"SplitLongMask":SplitLongMask,
|
||||
"FeatheredMask":FeatheredMask,
|
||||
@@ -471,15 +976,18 @@ NODE_CLASS_MAPPINGS = {
|
||||
"FaceToMask":FaceToMask,
|
||||
"AreaToMask":AreaToMask,
|
||||
"ImageCropByAlpha":ImageCropByAlpha,
|
||||
"VAELoaderConsistencyDecoder":VAELoader,
|
||||
"VAEDecodeConsistencyDecoder":VAEDecode,
|
||||
"ImagesPrompt_":ImagesPrompt,
|
||||
# "VAELoaderConsistencyDecoder":VAELoader,
|
||||
"SaveImageToLocal":SaveImageToLocal,
|
||||
"SaveImageAndMetadata_":SaveImageAndMetadata,
|
||||
"ComparingTwoFrames_":ComparingTwoFrames,
|
||||
# "VAEDecodeConsistencyDecoder":VAEDecode,
|
||||
"ScreenShare":ScreenShareNode,
|
||||
"FloatingVideo":FloatingVideo,
|
||||
"CLIPSeg_":CLIPSeg,
|
||||
"CombineMasks_":CombineMasks,
|
||||
"ChatGPTOpenAI":ChatGPTNode,
|
||||
"ShowTextForGPT":ShowTextForGPT,
|
||||
"CharacterInText":CharacterInText,
|
||||
"TextSplitByDelimiter":TextSplitByDelimiter,
|
||||
"SpeechRecognition":SpeechRecognition,
|
||||
"SpeechSynthesis":SpeechSynthesis,
|
||||
"Color":ColorInput,
|
||||
@@ -493,36 +1001,130 @@ NODE_CLASS_MAPPINGS = {
|
||||
"GetImageSize_":GetImageSize_,
|
||||
"SwitchByIndex":SwitchByIndex,
|
||||
"LimitNumber":LimitNumber,
|
||||
"LaMaInpainting":LaMaInpainting
|
||||
"OutlineMask":OutlineMask,
|
||||
"MaskListMerge_":MaskListMerge,
|
||||
"JoinWithDelimiter":JoinWithDelimiter,
|
||||
"Seed_":CreateSeedNode,
|
||||
"CkptNames_":CreateCkptNames,
|
||||
"SamplerNames_":CreateSampler_names,
|
||||
"LoraNames_":CreateLoraNames,
|
||||
"ApplyVisualStylePrompting_":ApplyVisualStylePrompting,
|
||||
"StyleAlignedReferenceSampler_": StyleAlignedReferenceSampler,
|
||||
"StyleAlignedSampleReferenceLatents_": StyleAlignedSampleReferenceLatents,
|
||||
"StyleAlignedBatchAlign_": StyleAlignedBatchAlign,
|
||||
"LoadVideoAndSegment_":LoadVideoAndSegment,
|
||||
"VideoCombine_Adv":VideoCombine_Adv,
|
||||
"ListSplit_":ListSplit,
|
||||
"MaskListReplace_":MaskListReplace,
|
||||
"ImageListReplace_":ImageListReplace,
|
||||
"VAEEncodeForInpaint_Frames":VAEEncodeForInpaint_Frames,
|
||||
"IncrementingListNode_":IncrementingListNode,
|
||||
"PreviewMask_":PreviewMask_,
|
||||
"LoadTripoSRModel_": LoadTripoSRModel,
|
||||
"TripoSRSampler_": TripoSRSampler,
|
||||
"SaveTripoSRMesh": SaveTripoSRMesh
|
||||
# "GamePal":GamePal
|
||||
}
|
||||
|
||||
# 一个包含节点友好/可读的标题的字典
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AppInfo":"AppInfo ♾️Mixlab",
|
||||
"ResizeImageMixlab":"ResizeImage ♾️Mixlab",
|
||||
"AppInfo":"App Info ♾️MixlabApp",
|
||||
"Color":"Color Input ♾️MixlabApp",
|
||||
"TextInput_":"Text Input ♾️MixlabApp",
|
||||
"FloatSlider":"Float Slider Input ♾️MixlabApp",
|
||||
"IntNumber":"Int Input ♾️MixlabApp",
|
||||
"ImagesPrompt_":"Images Input ♾️MixlabApp",
|
||||
"SaveImageAndMetadata_":"Save Image Output ♾️MixlabApp",
|
||||
"ComparingTwoFrames_":"Comparing Two Frames ♾️MixlabApp",
|
||||
"ResizeImageMixlab":"Resize Image ♾️Mixlab",
|
||||
"RandomPrompt": "Random Prompt ♾️Mixlab",
|
||||
"PromptImage":"Output Prompt and Image ♾️Mixlab",
|
||||
"SplitLongMask":"Splitting a long image into sections",
|
||||
"VAELoaderConsistencyDecoder":"Consistency Decoder Loader",
|
||||
"VAEDecodeConsistencyDecoder":"Consistency Decoder Decode",
|
||||
"ScreenShare":"ScreenShare ♾️Mixlab",
|
||||
"ScreenShare":"Screen Share ♾️Mixlab",
|
||||
"FloatingVideo":"FloatingVideo ♾️Mixlab",
|
||||
"ChatGPTOpenAI":"ChatGPT ♾️Mixlab",
|
||||
"ShowTextForGPT":"ShowTextForGPT ♾️Mixlab",
|
||||
"MergeLayers":"MergeLayers ♾️Mixlab",
|
||||
"ChatGPTOpenAI":"ChatGPT & Local LLM ♾️Mixlab",
|
||||
"ShowTextForGPT":"Show Text ♾️MixlabApp",
|
||||
"MergeLayers":"Merge Layers ♾️Mixlab",
|
||||
"SpeechSynthesis":"SpeechSynthesis ♾️Mixlab",
|
||||
"SpeechRecognition":"SpeechRecognition ♾️Mixlab",
|
||||
"3DImage":"3DImage ♾️Mixlab",
|
||||
"CompositeImages_":"Composite Images ♾️Mixlab",
|
||||
"DynamicDelayProcessor":"DynamicDelayByText ♾️Mixlab",
|
||||
"LaMaInpainting":"LaMaInpainting ♾️Mixlab",
|
||||
"PromptSlide":"PromptSlide ♾️Mixlab"
|
||||
|
||||
# "GamePal":"GamePal ♾️Mixlab"
|
||||
"PromptSlide":"Prompt Slide ♾️Mixlab",
|
||||
"PromptGenerate_Mix":"Prompt Generate ♾️Mixlab",
|
||||
"ChinesePrompt_Mix":"Chinese Prompt ♾️Mixlab",
|
||||
"GamePal":"GamePal ♾️Mixlab",
|
||||
"RembgNode_Mix":"Remove Background ♾️Mixlab",
|
||||
"LoraNames_":"LoraName ♾️Mixlab",
|
||||
"ApplyVisualStylePrompting_":"Apply VisualStyle Prompting ♾️Mixlab",
|
||||
"StyleAlignedReferenceSampler_": "StyleAligned Reference Sampler ♾️Mixlab",
|
||||
"StyleAlignedSampleReferenceLatents_": "StyleAligned Sample Reference Latents ♾️Mixlab",
|
||||
"StyleAlignedBatchAlign_": "StyleAligned Batch Align ♾️Mixlab",
|
||||
"LoadVideoAndSegment_":"Load Video And Segment ♾️Mixlab",
|
||||
"VideoCombine_Adv":"Video Combine ♾️Mixlab",
|
||||
"MaskListMerge_":"MaskList to Mask ♾️Mixlab",
|
||||
"ListSplit_":"Split List ♾️Mixlab",
|
||||
"MaskListReplace_":"MaskList Replace ♾️Mixlab",
|
||||
"ImageListReplace_":"ImageList Replace ♾️Mixlab",
|
||||
"SwitchByIndex":"List Switch By Index ♾️Mixlab",
|
||||
"GLIGENTextBoxApply_Advanced":"GLIGEN TextBox Apply ♾️Mixlab",
|
||||
"GridDisplayAndSave":"Grid Display And Save ♾️Mixlab",
|
||||
"GridInput":"Grid Input ♾️Mixlab",
|
||||
"GridOutput":"Grid Output ♾️Mixlab",
|
||||
"GetImageSize_":"Get Image Size ♾️Mixlab",
|
||||
"VAEEncodeForInpaint_Frames":"VAE Encode For Inpaint Frames ♾️Mixlab",
|
||||
"IncrementingListNode_":"Create Incrementing Number List ♾️Mixlab",
|
||||
"LoadImagesToBatch":"Load Images(base64) ♾️Mixlab",
|
||||
"PreviewMask_":"Preview Mask",
|
||||
"LoadTripoSRModel_": "Load TripoSR Model",
|
||||
"TripoSRSampler_": "TripoSR Sampler",
|
||||
"SaveTripoSRMesh": "Save TripoSR Mesh"
|
||||
}
|
||||
|
||||
# web ui的节点功能
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
print('--------------')
|
||||
print('\033[91m ### Mixlab Nodes: \033[93mLoaded\033[0m')
|
||||
print('--------------')
|
||||
|
||||
logging.info('--------------')
|
||||
logging.info('\033[91m ### Mixlab Nodes: \033[93mLoaded')
|
||||
# print('\033[91m ### Mixlab Nodes: \033[93mLoaded')
|
||||
|
||||
try:
|
||||
from .nodes.Lama import LaMaInpainting
|
||||
logging.info('LaMaInpainting.available {}'.format(LaMaInpainting.available))
|
||||
if LaMaInpainting.available:
|
||||
NODE_CLASS_MAPPINGS['LaMaInpainting']=LaMaInpainting
|
||||
except Exception as e:
|
||||
logging.info('LaMaInpainting.available False')
|
||||
|
||||
try:
|
||||
from .nodes.ClipInterrogator import ClipInterrogator
|
||||
logging.info('ClipInterrogator.available {}'.format(ClipInterrogator.available))
|
||||
if ClipInterrogator.available:
|
||||
NODE_CLASS_MAPPINGS['ClipInterrogator']=ClipInterrogator
|
||||
except Exception as e:
|
||||
logging.info('ClipInterrogator.available False')
|
||||
|
||||
try:
|
||||
from .nodes.TextGenerateNode import PromptGenerate,ChinesePrompt
|
||||
logging.info('PromptGenerate.available {}'.format(PromptGenerate.available))
|
||||
if PromptGenerate.available:
|
||||
NODE_CLASS_MAPPINGS['PromptGenerate_Mix']=PromptGenerate
|
||||
logging.info('ChinesePrompt.available {}'.format(ChinesePrompt.available))
|
||||
if ChinesePrompt.available:
|
||||
NODE_CLASS_MAPPINGS['ChinesePrompt_Mix']=ChinesePrompt
|
||||
except Exception as e:
|
||||
logging.info('TextGenerateNode.available False')
|
||||
|
||||
try:
|
||||
from .nodes.RembgNode import RembgNode_
|
||||
logging.info('RembgNode_.available {}'.format(RembgNode_.available))
|
||||
if RembgNode_.available:
|
||||
NODE_CLASS_MAPPINGS['RembgNode_Mix']=RembgNode_
|
||||
except Exception as e:
|
||||
logging.info('RembgNode_.available False' )
|
||||
|
||||
logging.info('\033[93m -------------- \033[0m')
|
||||
|
||||
|
After Width: | Height: | Size: 135 KiB |
|
After Width: | Height: | Size: 210 KiB |
|
After Width: | Height: | Size: 2.4 MiB |
|
After Width: | Height: | Size: 1.1 MiB |
|
Before Width: | Height: | Size: 784 KiB |
|
After Width: | Height: | Size: 75 KiB |
|
After Width: | Height: | Size: 63 KiB |
|
After Width: | Height: | Size: 477 KiB |
|
After Width: | Height: | Size: 965 KiB |
@@ -0,0 +1,30 @@
|
||||
Jony Ive
|
||||
Dieter Rams
|
||||
Philippe Starck
|
||||
Karim Rashid
|
||||
Yves Béhar
|
||||
Marc Newson
|
||||
Naoto Fukasawa
|
||||
Jonathan Adler
|
||||
Patricia Urquiola
|
||||
Ross Lovegrove
|
||||
Tom Dixon
|
||||
Jasper Morrison
|
||||
Charles Eames
|
||||
Ray Eames
|
||||
Achille Castiglioni
|
||||
Ron Arad
|
||||
Konstantin Grcic
|
||||
Marcel Wanders
|
||||
Maarten Baas
|
||||
Stefan Sagmeister
|
||||
Ingo Maurer
|
||||
Hella Jongerius
|
||||
Sam Hecht
|
||||
Kim Colin
|
||||
Jaime Hayon
|
||||
Michael Anastassiades
|
||||
Nendo
|
||||
Oki Sato
|
||||
Matali Crasset
|
||||
Tokujin Yoshioka
|
||||
@@ -0,0 +1,10 @@
|
||||
Chibi Anime Style
|
||||
Gakuen Anime Style
|
||||
Gekiga Anime Style
|
||||
Jidaimono Anime Style
|
||||
Kawaii Anime Style
|
||||
Mecha Anime Style
|
||||
Realistic Anime Style
|
||||
Semi-Realistic Anime Style
|
||||
Shoji Anime Style
|
||||
Kemonomimi Anime Style
|
||||
@@ -0,0 +1,23 @@
|
||||
GoPro
|
||||
Drone
|
||||
polaroid
|
||||
black and white film
|
||||
Kodachrome
|
||||
shot on 8mm
|
||||
shot on 16mm
|
||||
shot on 35mm
|
||||
Microscopic
|
||||
Fisheye Lens
|
||||
Wide Angle
|
||||
Ultra-Wide Angle
|
||||
Panorama
|
||||
Short Exposure
|
||||
Long Exposure
|
||||
Double Exposure
|
||||
f2.8
|
||||
Depth of Field
|
||||
Soft Focus
|
||||
Deep Focus
|
||||
Shallow Focus
|
||||
Vanishing Point
|
||||
Vantage Point
|
||||
@@ -0,0 +1,30 @@
|
||||
Elegant evening gown
|
||||
Casual jeans and t-shirt
|
||||
Formal black suit
|
||||
Stylish leather jacket
|
||||
Flowy bohemian dress
|
||||
Sporty tracksuit
|
||||
Chic little black dress
|
||||
Trendy ripped jeans
|
||||
Classic white button-down shirt
|
||||
Cozy oversized sweater
|
||||
Sophisticated tailored blazer
|
||||
Quirky patterned leggings
|
||||
Striped sailor top
|
||||
Polished knee-length skirt
|
||||
Vintage-inspired floral dress
|
||||
Edgy motorcycle jacket
|
||||
Preppy polo shirt
|
||||
Boho maxi skirt
|
||||
Professional pinstripe suit
|
||||
Relaxed denim shorts
|
||||
Glamorous sequined dress
|
||||
Athletic running shoes
|
||||
Formal bow tie
|
||||
Casual baseball cap
|
||||
Stylish fedora hat
|
||||
Warm woolen scarf
|
||||
Comfortable cotton socks
|
||||
Trendy ankle boots
|
||||
Cute summer sandals
|
||||
Cozy pajama set
|
||||
@@ -0,0 +1,30 @@
|
||||
Happy
|
||||
Sad
|
||||
Angry
|
||||
Surprised
|
||||
Excited
|
||||
Worried
|
||||
Confused
|
||||
Disgusted
|
||||
Amused
|
||||
Bored
|
||||
Curious
|
||||
Embarrassed
|
||||
Frustrated
|
||||
Nervous
|
||||
Pleased
|
||||
Relieved
|
||||
Shy
|
||||
Tired
|
||||
Serious
|
||||
Silly
|
||||
Proud
|
||||
Grumpy
|
||||
Smug
|
||||
Sarcastic
|
||||
Flirty
|
||||
Skeptical
|
||||
Shocked
|
||||
Blissful
|
||||
Envious
|
||||
Mischievous
|
||||
@@ -4761,19 +4761,29 @@
|
||||
],
|
||||
"https://github.com/shadowcz007/comfyui-mixlab-nodes": [
|
||||
[
|
||||
"GridOutput",
|
||||
"SplitImage",
|
||||
"PromptGenerate_Mix",
|
||||
"JoinWithDelimiter",
|
||||
"ChinesePrompt_Mix",
|
||||
"3DImage",
|
||||
"AppInfo",
|
||||
"IntNumber",
|
||||
"FloatSlider",
|
||||
"ResizeImage",
|
||||
"NoiseImage",
|
||||
"PromptImage",
|
||||
"SaveImageToLocal",
|
||||
"AreaToMask",
|
||||
"CLIPSeg_",
|
||||
"CharacterInText",
|
||||
"ChatGPTOpenAI",
|
||||
"Color",
|
||||
"CombineMasks_",
|
||||
"Seed_",
|
||||
"CkptNames_",
|
||||
"SamplerNames_",
|
||||
"LoraNames_",
|
||||
"EnhanceImage",
|
||||
"GradientImage",
|
||||
"FaceToMask",
|
||||
"FeatheredMask",
|
||||
"FloatingVideo",
|
||||
@@ -4783,8 +4793,11 @@
|
||||
"LoadImagesFromURL",
|
||||
"MergeLayers",
|
||||
"NewLayer",
|
||||
"CenterImage",
|
||||
"RandomPrompt",
|
||||
"PromptSlide",
|
||||
"PromptSimplification",
|
||||
"ClipInterrogator",
|
||||
"ScreenShare",
|
||||
"ShowLayer",
|
||||
"ShowTextForGPT",
|
||||
@@ -4796,12 +4809,11 @@
|
||||
"TextImage",
|
||||
"ResizeImageMixlab",
|
||||
"TransparentImage",
|
||||
"VAEDecodeConsistencyDecoder",
|
||||
"VAELoaderConsistencyDecoder",
|
||||
"TextToNumber",
|
||||
"TextInput_",
|
||||
"DynamicDelayProcessor",
|
||||
"LaMaInpainting"
|
||||
"LaMaInpainting",
|
||||
"Moondream"
|
||||
],
|
||||
{
|
||||
"title_aux": "comfyui-mixlab-nodes"
|
||||
@@ -5189,7 +5201,7 @@
|
||||
"title_aux": "ComfyUI Stable Video Diffusion"
|
||||
}
|
||||
],
|
||||
"https://github.com/thedyze/save-image-extended-comfyui": [
|
||||
"https://github.com/audioscavenger/save-image-extended-comfyui": [
|
||||
[
|
||||
"SaveImageExtended"
|
||||
],
|
||||
@@ -5693,4 +5705,4 @@
|
||||
"title_aux": "SDXLCustomAspectRatio"
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,16 @@
|
||||
Mood Lighting
|
||||
Moody Lighting
|
||||
Studio Lighting
|
||||
Cove Lighting
|
||||
Soft Lighting
|
||||
Hard Lighting
|
||||
Volumetric Lighting
|
||||
Low-Key Lighting
|
||||
High-Key Lighting
|
||||
Epic Light
|
||||
Rembrandt Lighting
|
||||
Contre-Jour
|
||||
Veiling Flare
|
||||
Crepuscular Rays
|
||||
Rays of Shimmering Light
|
||||
Godrays
|
||||
@@ -0,0 +1 @@
|
||||
{}
|
||||
@@ -0,0 +1,132 @@
|
||||
Aaron Siskind
|
||||
Alessio Albi
|
||||
Alfred Eisenstaedt
|
||||
Alfred Stieglitz
|
||||
Alyssa Monks
|
||||
André Kertész
|
||||
Andreas Gursky
|
||||
Andrew Wyeth
|
||||
Anne Geddes
|
||||
Annie Leibovitz
|
||||
Ansel Adams
|
||||
Arnold Newman
|
||||
August Sander
|
||||
Balthus
|
||||
Berenice Abbott
|
||||
Bill Brandt
|
||||
Bill Henson
|
||||
Brassaï (Gyula Halász)
|
||||
Brooke Shaden
|
||||
Bruce Davidson
|
||||
Bruce Weber
|
||||
Bunny Yeager
|
||||
Carleton Watkins
|
||||
Carrie Mae Weems
|
||||
Chuck Close
|
||||
Cindy Sherman
|
||||
Clarence H. White
|
||||
Claude Cahun
|
||||
Danny Lyon
|
||||
David LaChapelle
|
||||
Dawoud Bey
|
||||
Diane Arbus
|
||||
Don McCullin
|
||||
Dora Maar
|
||||
Dorothea Lange
|
||||
Duane Michals
|
||||
Eadweard Muybridge
|
||||
Edward Burtynsky
|
||||
Edward Curtis
|
||||
Edward Ruscha
|
||||
Edward Steichen
|
||||
Edward Weston
|
||||
Elliott Erwitt
|
||||
Ernst Haas
|
||||
Eugene Atget
|
||||
Fan Ho
|
||||
Francesca Woodman
|
||||
Frans Lanting
|
||||
Garry Winogrand
|
||||
Georges Melies
|
||||
Gerda Taro
|
||||
Gertrude Käsebier
|
||||
Gordon Parks
|
||||
Graciela Iturbide
|
||||
Gregory Crewdson
|
||||
Harold Edgerton
|
||||
Helen Levitt
|
||||
Helmut Newton
|
||||
Hendrik Kerstens
|
||||
Henri Cartier-Bresson
|
||||
Hugh Kretschmer
|
||||
Irving Penn
|
||||
Jacques Henri Lartigue
|
||||
James Nachtwey
|
||||
James Van Der Zee
|
||||
Jay Maisel
|
||||
Jerry Uelsmann
|
||||
Joel Peter Witkin
|
||||
Joel Sartore
|
||||
John Frederick William Herschel
|
||||
Josef Sudek
|
||||
Julia Margaret Cameron
|
||||
Karl Blossfeldt
|
||||
Larry Burrows
|
||||
László Moholy-Nagy (photography)
|
||||
Lee Jeffries
|
||||
Lewis Hine
|
||||
Lorna Simpson
|
||||
Lynsey Addario
|
||||
Margaret Bourke-White
|
||||
Mario Testino
|
||||
Martin Parr
|
||||
Martin Schoeller
|
||||
Mary Ellen Mark
|
||||
Mathew B. Brady
|
||||
Méret Oppenheim
|
||||
Meryl McMaster
|
||||
Mick Rock
|
||||
Miles Aldridge
|
||||
Minor Martin White
|
||||
Nan Goldin
|
||||
Nathan Wirth
|
||||
Olive Cotton
|
||||
Olivier Rousteing
|
||||
Patrick Demarchelier
|
||||
Paul Nicklen
|
||||
Paul Outerbridge
|
||||
Paul Strand
|
||||
Pete Souza
|
||||
Peter Dombrovskis
|
||||
Peter Henry Emerson
|
||||
Peter Lik
|
||||
Peter Lindbergh
|
||||
Philip-Lorca diCorcia
|
||||
Philippe Halsman
|
||||
Ralph Gibson
|
||||
Richard Avedon
|
||||
Robert Adams
|
||||
Robert Bechtle
|
||||
Robert Capa
|
||||
Robert Frank
|
||||
Robert Mapplethorpe
|
||||
Roger Fenton
|
||||
Ruth Bernhard
|
||||
Sally Mann
|
||||
Sebastião Salgado
|
||||
Shirin Neshat
|
||||
Stefan Gesell
|
||||
Steven Meisel
|
||||
Susan Meiselas
|
||||
Vivian Maier
|
||||
Vivian Maier
|
||||
Viviane Sassen
|
||||
Walker Evans
|
||||
Wes Anderson
|
||||
William Eggleston
|
||||
William Eugene Smith
|
||||
William Henry Fox Talbot
|
||||
Yinka Shonibare
|
||||
Yousuf Karsh
|
||||
Man Ray
|
||||
Robert Mapplethorpe
|
||||
@@ -0,0 +1,101 @@
|
||||
Doctor
|
||||
Teacher
|
||||
Engineer
|
||||
Lawyer
|
||||
Accountant
|
||||
Nurse
|
||||
Architect
|
||||
Chef
|
||||
Pilot
|
||||
Scientist
|
||||
Artist
|
||||
Writer
|
||||
Musician
|
||||
Actor
|
||||
Photographer
|
||||
Police officer
|
||||
Firefighter
|
||||
Dentist
|
||||
Pharmacist
|
||||
Veterinarian
|
||||
Electrician
|
||||
Plumber
|
||||
Carpenter
|
||||
Mechanic
|
||||
Farmer
|
||||
Astronaut
|
||||
Athlete
|
||||
Journalist
|
||||
Politician
|
||||
Economist
|
||||
Psychologist
|
||||
Social worker
|
||||
Librarian
|
||||
Translator
|
||||
Salesperson
|
||||
Entrepreneur
|
||||
Financial advisor
|
||||
Graphic designer
|
||||
Web developer
|
||||
Marketing manager
|
||||
Human resources manager
|
||||
Project manager
|
||||
Event planner
|
||||
Fashion designer
|
||||
Interior decorator
|
||||
Real estate agent
|
||||
Archaeologist
|
||||
Biologist
|
||||
Chemist
|
||||
Geologist
|
||||
Physicist
|
||||
Mathematician
|
||||
Historian
|
||||
Geographer
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
#MixCopilot
|
||||
@@ -0,0 +1,58 @@
|
||||
Residential space
|
||||
Apartment building
|
||||
Villa
|
||||
Bungalow
|
||||
Condominium
|
||||
Commercial space
|
||||
Shopping mall
|
||||
Supermarket
|
||||
Restaurant
|
||||
Store
|
||||
Market
|
||||
Office space
|
||||
Office building
|
||||
Office
|
||||
Meeting room
|
||||
Co-working space
|
||||
Educational space
|
||||
School
|
||||
University
|
||||
Training institution
|
||||
Library
|
||||
Laboratory
|
||||
Medical space
|
||||
Hospital
|
||||
Clinic
|
||||
Pharmacy
|
||||
Nursing home
|
||||
Rehabilitation center
|
||||
Cultural space
|
||||
Museum
|
||||
Library
|
||||
Theater
|
||||
Concert hall
|
||||
Gallery
|
||||
Sports space
|
||||
Sports stadium
|
||||
Gym
|
||||
Swimming pool
|
||||
Basketball court
|
||||
Football field
|
||||
Transportation space
|
||||
Airport
|
||||
Train station
|
||||
Subway station
|
||||
Bus stop
|
||||
Parking lot
|
||||
Public space
|
||||
Park
|
||||
Square
|
||||
Street
|
||||
Pedestrian street
|
||||
Community center
|
||||
Industrial space
|
||||
Factory
|
||||
Warehouse
|
||||
Production workshop
|
||||
Mine
|
||||
Power plant
|
||||
@@ -0,0 +1,135 @@
|
||||
Vintage
|
||||
Grain
|
||||
Sepia
|
||||
High Key
|
||||
Low Key
|
||||
High Dynamic Range
|
||||
Cross Process
|
||||
Radial Blur
|
||||
Infrared
|
||||
Lomo
|
||||
Photocopy
|
||||
Pencil Sketch
|
||||
Pop Art
|
||||
Orton
|
||||
Mosaic
|
||||
Selective Black and White
|
||||
Torn Paper
|
||||
Tilt-Shift
|
||||
Double Exposure
|
||||
Polaroid
|
||||
Liquid Ink
|
||||
Color Splash
|
||||
Sketch
|
||||
Water Drops
|
||||
Polarizer
|
||||
Chinese Painting
|
||||
Water Droplets
|
||||
Polarization
|
||||
Color Inversion
|
||||
Fish-eye
|
||||
Soft Focus
|
||||
Solarization
|
||||
Posterize
|
||||
Comic Book
|
||||
Duotone
|
||||
Gradient Map
|
||||
Edge Detection
|
||||
Oil Painting
|
||||
Reflection
|
||||
Mirror
|
||||
ASCII Art
|
||||
Glitch
|
||||
Time-Lapse
|
||||
Day to Night
|
||||
Surreal
|
||||
Black and White
|
||||
Sepia Tone
|
||||
Vintage Film
|
||||
Grainy Texture
|
||||
High Key Lighting
|
||||
Low Key Lighting
|
||||
Cross Processed Film
|
||||
Infrared Photography
|
||||
Photocopy
|
||||
Pencil Drawing
|
||||
Pop Art Filter
|
||||
Mosaic Filter
|
||||
Selective Desaturation
|
||||
Torn Paper
|
||||
Tilt-Shift Photography
|
||||
Double Exposure
|
||||
Polaroid Style Frame
|
||||
Water Drops Texture
|
||||
Polarizer
|
||||
Chinese Painting
|
||||
Water Droplets Texture
|
||||
Polarization
|
||||
Color Inversion
|
||||
Fish-eye Lens
|
||||
Soft Focus
|
||||
Solarize Filter
|
||||
Edge Detection
|
||||
Oil Painting
|
||||
Reflection
|
||||
Mirror Image
|
||||
Time-Lapse Photography
|
||||
Day to Night Transition
|
||||
Surreal Art Style
|
||||
Abstract Expressionism
|
||||
Acrylic Painting
|
||||
Anime
|
||||
Art Deco
|
||||
Biomorphic Abstraction
|
||||
Black and White Photograph
|
||||
Cartoon
|
||||
Charcoal Sketch
|
||||
Chibi Anime
|
||||
Chinese Painting
|
||||
Classicist Painting
|
||||
Collage
|
||||
Concept Art
|
||||
Cyberpunk
|
||||
Dada Art
|
||||
Digital Art
|
||||
Fantasy Art
|
||||
Fashion Art
|
||||
Fashion Sketch
|
||||
Fish-Eye lens Photograph
|
||||
Goth Art
|
||||
Graffiti
|
||||
Harlem Renaissance
|
||||
High Key Photograph
|
||||
Hyperrealist Pencil Sketch
|
||||
Impressionist Painting
|
||||
Josei Anime
|
||||
Long Exposure Photograph
|
||||
Low Key Photograph
|
||||
Macro Photograph
|
||||
Manga
|
||||
Metal Sculpture
|
||||
Mid Century Modern Illustration
|
||||
Mixed Media
|
||||
Modern Art
|
||||
Moe Anime
|
||||
Nihonga
|
||||
Origami
|
||||
Paper Mache
|
||||
Pen and Ink
|
||||
Pencil Sketch
|
||||
Photograph
|
||||
Photorealism
|
||||
Pinup Art
|
||||
Romanticist Painting
|
||||
Sci-Fi Art
|
||||
Semi Realistic Fantasy Art
|
||||
Semi Realistic Cyberpunk Art
|
||||
Shallow Depth of Field Photograph
|
||||
Steam Punk Art
|
||||
Stone Sculpture
|
||||
Superhero Comic
|
||||
Surrealist Art
|
||||
Tempura Painting
|
||||
Underground Comic
|
||||
Watercolor Painting
|
||||
Zulu Urban Art
|
||||
@@ -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
|
||||
)
|
||||
|
||||
%python_exec% -s -m pip install --upgrade --force llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
|
||||
|
||||
%python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
|
||||
|
||||
|
||||
) else (
|
||||
echo Installing with system Python
|
||||
for /f "delims=" %%i in (%requirements_txt%) do (
|
||||
|
||||
@@ -24,7 +24,7 @@ class SpeechRecognition:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/audio"
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -48,7 +48,7 @@ class SpeechSynthesis:
|
||||
OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/audio"
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
def run(self, text):
|
||||
# print(session_history)
|
||||
@@ -82,7 +82,7 @@ class GamePal:
|
||||
OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/audio"
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
def run(self, input_text,input_num,python_code):
|
||||
exec(python_code)
|
||||
|
||||
@@ -1,7 +1,37 @@
|
||||
import openai
|
||||
import time
|
||||
import urllib.error
|
||||
import re,json
|
||||
import re,json,os,string,random
|
||||
import folder_paths
|
||||
import hashlib
|
||||
import codecs,sys
|
||||
import importlib.util
|
||||
|
||||
|
||||
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()
|
||||
return unique_hash
|
||||
|
||||
def generate_random_string(length):
|
||||
letters = string.ascii_letters + string.digits
|
||||
return ''.join(random.choice(letters) for _ in range(length))
|
||||
|
||||
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("*")
|
||||
|
||||
# 判断是否是azure服务
|
||||
def is_azure_url(url):
|
||||
@@ -28,6 +58,108 @@ def openai_client(key,url):
|
||||
)
|
||||
return client
|
||||
|
||||
def ZhipuAI_client(key):
|
||||
|
||||
try:
|
||||
if is_installed('zhipuai')==False:
|
||||
import subprocess
|
||||
|
||||
# 安装
|
||||
print('#pip install zhipuai')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'zhipuai'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0:
|
||||
print("#install success")
|
||||
from zhipuai import ZhipuAI
|
||||
else:
|
||||
print("#install error")
|
||||
|
||||
else:
|
||||
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 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()
|
||||
|
||||
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
|
||||
|
||||
|
||||
|
||||
|
||||
def chat(client, model_name,messages ):
|
||||
@@ -36,10 +168,21 @@ def chat(client, model_name,messages ):
|
||||
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
|
||||
)
|
||||
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
|
||||
@@ -48,7 +191,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)
|
||||
@@ -71,18 +215,28 @@ class ChatGPTNode:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
model_list=llama_modes_list+[
|
||||
"gpt-3.5-turbo",
|
||||
"gpt-3.5-turbo-0125",
|
||||
"gpt-35-turbo",
|
||||
"gpt-3.5-turbo-16k",
|
||||
"gpt-3.5-turbo-16k-0613",
|
||||
"gpt-4-0613",
|
||||
"gpt-4-1106-preview",
|
||||
"glm-4"
|
||||
]
|
||||
return {
|
||||
"required": {
|
||||
"api_key":("KEY", {"default": "", "multiline": True}),
|
||||
"api_url":("URL", {"default": "", "multiline": True}),
|
||||
"api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
|
||||
"api_url":("URL", {"default": "", "multiline": True,"dynamicPrompts": False}),
|
||||
"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"],
|
||||
{"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}),
|
||||
},
|
||||
@@ -105,8 +259,8 @@ class ChatGPTNode:
|
||||
api_url,
|
||||
prompt,
|
||||
system_content,
|
||||
model,
|
||||
seed,context_size,unique_id = None, extra_pnginfo=None):
|
||||
model,
|
||||
seed,context_size,unique_id = None, extra_pnginfo=None):
|
||||
# print(api_key!='',api_url,prompt,system_content,model,seed)
|
||||
# 可以选择保留会话历史以维持上下文记忆
|
||||
# 或者在此处清除会话历史 self.session_history.clear()
|
||||
@@ -124,8 +278,16 @@ class ChatGPTNode:
|
||||
if is_azure_url(api_url):
|
||||
client=azure_client(api_key,api_url)
|
||||
else:
|
||||
client=openai_client(api_key,api_url)
|
||||
print('openai url')
|
||||
# 根据用户选择的模型,设置相应的接口和模型名称
|
||||
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')
|
||||
|
||||
# 把用户的提示添加到会话历史中
|
||||
# 调用API时传递整个会话历史
|
||||
@@ -168,7 +330,10 @@ class ShowTextForGPT:
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"forceInput": True,"dynamicPrompts": False}),
|
||||
}
|
||||
},
|
||||
"optional":{
|
||||
"output_dir": ("STRING",{"forceInput": True,"default": "","multiline": True,"dynamicPrompts": False}),
|
||||
}
|
||||
}
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
@@ -177,11 +342,65 @@ class ShowTextForGPT:
|
||||
OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/GPT"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
def run(self, text):
|
||||
# print(session_history)
|
||||
def run(self, text,output_dir=[""]):
|
||||
|
||||
# 类型纠正
|
||||
texts=[]
|
||||
for t in text:
|
||||
if not isinstance(t, str):
|
||||
t = str(t)
|
||||
texts.append(t)
|
||||
|
||||
text=texts
|
||||
|
||||
if len(output_dir)==1 and (output_dir[0]=='' or os.path.dirname(output_dir[0])==''):
|
||||
t='\n'.join(text)
|
||||
output_dir=[
|
||||
os.path.join(folder_paths.get_temp_directory(),
|
||||
get_unique_hash(t)+'.txt'
|
||||
)
|
||||
]
|
||||
elif len(output_dir)==1:
|
||||
base=os.path.basename(output_dir[0])
|
||||
t='\n'.join(text)
|
||||
if base=='' or os.path.splitext(base)[1]=='':
|
||||
base=get_unique_hash(t)+'.txt'
|
||||
output_dir=[
|
||||
os.path.join(output_dir[0],
|
||||
base
|
||||
)
|
||||
]
|
||||
# elif len(output_dir)>1:
|
||||
|
||||
|
||||
|
||||
if len(output_dir)==1 and len(text)>1:
|
||||
output_dir=[output_dir[0] for _ in range(len(text))]
|
||||
|
||||
for i in range(len(text)):
|
||||
|
||||
o_fp=output_dir[i]
|
||||
dirp=os.path.dirname(o_fp)
|
||||
if dirp=='':
|
||||
dirp=folder_paths.get_temp_directory()
|
||||
o_fp=os.path.join(folder_paths.get_temp_directory(),o_fp
|
||||
)
|
||||
|
||||
if not os.path.exists(dirp):
|
||||
os.mkdir(dirp)
|
||||
|
||||
if not os.path.splitext(o_fp)[1].lower()=='.txt':
|
||||
o_fp=o_fp+'.txt'
|
||||
|
||||
t=text[i]
|
||||
with open(o_fp, 'w') as file:
|
||||
file.write(t)
|
||||
|
||||
# print(text)
|
||||
return {"ui": {"text": text}, "result": (text,)}
|
||||
|
||||
|
||||
|
||||
class CharacterInText:
|
||||
@@ -207,11 +426,61 @@ 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)
|
||||
b=1 if character in text else 0
|
||||
b=1 if character.lower() in text.lower() else 0
|
||||
|
||||
return (b+start_index,)
|
||||
|
||||
class TextSplitByDelimiter:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True,"dynamicPrompts": False}),
|
||||
"delimiter":("STRING", {"multiline": False,"default":",","dynamicPrompts": False}),
|
||||
"start_index": ("INT", {
|
||||
"default": 0,
|
||||
"min": 0, #Minimum value
|
||||
"max": 1000, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"skip_every": ("INT", {
|
||||
"default": 0,
|
||||
"min": 0, #Minimum value
|
||||
"max": 10, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"max_count": ("INT", {
|
||||
"default": 10,
|
||||
"min": 1, #Minimum value
|
||||
"max": 1000, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "run"
|
||||
# OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
def run(self, text,delimiter,start_index,skip_every,max_count):
|
||||
|
||||
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,)
|
||||
|
||||
@@ -0,0 +1,283 @@
|
||||
import os,sys
|
||||
import folder_paths
|
||||
|
||||
from PIL import Image
|
||||
import importlib.util
|
||||
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import json
|
||||
import torch
|
||||
import random
|
||||
|
||||
|
||||
# from clip_interrogator import Config, Interrogator
|
||||
|
||||
|
||||
global _available
|
||||
_available=False
|
||||
|
||||
def is_installed(package):
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return spec is not None
|
||||
|
||||
|
||||
try:
|
||||
if is_installed('clip_interrogator')==False:
|
||||
import subprocess
|
||||
|
||||
# 安装
|
||||
print('#pip install clip-interrogator==0.6.0')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'clip-interrogator==0.6.0'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0:
|
||||
print("#install success")
|
||||
from clip_interrogator import Config, Interrogator
|
||||
_available=True
|
||||
else:
|
||||
print("#install error")
|
||||
|
||||
else:
|
||||
from clip_interrogator import Config, Interrogator
|
||||
_available=True
|
||||
|
||||
except:
|
||||
_available=False
|
||||
|
||||
try:
|
||||
from transformers import AutoProcessor, BlipForConditionalGeneration
|
||||
except:
|
||||
_available=False
|
||||
print('pls check transformers.__version__>=4.36.0:: AutoProcessor, BlipForConditionalGeneration')
|
||||
|
||||
|
||||
|
||||
def load_caption_model(model_path,config,t='blip-base'):
|
||||
dtype=torch.float16 if config.device == 'cuda' else torch.float32
|
||||
caption_model = BlipForConditionalGeneration.from_pretrained(model_path, torch_dtype=dtype)
|
||||
|
||||
caption_processor = AutoProcessor.from_pretrained(model_path)
|
||||
|
||||
caption_model.eval()
|
||||
if not config.caption_offload:
|
||||
caption_model = caption_model.to(config.device)
|
||||
|
||||
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")
|
||||
|
||||
|
||||
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'
|
||||
|
||||
|
||||
# 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 image_analysis_fn(ci,image):
|
||||
image = image.convert('RGB')
|
||||
image_features = ci.image_to_features(image)
|
||||
|
||||
top_mediums = ci.mediums.rank(image_features, 5)
|
||||
top_artists = ci.artists.rank(image_features, 5)
|
||||
top_movements = ci.movements.rank(image_features, 5)
|
||||
top_trendings = ci.trendings.rank(image_features, 5)
|
||||
top_flavors = ci.flavors.rank(image_features, 5)
|
||||
|
||||
medium_ranks = {medium: sim for medium, sim in zip(top_mediums, ci.similarities(image_features, top_mediums))}
|
||||
artist_ranks = {artist: sim for artist, sim in zip(top_artists, ci.similarities(image_features, top_artists))}
|
||||
movement_ranks = {movement: sim for movement, sim in zip(top_movements, ci.similarities(image_features, top_movements))}
|
||||
trending_ranks = {trending: sim for trending, sim in zip(top_trendings, ci.similarities(image_features, top_trendings))}
|
||||
flavor_ranks = {flavor: sim for flavor, sim in zip(top_flavors, ci.similarities(image_features, top_flavors))}
|
||||
|
||||
return {
|
||||
"medium_ranks":medium_ranks,
|
||||
"artist_ranks":artist_ranks,
|
||||
"movement_ranks":movement_ranks,
|
||||
"trending_ranks":trending_ranks,
|
||||
"flavor_ranks":flavor_ranks
|
||||
}
|
||||
|
||||
|
||||
def generate_sentences(data):
|
||||
sentences = []
|
||||
|
||||
# Get the length of data
|
||||
data_length = len(data)
|
||||
|
||||
# Use a recursive function to handle variable-length data
|
||||
def generate_recursive(index, current_sentence, current_score):
|
||||
# Check if recursion is complete
|
||||
if index == data_length:
|
||||
sentences.append({"sentence": current_sentence, "score": current_score})
|
||||
return
|
||||
|
||||
# Get the current level data
|
||||
current_data = data[index]
|
||||
|
||||
# Iterate through the current level data
|
||||
for phrase in current_data:
|
||||
sentence = current_sentence + ("," if current_sentence.strip() else "") + phrase
|
||||
score = current_score + current_data[phrase]
|
||||
generate_recursive(index + 1, sentence, score)
|
||||
|
||||
# Start recursive generation of sentences
|
||||
generate_recursive(0, "", 0)
|
||||
|
||||
# Sort the generated sentences by score in descending order
|
||||
sentences.sort(key=lambda x: x["score"], reverse=True)
|
||||
|
||||
def get_random_elements(elements, num):
|
||||
return random.sample(elements, num)
|
||||
|
||||
ps = get_random_elements(sentences, 5)
|
||||
ps = [s["sentence"] for s in sorted(ps, key=lambda x: x["score"], reverse=True)]
|
||||
|
||||
return ps
|
||||
|
||||
|
||||
|
||||
|
||||
def image_to_prompt(ci,image, mode):
|
||||
ci.config.chunk_size = 2048 if ci.config.clip_model_name == "ViT-L-14/openai" else 1024
|
||||
ci.config.flavor_intermediate_count = 2048 if ci.config.clip_model_name == "ViT-L-14/openai" else 1024
|
||||
image = image.convert('RGB')
|
||||
if mode == 'best':
|
||||
return ci.interrogate(image)
|
||||
elif mode == 'classic':
|
||||
return ci.interrogate_classic(image)
|
||||
elif mode == 'fast':
|
||||
return ci.interrogate_fast(image)
|
||||
elif mode == 'negative':
|
||||
return ci.interrogate_negative(image)
|
||||
|
||||
# image = Image.open(image_path).convert('RGB')
|
||||
# ci = Interrogator(Config(clip_model_name="ViT-L-14/openai"))
|
||||
# print(ci.interrogate(image))
|
||||
|
||||
|
||||
class ClipInterrogator:
|
||||
|
||||
global _available
|
||||
available=_available
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image": ("IMAGE",),
|
||||
"prompt_mode": (['fast','classic','best','negative'],),
|
||||
"image_analysis": (["off","on"],),
|
||||
},
|
||||
|
||||
# "optional":{
|
||||
# "output":("CLIPINTERROGATOR", {"multiline": True,"default": "", "dynamicPrompts": False})
|
||||
# },
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING","STRING",)
|
||||
RETURN_NAMES = ("prompt","random_samples",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,True,)
|
||||
global ci
|
||||
ci = None
|
||||
def run(self,image,prompt_mode,image_analysis):
|
||||
global ci
|
||||
|
||||
prompt_mode=prompt_mode[0]
|
||||
analysis=image_analysis[0]
|
||||
|
||||
prompt_result=[]
|
||||
analysis_result=[]
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(image)*(2 if analysis=='on' else 1))
|
||||
|
||||
if ci==None:
|
||||
config=Config(
|
||||
clip_model_name="ViT-L-14/openai",
|
||||
device="cuda" if torch.cuda.is_available() else "cpu",
|
||||
download_cache=True,
|
||||
clip_model_path=cache_path,
|
||||
cache_path=cache_path
|
||||
)
|
||||
config.apply_low_vram_defaults()
|
||||
|
||||
caption_model,caption_processor=load_caption_model(caption_model_path,config)
|
||||
|
||||
config.caption_model= caption_model
|
||||
config.caption_processor= caption_processor
|
||||
|
||||
ci = Interrogator(config)
|
||||
# else:
|
||||
# simple_lama.model.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
for i in range(len(image)):
|
||||
im=image[i]
|
||||
|
||||
im=tensor2pil(im)
|
||||
im=im.convert('RGB')
|
||||
|
||||
if analysis=='on':
|
||||
analysis_res=image_analysis_fn(ci,im)
|
||||
analysis_result.append( analysis_res )
|
||||
pbar.update(1)
|
||||
|
||||
prompt=image_to_prompt(ci,im,prompt_mode)
|
||||
pbar.update(1)
|
||||
prompt_result.append(prompt)
|
||||
|
||||
|
||||
# result.save("inpainted.png")
|
||||
if ci.config.clip_offload and not ci.clip_offloaded:
|
||||
ci.clip_model = ci.clip_model.to('cpu')
|
||||
ci.clip_offloaded = True
|
||||
|
||||
if ci.config.caption_offload and not ci.caption_offloaded:
|
||||
ci.caption_model = ci.caption_model.to('cpu')
|
||||
ci.caption_offloaded = True
|
||||
|
||||
# analysis_result=[]
|
||||
# items = app.graph.getNodeById(31).widgets[2].value["items"]
|
||||
|
||||
random_samples=[]
|
||||
|
||||
for r in analysis_result:
|
||||
random_sample = generate_sentences([r['medium_ranks'], r['artist_ranks'],r['movement_ranks'],r['trending_ranks'],r['flavor_ranks']])
|
||||
for s in random_sample:
|
||||
random_samples.append(s)
|
||||
# print(len(random_samples))
|
||||
# print('-----')
|
||||
# print( random_samples)
|
||||
return {
|
||||
"ui":{
|
||||
"prompt": prompt_result,
|
||||
"analysis":analysis_result,
|
||||
"random_samples":random_samples
|
||||
},
|
||||
"result": (prompt_result,random_samples,)}
|
||||
@@ -1,272 +0,0 @@
|
||||
#### Thanks:
|
||||
# [ComfyUI-CLIPSeg](https://github.com/biegert/ComfyUI-CLIPSeg/tree/main)
|
||||
|
||||
from transformers import CLIPSegProcessor, CLIPSegForImageSegmentation
|
||||
|
||||
from PIL import Image
|
||||
import torch
|
||||
import torchvision.transforms as T
|
||||
import numpy as np
|
||||
|
||||
from torchvision.transforms.functional import to_pil_image
|
||||
import matplotlib.pyplot as plt
|
||||
import matplotlib.cm as cm
|
||||
|
||||
|
||||
import cv2
|
||||
|
||||
from scipy.ndimage import gaussian_filter
|
||||
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import warnings,os
|
||||
warnings.filterwarnings("ignore", category=UserWarning, module="torch")
|
||||
warnings.filterwarnings("ignore", category=UserWarning, module="safetensors")
|
||||
|
||||
import folder_paths
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger('CLIPSeg nodes')
|
||||
|
||||
clipseg_model_dir = os.path.join(folder_paths.models_dir, "clipseg")
|
||||
|
||||
if not os.path.exists(clipseg_model_dir):
|
||||
print(f"## clipseg model not found: {clipseg_model_dir},pls download from https://huggingface.co/CIDAS/clipseg-rd64-refined/tree/main")
|
||||
clipseg_model_dir='CIDAS/clipseg-rd64-refined'
|
||||
|
||||
"""Helper methods for CLIPSeg nodes"""
|
||||
|
||||
# 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 tensor_to_numpy(tensor: torch.Tensor) -> np.ndarray:
|
||||
"""Convert a tensor to a numpy array and scale its values to 0-255."""
|
||||
array = tensor.numpy().squeeze()
|
||||
return (array * 255).astype(np.uint8)
|
||||
|
||||
def numpy_to_tensor(array: np.ndarray) -> torch.Tensor:
|
||||
"""Convert a numpy array to a tensor and scale its values from 0-255 to 0-1."""
|
||||
array = array.astype(np.float32) / 255.0
|
||||
return torch.from_numpy(array)[None,]
|
||||
|
||||
def apply_colormap(mask: torch.Tensor, colormap) -> np.ndarray:
|
||||
"""Apply a colormap to a tensor and convert it to a numpy array."""
|
||||
colored_mask = colormap(mask.numpy())[:, :, :3]
|
||||
return (colored_mask * 255).astype(np.uint8)
|
||||
|
||||
def resize_image(image: np.ndarray, dimensions: Tuple[int, int]) -> np.ndarray:
|
||||
"""Resize an image to the given dimensions using linear interpolation."""
|
||||
return cv2.resize(image, dimensions, interpolation=cv2.INTER_LINEAR)
|
||||
|
||||
def overlay_image(background: np.ndarray, foreground: np.ndarray, alpha: float) -> np.ndarray:
|
||||
"""Overlay the foreground image onto the background with a given opacity (alpha)."""
|
||||
return cv2.addWeighted(background, 1 - alpha, foreground, alpha, 0)
|
||||
|
||||
def dilate_mask(mask: torch.Tensor, dilation_factor: float) -> torch.Tensor:
|
||||
"""Dilate a mask using a square kernel with a given dilation factor."""
|
||||
kernel_size = int(dilation_factor * 2) + 1
|
||||
kernel = np.ones((kernel_size, kernel_size), np.uint8)
|
||||
mask_dilated = cv2.dilate(mask.numpy(), kernel, iterations=1)
|
||||
return torch.from_numpy(mask_dilated)
|
||||
|
||||
|
||||
|
||||
class CLIPSeg:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
"""
|
||||
Return a dictionary which contains config for all input fields.
|
||||
Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT".
|
||||
Input types "INT", "STRING" or "FLOAT" are special values for fields on the node.
|
||||
The type can be a list for selection.
|
||||
|
||||
Returns: `dict`:
|
||||
- Key input_fields_group (`string`): Can be either required, hidden or optional. A node class must have property `required`
|
||||
- Value input_fields (`dict`): Contains input fields config:
|
||||
* Key field_name (`string`): Name of a entry-point method's argument
|
||||
* Value field_config (`tuple`):
|
||||
+ First value is a string indicate the type of field or a list for selection.
|
||||
+ Secound value is a config for type "INT", "STRING" or "FLOAT".
|
||||
"""
|
||||
return {"required":
|
||||
{
|
||||
"image": ("IMAGE",),
|
||||
"text": ("STRING", {"multiline": False,"dynamicPrompts": False}),
|
||||
|
||||
},
|
||||
"optional":
|
||||
{
|
||||
"blur": ("FLOAT", {"min": 0, "max": 15, "step": 0.1, "default": 7}),
|
||||
"threshold": ("FLOAT", {"min": 0, "max": 1, "step": 0.05, "default": 0.4}),
|
||||
"dilation_factor": ("INT", {"min": 0, "max": 10, "step": 1, "default": 4}),
|
||||
}
|
||||
}
|
||||
|
||||
CATEGORY = "♾️Mixlab/mask"
|
||||
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
|
||||
RETURN_NAMES = ("Mask","Heatmap Mask", "BW Mask")
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (False,False,False,)
|
||||
|
||||
FUNCTION = "segment_image"
|
||||
def segment_image(self, image: torch.Tensor, text: str, blur: float, threshold: float, dilation_factor: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Create a segmentation mask from an image and a text prompt using CLIPSeg.
|
||||
|
||||
Args:
|
||||
image (torch.Tensor): The image to segment.
|
||||
text (str): The text prompt to use for segmentation.
|
||||
blur (float): How much to blur the segmentation mask.
|
||||
threshold (float): The threshold to use for binarizing the segmentation mask.
|
||||
dilation_factor (int): How much to dilate the segmentation mask.
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: The segmentation mask, the heatmap mask, and the binarized mask.
|
||||
"""
|
||||
|
||||
# Convert the Tensor to a PIL image
|
||||
image_np = image.numpy().squeeze() # Remove the first dimension (batch size of 1)
|
||||
# Convert the numpy array back to the original range (0-255) and data type (uint8)
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
# Create a PIL image from the numpy array
|
||||
i = Image.fromarray(image_np, mode="RGB")
|
||||
|
||||
processor = CLIPSegProcessor.from_pretrained(clipseg_model_dir)
|
||||
model = CLIPSegForImageSegmentation.from_pretrained(clipseg_model_dir)
|
||||
|
||||
prompt = text
|
||||
|
||||
input_prc = processor(text=prompt, images=i, padding="max_length", return_tensors="pt")
|
||||
|
||||
# Predict the segemntation mask
|
||||
with torch.no_grad():
|
||||
outputs = model(**input_prc)
|
||||
|
||||
tensor = torch.sigmoid(outputs[0]) # get the mask
|
||||
|
||||
# Apply a threshold to the original tensor to cut off low values
|
||||
thresh = threshold
|
||||
tensor_thresholded = torch.where(tensor > thresh, tensor, torch.tensor(0, dtype=torch.float))
|
||||
|
||||
# Apply Gaussian blur to the thresholded tensor
|
||||
sigma = blur
|
||||
tensor_smoothed = gaussian_filter(tensor_thresholded.numpy(), sigma=sigma)
|
||||
tensor_smoothed = torch.from_numpy(tensor_smoothed)
|
||||
|
||||
# Normalize the smoothed tensor to [0, 1]
|
||||
mask_normalized = (tensor_smoothed - tensor_smoothed.min()) / (tensor_smoothed.max() - tensor_smoothed.min())
|
||||
|
||||
# Dilate the normalized mask
|
||||
mask_dilated = dilate_mask(mask_normalized, dilation_factor)
|
||||
|
||||
# Convert the mask to a heatmap and a binary mask
|
||||
heatmap = apply_colormap(mask_dilated, cm.viridis)
|
||||
binary_mask = apply_colormap(mask_dilated, cm.Greys_r)
|
||||
|
||||
# Overlay the heatmap and binary mask on the original image
|
||||
dimensions = (image_np.shape[1], image_np.shape[0])
|
||||
heatmap_resized = resize_image(heatmap, dimensions)
|
||||
binary_mask_resized = resize_image(binary_mask, dimensions)
|
||||
|
||||
alpha_heatmap, alpha_binary = 0.5, 1
|
||||
overlay_heatmap = overlay_image(image_np, heatmap_resized, alpha_heatmap)
|
||||
overlay_binary = overlay_image(image_np, binary_mask_resized, alpha_binary)
|
||||
|
||||
# Convert the numpy arrays to tensors
|
||||
image_out_heatmap = numpy_to_tensor(overlay_heatmap)
|
||||
image_out_binary = numpy_to_tensor(overlay_binary)
|
||||
|
||||
# Save or display the resulting binary mask
|
||||
binary_mask_image = Image.fromarray(binary_mask_resized[..., 0])
|
||||
|
||||
# convert PIL image to numpy array
|
||||
tensor_bw = binary_mask_image.convert("L")
|
||||
tensor_bw=pil2tensor(tensor_bw)
|
||||
# tensor_bw = np.array(tensor_bw).astype(np.float32) / 255.0
|
||||
# tensor_bw = torch.from_numpy(tensor_bw)[None,]
|
||||
# tensor_bw = tensor_bw.squeeze(0)[..., 0]
|
||||
|
||||
return (tensor_bw, image_out_heatmap, image_out_binary,)
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
class CombineMasks:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{
|
||||
"input_image": ("IMAGE", ),
|
||||
"mask_1": ("MASK", ),
|
||||
"mask_2": ("MASK", ),
|
||||
},
|
||||
"optional":
|
||||
{
|
||||
"mask_3": ("MASK",),
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "♾️Mixlab/mask"
|
||||
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
|
||||
RETURN_NAMES = ("Combined Mask","Heatmap Mask", "BW Mask")
|
||||
|
||||
FUNCTION = "combine_masks"
|
||||
|
||||
def combine_masks(self, input_image: torch.Tensor, mask_1: torch.Tensor, mask_2: torch.Tensor, mask_3: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""A method that combines two or three masks into one mask. Takes in tensors and returns the mask as a tensor, as well as the heatmap and binary mask as tensors."""
|
||||
|
||||
# Combine masks
|
||||
if mask_1 is not None:
|
||||
mask_1 = mask_1.squeeze()
|
||||
if mask_2 is not None:
|
||||
mask_2 = mask_2.squeeze()
|
||||
if mask_3 is not None:
|
||||
mask_3 = mask_3.squeeze()
|
||||
|
||||
print(mask_1.shape,mask_2.shape , mask_3.shape)
|
||||
combined_mask = mask_1 + mask_2 + mask_3 if mask_3 is not None else mask_1 + mask_2
|
||||
# print(combined_mask)
|
||||
|
||||
# Convert image and masks to numpy arrays
|
||||
image_np = tensor_to_numpy(input_image)
|
||||
heatmap = apply_colormap(combined_mask, cm.viridis)
|
||||
binary_mask = apply_colormap(combined_mask, cm.Greys_r)
|
||||
|
||||
# Resize heatmap and binary mask to match the original image dimensions
|
||||
dimensions = (image_np.shape[1], image_np.shape[0])
|
||||
# print('heatmap',heatmap)
|
||||
if dimensions is None or dimensions[0] == 0 or dimensions[1] == 0:
|
||||
raise ValueError("Invalid dimensions")
|
||||
|
||||
heatmap_resized = resize_image(heatmap, dimensions)
|
||||
binary_mask_resized = resize_image(binary_mask, dimensions)
|
||||
|
||||
# Overlay the heatmap and binary mask onto the original image
|
||||
alpha_heatmap, alpha_binary = 0.5, 1
|
||||
overlay_heatmap = overlay_image(image_np, heatmap_resized, alpha_heatmap)
|
||||
overlay_binary = overlay_image(image_np, binary_mask_resized, alpha_binary)
|
||||
|
||||
# Convert overlays to tensors
|
||||
image_out_heatmap = numpy_to_tensor(overlay_heatmap)
|
||||
image_out_binary = numpy_to_tensor(overlay_binary)
|
||||
|
||||
return combined_mask, image_out_heatmap, image_out_binary
|
||||
|
||||
# A dictionary that contains all nodes you want to export with their names
|
||||
# NOTE: names should be globally unique
|
||||
# NODE_CLASS_MAPPINGS = {
|
||||
# "CLIPSeg": CLIPSeg,
|
||||
# "CombineSegMasks": CombineMasks,
|
||||
# }
|
||||
@@ -1,15 +1,54 @@
|
||||
import os
|
||||
import os,sys
|
||||
import folder_paths
|
||||
from simple_lama_inpainting import SimpleLama
|
||||
from PIL import Image
|
||||
|
||||
from PIL import Image
|
||||
import importlib.util
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
|
||||
global _available
|
||||
_available=False
|
||||
|
||||
|
||||
llma_model_path=os.path.join(folder_paths.models_dir, "lama/big-lama.pt")
|
||||
def is_installed(package):
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return spec is not None
|
||||
|
||||
if is_installed('simple_lama_inpainting')==False:
|
||||
import subprocess
|
||||
from packaging import version
|
||||
|
||||
if version.parse(torch.__version__)>=version.parse('2.1'):
|
||||
# 安装
|
||||
print('#pip install simple_lama_inpainting')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'simple_lama_inpainting'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0:
|
||||
print("#install success")
|
||||
from simple_lama_inpainting import SimpleLama
|
||||
_available=True
|
||||
else:
|
||||
print("#install error")
|
||||
else:
|
||||
print('#pls check your torch version >= 2.1')
|
||||
|
||||
else:
|
||||
from simple_lama_inpainting import SimpleLama
|
||||
_available=True
|
||||
|
||||
|
||||
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")
|
||||
@@ -38,6 +77,8 @@ def pil2tensor(image):
|
||||
|
||||
|
||||
class LaMaInpainting:
|
||||
global _available
|
||||
available=_available
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
@@ -52,7 +93,7 @@ class LaMaInpainting:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/image"
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
@@ -0,0 +1,274 @@
|
||||
|
||||
import scipy.ndimage
|
||||
import torch
|
||||
|
||||
import numpy as np
|
||||
# from PIL import Image, ImageDraw
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
from comfy.cli_args import args
|
||||
import cv2,os
|
||||
from nodes import MAX_RESOLUTION, SaveImage, common_ksampler
|
||||
import folder_paths,random
|
||||
|
||||
# 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 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],
|
||||
[1, 1, 1],
|
||||
[c, 1, c]])
|
||||
mask = mask.reshape((-1, mask.shape[-2], mask.shape[-1]))
|
||||
out = []
|
||||
for m in mask:
|
||||
output = m.numpy()
|
||||
for _ in range(abs(expand)):
|
||||
if expand < 0:
|
||||
output = scipy.ndimage.grey_erosion(output, footprint=kernel)
|
||||
else:
|
||||
output = scipy.ndimage.grey_dilation(output, footprint=kernel)
|
||||
output = torch.from_numpy(output)
|
||||
out.append(output)
|
||||
return torch.stack(out, dim=0)
|
||||
|
||||
def combine(destination, source, x, y):
|
||||
output = destination.reshape((-1, destination.shape[-2], destination.shape[-1])).clone()
|
||||
source = source.reshape((-1, source.shape[-2], source.shape[-1]))
|
||||
|
||||
left, top = (x, y,)
|
||||
right, bottom = (min(left + source.shape[-1], destination.shape[-1]), min(top + source.shape[-2], destination.shape[-2]))
|
||||
visible_width, visible_height = (right - left, bottom - top,)
|
||||
|
||||
source_portion = source[:, :visible_height, :visible_width]
|
||||
destination_portion = destination[:, top:bottom, left:right]
|
||||
|
||||
#operation == "subtract":
|
||||
output[:, top:bottom, left:right] = destination_portion - source_portion
|
||||
|
||||
output = torch.clamp(output, 0.0, 1.0)
|
||||
|
||||
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
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mask": ("MASK",),
|
||||
"outline_width":("INT", {"default": 10,"min": 1, "max": MAX_RESOLUTION, "step": 1}),
|
||||
"tapered_corners": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('MASK',)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
# 运行的函数
|
||||
def run(self, mask, outline_width, tapered_corners):
|
||||
|
||||
m1=grow(mask,outline_width,tapered_corners)
|
||||
m2=grow(mask,-outline_width,tapered_corners)
|
||||
|
||||
m3=combine(m1,m2,0,0)
|
||||
|
||||
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:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mask": ("MASK",),
|
||||
"start_offset":("INT", {"default": 1,
|
||||
"min": -150,
|
||||
"max": 150,
|
||||
"step": 1,
|
||||
"display": "slider"}),
|
||||
"feathering_weight":("FLOAT", {"default": 0.1,
|
||||
"min": 0.0,
|
||||
"max": 1,
|
||||
"step": 0.1,
|
||||
"display": "slider"})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('MASK',)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
# 运行的函数
|
||||
def run(self,mask,start_offset, feathering_weight):
|
||||
# print(mask.shape,mask.size())
|
||||
|
||||
num,_,_=mask.size()
|
||||
|
||||
masks=[]
|
||||
|
||||
for i in range(num):
|
||||
mm=mask[i]
|
||||
image=tensor2pil(mm)
|
||||
|
||||
# Open the image using PIL
|
||||
image = image.convert("L")
|
||||
if start_offset>0:
|
||||
image=ImageOps.invert(image)
|
||||
|
||||
# Convert the image to a numpy array
|
||||
image_np = np.array(image)
|
||||
|
||||
# Use Canny edge detection to get black contours
|
||||
edges = cv2.Canny(image_np, 30, 150)
|
||||
|
||||
for i in range(0,abs(start_offset)):
|
||||
# int(100*feathering_weight)
|
||||
a=int(abs(start_offset)*0.1*i)
|
||||
# Dilate the black contours to make them wider
|
||||
kernel = np.ones((a, a), np.uint8)
|
||||
|
||||
dilated_edges = cv2.dilate(edges, kernel, iterations=1)
|
||||
# dilated_edges = cv2.erode(edges, kernel, iterations=1)
|
||||
# Smooth the dilated edges using Gaussian blur
|
||||
smoothed_edges = cv2.GaussianBlur(dilated_edges, (5, 5), 0)
|
||||
|
||||
# Adjust the feathering weight
|
||||
feathering_weight = max(0, min(feathering_weight, 1))
|
||||
|
||||
# Blend the smoothed edges with the original image to achieve feathering effect
|
||||
image_np = cv2.addWeighted(image_np, 1, smoothed_edges, feathering_weight, feathering_weight)
|
||||
|
||||
# Convert the result back to PIL image
|
||||
result_image = Image.fromarray(np.uint8(image_np))
|
||||
result_image=result_image.convert("L")
|
||||
|
||||
if start_offset>0:
|
||||
result_image=ImageOps.invert(result_image)
|
||||
|
||||
result_image=result_image.convert("L")
|
||||
mt=pil2tensor(result_image)
|
||||
masks.append(mt)
|
||||
|
||||
# print( mt.size())
|
||||
return (masks,)
|
||||
@@ -1,7 +1,15 @@
|
||||
import random
|
||||
import comfy.utils
|
||||
import json
|
||||
import os
|
||||
import numpy as np
|
||||
from urllib import request, parse
|
||||
import folder_paths
|
||||
from PIL import Image, ImageOps,ImageFilter,ImageEnhance,ImageDraw,ImageSequence, ImageFont
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
|
||||
import hashlib
|
||||
import requests
|
||||
import json
|
||||
|
||||
|
||||
# def queue_prompt(prompt_workflow):
|
||||
@@ -10,6 +18,75 @@ from urllib import request, parse
|
||||
# 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_files_with_extension(directory, extension):
|
||||
|
||||
file_list = []
|
||||
for root, dirs, files in os.walk(directory):
|
||||
for file in files:
|
||||
if file.endswith(extension):
|
||||
file_name = os.path.splitext(file)[0]
|
||||
file_list.append(file_name)
|
||||
return file_list
|
||||
|
||||
def join_with_(text_list,delimiter):
|
||||
joined_text = delimiter.join(text_list)
|
||||
return joined_text
|
||||
|
||||
|
||||
|
||||
def load_json(file_path):
|
||||
try:
|
||||
with open(file_path, 'r') as json_file:
|
||||
data = json.load(json_file)
|
||||
return data
|
||||
except FileNotFoundError:
|
||||
print(f"File not found: {file_path}")
|
||||
return None
|
||||
except json.JSONDecodeError:
|
||||
print(f"Error decoding JSON in file: {file_path}")
|
||||
return None
|
||||
|
||||
def save_json(data_dict, file_path):
|
||||
try:
|
||||
with open(file_path, 'w') as json_file:
|
||||
json.dump(data_dict, json_file, indent=4)
|
||||
print(f"Data saved to {file_path}")
|
||||
except Exception as e:
|
||||
print(f"Error saving JSON to file: {e}")
|
||||
|
||||
# pysss的lora加载器
|
||||
# def get_model_version_info(hash_value):
|
||||
# # http://127.0.0.1:1082
|
||||
# proxies = {'http': 'http://127.0.0.1:1082', 'https': 'https://127.0.0.1:1082'}
|
||||
# api_url = f"https://civitai.com/api/v1/model-versions/by-hash/{hash_value}"
|
||||
# print(api_url)
|
||||
# response = requests.get(api_url,proxies=proxies, verify=False)
|
||||
|
||||
# if response.status_code == 200:
|
||||
# return response.json()
|
||||
# else:
|
||||
# return None
|
||||
|
||||
# def calculate_sha256(file_path):
|
||||
# sha256_hash = hashlib.sha256()
|
||||
# with open(file_path, "rb") as f:
|
||||
# for chunk in iter(lambda: f.read(4096), b""):
|
||||
# sha256_hash.update(chunk)
|
||||
# return sha256_hash.hexdigest()
|
||||
|
||||
|
||||
|
||||
|
||||
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("*")
|
||||
|
||||
|
||||
default_prompt1='''Swing
|
||||
Slide
|
||||
@@ -45,11 +122,170 @@ default_prompt1='''Swing
|
||||
default_prompt1="\n".join([p.strip() for p in default_prompt1.split('\n') if p.strip()!=''])
|
||||
|
||||
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
def addWeight(text, weight=1):
|
||||
if weight == 1:
|
||||
return text
|
||||
else:
|
||||
return f"({text}:{round(weight,2)})"
|
||||
return f"({text}:{round(weight,3)})"
|
||||
|
||||
def prompt_delete_words(sentence, new_words_length):
|
||||
# 使用逗号分割句子,并去除空格
|
||||
words = [word.strip() for word in sentence.split(",")]
|
||||
|
||||
# 计算需要删除的单词数量
|
||||
num_to_delete = len(words) - new_words_length
|
||||
|
||||
words_to=[w for w in words]
|
||||
|
||||
# 逐个删除单词并存储在新列表中
|
||||
new_words = []
|
||||
for i in range(len(words)):
|
||||
if num_to_delete > 0:
|
||||
num_to_delete -= 1
|
||||
else:
|
||||
words_to.pop()
|
||||
if len(words_to)>0:
|
||||
new_words.append(", ".join(words_to))
|
||||
|
||||
return new_words
|
||||
|
||||
# # 测试方法
|
||||
# sentence = "a computer, a glass tablet with a keyboard on a dark background, 3d illustration, reflection, cgi 8k, clear glass, archaic, cut-away, white outline"
|
||||
# new_words_length = 5
|
||||
# result = prompt_delete_words(sentence, new_words_length)
|
||||
# print(result)
|
||||
|
||||
class PromptImage:
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_output_directory()
|
||||
self.type = "output"
|
||||
self.prefix_append = "PromptImage"
|
||||
self.compress_level = 4
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"prompts": ("STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": '',
|
||||
"dynamicPrompts": False
|
||||
}),
|
||||
|
||||
"images": ("IMAGE",{"default": None}),
|
||||
"save_to_image": (["enable", "disable"],),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
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])
|
||||
|
||||
results = list()
|
||||
|
||||
save_to_image=save_to_image[0]=='enable'
|
||||
|
||||
for index in range(len(images)):
|
||||
res=[]
|
||||
imgs=images[index]
|
||||
|
||||
for image in imgs:
|
||||
img=tensor2pil(image)
|
||||
|
||||
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)
|
||||
res.append({
|
||||
"filename": file,
|
||||
"subfolder": subfolder,
|
||||
"type": self.type
|
||||
})
|
||||
counter += 1
|
||||
results.append(res)
|
||||
|
||||
return { "ui": { "_images": results,"prompts":prompts } }
|
||||
|
||||
|
||||
|
||||
|
||||
class PromptSimplification:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": '',
|
||||
"dynamicPrompts": False
|
||||
}),
|
||||
|
||||
"length":("INT", {"default": 5, "min": 1,"max":100, "step": 1, "display": "number"}),
|
||||
|
||||
# "min_value":("FLOAT", {
|
||||
# "default": -2,
|
||||
# "min": -10,
|
||||
# "max": 0xffffffffffffffff,
|
||||
# "step": 0.01,
|
||||
# "display": "number"
|
||||
# }),
|
||||
# "max_value":("FLOAT", {
|
||||
# "default": 2,
|
||||
# "min": -10,
|
||||
# "max": 0xffffffffffffffff,
|
||||
# "step": 0.01,
|
||||
# "display": "number"
|
||||
# }),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("prompts",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
OUTPUT_NODE = True
|
||||
|
||||
# 运行的函数
|
||||
def run(self,prompt,length):
|
||||
length=length[0]
|
||||
result=[]
|
||||
for p in prompt:
|
||||
nps=prompt_delete_words(p,length)
|
||||
for n in nps:
|
||||
result.append(n)
|
||||
|
||||
result= [elem.strip() for elem in result if elem.strip()]
|
||||
|
||||
return {"ui": {"prompts": result}, "result": (result,)}
|
||||
|
||||
|
||||
|
||||
@@ -62,7 +298,8 @@ class PromptSlide:
|
||||
"prompt_keyword": ("STRING",
|
||||
{
|
||||
"multiline": False,
|
||||
"default": ''
|
||||
"default": '',
|
||||
"dynamicPrompts": False
|
||||
}),
|
||||
|
||||
"weight":("FLOAT", {"default": 1, "min": -3,"max": 3,
|
||||
@@ -92,7 +329,7 @@ class PromptSlide:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/prompt"
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -135,6 +372,10 @@ class RandomPrompt:
|
||||
"default": 'sticker, Cartoon, ``'
|
||||
}),
|
||||
"random_sample": (["enable", "disable"],),
|
||||
# "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
|
||||
},
|
||||
"optional":{
|
||||
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -144,14 +385,14 @@ class RandomPrompt:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/prompt"
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
OUTPUT_NODE = True
|
||||
|
||||
|
||||
# 运行的函数
|
||||
def run(self,max_count,mutable_prompt,immutable_prompt,random_sample):
|
||||
def run(self,max_count,mutable_prompt,immutable_prompt,random_sample,seed=0):
|
||||
# print('#运行的函数',mutable_prompt,immutable_prompt,max_count,random_sample)
|
||||
|
||||
# Split the text into an array of words
|
||||
@@ -172,7 +413,10 @@ class RandomPrompt:
|
||||
for w2 in words2:
|
||||
w2=w2.strip()
|
||||
if '``' not in w2:
|
||||
w2=w2+',``'
|
||||
if w2=="":
|
||||
w2='``'
|
||||
else:
|
||||
w2=w2+',``'
|
||||
if w1!='' and w2!='':
|
||||
prompts.append(w2.replace('``', w1))
|
||||
pbar.update(1)
|
||||
@@ -186,69 +430,258 @@ class RandomPrompt:
|
||||
else:
|
||||
prompts = prompts[:min(max_count,len(prompts))]
|
||||
|
||||
prompts= [elem.strip() for elem in prompts if elem.strip()]
|
||||
|
||||
# return (new_prompt)
|
||||
return {"ui": {"prompts": prompts}, "result": (prompts,)}
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# class RunWorkflow:
|
||||
# class LoraPrompt:
|
||||
# @classmethod
|
||||
# def INPUT_TYPES(s):
|
||||
# return {
|
||||
# "required": {
|
||||
# "workflow": ("STRING", {
|
||||
# "multiline": False,
|
||||
# "default": ''
|
||||
# }),
|
||||
# "prompt": ("STRING", {
|
||||
# "multiline": False,
|
||||
# "default": ''
|
||||
# }),
|
||||
# "image": ("IMAGE",),
|
||||
# "input_node": ("STRING", {
|
||||
# "multiline": False,
|
||||
# "default": ''
|
||||
# }),
|
||||
# "output_node": ("STRING", {
|
||||
# "multiline": False,
|
||||
# "default": ''
|
||||
# }),
|
||||
# "lora_name":(sorted(folder_paths.get_filename_list("loras"), key=str.lower),),
|
||||
# "weight": ("FLOAT", {"default": 1, "min": -2, "max": 2,"step":0.01 ,"display": "slider"}),
|
||||
# "force_update": ("BOOLEAN", {"default": False}),
|
||||
# },
|
||||
|
||||
# }
|
||||
|
||||
|
||||
|
||||
# RETURN_TYPES = ("IMAGE","STRING",)
|
||||
# RETURN_TYPES = ("STRING","STRING",any_type)
|
||||
# RETURN_NAMES = ("lora_name","prompt","tags",)
|
||||
|
||||
# FUNCTION = "run"
|
||||
|
||||
# CATEGORY = "♾️Mixlab/workflow"
|
||||
|
||||
# OUTPUT_IS_LIST = (True,)
|
||||
# OUTPUT_NODE = True
|
||||
# CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
# OUTPUT_IS_LIST = (False,False,True,)
|
||||
# # OUTPUT_NODE = True
|
||||
|
||||
# # 运行的函数
|
||||
# def run(self,workflow,prompt,image,input_node,output_node):
|
||||
# print('#运行的函数',prompt,image,input_node,output_node)
|
||||
# workflow=json.loads(workflow)
|
||||
# input_node=input_node.split(".")
|
||||
# workflow[input_node[0]][input_node[1]][input_node[2]]=prompt
|
||||
# def run(self,lora_name,weight,force_update=False):
|
||||
|
||||
# # print('##LoraPrompt',__file__)
|
||||
# # 从本地数据库读取
|
||||
# json_tags_path = os.path.join(os.path.dirname(os.path.dirname(__file__)),r'data/loras_tags.json')
|
||||
|
||||
# workflow_new={}
|
||||
# # 遍历,seed设为随机
|
||||
# for key, value in workflow.items():
|
||||
# if 'inputs' in value:
|
||||
# if 'seed' in value['inputs']:
|
||||
# value['inputs']['seed']= random.randint(1, 18446744073709551614)
|
||||
# workflow_new[key]=value
|
||||
# if not os.path.exists(json_tags_path):
|
||||
# save_json({},json_tags_path)
|
||||
|
||||
# queue_prompt(workflow_new)
|
||||
# print('#运行的函数',workflow_new[input_node[0]])
|
||||
# lora_tags = load_json(json_tags_path)
|
||||
# output_tags = lora_tags.get(lora_name, None) if lora_tags is not None else None
|
||||
# if output_tags is not None:
|
||||
# output_tags = ",".join(output_tags)
|
||||
# print("trainedWords:",output_tags)
|
||||
# else:
|
||||
# output_tags = ""
|
||||
|
||||
# # return (new_prompt)
|
||||
# return {"ui":{"images": []},"result": ([image],['text'],)}
|
||||
|
||||
# lora_path = folder_paths.get_full_path("loras", lora_name)
|
||||
# if output_tags == "" or force_update:
|
||||
# print("calculating lora hash")
|
||||
# LORAsha256 = calculate_sha256(lora_path)
|
||||
# print("requesting infos")
|
||||
# model_info = get_model_version_info(LORAsha256)
|
||||
# if model_info is not None:
|
||||
# if "trainedWords" in model_info:
|
||||
# print("tags found!")
|
||||
# if lora_tags is None:
|
||||
# lora_tags = {}
|
||||
# lora_tags[lora_name] = model_info["trainedWords"]
|
||||
# save_json(lora_tags,json_tags_path)
|
||||
# output_tags = ",".join(model_info["trainedWords"])
|
||||
# print("trainedWords:",output_tags)
|
||||
# else:
|
||||
# print("No informations found.")
|
||||
# if lora_tags is None:
|
||||
# lora_tags = {}
|
||||
# lora_tags[lora_name] = []
|
||||
# save_json(lora_tags,json_tags_path)
|
||||
|
||||
|
||||
# weight = round(weight, 3)
|
||||
# prompt=[]
|
||||
# for p in output_tags.split(','):
|
||||
|
||||
# if weight!=1:
|
||||
# prompt.append('('+p+':'+str(weight)+')')
|
||||
# else:
|
||||
# prompt.append(p)
|
||||
|
||||
# prompt=",".join(prompt)
|
||||
|
||||
# return (lora_name,prompt,output_tags.split(','),)
|
||||
|
||||
|
||||
|
||||
class EmbeddingPrompt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"embedding":(folder_paths.get_filename_list("embeddings"),),
|
||||
"weight": ("FLOAT", {"default": 1, "min": -2, "max": 2,"step":0.01 ,"display": "slider"}),
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
# OUTPUT_NODE = True
|
||||
|
||||
# 运行的函数
|
||||
def run(self,embedding,weight):
|
||||
weight = round(weight, 3)
|
||||
prompt='embedding:'+embedding
|
||||
if weight!=1:
|
||||
prompt='('+prompt+':'+str(weight)+')'
|
||||
prompt=" "+prompt+' '
|
||||
# return (new_prompt)
|
||||
return (prompt,)
|
||||
|
||||
# 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
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"text_list": (any_type,),
|
||||
"delimiter":(["newline","comma","backslash","space"],),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
INPUT_IS_LIST = True # 当true的时候,输入时list,当false的时候,如果输入是list,则会自动包一层for循环调用
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self,text_list,delimiter):
|
||||
delimiter=delimiter[0]
|
||||
if delimiter =='newline':
|
||||
delimiter='\n'
|
||||
elif delimiter=='comma':
|
||||
delimiter=','
|
||||
elif delimiter=='backslash':
|
||||
delimiter='\\'
|
||||
elif delimiter=='space':
|
||||
delimiter=' '
|
||||
t=''
|
||||
if isinstance(text_list, list):
|
||||
t=join_with_(text_list,delimiter)
|
||||
return (t,)
|
||||
@@ -0,0 +1,708 @@
|
||||
import os,sys
|
||||
import folder_paths
|
||||
|
||||
from PIL import Image
|
||||
import importlib.util
|
||||
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from huggingface_hub import hf_hub_download
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torchvision.transforms.functional import normalize
|
||||
# BRIA-RMBG-1.4 / briarmbg.py
|
||||
class REBNCONV(nn.Module):
|
||||
def __init__(self,in_ch=3,out_ch=3,dirate=1,stride=1):
|
||||
super(REBNCONV,self).__init__()
|
||||
|
||||
self.conv_s1 = nn.Conv2d(in_ch,out_ch,3,padding=1*dirate,dilation=1*dirate,stride=stride)
|
||||
self.bn_s1 = nn.BatchNorm2d(out_ch)
|
||||
self.relu_s1 = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
xout = self.relu_s1(self.bn_s1(self.conv_s1(hx)))
|
||||
|
||||
return xout
|
||||
|
||||
## upsample tensor 'src' to have the same spatial size with tensor 'tar'
|
||||
def _upsample_like(src,tar):
|
||||
|
||||
src = F.interpolate(src,size=tar.shape[2:],mode='bilinear')
|
||||
|
||||
return src
|
||||
|
||||
|
||||
### RSU-7 ###
|
||||
class RSU7(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3, img_size=512):
|
||||
super(RSU7,self).__init__()
|
||||
|
||||
self.in_ch = in_ch
|
||||
self.mid_ch = mid_ch
|
||||
self.out_ch = out_ch
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1) ## 1 -> 1/2
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool5 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
|
||||
self.rebnconv7 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
|
||||
self.rebnconv6d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
b, c, h, w = x.shape
|
||||
|
||||
hx = x
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
hx = self.pool3(hx3)
|
||||
|
||||
hx4 = self.rebnconv4(hx)
|
||||
hx = self.pool4(hx4)
|
||||
|
||||
hx5 = self.rebnconv5(hx)
|
||||
hx = self.pool5(hx5)
|
||||
|
||||
hx6 = self.rebnconv6(hx)
|
||||
|
||||
hx7 = self.rebnconv7(hx6)
|
||||
|
||||
hx6d = self.rebnconv6d(torch.cat((hx7,hx6),1))
|
||||
hx6dup = _upsample_like(hx6d,hx5)
|
||||
|
||||
hx5d = self.rebnconv5d(torch.cat((hx6dup,hx5),1))
|
||||
hx5dup = _upsample_like(hx5d,hx4)
|
||||
|
||||
hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
|
||||
hx4dup = _upsample_like(hx4d,hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
|
||||
### RSU-6 ###
|
||||
class RSU6(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU6,self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
|
||||
self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
|
||||
self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
hx = self.pool3(hx3)
|
||||
|
||||
hx4 = self.rebnconv4(hx)
|
||||
hx = self.pool4(hx4)
|
||||
|
||||
hx5 = self.rebnconv5(hx)
|
||||
|
||||
hx6 = self.rebnconv6(hx5)
|
||||
|
||||
|
||||
hx5d = self.rebnconv5d(torch.cat((hx6,hx5),1))
|
||||
hx5dup = _upsample_like(hx5d,hx4)
|
||||
|
||||
hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
|
||||
hx4dup = _upsample_like(hx4d,hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
### RSU-5 ###
|
||||
class RSU5(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU5,self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
|
||||
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
|
||||
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
hx = self.pool3(hx3)
|
||||
|
||||
hx4 = self.rebnconv4(hx)
|
||||
|
||||
hx5 = self.rebnconv5(hx4)
|
||||
|
||||
hx4d = self.rebnconv4d(torch.cat((hx5,hx4),1))
|
||||
hx4dup = _upsample_like(hx4d,hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
### RSU-4 ###
|
||||
class RSU4(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU4,self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
|
||||
hx4 = self.rebnconv4(hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
### RSU-4F ###
|
||||
class RSU4F(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU4F,self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=4)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=8)
|
||||
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=4)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=2)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx2 = self.rebnconv2(hx1)
|
||||
hx3 = self.rebnconv3(hx2)
|
||||
|
||||
hx4 = self.rebnconv4(hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3d,hx2),1))
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2d,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
|
||||
class myrebnconv(nn.Module):
|
||||
def __init__(self, in_ch=3,
|
||||
out_ch=1,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1,
|
||||
dilation=1,
|
||||
groups=1):
|
||||
super(myrebnconv,self).__init__()
|
||||
|
||||
self.conv = nn.Conv2d(in_ch,
|
||||
out_ch,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=groups)
|
||||
self.bn = nn.BatchNorm2d(out_ch)
|
||||
self.rl = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self,x):
|
||||
return self.rl(self.bn(self.conv(x)))
|
||||
|
||||
|
||||
class BriaRMBG(nn.Module):
|
||||
|
||||
def __init__(self,in_ch=3,out_ch=1):
|
||||
super(BriaRMBG,self).__init__()
|
||||
|
||||
self.conv_in = nn.Conv2d(in_ch,64,3,stride=2,padding=1)
|
||||
self.pool_in = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage1 = RSU7(64,32,64)
|
||||
self.pool12 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage2 = RSU6(64,32,128)
|
||||
self.pool23 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage3 = RSU5(128,64,256)
|
||||
self.pool34 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage4 = RSU4(256,128,512)
|
||||
self.pool45 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage5 = RSU4F(512,256,512)
|
||||
self.pool56 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage6 = RSU4F(512,256,512)
|
||||
|
||||
# decoder
|
||||
self.stage5d = RSU4F(1024,256,512)
|
||||
self.stage4d = RSU4(1024,128,256)
|
||||
self.stage3d = RSU5(512,64,128)
|
||||
self.stage2d = RSU6(256,32,64)
|
||||
self.stage1d = RSU7(128,16,64)
|
||||
|
||||
self.side1 = nn.Conv2d(64,out_ch,3,padding=1)
|
||||
self.side2 = nn.Conv2d(64,out_ch,3,padding=1)
|
||||
self.side3 = nn.Conv2d(128,out_ch,3,padding=1)
|
||||
self.side4 = nn.Conv2d(256,out_ch,3,padding=1)
|
||||
self.side5 = nn.Conv2d(512,out_ch,3,padding=1)
|
||||
self.side6 = nn.Conv2d(512,out_ch,3,padding=1)
|
||||
|
||||
# self.outconv = nn.Conv2d(6*out_ch,out_ch,1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.conv_in(hx)
|
||||
#hx = self.pool_in(hxin)
|
||||
|
||||
#stage 1
|
||||
hx1 = self.stage1(hxin)
|
||||
hx = self.pool12(hx1)
|
||||
|
||||
#stage 2
|
||||
hx2 = self.stage2(hx)
|
||||
hx = self.pool23(hx2)
|
||||
|
||||
#stage 3
|
||||
hx3 = self.stage3(hx)
|
||||
hx = self.pool34(hx3)
|
||||
|
||||
#stage 4
|
||||
hx4 = self.stage4(hx)
|
||||
hx = self.pool45(hx4)
|
||||
|
||||
#stage 5
|
||||
hx5 = self.stage5(hx)
|
||||
hx = self.pool56(hx5)
|
||||
|
||||
#stage 6
|
||||
hx6 = self.stage6(hx)
|
||||
hx6up = _upsample_like(hx6,hx5)
|
||||
|
||||
#-------------------- decoder --------------------
|
||||
hx5d = self.stage5d(torch.cat((hx6up,hx5),1))
|
||||
hx5dup = _upsample_like(hx5d,hx4)
|
||||
|
||||
hx4d = self.stage4d(torch.cat((hx5dup,hx4),1))
|
||||
hx4dup = _upsample_like(hx4d,hx3)
|
||||
|
||||
hx3d = self.stage3d(torch.cat((hx4dup,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.stage2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.stage1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
|
||||
#side output
|
||||
d1 = self.side1(hx1d)
|
||||
d1 = _upsample_like(d1,x)
|
||||
|
||||
d2 = self.side2(hx2d)
|
||||
d2 = _upsample_like(d2,x)
|
||||
|
||||
d3 = self.side3(hx3d)
|
||||
d3 = _upsample_like(d3,x)
|
||||
|
||||
d4 = self.side4(hx4d)
|
||||
d4 = _upsample_like(d4,x)
|
||||
|
||||
d5 = self.side5(hx5d)
|
||||
d5 = _upsample_like(d5,x)
|
||||
|
||||
d6 = self.side6(hx6)
|
||||
d6 = _upsample_like(d6,x)
|
||||
|
||||
return [F.sigmoid(d1), F.sigmoid(d2), F.sigmoid(d3), F.sigmoid(d4), F.sigmoid(d5), F.sigmoid(d6)],[hx1d,hx2d,hx3d,hx4d,hx5d,hx6]
|
||||
|
||||
|
||||
|
||||
|
||||
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=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)
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return spec is not None
|
||||
|
||||
|
||||
try:
|
||||
if is_installed('rembg')==False:
|
||||
import subprocess
|
||||
|
||||
# 安装
|
||||
print('#pip install rembg[gpu]')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'rembg[gpu]'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0:
|
||||
print("#install success")
|
||||
from rembg import new_session, remove
|
||||
_available=True
|
||||
else:
|
||||
print("#install error")
|
||||
|
||||
else:
|
||||
from rembg import new_session, remove
|
||||
_available=True
|
||||
|
||||
except:
|
||||
_available=False
|
||||
|
||||
|
||||
def run_briarmbg(images=[]):
|
||||
mroot=U2NET_HOME
|
||||
m=os.path.join(mroot,'briarmbg.pth')
|
||||
if os.path.exists(m)==False:
|
||||
# 下载
|
||||
m1=hf_hub_download("briaai/RMBG-1.4",
|
||||
local_dir=mroot,
|
||||
filename='model.pth',
|
||||
local_dir_use_symlinks=False,
|
||||
endpoint='https://hf-mirror.com')
|
||||
os.rename(m1, m)
|
||||
|
||||
net=BriaRMBG()
|
||||
if torch.cuda.is_available():
|
||||
net.load_state_dict(torch.load(m))
|
||||
net=net.cuda()
|
||||
else:
|
||||
net.load_state_dict(torch.load(m,map_location="cpu"))
|
||||
net.eval()
|
||||
|
||||
masks=[]
|
||||
rgba_images=[]
|
||||
rgb_images=[]
|
||||
for orig_image in images:
|
||||
|
||||
w,h = orig_im_size = orig_image.size
|
||||
|
||||
image = orig_image.convert('RGB')
|
||||
model_input_size = (1024, 1024)
|
||||
image = image.resize(model_input_size, Image.BILINEAR)
|
||||
|
||||
im_np = np.array(image)
|
||||
im_tensor = torch.tensor(im_np, dtype=torch.float32).permute(2,0,1)
|
||||
im_tensor = torch.unsqueeze(im_tensor,0)
|
||||
im_tensor = torch.divide(im_tensor,255.0)
|
||||
im_tensor = normalize(im_tensor,[0.5,0.5,0.5],[1.0,1.0,1.0])
|
||||
if torch.cuda.is_available():
|
||||
im_tensor=im_tensor.cuda()
|
||||
|
||||
result=net(im_tensor)
|
||||
result = torch.squeeze(F.interpolate(result[0][0], size=(h,w), mode='bilinear') ,0)
|
||||
ma = torch.max(result)
|
||||
mi = torch.min(result)
|
||||
result = (result-mi)/(ma-mi)
|
||||
im_array = (result*255).cpu().data.numpy().astype(np.uint8)
|
||||
mask = Image.fromarray(np.squeeze(im_array))
|
||||
# mask.save('test.png')
|
||||
# mask=tensor2pil(result)
|
||||
mask=mask.convert('L')
|
||||
|
||||
masks.append(mask)
|
||||
|
||||
# rgba图
|
||||
image_rgba =orig_image.convert("RGBA")
|
||||
image_rgba.putalpha(mask)
|
||||
rgba_images.append(image_rgba)
|
||||
|
||||
#rgb
|
||||
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)
|
||||
return (masks,rgba_images,rgb_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=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)
|
||||
# mask=mask.convert('L')
|
||||
# masks.append(mask)
|
||||
if model_name=="u2net_cloth_seg":
|
||||
width, original_height = mask.size
|
||||
num_slices = original_height // img.height
|
||||
for i in range(num_slices):
|
||||
top = i * img.height
|
||||
bottom = (i + 1) * img.height
|
||||
slice_image = mask.crop((0, top, width, bottom))
|
||||
slice_mask=slice_image.convert('L')
|
||||
masks.append(slice_mask)
|
||||
|
||||
# rgba图
|
||||
image_rgba = img.convert("RGBA")
|
||||
image_rgba.putalpha(slice_mask)
|
||||
rgba_images.append(image_rgba)
|
||||
|
||||
#rgb
|
||||
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)
|
||||
|
||||
else:
|
||||
mask=mask.convert('L')
|
||||
# mask.save(output_path)
|
||||
masks.append(mask)
|
||||
|
||||
# rgba图
|
||||
image_rgba = img.convert("RGBA")
|
||||
image_rgba.putalpha(mask)
|
||||
rgba_images.append(image_rgba)
|
||||
|
||||
#rgb
|
||||
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)
|
||||
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
return (masks,rgba_images,rgb_images)
|
||||
|
||||
|
||||
# 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)
|
||||
|
||||
|
||||
class RembgNode_:
|
||||
|
||||
global _available
|
||||
available=_available
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image": ("IMAGE",),
|
||||
"model_name": (get_rembg_models(U2NET_HOME),),
|
||||
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK","IMAGE","RGBA",)
|
||||
RETURN_NAMES = ("masks","images","RGBAs")
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,True,True,)
|
||||
|
||||
def run(self,image,model_name):
|
||||
# 兼容list输入和batch输入
|
||||
|
||||
model_name=model_name[0]
|
||||
|
||||
images=[]
|
||||
|
||||
for ims in image:
|
||||
for im in ims:
|
||||
im=tensor2pil(im)
|
||||
images.append(im)
|
||||
|
||||
if model_name=='briarmbg':
|
||||
masks,rgba_images,rgb_images=run_briarmbg(images)
|
||||
else:
|
||||
masks,rgba_images,rgb_images=run_rembg(model_name,images, comfy.utils.ProgressBar(len(images) ))
|
||||
|
||||
masks=[pil2tensor(m) for m in masks]
|
||||
|
||||
rgba_images=[pil2tensor(m) for m in rgba_images]
|
||||
|
||||
rgb_images=[pil2tensor(m) for m in rgb_images]
|
||||
|
||||
return (masks,rgb_images,rgba_images,)
|
||||
@@ -93,7 +93,7 @@ class ScreenShareNode:
|
||||
RETURN_NAMES = ("IMAGE","PROMPT","FLOAT","INT")
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/image"
|
||||
CATEGORY = "♾️Mixlab/Screen"
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (False,False,False,False)
|
||||
@@ -118,7 +118,7 @@ class FloatingVideo:
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/image"
|
||||
CATEGORY = "♾️Mixlab/Screen"
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
# OUTPUT_IS_LIST = (False,False,)
|
||||
|
||||
@@ -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.reverse(), ),
|
||||
"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,)
|
||||
|
||||
|
||||
@@ -0,0 +1,437 @@
|
||||
from transformers import pipeline, set_seed,AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
import random
|
||||
import re
|
||||
|
||||
import os,sys
|
||||
import folder_paths
|
||||
|
||||
# from PIL import Image
|
||||
import importlib.util
|
||||
|
||||
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")
|
||||
|
||||
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(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)
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return spec is not None
|
||||
|
||||
|
||||
try:
|
||||
if is_installed('sentencepiece')==False:
|
||||
import subprocess
|
||||
|
||||
# 安装
|
||||
print('#pip install sentencepiece')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'sentencepiece'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0 and is_installed('sentencepiece'):
|
||||
print("#install success")
|
||||
_available=True
|
||||
else:
|
||||
print("#install error")
|
||||
_available=False
|
||||
|
||||
else:
|
||||
_available=True
|
||||
|
||||
except:
|
||||
_available=False
|
||||
|
||||
|
||||
|
||||
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)
|
||||
sequences = zh_en_model.generate(**encoded)
|
||||
return zh_en_tokenizer.batch_decode(sequences, skip_special_tokens=True)[0]
|
||||
|
||||
# input = "青春不能回头,所以青春没有终点。 ——《火影忍者》"
|
||||
# print(input, translate(input))
|
||||
|
||||
|
||||
def text_generate(text_pipe,input,seed=None):
|
||||
|
||||
if seed==None:
|
||||
seed = random.randint(100, 1000000)
|
||||
|
||||
set_seed(seed)
|
||||
|
||||
for count in range(6):
|
||||
sequences = text_pipe(input, max_length=random.randint(60, 90), num_return_sequences=8)
|
||||
list = []
|
||||
for sequence in sequences:
|
||||
line = sequence['generated_text'].strip()
|
||||
if line != input and len(line) > (len(input) + 4) and line.endswith((":", "-", "—")) is False:
|
||||
list.append(line)
|
||||
|
||||
result = "\n".join(list)
|
||||
result = re.sub('[^ ]+\.[^ ]+','', result)
|
||||
result = result.replace("<", "").replace(">", "")
|
||||
if result != "":
|
||||
return result
|
||||
if count == 5:
|
||||
return result
|
||||
|
||||
# input = "Youth can't turn back, so there's no end to youth."
|
||||
# print(input, text_generate(input))
|
||||
|
||||
|
||||
import re
|
||||
|
||||
def correct_prompt_syntax(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()
|
||||
|
||||
# 处理空元素
|
||||
if not element:
|
||||
continue
|
||||
|
||||
# 检查并处理圆括号、方括号、尖括号
|
||||
if element[0] in '([':
|
||||
corrected_element = balance_brackets(element, '(', ')') if element[0] == '(' else balance_brackets(element, '[', ']')
|
||||
elif element[0] == '<':
|
||||
corrected_element = balance_brackets(element, '<', '>')
|
||||
else:
|
||||
# 删除开头的右括号或右方括号
|
||||
corrected_element = element.lstrip(')]')
|
||||
|
||||
corrected_elements.append(corrected_element)
|
||||
|
||||
# 重组修正后的prompt
|
||||
return ','.join(corrected_elements)
|
||||
|
||||
|
||||
# # 示例使用
|
||||
# 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:
|
||||
|
||||
global _available
|
||||
available=_available
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"text": ("STRING",{"multiline": True,"default": "", "dynamicPrompts": False}),
|
||||
"generation": (["on","off"],{"default": "off"}),
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
|
||||
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("prompt",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
global text_pipe,zh_en_model,zh_en_tokenizer
|
||||
|
||||
text_pipe= None
|
||||
zh_en_model=None
|
||||
zh_en_tokenizer=None
|
||||
|
||||
def run(self,text,seed,generation):
|
||||
|
||||
|
||||
seed=seed[0]
|
||||
generation=generation[0]
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(text)+1)
|
||||
texts = [correct_prompt_syntax(t) for t in 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")
|
||||
# zh_en_tokenizer.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
text_pipe=pipeline('text-generation', model=text_generator_model_path,device="cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
# text_pipe.model.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
prompt_result=[]
|
||||
|
||||
# print('zh_en_model device',zh_en_model.device,text_pipe.model.device,torch.cuda.current_device() )
|
||||
en_texts=[]
|
||||
|
||||
for t in texts:
|
||||
if t:
|
||||
# translated_text = translated_word = translate(zh_en_tokenizer,zh_en_model,str(t))
|
||||
parser = Lark(grammar, start="start", parser="lalr", transformer=ChinesePromptTranslate())
|
||||
# print('t',t)
|
||||
result = parser.parse(t).children
|
||||
# print('en_result',result)
|
||||
# en_text=translate(zh_en_tokenizer,zh_en_model,text_without_syntax)
|
||||
en_texts.append(result[0])
|
||||
|
||||
zh_en_model.to('cpu')
|
||||
print("test en_text",en_texts)
|
||||
# en_text.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
pbar.update(1)
|
||||
for t in en_texts:
|
||||
if generation=='on':
|
||||
prompt =text_generate(text_pipe,t,seed)
|
||||
# 多条,还是单条
|
||||
lines = prompt.split("\n")
|
||||
longest_line = max(lines, key=len)
|
||||
# print(longest_line)
|
||||
prompt_result.append(longest_line)
|
||||
else:
|
||||
prompt_result.append(t)
|
||||
pbar.update(1)
|
||||
|
||||
text_pipe.model.to('cpu')
|
||||
|
||||
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
|
||||
},
|
||||
"result": (prompt_result,)}
|
||||
|
||||
|
||||
|
||||
|
||||
class PromptGenerate:
|
||||
|
||||
global _available
|
||||
available=_available
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"text": ("STRING",{"multiline": True,"default": "", "dynamicPrompts": False}),
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"multiple": (["off","on"],),
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("prompt",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
global text_pipe
|
||||
|
||||
text_pipe= None
|
||||
#
|
||||
|
||||
def run(self,text,multiple,seed):
|
||||
global text_pipe
|
||||
|
||||
seed=seed[0]
|
||||
|
||||
multiple=multiple[0]
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(text))
|
||||
|
||||
text_pipe=pipeline('text-generation', model=text_generator_model_path,device="cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
prompt_result=[]
|
||||
|
||||
for t in text:
|
||||
prompt =text_generate(text_pipe,t,seed)
|
||||
prompt = prompt.split("\n")
|
||||
if multiple=='off':
|
||||
prompt = [max(prompt, key=len)]
|
||||
|
||||
for p in prompt:
|
||||
prompt_result.append(p)
|
||||
pbar.update(1)
|
||||
|
||||
text_pipe.model.to('cpu')
|
||||
|
||||
return {
|
||||
"ui":{
|
||||
"prompt": prompt_result
|
||||
},
|
||||
"result": (prompt_result,)}
|
||||
@@ -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}}
|
||||
|
||||
|
||||
|
||||
@@ -1,10 +1,70 @@
|
||||
import os
|
||||
import re,random
|
||||
import os,platform
|
||||
import re,random,json
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
# FONT_PATH= os.path.abspath(os.path.join(os.path.dirname(__file__),'../assets/王汉宗颜楷体繁.ttf'))
|
||||
import folder_paths
|
||||
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):
|
||||
return [], {}
|
||||
|
||||
if excluded_dir_names is None:
|
||||
excluded_dir_names = []
|
||||
|
||||
result = []
|
||||
dirs = {directory: os.path.getmtime(directory)}
|
||||
for dirpath, subdirs, filenames in os.walk(directory, followlinks=True, topdown=True):
|
||||
subdirs[:] = [d for d in subdirs if d not in excluded_dir_names]
|
||||
for file_name in filenames:
|
||||
relative_path = os.path.relpath(os.path.join(dirpath, file_name), directory)
|
||||
result.append(relative_path)
|
||||
for d in subdirs:
|
||||
path = os.path.join(dirpath, d)
|
||||
dirs[path] = os.path.getmtime(path)
|
||||
return result, dirs
|
||||
|
||||
def filter_files_extensions(files, extensions):
|
||||
return sorted(list(filter(lambda a: os.path.splitext(a)[-1].lower() in extensions or len(extensions) == 0, files)))
|
||||
|
||||
|
||||
def get_system_font_path():
|
||||
ps=[]
|
||||
system = platform.system()
|
||||
if system == "Windows":
|
||||
ps.append(os.path.join(os.environ["WINDIR"], "Fonts"))
|
||||
elif system == "Darwin":
|
||||
ps.append(os.path.join("/Library", "Fonts"))
|
||||
elif system == "Linux":
|
||||
ps.append(os.path.join("/usr", "share", "fonts"))
|
||||
ps.append(os.path.join("/usr", "local", "share", "fonts"))
|
||||
ps=[p for p in ps if os.path.exists(p)]
|
||||
file_paths=[]
|
||||
for f in ps:
|
||||
result, dirs=recursive_search(f)
|
||||
for r in result:
|
||||
file_paths.append(r)
|
||||
file_paths=filter_files_extensions(file_paths,[".otf", ".ttf"])
|
||||
|
||||
return file_paths
|
||||
|
||||
|
||||
|
||||
# import json
|
||||
# import hashlib
|
||||
@@ -34,13 +94,13 @@ def create_temp_file(image):
|
||||
) = folder_paths.get_save_image_path('tmp', output_dir)
|
||||
|
||||
|
||||
image=tensor2pil(image)
|
||||
im=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)
|
||||
im.save(image_path,compress_level=4)
|
||||
|
||||
return [{
|
||||
"filename": image_file,
|
||||
@@ -60,14 +120,14 @@ def get_font_files(directory):
|
||||
|
||||
# 尝试获取系统字体
|
||||
try:
|
||||
font_paths = fm.findSystemFonts()
|
||||
for path in font_paths:
|
||||
font_paths = get_system_font_path()
|
||||
for file in font_paths:
|
||||
try:
|
||||
font_prop = fm.FontProperties(fname=path)
|
||||
font_name = font_prop.get_name()
|
||||
font_files[font_name] = path
|
||||
font_name = os.path.splitext(file)[0]
|
||||
font_path = file
|
||||
font_files[font_name] = os.path.abspath(font_path)
|
||||
except Exception as e:
|
||||
print(f"Error processing font {path}: {e}")
|
||||
print(f"Error processing font {file}: {e}")
|
||||
except Exception as e:
|
||||
print(f"Error finding system fonts: {e}")
|
||||
|
||||
@@ -79,11 +139,25 @@ font_files = get_font_files(r_directory)
|
||||
# print(font_files)
|
||||
|
||||
|
||||
def flatten_list(nested_list):
|
||||
flat_list = []
|
||||
for item in nested_list:
|
||||
if isinstance(item, list):
|
||||
flat_list.extend(flatten_list(item))
|
||||
else:
|
||||
if torch.is_tensor(item):
|
||||
print('item.shape',item.shape)
|
||||
for i in range(item.shape[0]):
|
||||
flat_list.append(item[i:i + 1, ...])
|
||||
else:
|
||||
flat_list.append(item)
|
||||
return flat_list
|
||||
|
||||
|
||||
class ColorInput:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
|
||||
"color":("TCOLOR",),
|
||||
},
|
||||
}
|
||||
@@ -93,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,)
|
||||
@@ -122,7 +196,7 @@ class FontInput:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -137,14 +211,17 @@ class TextToNumber:
|
||||
return {"required": {
|
||||
"text": ("STRING",{"multiline": False,"default": "1"}),
|
||||
"random_number": (["enable", "disable"],),
|
||||
"number":("INT", {
|
||||
"default": 0,
|
||||
"min": 0, #Minimum value
|
||||
"max_num":("INT", {
|
||||
"default": 10,
|
||||
"min":2, #Minimum value
|
||||
"max": 10000000000, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
},
|
||||
"optional":{
|
||||
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
@@ -152,12 +229,12 @@ class TextToNumber:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self,text,random_number,number):
|
||||
def run(self,text,random_number,max_num,seed=0):
|
||||
|
||||
numbers = re.findall(r'\d+', text)
|
||||
result=0
|
||||
@@ -166,7 +243,7 @@ class TextToNumber:
|
||||
# print(result)
|
||||
|
||||
if random_number=='enable' and result>0:
|
||||
result= random.randint(1, 10000000000)
|
||||
result= random.randint(1, max_num)
|
||||
return {"ui": {"text": [text],"num":[result]}, "result": (result,)}
|
||||
|
||||
|
||||
@@ -178,7 +255,7 @@ class FloatSlider:
|
||||
"number":("FLOAT", {
|
||||
"default": 0,
|
||||
"min": 0, #Minimum value
|
||||
"max": 1, #Maximum value
|
||||
"max": 0xffffffffffffffff, #Maximum value
|
||||
"step": 0.001, #Slider's step
|
||||
"display": "slider" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
@@ -207,21 +284,20 @@ class FloatSlider:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
|
||||
RETURN_NAMES = ('FLOAT',)
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self,number,min_value,max_value,step):
|
||||
def run(self, number, min_value, max_value, step):
|
||||
if number < min_value:
|
||||
number= min_value
|
||||
number = min_value
|
||||
elif number > max_value:
|
||||
number= max_value
|
||||
return (number,)
|
||||
|
||||
number = max_value
|
||||
return (number,)
|
||||
|
||||
class IntNumber:
|
||||
@classmethod
|
||||
@@ -252,7 +328,7 @@ class IntNumber:
|
||||
"default": 1,
|
||||
"min": -0xffffffffffffffff,
|
||||
"max": 0xffffffffffffffff,
|
||||
"step":1,
|
||||
"step":1,
|
||||
"display": "number"
|
||||
}),
|
||||
},
|
||||
@@ -262,7 +338,7 @@ class IntNumber:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -279,11 +355,18 @@ class MultiplicationNode:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"numberA":(any_type,),
|
||||
"numberB":("FLOAT", {
|
||||
"default": 0,
|
||||
"min": -1, #Minimum value
|
||||
"multiply_by":("FLOAT", {
|
||||
"default": 1,
|
||||
"min": -2, #Minimum value
|
||||
"max": 0xffffffffffffffff,
|
||||
"step": 0.1, #Slider's step
|
||||
"step": 0.01, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"add_by":("FLOAT", {
|
||||
"default": 0,
|
||||
"min": -2000, #Minimum value
|
||||
"max": 0xffffffffffffffff,
|
||||
"step": 0.01, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
})
|
||||
},
|
||||
@@ -293,21 +376,21 @@ class MultiplicationNode:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,False,)
|
||||
|
||||
def run(self,numberA,numberB):
|
||||
b=int(numberA*numberB)
|
||||
a=float(numberA*numberB)
|
||||
def run(self,numberA,multiply_by,add_by):
|
||||
b=int(numberA*multiply_by+add_by)
|
||||
a=float(numberA*multiply_by+add_by)
|
||||
return (a,b,)
|
||||
|
||||
class TextInput:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"text": ("STRING",{"multiline": True,"default": ""}),
|
||||
"text": ("STRING",{"multiline": True,"default": ""})
|
||||
},
|
||||
}
|
||||
|
||||
@@ -315,7 +398,7 @@ class TextInput:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -324,6 +407,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
|
||||
@@ -356,7 +494,7 @@ class DynamicDelayProcessor:
|
||||
},
|
||||
"optional":{
|
||||
"any_input":(any_type,),
|
||||
"delay_by_text":("STRING",{"multiline":True,}),
|
||||
"delay_by_text":("STRING",{"multiline":True,"dynamicPrompts": False,}),
|
||||
"words_per_seconds":("FLOAT",{ "default":1.50,"min": 0.0,"max": 1000.00,"display":"Chars per second?"}),
|
||||
"replace_output": (["disable","enable"],),
|
||||
"replace_value":("INT",{ "default":-1,"min": 0,"max": 1000000,"display":"Replacement value"})
|
||||
@@ -389,7 +527,7 @@ class DynamicDelayProcessor:
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ('output',)
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
def run(self,any_input,delay_seconds,delay_by_text,words_per_seconds,replace_output,replace_value):
|
||||
# print(f"Delay text:",delay_by_text )
|
||||
# 获取开始时间戳
|
||||
@@ -423,12 +561,12 @@ class AppInfo:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"name": ("STRING",{"multiline": False,"default": "Mixlab-App","dynamicPrompts": False}),
|
||||
"image": ("IMAGE",),
|
||||
"input_ids":("STRING",{"multiline": True,"default": "\n".join(["1","2","3"]),"dynamicPrompts": False}),
|
||||
"output_ids":("STRING",{"multiline": True,"default": "\n".join(["5","9"]),"dynamicPrompts": False}),
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"IMAGE": ("IMAGE",),
|
||||
"description":("STRING",{"multiline": True,"default": "","dynamicPrompts": False}),
|
||||
"version":("INT", {
|
||||
"default": 1,
|
||||
@@ -439,60 +577,58 @@ class AppInfo:
|
||||
}),
|
||||
"share_prefix":("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
|
||||
"link":("STRING",{"multiline": False,"default": "https://","dynamicPrompts": False}),
|
||||
"category":("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
|
||||
"auto_save": (["enable","disable"],),
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("IMAGE",)
|
||||
RETURN_TYPES = ()
|
||||
# RETURN_NAMES = ("IMAGE",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
# OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self,name,image,input_ids,output_ids,description,version,share_prefix,link):
|
||||
def run(self,name,input_ids,output_ids,IMAGE,description,version,share_prefix,link,category,auto_save):
|
||||
name=name[0]
|
||||
|
||||
im=None
|
||||
if IMAGE:
|
||||
im=IMAGE[0][0]
|
||||
#TODO batch 的方式需要处理
|
||||
im=create_temp_file(im)
|
||||
# image [img,] img[batch,w,h,a] 列表里面是batch,
|
||||
|
||||
im=create_temp_file(image)
|
||||
input_ids=input_ids[0]
|
||||
output_ids=output_ids[0]
|
||||
description=description[0]
|
||||
version=version[0]
|
||||
share_prefix=share_prefix[0]
|
||||
link=link[0]
|
||||
category=category[0]
|
||||
|
||||
# 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]}, "result": (image,)}
|
||||
return {"ui": {"json": [name,im,input_ids,output_ids,description,version,share_prefix,link,category]}, "result": ()}
|
||||
|
||||
|
||||
|
||||
|
||||
class GetImageSize_:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT", "INT")
|
||||
RETURN_NAMES = ("width", "height")
|
||||
|
||||
FUNCTION = "get_size"
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
|
||||
def get_size(self, image):
|
||||
_, height, width, _ = image.shape
|
||||
return (width, height)
|
||||
|
||||
|
||||
|
||||
class SwitchByIndex:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"A":(any_type,),
|
||||
"B":(any_type,),
|
||||
"optional":{
|
||||
"A":(any_type,),
|
||||
"B":(any_type,),
|
||||
},
|
||||
"required": {
|
||||
"index":("INT", {
|
||||
"default": -1,
|
||||
"min": -1,
|
||||
@@ -500,33 +636,76 @@ class SwitchByIndex:
|
||||
"step": 1,
|
||||
"display": "number"
|
||||
}),
|
||||
"flat": (['off',"on"],),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ("C",)
|
||||
RETURN_TYPES = (any_type,"INT",)
|
||||
RETURN_NAMES = ("list", "count",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
OUTPUT_IS_LIST = (True, False,)
|
||||
|
||||
def run(self, A=[],B=[],index=-1,flat='on'):
|
||||
|
||||
flat=flat[0]
|
||||
|
||||
def run(self, A,B,index):
|
||||
C=[]
|
||||
index=index[0]
|
||||
for a in A:
|
||||
|
||||
for a in A:
|
||||
C.append(a)
|
||||
for b in B:
|
||||
C.append(b)
|
||||
|
||||
if flat=='on':
|
||||
C=flatten_list(C)
|
||||
|
||||
if index>-1:
|
||||
try:
|
||||
C=[C[index]]
|
||||
except Exception as e:
|
||||
C=[]
|
||||
return (C,)
|
||||
C=[C[-1]] #最后一个
|
||||
|
||||
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:
|
||||
@@ -557,7 +736,7 @@ class LimitNumber:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -579,3 +758,230 @@ class LimitNumber:
|
||||
|
||||
return (nn,)
|
||||
|
||||
|
||||
|
||||
class ListStatistics:
|
||||
@staticmethod
|
||||
def count_types(lst):
|
||||
type_count = {}
|
||||
|
||||
for item in lst:
|
||||
item_type = type(item).__name__
|
||||
if item_type not in type_count:
|
||||
type_count[item_type] = []
|
||||
|
||||
if item_type in ['dict', 'str', 'int', 'float']:
|
||||
type_count[item_type].append(item)
|
||||
|
||||
return type_count
|
||||
|
||||
# # 示例列表
|
||||
# my_list = [1, 'hello', {'name': 'John'}, 3.14, {'age': 25}, 'world', 10]
|
||||
|
||||
# # 创建ListStatistics对象
|
||||
# list_stats = ListStatistics()
|
||||
|
||||
# # 调用count_types方法进行统计
|
||||
# result = list_stats.count_types(my_list)
|
||||
|
||||
# # 输出结果
|
||||
# for item_type, values in result.items():
|
||||
# print(item_type + ':')
|
||||
# for value in values:
|
||||
# print(value)
|
||||
# print('---')
|
||||
|
||||
class TESTNODE_:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"ANY":(any_type,),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Test"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self,ANY):
|
||||
print(type(ANY))
|
||||
try:
|
||||
print(ANY[0].shape)
|
||||
img= tensor2pil(ANY[0])
|
||||
print(img.size)
|
||||
except:
|
||||
print('')
|
||||
|
||||
# data=ANY
|
||||
list_stats = ListStatistics()
|
||||
|
||||
# 调用count_types方法进行统计
|
||||
result = list_stats.count_types(ANY)
|
||||
|
||||
|
||||
# 假设我们有一个模块文件名为 my_module.py,它位于 'importables' 目录下
|
||||
module_path = os.path.join(os.path.dirname(__file__),'test.py')
|
||||
|
||||
# 使用 spec_from_file_location 获取模块的元数据(名称、定义等)
|
||||
spec = importlib.util.spec_from_file_location('test', module_path)
|
||||
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
functions = getattr(module, 'run') # 获取函数
|
||||
|
||||
functions(ANY)
|
||||
|
||||
|
||||
return {"ui": {"data": result,"type":[str(type(ANY[0]))]}, "result": (ANY,)}
|
||||
|
||||
|
||||
class TESTNODE_TOKEN:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"text":("STRING", {"forceInput": True,}),
|
||||
"clip": ("CLIP", )
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Test"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self,text,clip=None):
|
||||
# print(text)
|
||||
|
||||
tokens = clip.tokenize(text)
|
||||
|
||||
tokens=[v for v in tokens.values()][0][0]
|
||||
|
||||
tokens=json.dumps(tokens)
|
||||
|
||||
return (tokens,)
|
||||
|
||||
|
||||
|
||||
class CreateSeedNode:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
RETURN_NAMES = ("seed",)
|
||||
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, seed):
|
||||
return (seed,)
|
||||
|
||||
|
||||
class CreateCkptNames:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"ckpt_names": ("STRING",{"multiline": True,"default": "\n".join(folder_paths.get_filename_list("checkpoints")),"dynamicPrompts": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ("ckpt_names",)
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
# OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, ckpt_names):
|
||||
ckpt_names=ckpt_names.split('\n')
|
||||
ckpt_names = [name for name in ckpt_names if name.strip()]
|
||||
return (ckpt_names,)
|
||||
|
||||
|
||||
class CreateLoraNames:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"lora_names": ("STRING",{"multiline": True,"default": "\n".join(folder_paths.get_filename_list("loras")),"dynamicPrompts": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,"STRING",)
|
||||
RETURN_NAMES = ("lora_names","prompt",)
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (True,True,)
|
||||
|
||||
# OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, lora_names):
|
||||
lora_names=lora_names.split('\n')
|
||||
lora_names = [name for name in lora_names if name.strip()]
|
||||
prompts=[os.path.splitext(n)[0] for n in lora_names]
|
||||
return (lora_names,prompts,)
|
||||
|
||||
|
||||
|
||||
class CreateSampler_names:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"sampler_names": ("STRING",{"multiline": True,"default": "\n".join(comfy.samplers.KSampler.SAMPLERS),"dynamicPrompts": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ("sampler_names",)
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
# OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, sampler_names):
|
||||
sampler_names=sampler_names.split('\n')
|
||||
sampler_names = [name for name in sampler_names if name.strip()]
|
||||
return (sampler_names,)
|
||||
@@ -1,179 +0,0 @@
|
||||
# https://github.com/openai/consistencydecoder/blob/main/consistencydecoder/__init__.py
|
||||
|
||||
import folder_paths
|
||||
from comfy import model_management
|
||||
|
||||
import math
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
|
||||
class ConsistencyDecoderWrapper:
|
||||
def __init__(self, decoder):
|
||||
self.decoder = decoder
|
||||
def decode(self, x):
|
||||
return self.decoder(x)
|
||||
|
||||
def _extract_into_tensor(arr, timesteps, broadcast_shape):
|
||||
|
||||
# from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L895 """
|
||||
res = arr[timesteps].float()
|
||||
dims_to_append = len(broadcast_shape) - len(res.shape)
|
||||
return res[(...,) + (None,) * dims_to_append]
|
||||
|
||||
|
||||
def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
|
||||
# from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L45
|
||||
betas = []
|
||||
for i in range(num_diffusion_timesteps):
|
||||
t1 = i / num_diffusion_timesteps
|
||||
t2 = (i + 1) / num_diffusion_timesteps
|
||||
betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
|
||||
return torch.tensor(betas)
|
||||
|
||||
class ConsistencyDecoder:
|
||||
def __init__(self, device="cuda:0", download_target=""):
|
||||
self.n_distilled_steps = 64
|
||||
# download_target = _download("https://openaipublic.azureedge.net/diff-vae/c9cebd3132dd9c42936d803e33424145a748843c8f716c0814838bdc8a2fe7cb/decoder.pt", download_root)
|
||||
self.ckpt = torch.jit.load(download_target).to(device)
|
||||
self.device = device
|
||||
sigma_data = 0.5
|
||||
betas = betas_for_alpha_bar(
|
||||
1024, lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
|
||||
).to(device)
|
||||
alphas = 1.0 - betas
|
||||
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||
self.sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
|
||||
self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
|
||||
sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / alphas_cumprod)
|
||||
sigmas = torch.sqrt(1.0 / alphas_cumprod - 1)
|
||||
self.c_skip = (
|
||||
sqrt_recip_alphas_cumprod
|
||||
* sigma_data**2
|
||||
/ (sigmas**2 + sigma_data**2)
|
||||
)
|
||||
self.c_out = sigmas * sigma_data / (sigmas**2 + sigma_data**2) ** 0.5
|
||||
self.c_in = sqrt_recip_alphas_cumprod / (sigmas**2 + sigma_data**2) ** 0.5
|
||||
|
||||
@staticmethod
|
||||
def round_timesteps(
|
||||
timesteps, total_timesteps, n_distilled_steps, truncate_start=True
|
||||
):
|
||||
with torch.no_grad():
|
||||
space = torch.div(total_timesteps, n_distilled_steps, rounding_mode="floor")
|
||||
rounded_timesteps = (
|
||||
torch.div(timesteps, space, rounding_mode="floor") + 1
|
||||
) * space
|
||||
if truncate_start:
|
||||
rounded_timesteps[rounded_timesteps == total_timesteps] -= space
|
||||
else:
|
||||
rounded_timesteps[rounded_timesteps == total_timesteps] -= space
|
||||
rounded_timesteps[rounded_timesteps == 0] += space
|
||||
return rounded_timesteps
|
||||
|
||||
@staticmethod
|
||||
def ldm_transform_latent(z, extra_scale_factor=1):
|
||||
channel_means = [0.38862467, 0.02253063, 0.07381133, -0.0171294]
|
||||
channel_stds = [0.9654121, 1.0440036, 0.76147926, 0.77022034]
|
||||
|
||||
if len(z.shape) != 4:
|
||||
raise ValueError()
|
||||
|
||||
z = z * 0.18215
|
||||
channels = [z[:, i] for i in range(z.shape[1])]
|
||||
|
||||
channels = [
|
||||
extra_scale_factor * (c - channel_means[i]) / channel_stds[i]
|
||||
for i, c in enumerate(channels)
|
||||
]
|
||||
return torch.stack(channels, dim=1)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
features: torch.Tensor,
|
||||
schedule=[1.0, 0.5],
|
||||
):
|
||||
features = self.ldm_transform_latent(features)
|
||||
|
||||
ts = self.round_timesteps(
|
||||
torch.arange(0, 1024),
|
||||
1024,
|
||||
self.n_distilled_steps,
|
||||
truncate_start=False,
|
||||
)
|
||||
shape = (
|
||||
features.size(0),
|
||||
3,
|
||||
8 * features.size(2),
|
||||
8 * features.size(3),
|
||||
)
|
||||
|
||||
x_start = torch.zeros(shape, device=features.device, dtype=features.dtype)
|
||||
schedule_timesteps = [int((1024 - 1) * s) for s in schedule]
|
||||
for i in schedule_timesteps:
|
||||
t = ts[i].item()
|
||||
t_ = torch.tensor([t] * features.shape[0]).to(self.device)
|
||||
noise = torch.randn_like(x_start)
|
||||
|
||||
x_start = (
|
||||
_extract_into_tensor(self.sqrt_alphas_cumprod, t_, x_start.shape)
|
||||
* x_start
|
||||
+ _extract_into_tensor(
|
||||
self.sqrt_one_minus_alphas_cumprod, t_, x_start.shape
|
||||
)
|
||||
* noise
|
||||
)
|
||||
c_in = _extract_into_tensor(self.c_in, t_, x_start.shape)
|
||||
model_output = self.ckpt(c_in * x_start, t_, features=features)
|
||||
B, C = x_start.shape[:2]
|
||||
model_output, _ = torch.split(model_output, C, dim=1)
|
||||
pred_xstart = (
|
||||
_extract_into_tensor(self.c_out, t_, x_start.shape) * model_output
|
||||
+ _extract_into_tensor(self.c_skip, t_, x_start.shape) * x_start
|
||||
).clamp(-1, 1)
|
||||
x_start = pred_xstart
|
||||
return x_start
|
||||
|
||||
|
||||
|
||||
class VAELoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "vae_name": (folder_paths.get_filename_list("vae"), )}}
|
||||
RETURN_TYPES = ("VAE",)
|
||||
FUNCTION = "load_vae"
|
||||
|
||||
CATEGORY = "♾️Mixlab/_test"
|
||||
|
||||
#TODO: scale factor?
|
||||
def load_vae(self, vae_name):
|
||||
vae_path = folder_paths.get_full_path("vae", vae_name)
|
||||
device = 'cuda:0'
|
||||
# print('device',device)
|
||||
consistencyDecoder = ConsistencyDecoder(device=device,
|
||||
download_target=vae_path) # Model size: 2.49 GB
|
||||
vae = ConsistencyDecoderWrapper(consistencyDecoder)
|
||||
return (vae,)
|
||||
|
||||
|
||||
class VAEDecode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "samples": ("LATENT", ), "vae": ("VAE", )}}
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "decode"
|
||||
|
||||
CATEGORY = "♾️Mixlab/_test"
|
||||
|
||||
def decode(self, vae, samples):
|
||||
image = vae.decode(samples["samples"].to("cuda:0"))
|
||||
image = image[0].cpu().numpy()
|
||||
image = (image + 1.0) * 127.5
|
||||
image = image.clip(0, 255).astype(np.uint8)
|
||||
image = Image.fromarray(image.transpose(1, 2, 0))
|
||||
image = image.convert("RGB")
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[None,]
|
||||
return (image, )
|
||||
@@ -0,0 +1,693 @@
|
||||
import os
|
||||
import hashlib
|
||||
import json
|
||||
import subprocess
|
||||
import shutil
|
||||
import re
|
||||
import time,math
|
||||
import numpy as np
|
||||
from typing import List
|
||||
import torch
|
||||
from PIL import Image, ImageOps
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
import cv2,random,string
|
||||
from pathlib import Path
|
||||
|
||||
import folder_paths
|
||||
from comfy.k_diffusion.utils import FolderOfImages
|
||||
from comfy.utils import common_upscale
|
||||
|
||||
|
||||
|
||||
|
||||
def generate_folder_name(directory,video_path):
|
||||
# Get the directory and filename from the video path
|
||||
_, filename = os.path.split(video_path)
|
||||
# Generate a random string of lowercase letters and digits
|
||||
random_string = ''.join(random.choices(string.ascii_lowercase + string.digits, k=8))
|
||||
# Create the folder name by combining the random string and the filename
|
||||
folder_name = random_string + '_' + filename
|
||||
# Create the full folder path by joining the directory and the folder name
|
||||
folder_path = os.path.join(directory, folder_name)
|
||||
return folder_path
|
||||
|
||||
def create_folder(directory,video_path):
|
||||
folder_path = generate_folder_name(directory,video_path)
|
||||
os.makedirs(folder_path)
|
||||
return folder_path
|
||||
|
||||
|
||||
def split_video(video_path, video_segment_frames, transition_frames, output_dir):
|
||||
# 读取视频文件
|
||||
video_capture = cv2.VideoCapture(video_path)
|
||||
|
||||
# 获取视频的总帧数和帧率
|
||||
total_frames = int(video_capture.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
fps = video_capture.get(cv2.CAP_PROP_FPS)
|
||||
|
||||
# 计算每个视频片段的总帧数,包括过渡帧
|
||||
segment_total_frames = video_segment_frames + transition_frames
|
||||
|
||||
# 计算可以分割的片段数量,向上取整
|
||||
num_segments = (total_frames + transition_frames - 1) // segment_total_frames
|
||||
|
||||
vs=[]
|
||||
# 计算每个片段的起始帧和结束帧
|
||||
start_frame = 0
|
||||
for i in range(num_segments):
|
||||
# 计算当前片段的结束帧,注意最后一个片段可能没有过渡帧
|
||||
end_frame = min(start_frame + segment_total_frames, total_frames)
|
||||
|
||||
# 打印当前片段的起始帧和结束帧
|
||||
print(f"Segment {i+1}: Start Frame {start_frame}, End Frame {end_frame}")
|
||||
|
||||
# 保存当前片段为一个视频文件
|
||||
segment_video_path = f"{output_dir}/segment_{i+1}.avi"
|
||||
|
||||
fourcc = cv2.VideoWriter_fourcc(*'XVID')
|
||||
segment_video = cv2.VideoWriter(segment_video_path, fourcc, fps, (int(video_capture.get(cv2.CAP_PROP_FRAME_WIDTH)),
|
||||
int(video_capture.get(cv2.CAP_PROP_FRAME_HEIGHT))))
|
||||
|
||||
for frame_num in range(start_frame, end_frame):
|
||||
ret, frame = video_capture.read()
|
||||
if ret:
|
||||
segment_video.write(frame)
|
||||
else:
|
||||
break # 如果读取失败,则退出循环
|
||||
|
||||
# 更新起始帧为下一个片段的起始位置
|
||||
start_frame = end_frame + transition_frames
|
||||
vs.append(segment_video_path)
|
||||
|
||||
# 释放视频捕获对象
|
||||
video_capture.release()
|
||||
# print(vs)
|
||||
return (vs,total_frames,fps)
|
||||
|
||||
|
||||
folder_paths.folder_names_and_paths["video_formats"] = (
|
||||
[
|
||||
os.path.join(os.path.dirname(os.path.abspath(__file__)), ".", "video_formats"),
|
||||
],
|
||||
[".json"]
|
||||
)
|
||||
|
||||
ffmpeg_path = shutil.which("ffmpeg")
|
||||
if ffmpeg_path is None:
|
||||
print("ffmpeg could not be found. Using ffmpeg from imageio-ffmpeg.")
|
||||
from imageio_ffmpeg import get_ffmpeg_exe
|
||||
try:
|
||||
ffmpeg_path = get_ffmpeg_exe()
|
||||
except:
|
||||
print("ffmpeg could not be found. Outputs that require it have been disabled")
|
||||
|
||||
# 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 count_files(directory):
|
||||
count = 0
|
||||
for root, dirs, files in os.walk(directory):
|
||||
count += len(files)
|
||||
return count
|
||||
|
||||
def create_temp_file(image):
|
||||
output_dir = folder_paths.get_temp_directory()
|
||||
|
||||
c=count_files(output_dir)
|
||||
|
||||
(
|
||||
full_output_folder,
|
||||
filename,
|
||||
counter,
|
||||
subfolder,
|
||||
_,
|
||||
) = folder_paths.get_save_image_path('temp_', output_dir)
|
||||
|
||||
|
||||
image=tensor2pil(image)
|
||||
|
||||
image_file = f"{filename}_{c}_{counter:05}.png"
|
||||
|
||||
image_path=os.path.join(full_output_folder, image_file)
|
||||
|
||||
image.save(image_path,compress_level=4)
|
||||
|
||||
return [{
|
||||
"filename": image_file,
|
||||
"subfolder": subfolder,
|
||||
"type": "temp"
|
||||
}]
|
||||
|
||||
|
||||
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
|
||||
|
||||
# images = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]
|
||||
# chunk_size = 3
|
||||
# transition_size = 1
|
||||
|
||||
# result = split_list(images, chunk_size, transition_size)
|
||||
# print(result)
|
||||
|
||||
class ImageListReplace:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"images": ("IMAGE",),
|
||||
"start_index":("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"end_index":("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"invert": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional":{
|
||||
"image_replace": ("IMAGE",),
|
||||
"images_replace": ("IMAGE",),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE","IMAGE",)
|
||||
RETURN_NAMES = ("images","select_images",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,True,)
|
||||
|
||||
def run(self, images,start_index=[0],end_index=[0],invert=[False],image_replace=None,images_replace=None):
|
||||
start_index=start_index[0]
|
||||
end_index=end_index[0]
|
||||
invert=invert[0]
|
||||
|
||||
image_rs=[]
|
||||
|
||||
if image_replace!=None:
|
||||
for i in range(end_index-start_index+1):
|
||||
image_rs.append(image_replace[0])
|
||||
|
||||
if images_replace!=None:
|
||||
image_rs=images_replace
|
||||
|
||||
# 如果image replace 为空
|
||||
if image_replace==None and images_replace==None:
|
||||
# print('如果image replace 为空',images[0])
|
||||
# [[tensor(
|
||||
# tensor([[[[0.
|
||||
first_image=tensor2pil(images[0][0])
|
||||
width, height = first_image.size
|
||||
image_replace=Image.new("RGB", (width, height), (0, 0, 0))
|
||||
image_replace=pil2tensor(image_replace)
|
||||
for i in range(end_index-start_index+1):
|
||||
image_rs.append(image_replace)
|
||||
|
||||
|
||||
new_images=[]
|
||||
select_images=[]
|
||||
k=0
|
||||
for i in range(len(images)):
|
||||
if i>=start_index and i<=end_index:
|
||||
if invert:
|
||||
new_images.append(images[i])
|
||||
else:
|
||||
new_images.append(image_rs[k])
|
||||
select_images.append(images[i])
|
||||
k+=1
|
||||
else:
|
||||
if invert:
|
||||
new_images.append(image_rs[k])
|
||||
select_images.append(images[i])
|
||||
k+=1
|
||||
else:
|
||||
new_images.append(images[i])
|
||||
|
||||
imss=[]
|
||||
# print(len(images))
|
||||
for i in range(len(images)):
|
||||
t=images[i][0]
|
||||
t=tensor2pil(t)
|
||||
t = t.convert("RGB")
|
||||
original_width, original_height = t.size
|
||||
scale = 300 / original_width
|
||||
new_height = int(original_height * scale)
|
||||
t = t.resize((300, new_height))
|
||||
|
||||
ims=create_temp_file(pil2tensor(t))
|
||||
imss.append(ims[0])
|
||||
|
||||
# image_replace=create_temp_file(image_replace)
|
||||
|
||||
return {"ui":{"_images": imss},"result": (new_images,select_images,)}
|
||||
|
||||
# The code is based on ComfyUI-VideoHelperSuite modification.
|
||||
class LoadVideoAndSegment:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
video_extensions = ['webm', 'mp4', 'mkv', 'gif']
|
||||
input_dir = folder_paths.get_input_directory()
|
||||
files = []
|
||||
for f in os.listdir(input_dir):
|
||||
if os.path.isfile(os.path.join(input_dir, f)):
|
||||
file_parts = f.split('.')
|
||||
if len(file_parts) > 1 and (file_parts[-1] in video_extensions):
|
||||
files.append(f)
|
||||
return {"required": {
|
||||
"video": (sorted(files), {"video_upload": True}),
|
||||
"video_segment_frames": ("INT", {"default": 10, "min": 1, "step": 1}),
|
||||
"transition_frames": ("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
},}
|
||||
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
RETURN_TYPES = ("SCENE_VIDEO","INT", "INT","INT",)
|
||||
RETURN_NAMES = ("scenes_video","scenes_count","frame_count","fps",)
|
||||
FUNCTION = "load_video"
|
||||
OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (True,False,False,False,)
|
||||
|
||||
|
||||
def is_gif(self, filename):
|
||||
file_parts = filename.split('.')
|
||||
return len(file_parts) > 1 and file_parts[-1] == "gif"
|
||||
|
||||
def load_video_cv_fallback(self, video, frame_load_cap, skip_first_frames):
|
||||
try:
|
||||
video_cap = cv2.VideoCapture(folder_paths.get_annotated_filepath(video))
|
||||
if not video_cap.isOpened():
|
||||
raise ValueError(f"{video} could not be loaded with cv fallback.")
|
||||
# set video_cap to look at start_index frame
|
||||
images = []
|
||||
total_frame_count = 0
|
||||
frames_added = 0
|
||||
base_frame_time = 1/video_cap.get(cv2.CAP_PROP_FPS)
|
||||
|
||||
target_frame_time = base_frame_time
|
||||
|
||||
time_offset=0.0
|
||||
while video_cap.isOpened():
|
||||
if time_offset < target_frame_time:
|
||||
is_returned, frame = video_cap.read()
|
||||
# if didn't return frame, video has ended
|
||||
if not is_returned:
|
||||
break
|
||||
time_offset += base_frame_time
|
||||
if time_offset < target_frame_time:
|
||||
continue
|
||||
time_offset -= target_frame_time
|
||||
# if not at start_index, skip doing anything with frame
|
||||
total_frame_count += 1
|
||||
if total_frame_count <= skip_first_frames:
|
||||
continue
|
||||
# TODO: do whatever operations need to happen, like force_size, etc
|
||||
|
||||
# opencv loads images in BGR format (yuck), so need to convert to RGB for ComfyUI use
|
||||
# follow up: can videos ever have an alpha channel?
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
# convert frame to comfyui's expected format (taken from comfy's load image code)
|
||||
image = Image.fromarray(frame)
|
||||
image = ImageOps.exif_transpose(image)
|
||||
image = np.array(image, dtype=np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[None,]
|
||||
images.append(image)
|
||||
frames_added += 1
|
||||
# if cap exists and we've reached it, stop processing frames
|
||||
if frame_load_cap > 0 and frames_added >= frame_load_cap:
|
||||
break
|
||||
finally:
|
||||
video_cap.release()
|
||||
images = torch.cat(images, dim=0)
|
||||
|
||||
return (images, frames_added)
|
||||
|
||||
def load_video(self, video,video_segment_frames,transition_frames ):
|
||||
|
||||
video_path = folder_paths.get_annotated_filepath(video)
|
||||
|
||||
# check if video is a gif - will need to use cv fallback to read frames
|
||||
# use cv fallback if ffmpeg not installed or gif
|
||||
# if ffmpeg_path is None:
|
||||
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
|
||||
# otherwise, continue with ffmpeg
|
||||
|
||||
# args_dummy = [ffmpeg_path, "-i", video_path, "-f", "null", "-"]
|
||||
# try:
|
||||
# with subprocess.Popen(args_dummy, stdout=subprocess.DEVNULL, stderr=subprocess.PIPE) as proc:
|
||||
# for line in proc.stderr.readlines():
|
||||
# match = re.search(", ([1-9]|\\d{2,})x(\\d+)",line.decode('utf-8'))
|
||||
# if match is not None:
|
||||
# size = [int(match.group(1)), int(match.group(2))]
|
||||
# break
|
||||
# except Exception as e:
|
||||
# print(f"Retrying with opencv due to ffmpeg error: {e}")
|
||||
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
|
||||
# args_all_frames = [ffmpeg_path, "-i", video_path, "-v", "error",
|
||||
# "-pix_fmt", "rgb24"]
|
||||
|
||||
# vfilters = []
|
||||
|
||||
# if skip_first_frames > 0:
|
||||
# vfilters.append(f"select=gt(n\\,{skip_first_frames-1})")
|
||||
# if frame_load_cap > 0:
|
||||
# vfilters.append(f"select=gt({frame_load_cap}\\,n)")
|
||||
# #manually calculate aspect ratio to ensure reads remain aligned
|
||||
|
||||
# if len(vfilters) > 0:
|
||||
# args_all_frames += ["-vf", ",".join(vfilters)]
|
||||
|
||||
# args_all_frames += ["-f", "rawvideo", "-"]
|
||||
# images = []
|
||||
# try:
|
||||
# with subprocess.Popen(args_all_frames, stdout=subprocess.PIPE) as proc:
|
||||
# #Manually buffer enough bytes for an image
|
||||
# bpi = size[0]*size[1]*3
|
||||
# current_bytes = bytearray(bpi)
|
||||
# current_offset=0
|
||||
# while True:
|
||||
# bytes_read = proc.stdout.read(bpi - current_offset)
|
||||
# if bytes_read is None:#sleep to wait for more data
|
||||
# time.sleep(.2)
|
||||
# continue
|
||||
# if len(bytes_read) == 0:#EOF
|
||||
# break
|
||||
# current_bytes[current_offset:len(bytes_read)] = bytes_read
|
||||
# current_offset+=len(bytes_read)
|
||||
# if current_offset == bpi:
|
||||
# images.append(np.array(current_bytes, dtype=np.float32).reshape(size[1], size[0], 3) / 255.0)
|
||||
# current_offset = 0
|
||||
# except Exception as e:
|
||||
# print(f"Retrying with opencv due to ffmpeg error: {e}")
|
||||
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
|
||||
|
||||
# imgs=split_list(images,video_segment_frames,transition_frames)
|
||||
|
||||
# temp path
|
||||
tp=folder_paths.get_temp_directory()
|
||||
basename = os.path.basename(video_path) # 获取文件名
|
||||
name_without_extension = os.path.splitext(basename)[0] # 去掉文件后缀
|
||||
|
||||
folder_path = create_folder(tp,name_without_extension)
|
||||
|
||||
|
||||
# 导出的数据
|
||||
scenes_video,total_frames,fps=split_video(video_path,video_segment_frames,
|
||||
transition_frames,folder_path)
|
||||
|
||||
|
||||
# imgs=[torch.from_numpy(np.stack(im)) for im in imgs]
|
||||
|
||||
# images = torch.from_numpy(np.stack(images))
|
||||
|
||||
return (scenes_video,len(scenes_video), total_frames,fps,)
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, video, **kwargs):
|
||||
image_path = folder_paths.get_annotated_filepath(video)
|
||||
m = hashlib.sha256()
|
||||
with open(image_path, 'rb') as f:
|
||||
m.update(f.read())
|
||||
return m.digest().hex()
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(s, video, **kwargs):
|
||||
if not folder_paths.exists_annotated_filepath(video):
|
||||
return "Invalid image file: {}".format(video)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
# The code is based on ComfyUI-VideoHelperSuite modification.
|
||||
class VideoCombine_Adv:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
#Hide ffmpeg formats if ffmpeg isn't available
|
||||
if ffmpeg_path is not None:
|
||||
ffmpeg_formats = ["video/"+x[:-5] for x in folder_paths.get_filename_list("video_formats")]
|
||||
else:
|
||||
ffmpeg_formats = []
|
||||
# ffmpeg_formats =["video/"+x for x in ['webm', 'mp4', 'mkv']]
|
||||
return {
|
||||
"required": {
|
||||
"image_batch": ("IMAGE",),
|
||||
"frame_rate": (
|
||||
"INT",
|
||||
{"default": 8, "min": 1, "step": 1},
|
||||
),
|
||||
"loop_count": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}),
|
||||
"filename_prefix": ("STRING", {"default": "Comfyui"}),
|
||||
"format": (["image/gif", "image/webp"] + ffmpeg_formats,),
|
||||
"pingpong": ("BOOLEAN", {"default": False}),
|
||||
"save_image": ("BOOLEAN", {"default": True}),
|
||||
"metadata": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"hidden": {
|
||||
"prompt": "PROMPT",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO",
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
FUNCTION = "run"
|
||||
|
||||
def save_with_tempfile(self, args, metadata, file_path, frames, env):
|
||||
#Ensure temp directory exists
|
||||
os.makedirs(folder_paths.get_temp_directory(), exist_ok=True)
|
||||
|
||||
metadata_path = os.path.join(folder_paths.get_temp_directory(), "metadata.txt")
|
||||
#metadata from file should escape = ; # \ and newline
|
||||
#From my testing, though, only backslashes need escapes and = in particular causes problems
|
||||
#It is likely better to prioritize future compatibility with containers that don't support
|
||||
#or shouldn't use the comment tag for embedding metadata
|
||||
metadata = metadata.replace("\\","\\\\")
|
||||
metadata = metadata.replace(";","\\;")
|
||||
metadata = metadata.replace("#","\\#")
|
||||
#metadata = metadata.replace("=","\\=")
|
||||
metadata = metadata.replace("\n","\\\n")
|
||||
with open(metadata_path, "w") as f:
|
||||
f.write(";FFMETADATA1\n")
|
||||
f.write(metadata)
|
||||
args = args[:1] + ["-i", metadata_path] + args[1:] + [file_path]
|
||||
with subprocess.Popen(args, stdin=subprocess.PIPE, env=env) as proc:
|
||||
for frame in frames:
|
||||
proc.stdin.write(frame.tobytes())
|
||||
|
||||
def run(
|
||||
self,
|
||||
image_batch,
|
||||
frame_rate: int,
|
||||
loop_count: int,
|
||||
filename_prefix="AnimateDiff",
|
||||
format="image/gif",
|
||||
pingpong=False,
|
||||
save_image=True,
|
||||
metadata=False,
|
||||
prompt=None,
|
||||
extra_pnginfo=None,
|
||||
):
|
||||
images=image_batch
|
||||
|
||||
frames: List[Image.Image] = []
|
||||
for image in images:
|
||||
img = 255.0 * image.cpu().numpy()
|
||||
img = Image.fromarray(np.clip(img, 0, 255).astype(np.uint8))
|
||||
# resize 保证
|
||||
# 检查图像的高度是否是2的倍数,如果不是,则调整高度
|
||||
if img.height % 2 != 0:
|
||||
img = img.resize((img.width, img.height + 1))
|
||||
|
||||
# 检查图像的宽度是否是2的倍数,如果不是,则调整宽度
|
||||
if img.width % 2 != 0:
|
||||
img = img.resize((img.width + 1, img.height))
|
||||
|
||||
frames.append(img)
|
||||
|
||||
# get output information
|
||||
output_dir = (
|
||||
folder_paths.get_output_directory()
|
||||
if save_image
|
||||
else folder_paths.get_temp_directory()
|
||||
)
|
||||
(
|
||||
full_output_folder,
|
||||
filename,
|
||||
counter,
|
||||
subfolder,
|
||||
_,
|
||||
) = folder_paths.get_save_image_path(filename_prefix, output_dir)
|
||||
|
||||
metadata = PngInfo()
|
||||
video_metadata = {}
|
||||
if prompt is not None:
|
||||
metadata.add_text("prompt", json.dumps(prompt))
|
||||
video_metadata["prompt"] = prompt
|
||||
if extra_pnginfo is not None:
|
||||
for x in extra_pnginfo:
|
||||
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
|
||||
video_metadata[x] = extra_pnginfo[x]
|
||||
|
||||
# 取消保存metadata
|
||||
if metadata==False:
|
||||
metadata = PngInfo()
|
||||
|
||||
# save first frame as png to keep metadata
|
||||
file = f"{filename}_{counter:05}_.png"
|
||||
file_path = os.path.join(full_output_folder, file)
|
||||
frames[0].save(
|
||||
file_path,
|
||||
pnginfo=metadata,
|
||||
compress_level=4,
|
||||
)
|
||||
if pingpong:
|
||||
frames = frames + frames[-2:0:-1]
|
||||
|
||||
format_type, format_ext = format.split("/")
|
||||
file = f"{filename}_{counter:05}_.{format_ext}"
|
||||
file_path = os.path.join(full_output_folder, file)
|
||||
if format_type == "image":
|
||||
# Use pillow directly to save an animated image
|
||||
frames[0].save(
|
||||
file_path,
|
||||
format=format_ext.upper(),
|
||||
save_all=True,
|
||||
append_images=frames[1:],
|
||||
duration=round(1000 / frame_rate),
|
||||
loop=loop_count,
|
||||
compress_level=4,
|
||||
)
|
||||
else:
|
||||
# Use ffmpeg to save a video
|
||||
if ffmpeg_path is None:
|
||||
#Should never be reachable
|
||||
raise ProcessLookupError("Could not find ffmpeg")
|
||||
|
||||
video_format_path = folder_paths.get_full_path("video_formats", format_ext + ".json")
|
||||
with open(video_format_path, 'r') as stream:
|
||||
video_format = json.load(stream)
|
||||
file = f"{filename}_{counter:05}_.{video_format['extension']}"
|
||||
file_path = os.path.join(full_output_folder, file)
|
||||
dimensions = f"{frames[0].width}x{frames[0].height}"
|
||||
metadata_args = ["-metadata", "comment=" + json.dumps(video_metadata)]
|
||||
args = [ffmpeg_path, "-v", "error", "-f", "rawvideo", "-pix_fmt", "rgb24",
|
||||
"-s", dimensions, "-r", str(frame_rate), "-i", "-"] \
|
||||
+ video_format['main_pass']
|
||||
# On linux, max arg length is Pagesize * 32 -> 131072
|
||||
# On windows, this around 32767 but seems to vary wildly by > 500
|
||||
# in a manor not solely related to other arguments
|
||||
if os.name == 'posix':
|
||||
max_arg_length = 4096*32
|
||||
else:
|
||||
max_arg_length = 32767 - len(" ".join(args + [metadata_args[0]] + [file_path])) - 1
|
||||
#test max limit
|
||||
#metadata_args[1] = metadata_args[1] + "a"*(max_arg_length - len(metadata_args[1])-1)
|
||||
|
||||
env=os.environ.copy()
|
||||
if "environment" in video_format:
|
||||
env.update(video_format["environment"])
|
||||
if len(metadata_args[1]) >= max_arg_length:
|
||||
print(f"Using fallback file for extremely long metadata: {len(metadata_args[1])}/{max_arg_length}")
|
||||
self.save_with_tempfile(args, metadata_args[1], file_path, frames, env)
|
||||
else:
|
||||
try:
|
||||
with subprocess.Popen(args + metadata_args + [file_path],
|
||||
stdin=subprocess.PIPE, env=env) as proc:
|
||||
for frame in frames:
|
||||
proc.stdin.write(frame.tobytes())
|
||||
except FileNotFoundError as e:
|
||||
if "winerror" in dir(e) and e.winerror == 206:
|
||||
print("Metadata was too long. Retrying with fallback file")
|
||||
self.save_with_tempfile(args, metadata_args[1], file_path, frames, env)
|
||||
else:
|
||||
raise
|
||||
except OSError as e:
|
||||
if "errno" in dir(e) and e.errno == 7:
|
||||
print("Metadata was too long. Retrying with fallback file")
|
||||
self.save_with_tempfile(args, metadata_args[1], file_path, frames, env)
|
||||
else:
|
||||
raise
|
||||
|
||||
previews = [
|
||||
{
|
||||
"filename": file,
|
||||
"subfolder": subfolder,
|
||||
"type": "output" if save_image else "temp",
|
||||
"format": format,
|
||||
}
|
||||
]
|
||||
return {"ui": {"gifs": previews}}
|
||||
|
||||
|
||||
class VAEEncodeForInpaint_Frames:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"vae": ("VAE", ),
|
||||
"images": ("IMAGE", ),
|
||||
"masks": ("MASK", ),
|
||||
"grow_mask_by": ("INT", {"default": 6, "min": 0, "max": 64, "step": 1}),
|
||||
}}
|
||||
|
||||
FUNCTION = "encode"
|
||||
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
RETURN_NAMES = ("LATENT",)
|
||||
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
|
||||
def encode(self, vae, images, masks, grow_mask_by=[6]):
|
||||
vae=vae[0]
|
||||
grow_mask_by=grow_mask_by[0]
|
||||
|
||||
result=[]
|
||||
|
||||
for i in range(len(images)):
|
||||
pixels=images[i]
|
||||
mask=masks[i]
|
||||
|
||||
|
||||
x = (pixels.shape[1] // 8) * 8
|
||||
y = (pixels.shape[2] // 8) * 8
|
||||
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(pixels.shape[1], pixels.shape[2]), mode="bilinear")
|
||||
|
||||
pixels = pixels.clone()
|
||||
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,:]
|
||||
mask = mask[:,:,x_offset:x + x_offset, y_offset:y + y_offset]
|
||||
|
||||
#grow mask by a few pixels to keep things seamless in latent space
|
||||
if grow_mask_by == 0:
|
||||
mask_erosion = mask
|
||||
else:
|
||||
kernel_tensor = torch.ones((1, 1, grow_mask_by, grow_mask_by))
|
||||
padding = math.ceil((grow_mask_by - 1) / 2)
|
||||
|
||||
mask_erosion = torch.clamp(torch.nn.functional.conv2d(mask.round(), kernel_tensor, padding=padding), 0, 1)
|
||||
|
||||
m = (1.0 - mask.round()).squeeze(1)
|
||||
for i in range(3):
|
||||
pixels[:,:,:,i] -= 0.5
|
||||
pixels[:,:,:,i] *= m
|
||||
pixels[:,:,:,i] += 0.5
|
||||
t = vae.encode(pixels)
|
||||
|
||||
result.append({"samples":t, "noise_mask": (mask_erosion[:,:,:x,:y].round())})
|
||||
|
||||
|
||||
return (result, )
|
||||
@@ -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,12 @@
|
||||
from VoiceStreamAI.asr.whisper_asr import WhisperASR
|
||||
from VoiceStreamAI.asr.faster_whisper_asr import FasterWhisperASR
|
||||
|
||||
class ASRFactory:
|
||||
@staticmethod
|
||||
def create_asr_pipeline(type, **kwargs):
|
||||
if type == "whisper":
|
||||
return WhisperASR(**kwargs)
|
||||
if type == "faster_whisper":
|
||||
return FasterWhisperASR(**kwargs)
|
||||
else:
|
||||
raise ValueError(f"Unknown ASR pipeline type: {type}")
|
||||
@@ -0,0 +1,9 @@
|
||||
class ASRInterface:
|
||||
async def transcribe(self, client):
|
||||
"""
|
||||
Transcribe the given audio data.
|
||||
|
||||
:param client: The client object with all the member variables including the buffer
|
||||
:return: The transcription structure, see for example the faster_whisper_asr.py file.
|
||||
"""
|
||||
raise NotImplementedError("This method should be implemented by subclasses.")
|
||||
@@ -0,0 +1,142 @@
|
||||
import os
|
||||
from faster_whisper import WhisperModel
|
||||
|
||||
from VoiceStreamAI.asr.asr_interface import ASRInterface
|
||||
from VoiceStreamAI.audio_utils import save_audio_to_file
|
||||
|
||||
import folder_paths
|
||||
|
||||
language_codes = {
|
||||
"afrikaans": "af",
|
||||
"amharic": "am",
|
||||
"arabic": "ar",
|
||||
"assamese": "as",
|
||||
"azerbaijani": "az",
|
||||
"bashkir": "ba",
|
||||
"belarusian": "be",
|
||||
"bulgarian": "bg",
|
||||
"bengali": "bn",
|
||||
"tibetan": "bo",
|
||||
"breton": "br",
|
||||
"bosnian": "bs",
|
||||
"catalan": "ca",
|
||||
"czech": "cs",
|
||||
"welsh": "cy",
|
||||
"danish": "da",
|
||||
"german": "de",
|
||||
"greek": "el",
|
||||
"english": "en",
|
||||
"spanish": "es",
|
||||
"estonian": "et",
|
||||
"basque": "eu",
|
||||
"persian": "fa",
|
||||
"finnish": "fi",
|
||||
"faroese": "fo",
|
||||
"french": "fr",
|
||||
"galician": "gl",
|
||||
"gujarati": "gu",
|
||||
"hausa": "ha",
|
||||
"hawaiian": "haw",
|
||||
"hebrew": "he",
|
||||
"hindi": "hi",
|
||||
"croatian": "hr",
|
||||
"haitian": "ht",
|
||||
"hungarian": "hu",
|
||||
"armenian": "hy",
|
||||
"indonesian": "id",
|
||||
"icelandic": "is",
|
||||
"italian": "it",
|
||||
"japanese": "ja",
|
||||
"javanese": "jw",
|
||||
"georgian": "ka",
|
||||
"kazakh": "kk",
|
||||
"khmer": "km",
|
||||
"kannada": "kn",
|
||||
"korean": "ko",
|
||||
"latin": "la",
|
||||
"luxembourgish": "lb",
|
||||
"lingala": "ln",
|
||||
"lao": "lo",
|
||||
"lithuanian": "lt",
|
||||
"latvian": "lv",
|
||||
"malagasy": "mg",
|
||||
"maori": "mi",
|
||||
"macedonian": "mk",
|
||||
"malayalam": "ml",
|
||||
"mongolian": "mn",
|
||||
"marathi": "mr",
|
||||
"malay": "ms",
|
||||
"maltese": "mt",
|
||||
"burmese": "my",
|
||||
"nepali": "ne",
|
||||
"dutch": "nl",
|
||||
"norwegian nynorsk": "nn",
|
||||
"norwegian": "no",
|
||||
"occitan": "oc",
|
||||
"punjabi": "pa",
|
||||
"polish": "pl",
|
||||
"pashto": "ps",
|
||||
"portuguese": "pt",
|
||||
"romanian": "ro",
|
||||
"russian": "ru",
|
||||
"sanskrit": "sa",
|
||||
"sindhi": "sd",
|
||||
"sinhalese": "si",
|
||||
"slovak": "sk",
|
||||
"slovenian": "sl",
|
||||
"shona": "sn",
|
||||
"somali": "so",
|
||||
"albanian": "sq",
|
||||
"serbian": "sr",
|
||||
"sundanese": "su",
|
||||
"swedish": "sv",
|
||||
"swahili": "sw",
|
||||
"tamil": "ta",
|
||||
"telugu": "te",
|
||||
"tajik": "tg",
|
||||
"thai": "th",
|
||||
"turkmen": "tk",
|
||||
"tagalog": "tl",
|
||||
"turkish": "tr",
|
||||
"tatar": "tt",
|
||||
"ukrainian": "uk",
|
||||
"urdu": "ur",
|
||||
"uzbek": "uz",
|
||||
"vietnamese": "vi",
|
||||
"yiddish": "yi",
|
||||
"yoruba": "yo",
|
||||
"chinese": "zh",
|
||||
"cantonese": "yue",
|
||||
}
|
||||
|
||||
|
||||
class FasterWhisperASR(ASRInterface):
|
||||
def __init__(self, **kwargs):
|
||||
model_size = kwargs.get('model_size', "large-v3")
|
||||
device = kwargs.get('device', "cuda")
|
||||
model_root = os.path.join(folder_paths.models_dir, "whisper")
|
||||
# Run on GPU with FP16
|
||||
self.asr_pipeline = WhisperModel(model_size, device=device, compute_type="float16",download_root=model_root)
|
||||
|
||||
async def transcribe(self, client):
|
||||
file_path = await save_audio_to_file(client.scratch_buffer, client.get_file_name())
|
||||
|
||||
language = None if client.config['language'] is None else language_codes.get(client.config['language'].lower())
|
||||
segments, info = self.asr_pipeline.transcribe(file_path, word_timestamps=True, language=language)
|
||||
|
||||
segments = list(segments) # The transcription will actually run here.
|
||||
os.remove(file_path)
|
||||
|
||||
flattened_words = [word for segment in segments for word in segment.words]
|
||||
|
||||
to_return = {
|
||||
"language": info.language,
|
||||
"language_probability": info.language_probability,
|
||||
"text": ' '.join([s.text.strip() for s in segments]),
|
||||
"words":
|
||||
[
|
||||
{"word": w.word, "start": w.start, "end": w.end, "probability":w.probability} for w in flattened_words
|
||||
]
|
||||
}
|
||||
return to_return
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
from transformers import pipeline
|
||||
from VoiceStreamAI.asr.asr_interface import ASRInterface
|
||||
from VoiceStreamAI.audio_utils import save_audio_to_file
|
||||
import os
|
||||
|
||||
class WhisperASR(ASRInterface):
|
||||
def __init__(self, **kwargs):
|
||||
model_name = kwargs.get('model_name', "openai/whisper-large-v3")
|
||||
self.asr_pipeline = pipeline("automatic-speech-recognition", model=model_name)
|
||||
|
||||
async def transcribe(self, client):
|
||||
file_path = await save_audio_to_file(client.scratch_buffer, client.get_file_name())
|
||||
|
||||
if client.config['language'] is not None:
|
||||
to_return = self.asr_pipeline(file_path, generate_kwargs={"language": client.config['language']})['text']
|
||||
else:
|
||||
to_return = self.asr_pipeline(file_path)['text']
|
||||
|
||||
os.remove(file_path)
|
||||
|
||||
to_return = {
|
||||
"language": "UNSUPPORTED_BY_HUGGINGFACE_WHISPER",
|
||||
"language_probability": None,
|
||||
"text": to_return.strip(),
|
||||
"words": "UNSUPPORTED_BY_HUGGINGFACE_WHISPER"
|
||||
}
|
||||
return to_return
|
||||
@@ -0,0 +1,26 @@
|
||||
import wave
|
||||
import os
|
||||
|
||||
async def save_audio_to_file(audio_data, file_name, audio_dir="audio_files", audio_format="wav"):
|
||||
"""
|
||||
Saves the audio data to a file.
|
||||
|
||||
:param client_id: Unique identifier for the client.
|
||||
:param audio_data: The audio data to save.
|
||||
:param file_counters: Dictionary to keep track of file counts for each client.
|
||||
:param audio_dir: Directory where audio files will be saved.
|
||||
:param audio_format: Format of the audio file.
|
||||
:return: Path to the saved audio file.
|
||||
"""
|
||||
|
||||
os.makedirs(audio_dir, exist_ok=True)
|
||||
|
||||
file_path = os.path.join(audio_dir, file_name)
|
||||
|
||||
with wave.open(file_path, 'wb') as wav_file:
|
||||
wav_file.setnchannels(1) # Assuming mono audio
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(16000)
|
||||
wav_file.writeframes(audio_data)
|
||||
|
||||
return file_path
|
||||
@@ -0,0 +1,142 @@
|
||||
import os
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
|
||||
from VoiceStreamAI.buffering_strategy.buffering_strategy_interface import BufferingStrategyInterface
|
||||
from openai import OpenAI
|
||||
|
||||
class SilenceAtEndOfChunk(BufferingStrategyInterface):
|
||||
"""
|
||||
A buffering strategy that processes audio at the end of each chunk with silence detection.
|
||||
|
||||
This class is responsible for handling audio chunks, detecting silence at the end of each chunk,
|
||||
and initiating the transcription process for the chunk.
|
||||
|
||||
Attributes:
|
||||
client (Client): The client instance associated with this buffering strategy.
|
||||
chunk_length_seconds (float): Length of each audio chunk in seconds.
|
||||
chunk_offset_seconds (float): Offset time in seconds to be considered for processing audio chunks.
|
||||
"""
|
||||
|
||||
def __init__(self, client, **kwargs):
|
||||
"""
|
||||
Initialize the SilenceAtEndOfChunk buffering strategy.
|
||||
|
||||
Args:
|
||||
client (Client): The client instance associated with this buffering strategy.
|
||||
**kwargs: Additional keyword arguments, including 'chunk_length_seconds' and 'chunk_offset_seconds'.
|
||||
"""
|
||||
self.client = client
|
||||
|
||||
self.chunk_length_seconds = os.environ.get('BUFFERING_CHUNK_LENGTH_SECONDS')
|
||||
if not self.chunk_length_seconds:
|
||||
self.chunk_length_seconds = kwargs.get('chunk_length_seconds')
|
||||
self.chunk_length_seconds = float(self.chunk_length_seconds)
|
||||
|
||||
self.chunk_offset_seconds = os.environ.get('BUFFERING_CHUNK_OFFSET_SECONDS')
|
||||
if not self.chunk_offset_seconds:
|
||||
self.chunk_offset_seconds = kwargs.get('chunk_offset_seconds')
|
||||
self.chunk_offset_seconds = float(self.chunk_offset_seconds)
|
||||
|
||||
self.error_if_not_realtime = os.environ.get('ERROR_IF_NOT_REALTIME')
|
||||
if not self.error_if_not_realtime:
|
||||
self.error_if_not_realtime = kwargs.get('error_if_not_realtime', False)
|
||||
|
||||
self.processing_flag = False
|
||||
|
||||
self.messages=[]
|
||||
|
||||
def process_audio(self, websocket, vad_pipeline, asr_pipeline,llm_port):
|
||||
"""
|
||||
Process audio chunks by checking their length and scheduling asynchronous processing.
|
||||
|
||||
This method checks if the length of the audio buffer exceeds the chunk length and, if so,
|
||||
it schedules asynchronous processing of the audio.
|
||||
|
||||
Args:
|
||||
websocket (Websocket): The WebSocket connection for sending transcriptions.
|
||||
vad_pipeline: The voice activity detection pipeline.
|
||||
asr_pipeline: The automatic speech recognition pipeline.
|
||||
"""
|
||||
chunk_length_in_bytes = self.chunk_length_seconds * self.client.sampling_rate * self.client.samples_width
|
||||
if len(self.client.buffer) > chunk_length_in_bytes:
|
||||
if self.processing_flag:
|
||||
exit("Error in realtime processing: tried processing a new chunk while the previous one was still being processed")
|
||||
|
||||
self.client.scratch_buffer += self.client.buffer
|
||||
self.client.buffer.clear()
|
||||
self.processing_flag = True
|
||||
# Schedule the processing in a separate task
|
||||
asyncio.create_task(self.process_audio_async(websocket, vad_pipeline, asr_pipeline,llm_port))
|
||||
|
||||
async def process_audio_async(self, websocket, vad_pipeline, asr_pipeline,llm_port):
|
||||
"""
|
||||
Asynchronously process audio for activity detection and transcription.
|
||||
|
||||
This method performs heavy processing, including voice activity detection and transcription of
|
||||
the audio data. It sends the transcription results through the WebSocket connection.
|
||||
|
||||
Args:
|
||||
websocket (Websocket): The WebSocket connection for sending transcriptions.
|
||||
vad_pipeline: The voice activity detection pipeline.
|
||||
asr_pipeline: The automatic speech recognition pipeline.
|
||||
"""
|
||||
start = time.time()
|
||||
vad_results = await vad_pipeline.detect_activity(self.client)
|
||||
|
||||
if len(vad_results) == 0:
|
||||
self.client.scratch_buffer.clear()
|
||||
self.client.buffer.clear()
|
||||
self.processing_flag = False
|
||||
return
|
||||
|
||||
last_segment_should_end_before = ((len(self.client.scratch_buffer) / (self.client.sampling_rate * self.client.samples_width)) - self.chunk_offset_seconds)
|
||||
if vad_results[-1]['end'] < last_segment_should_end_before:
|
||||
transcription = await asr_pipeline.transcribe(self.client)
|
||||
if transcription['text'] != '':
|
||||
end = time.time()
|
||||
transcription['processing_time'] = end - start
|
||||
|
||||
transcription['status']="chat_start"
|
||||
|
||||
json_transcription = json.dumps(transcription)
|
||||
|
||||
await websocket.send(json_transcription)
|
||||
|
||||
# Point to the local server
|
||||
client = OpenAI(base_url=f"http://localhost:{llm_port}/v1", api_key="lm-studio")
|
||||
|
||||
|
||||
messages=[
|
||||
{"role": "system", "content": "You are a friendly and engaging AI designed to interact with users in a conversational manner. Your personality is that of a sophisticated and polite young professional who is both a designer and a programmer. You are well-mannered, articulate, and possess a good sense of humor. Your goal is to provide helpful and insightful responses while maintaining a pleasant and enjoyable conversation. Be sure to use your knowledge in design and programming to enrich the dialogue and offer relevant advice or information when appropriate. Always be respectful and considerate of the user's feelings and perspectives. Additionally, you are fluent in both English and Chinese, and can seamlessly switch between the two languages to best assist users."},
|
||||
]+self.messages[-10:0]+[{"role": "user", "content":transcription['text']}]
|
||||
# print('#messages',messages)
|
||||
|
||||
completion = client.chat.completions.create(
|
||||
model="model-identifier",
|
||||
messages=messages,
|
||||
temperature=0.7,
|
||||
)
|
||||
|
||||
transcription['asistant'] = completion.choices[0].message.content
|
||||
|
||||
transcription['status']="chat_end"
|
||||
|
||||
json_transcription = json.dumps(transcription)
|
||||
|
||||
self.messages.append({
|
||||
"role": "user",
|
||||
"content":transcription['text']})
|
||||
|
||||
self.messages.append({
|
||||
"role": "asistant",
|
||||
"content": transcription['asistant']
|
||||
})
|
||||
# print('#messages',completion.choices[0].message.content)
|
||||
|
||||
await websocket.send(json_transcription)
|
||||
self.client.scratch_buffer.clear()
|
||||
self.client.increment_file_counter()
|
||||
|
||||
self.processing_flag = False
|
||||
@@ -0,0 +1,41 @@
|
||||
from VoiceStreamAI.buffering_strategy.buffering_strategies import SilenceAtEndOfChunk
|
||||
|
||||
class BufferingStrategyFactory:
|
||||
"""
|
||||
A factory class for creating instances of different buffering strategies.
|
||||
|
||||
This factory provides a centralized way to instantiate various buffering strategies
|
||||
based on the type specified. It abstracts the creation logic, making it easier to
|
||||
manage and extend with new buffering strategy types.
|
||||
|
||||
Methods:
|
||||
create_buffering_strategy: Creates and returns an instance of a specified buffering strategy.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def create_buffering_strategy(type, client, **kwargs):
|
||||
"""
|
||||
Creates an instance of a buffering strategy based on the specified type.
|
||||
|
||||
This method acts as a factory for creating buffering strategy objects. It returns
|
||||
an instance of the strategy corresponding to the given type. If the type is not
|
||||
recognized, it raises a ValueError.
|
||||
|
||||
Args:
|
||||
type (str): The type of buffering strategy to create. Currently supports 'silence_at_end_of_chunk'.
|
||||
client (Client): The client instance to be associated with the buffering strategy.
|
||||
**kwargs: Additional keyword arguments specific to the buffering strategy being created.
|
||||
|
||||
Returns:
|
||||
An instance of the specified buffering strategy.
|
||||
|
||||
Raises:
|
||||
ValueError: If the specified type is not recognized or supported.
|
||||
|
||||
Example:
|
||||
strategy = BufferingStrategyFactory.create_buffering_strategy("silence_at_end_of_chunk", client)
|
||||
"""
|
||||
if type == "silence_at_end_of_chunk":
|
||||
return SilenceAtEndOfChunk(client, **kwargs)
|
||||
else:
|
||||
raise ValueError(f"Unknown buffering strategy type: {type}")
|
||||
@@ -0,0 +1,31 @@
|
||||
class BufferingStrategyInterface:
|
||||
"""
|
||||
An interface class for buffering strategies in audio processing systems.
|
||||
|
||||
This class defines the structure for buffering strategies used in handling
|
||||
and processing audio data. It serves as a template for creating custom buffering
|
||||
strategies that fit specific requirements of an audio processing pipeline.
|
||||
|
||||
Subclasses should implement the methods defined in this interface to ensure
|
||||
consistency and compatibility with the system's audio processing framework.
|
||||
|
||||
Methods:
|
||||
process_audio: Process audio data. This method should be implemented by subclasses.
|
||||
"""
|
||||
|
||||
def process_audio(self, websocket, vad_pipeline, asr_pipeline):
|
||||
"""
|
||||
Process audio data using the given WebSocket connection, VAD pipeline, and ASR pipeline.
|
||||
|
||||
This method is intended to be overridden in subclasses to provide specific logic
|
||||
for handling and processing audio data in different buffering strategies.
|
||||
|
||||
Args:
|
||||
websocket (Websocket): The WebSocket connection for communication with clients.
|
||||
vad_pipeline: The Voice Activity Detection (VAD) pipeline used for detecting speech in the audio.
|
||||
asr_pipeline: The Automatic Speech Recognition (ASR) pipeline used for transcribing speech in the audio.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: If the method is not implemented in the subclass.
|
||||
"""
|
||||
raise NotImplementedError("This method should be implemented by subclasses.")
|
||||
@@ -0,0 +1,54 @@
|
||||
from VoiceStreamAI.buffering_strategy.buffering_strategy_factory import BufferingStrategyFactory
|
||||
|
||||
class Client:
|
||||
"""
|
||||
Represents a client connected to the VoiceStreamAI server.
|
||||
|
||||
This class maintains the state for each connected client, including their
|
||||
unique identifier, audio buffer, configuration, and a counter for processed audio files.
|
||||
|
||||
Attributes:
|
||||
client_id (str): A unique identifier for the client.
|
||||
buffer (bytearray): A buffer to store incoming audio data.
|
||||
config (dict): Configuration settings for the client, like chunk length and offset.
|
||||
file_counter (int): Counter for the number of audio files processed.
|
||||
total_samples (int): Total number of audio samples received from this client.
|
||||
sampling_rate (int): The sampling rate of the audio data in Hz.
|
||||
samples_width (int): The width of each audio sample in bits.
|
||||
"""
|
||||
def __init__(self, client_id, sampling_rate, samples_width):
|
||||
self.client_id = client_id
|
||||
self.buffer = bytearray()
|
||||
self.scratch_buffer = bytearray()
|
||||
self.config = {"language": None,
|
||||
"processing_strategy": "silence_at_end_of_chunk",
|
||||
"processing_args": {
|
||||
"chunk_length_seconds": 5,
|
||||
"chunk_offset_seconds": 0.1
|
||||
}
|
||||
}
|
||||
self.file_counter = 0
|
||||
self.total_samples = 0
|
||||
self.sampling_rate = sampling_rate
|
||||
self.samples_width = samples_width
|
||||
self.buffering_strategy = BufferingStrategyFactory.create_buffering_strategy(self.config['processing_strategy'], self, **self.config['processing_args'])
|
||||
|
||||
def update_config(self, config_data):
|
||||
self.config.update(config_data)
|
||||
self.buffering_strategy = BufferingStrategyFactory.create_buffering_strategy(self.config['processing_strategy'], self, **self.config['processing_args'])
|
||||
|
||||
def append_audio_data(self, audio_data):
|
||||
self.buffer.extend(audio_data)
|
||||
self.total_samples += len(audio_data) / self.samples_width
|
||||
|
||||
def clear_buffer(self):
|
||||
self.buffer.clear()
|
||||
|
||||
def increment_file_counter(self):
|
||||
self.file_counter += 1
|
||||
|
||||
def get_file_name(self):
|
||||
return f"{self.client_id}_{self.file_counter}.wav"
|
||||
|
||||
def process_audio(self, websocket, vad_pipeline, asr_pipeline,llm_port):
|
||||
self.buffering_strategy.process_audio(websocket, vad_pipeline, asr_pipeline,llm_port)
|
||||
@@ -0,0 +1,54 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
# 获取当前文件的绝对路径
|
||||
current_file_path = os.path.abspath(__file__)
|
||||
|
||||
# 获取当前文件的目录
|
||||
current_directory = os.path.dirname(current_file_path)
|
||||
sys.path.append(str(Path(current_directory).parent))
|
||||
# print("sys.path", current_directory)
|
||||
|
||||
|
||||
from VoiceStreamAI.server import Server
|
||||
from VoiceStreamAI.asr.asr_factory import ASRFactory
|
||||
from VoiceStreamAI.vad.vad_factory import VADFactory
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="VoiceStreamAI Server: Real-time audio transcription using self-hosted Whisper and WebSocket")
|
||||
parser.add_argument("--vad-type", type=str, default="pyannote", help="Type of VAD pipeline to use (e.g., 'pyannote')")
|
||||
parser.add_argument("--vad-args", type=str, default='{"auth_token": "huggingface_token"}', help="JSON string of additional arguments for VAD pipeline")
|
||||
parser.add_argument("--asr-type", type=str, default="faster_whisper", help="Type of ASR pipeline to use (e.g., 'whisper')")
|
||||
parser.add_argument("--asr-args", type=str, default='{"model_size": "large-v3"}', help="JSON string of additional arguments for ASR pipeline")
|
||||
parser.add_argument("--host", type=str, default="127.0.0.1", help="Host for the WebSocket server")
|
||||
parser.add_argument("--port", type=int, default=8765, help="Port for the WebSocket server")
|
||||
parser.add_argument("--certfile", type=str, default=None, help="The path to the SSL certificate (cert file) if using secure websockets")
|
||||
parser.add_argument("--keyfile", type=str, default=None, help="The path to the SSL key file if using secure websockets")
|
||||
return parser.parse_args()
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
try:
|
||||
vad_args = json.loads(args.vad_args)
|
||||
asr_args = json.loads(args.asr_args)
|
||||
except json.JSONDecodeError as e:
|
||||
print(f"Error parsing JSON arguments: {e}")
|
||||
return
|
||||
|
||||
vad_pipeline = VADFactory.create_vad_pipeline(args.vad_type, **vad_args)
|
||||
asr_pipeline = ASRFactory.create_asr_pipeline(args.asr_type, **asr_args)
|
||||
|
||||
server = Server(vad_pipeline, asr_pipeline, host=args.host, port=args.port, sampling_rate=16000, samples_width=2, certfile=args.certfile, keyfile=args.keyfile)
|
||||
|
||||
asyncio.get_event_loop().run_until_complete(server.start())
|
||||
asyncio.get_event_loop().run_forever()
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,7 @@
|
||||
websockets
|
||||
speechbrain
|
||||
pyannote-audio
|
||||
asyncio
|
||||
sentence-transformers
|
||||
transformers
|
||||
faster-whisper
|
||||
@@ -0,0 +1,88 @@
|
||||
import websockets
|
||||
import uuid
|
||||
import json
|
||||
import asyncio
|
||||
import ssl
|
||||
|
||||
from VoiceStreamAI.audio_utils import save_audio_to_file
|
||||
from VoiceStreamAI.client import Client
|
||||
|
||||
class Server:
|
||||
"""
|
||||
Represents the WebSocket server for handling real-time audio transcription.
|
||||
|
||||
This class manages WebSocket connections, processes incoming audio data,
|
||||
and interacts with VAD and ASR pipelines for voice activity detection and
|
||||
speech recognition.
|
||||
|
||||
Attributes:
|
||||
vad_pipeline: An instance of a voice activity detection pipeline.
|
||||
asr_pipeline: An instance of an automatic speech recognition pipeline.
|
||||
host (str): Host address of the server.
|
||||
port (int): Port on which the server listens.
|
||||
sampling_rate (int): The sampling rate of audio data in Hz.
|
||||
samples_width (int): The width of each audio sample in bits.
|
||||
connected_clients (dict): A dictionary mapping client IDs to Client objects.
|
||||
"""
|
||||
def __init__(self, vad_pipeline, asr_pipeline, host='localhost', port=8765, sampling_rate=16000, samples_width=2, certfile = None, keyfile = None,llm_port=9000):
|
||||
self.vad_pipeline = vad_pipeline
|
||||
self.asr_pipeline = asr_pipeline
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.sampling_rate = sampling_rate
|
||||
self.samples_width = samples_width
|
||||
self.certfile = certfile
|
||||
self.keyfile = keyfile
|
||||
self.connected_clients = {}
|
||||
|
||||
self.llm_port=llm_port
|
||||
|
||||
async def handle_audio(self, client, websocket):
|
||||
while True:
|
||||
message = await websocket.recv()
|
||||
|
||||
if isinstance(message, bytes):
|
||||
client.append_audio_data(message)
|
||||
elif isinstance(message, str):
|
||||
config = json.loads(message)
|
||||
if config.get('type') == 'config':
|
||||
client.update_config(config['data'])
|
||||
continue
|
||||
else:
|
||||
print(f"Unexpected message type from {client.client_id}")
|
||||
|
||||
# this is synchronous, any async operation is in BufferingStrategy
|
||||
client.process_audio(websocket, self.vad_pipeline, self.asr_pipeline,self.llm_port)
|
||||
|
||||
|
||||
async def handle_websocket(self, websocket, path):
|
||||
client_id = str(uuid.uuid4())
|
||||
client = Client(client_id, self.sampling_rate, self.samples_width)
|
||||
self.connected_clients[client_id] = client
|
||||
|
||||
print(f"Client {client_id} connected")
|
||||
|
||||
try:
|
||||
await self.handle_audio(client, websocket)
|
||||
except websockets.ConnectionClosed as e:
|
||||
print(f"Connection with {client_id} closed: {e}")
|
||||
finally:
|
||||
del self.connected_clients[client_id]
|
||||
|
||||
def start(self):
|
||||
if self.certfile:
|
||||
# Create an SSL context to enforce encrypted connections
|
||||
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||||
|
||||
# Load your server's certificate and private key
|
||||
# Replace 'your_cert_path.pem' and 'your_key_path.pem' with the actual paths to your files
|
||||
ssl_context.load_cert_chain(certfile=self.certfile, keyfile=self.keyfile)
|
||||
|
||||
print(f"WebSocket server ready to accept secure connections on {self.host}:{self.port}")
|
||||
|
||||
# Pass the SSL context to the serve function along with the host and port
|
||||
# Ensure the secure flag is set to True if using a secure WebSocket protocol (wss://)
|
||||
return websockets.serve(self.handle_websocket, self.host, self.port, ssl=ssl_context)
|
||||
else:
|
||||
print(f"WebSocket server ready to accept secure connections on {self.host}:{self.port}")
|
||||
return websockets.serve(self.handle_websocket, self.host, self.port)
|
||||
@@ -0,0 +1,50 @@
|
||||
from os import remove
|
||||
import os
|
||||
|
||||
from pyannote.core import Segment
|
||||
from pyannote.audio import Model
|
||||
from pyannote.audio.pipelines import VoiceActivityDetection
|
||||
|
||||
from VoiceStreamAI.vad.vad_interface import VADInterface
|
||||
from VoiceStreamAI.audio_utils import save_audio_to_file
|
||||
|
||||
|
||||
class PyannoteVAD(VADInterface):
|
||||
"""
|
||||
Pyannote-based implementation of the VADInterface.
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
"""
|
||||
Initializes Pyannote's VAD pipeline.
|
||||
|
||||
Args:
|
||||
model_name (str): The model name for Pyannote.
|
||||
auth_token (str, optional): Authentication token for Hugging Face.
|
||||
"""
|
||||
|
||||
model_name = kwargs.get('model_name', "pyannote/segmentation")
|
||||
|
||||
auth_token = os.environ.get('PYANNOTE_AUTH_TOKEN')
|
||||
if not auth_token:
|
||||
auth_token = kwargs.get('auth_token')
|
||||
|
||||
if auth_token is None:
|
||||
raise ValueError("Missing required env var in PYANNOTE_AUTH_TOKEN or argument in --vad-args: 'auth_token'")
|
||||
|
||||
pyannote_args = kwargs.get('pyannote_args', {"onset": 0.5, "offset": 0.5, "min_duration_on": 0.3, "min_duration_off": 0.3})
|
||||
self.model = Model.from_pretrained(model_name, use_auth_token=auth_token)
|
||||
self.vad_pipeline = VoiceActivityDetection(segmentation=self.model)
|
||||
self.vad_pipeline.instantiate(pyannote_args)
|
||||
|
||||
async def detect_activity(self, client):
|
||||
audio_file_path = await save_audio_to_file(client.scratch_buffer, client.get_file_name())
|
||||
vad_results = self.vad_pipeline(audio_file_path)
|
||||
remove(audio_file_path)
|
||||
vad_segments = []
|
||||
if len(vad_results) > 0:
|
||||
vad_segments = [
|
||||
{"start": segment.start, "end": segment.end, "confidence": 1.0}
|
||||
for segment in vad_results.itersegments()
|
||||
]
|
||||
return vad_segments
|
||||
@@ -0,0 +1,23 @@
|
||||
from VoiceStreamAI.vad.pyannote_vad import PyannoteVAD
|
||||
|
||||
class VADFactory:
|
||||
"""
|
||||
Factory for creating instances of VAD systems.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def create_vad_pipeline(type, **kwargs):
|
||||
"""
|
||||
Creates a VAD pipeline based on the specified type.
|
||||
|
||||
Args:
|
||||
type (str): The type of VAD pipeline to create (e.g., 'pyannote').
|
||||
kwargs: Additional arguments for the VAD pipeline creation.
|
||||
|
||||
Returns:
|
||||
VADInterface: An instance of a class that implements VADInterface.
|
||||
"""
|
||||
if type == "pyannote":
|
||||
return PyannoteVAD(**kwargs)
|
||||
else:
|
||||
raise ValueError(f"Unknown VAD pipeline type: {type}")
|
||||
@@ -0,0 +1,16 @@
|
||||
class VADInterface:
|
||||
"""
|
||||
Interface for voice activity detection (VAD) systems.
|
||||
"""
|
||||
|
||||
async def detect_activity(self, client):
|
||||
"""
|
||||
Detects voice activity in the given audio data.
|
||||
|
||||
Args:
|
||||
client (src.Client): The client to detect on
|
||||
|
||||
Returns:
|
||||
List: VAD result, a list of objects containing "start", "end", "confidence"
|
||||
"""
|
||||
raise NotImplementedError("This method should be implemented by subclasses.")
|
||||
@@ -0,0 +1,8 @@
|
||||
import folder_paths
|
||||
|
||||
# 外挂一个文件,用来编写新的节点
|
||||
def run(v):
|
||||
|
||||
output_dir = folder_paths.get_temp_directory()
|
||||
|
||||
print('1323',v,output_dir)
|
||||
@@ -0,0 +1,38 @@
|
||||
cond_image_size: 512
|
||||
|
||||
image_tokenizer_cls: tsr.models.tokenizers.image.DINOSingleImageTokenizer
|
||||
image_tokenizer:
|
||||
pretrained_model_name_or_path: "facebook/dino-vitb16"
|
||||
|
||||
tokenizer_cls: tsr.models.tokenizers.triplane.Triplane1DTokenizer
|
||||
tokenizer:
|
||||
plane_size: 32
|
||||
num_channels: 1024
|
||||
|
||||
backbone_cls: tsr.models.transformer.transformer_1d.Transformer1D
|
||||
backbone:
|
||||
in_channels: ${tokenizer.num_channels}
|
||||
num_attention_heads: 16
|
||||
attention_head_dim: 64
|
||||
num_layers: 16
|
||||
cross_attention_dim: 768
|
||||
|
||||
post_processor_cls: tsr.models.network_utils.TriplaneUpsampleNetwork
|
||||
post_processor:
|
||||
in_channels: 1024
|
||||
out_channels: 40
|
||||
|
||||
decoder_cls: tsr.models.network_utils.NeRFMLP
|
||||
decoder:
|
||||
in_channels: 120 # 3 * 40
|
||||
n_neurons: 64
|
||||
n_hidden_layers: 9
|
||||
activation: silu
|
||||
|
||||
renderer_cls: tsr.models.nerf_renderer.TriplaneNeRFRenderer
|
||||
renderer:
|
||||
radius: 0.87 # slightly larger than 0.5 * sqrt(3)
|
||||
feature_reduction: concat
|
||||
density_activation: exp
|
||||
density_bias: -1.0
|
||||
num_samples_per_ray: 128
|
||||
@@ -0,0 +1,51 @@
|
||||
from typing import Callable, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from skimage import measure
|
||||
|
||||
|
||||
class IsosurfaceHelper(nn.Module):
|
||||
points_range: Tuple[float, float] = (0, 1)
|
||||
|
||||
@property
|
||||
def grid_vertices(self) -> torch.FloatTensor:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class MarchingCubeHelper(IsosurfaceHelper):
|
||||
def __init__(self, resolution: int) -> None:
|
||||
super().__init__()
|
||||
self.resolution = resolution
|
||||
#self.mc_func: Callable = marching_cubes
|
||||
self._grid_vertices: Optional[torch.FloatTensor] = None
|
||||
|
||||
@property
|
||||
def grid_vertices(self) -> torch.FloatTensor:
|
||||
if self._grid_vertices is None:
|
||||
# keep the vertices on CPU so that we can support very large resolution
|
||||
x, y, z = (
|
||||
torch.linspace(*self.points_range, self.resolution),
|
||||
torch.linspace(*self.points_range, self.resolution),
|
||||
torch.linspace(*self.points_range, self.resolution),
|
||||
)
|
||||
x, y, z = torch.meshgrid(x, y, z, indexing="ij")
|
||||
verts = torch.cat(
|
||||
[x.reshape(-1, 1), y.reshape(-1, 1), z.reshape(-1, 1)], dim=-1
|
||||
).reshape(-1, 3)
|
||||
self._grid_vertices = verts
|
||||
return self._grid_vertices
|
||||
|
||||
def forward(
|
||||
self,
|
||||
level: torch.FloatTensor,
|
||||
) -> Tuple[torch.FloatTensor, torch.LongTensor]:
|
||||
level = -level.view(self.resolution, self.resolution, self.resolution)
|
||||
v_pos, t_pos_idx, _, __ = measure.marching_cubes((level.detach().cpu() if level.is_cuda else level.detach()).numpy(), 0.0) #self.mc_func(level.detach(), 0.0)
|
||||
v_pos = torch.from_numpy(v_pos.copy()).type(torch.FloatTensor).to(level.device)
|
||||
t_pos_idx = torch.from_numpy(t_pos_idx.copy()).type(torch.LongTensor).to(level.device)
|
||||
v_pos = v_pos[..., [0, 1, 2]]
|
||||
t_pos_idx = t_pos_idx[..., [1, 0, 2]]
|
||||
v_pos = v_pos / (self.resolution - 1.0)
|
||||
return v_pos, t_pos_idx
|
||||
@@ -0,0 +1,180 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange, reduce
|
||||
|
||||
from ..utils import (
|
||||
BaseModule,
|
||||
chunk_batch,
|
||||
get_activation,
|
||||
rays_intersect_bbox,
|
||||
scale_tensor,
|
||||
)
|
||||
|
||||
|
||||
class TriplaneNeRFRenderer(BaseModule):
|
||||
@dataclass
|
||||
class Config(BaseModule.Config):
|
||||
radius: float
|
||||
|
||||
feature_reduction: str = "concat"
|
||||
density_activation: str = "trunc_exp"
|
||||
density_bias: float = -1.0
|
||||
color_activation: str = "sigmoid"
|
||||
num_samples_per_ray: int = 128
|
||||
randomized: bool = False
|
||||
|
||||
cfg: Config
|
||||
|
||||
def configure(self) -> None:
|
||||
assert self.cfg.feature_reduction in ["concat", "mean"]
|
||||
self.chunk_size = 0
|
||||
|
||||
def set_chunk_size(self, chunk_size: int):
|
||||
assert (
|
||||
chunk_size >= 0
|
||||
), "chunk_size must be a non-negative integer (0 for no chunking)."
|
||||
self.chunk_size = chunk_size
|
||||
|
||||
def query_triplane(
|
||||
self,
|
||||
decoder: torch.nn.Module,
|
||||
positions: torch.Tensor,
|
||||
triplane: torch.Tensor,
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
input_shape = positions.shape[:-1]
|
||||
positions = positions.view(-1, 3)
|
||||
|
||||
# positions in (-radius, radius)
|
||||
# normalized to (-1, 1) for grid sample
|
||||
positions = scale_tensor(
|
||||
positions, (-self.cfg.radius, self.cfg.radius), (-1, 1)
|
||||
)
|
||||
|
||||
def _query_chunk(x):
|
||||
indices2D: torch.Tensor = torch.stack(
|
||||
(x[..., [0, 1]], x[..., [0, 2]], x[..., [1, 2]]),
|
||||
dim=-3,
|
||||
)
|
||||
out: torch.Tensor = F.grid_sample(
|
||||
rearrange(triplane, "Np Cp Hp Wp -> Np Cp Hp Wp", Np=3),
|
||||
rearrange(indices2D, "Np N Nd -> Np () N Nd", Np=3),
|
||||
align_corners=False,
|
||||
mode="bilinear",
|
||||
)
|
||||
if self.cfg.feature_reduction == "concat":
|
||||
out = rearrange(out, "Np Cp () N -> N (Np Cp)", Np=3)
|
||||
elif self.cfg.feature_reduction == "mean":
|
||||
out = reduce(out, "Np Cp () N -> N Cp", Np=3, reduction="mean")
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
net_out: Dict[str, torch.Tensor] = decoder(out)
|
||||
return net_out
|
||||
|
||||
if self.chunk_size > 0:
|
||||
net_out = chunk_batch(_query_chunk, self.chunk_size, positions)
|
||||
else:
|
||||
net_out = _query_chunk(positions)
|
||||
|
||||
net_out["density_act"] = get_activation(self.cfg.density_activation)(
|
||||
net_out["density"] + self.cfg.density_bias
|
||||
)
|
||||
net_out["color"] = get_activation(self.cfg.color_activation)(
|
||||
net_out["features"]
|
||||
)
|
||||
|
||||
net_out = {k: v.view(*input_shape, -1) for k, v in net_out.items()}
|
||||
|
||||
return net_out
|
||||
|
||||
def _forward(
|
||||
self,
|
||||
decoder: torch.nn.Module,
|
||||
triplane: torch.Tensor,
|
||||
rays_o: torch.Tensor,
|
||||
rays_d: torch.Tensor,
|
||||
**kwargs,
|
||||
):
|
||||
rays_shape = rays_o.shape[:-1]
|
||||
rays_o = rays_o.view(-1, 3)
|
||||
rays_d = rays_d.view(-1, 3)
|
||||
n_rays = rays_o.shape[0]
|
||||
|
||||
t_near, t_far, rays_valid = rays_intersect_bbox(rays_o, rays_d, self.cfg.radius)
|
||||
t_near, t_far = t_near[rays_valid], t_far[rays_valid]
|
||||
|
||||
t_vals = torch.linspace(
|
||||
0, 1, self.cfg.num_samples_per_ray + 1, device=triplane.device
|
||||
)
|
||||
t_mid = (t_vals[:-1] + t_vals[1:]) / 2.0
|
||||
z_vals = t_near * (1 - t_mid[None]) + t_far * t_mid[None] # (N_rays, N_samples)
|
||||
|
||||
xyz = (
|
||||
rays_o[:, None, :] + z_vals[..., None] * rays_d[..., None, :]
|
||||
) # (N_rays, N_sample, 3)
|
||||
|
||||
mlp_out = self.query_triplane(
|
||||
decoder=decoder,
|
||||
positions=xyz,
|
||||
triplane=triplane,
|
||||
)
|
||||
|
||||
eps = 1e-10
|
||||
# deltas = z_vals[:, 1:] - z_vals[:, :-1] # (N_rays, N_samples)
|
||||
deltas = t_vals[1:] - t_vals[:-1] # (N_rays, N_samples)
|
||||
alpha = 1 - torch.exp(
|
||||
-deltas * mlp_out["density_act"][..., 0]
|
||||
) # (N_rays, N_samples)
|
||||
accum_prod = torch.cat(
|
||||
[
|
||||
torch.ones_like(alpha[:, :1]),
|
||||
torch.cumprod(1 - alpha[:, :-1] + eps, dim=-1),
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
weights = alpha * accum_prod # (N_rays, N_samples)
|
||||
comp_rgb_ = (weights[..., None] * mlp_out["color"]).sum(dim=-2) # (N_rays, 3)
|
||||
opacity_ = weights.sum(dim=-1) # (N_rays)
|
||||
|
||||
comp_rgb = torch.zeros(
|
||||
n_rays, 3, dtype=comp_rgb_.dtype, device=comp_rgb_.device
|
||||
)
|
||||
opacity = torch.zeros(n_rays, dtype=opacity_.dtype, device=opacity_.device)
|
||||
comp_rgb[rays_valid] = comp_rgb_
|
||||
opacity[rays_valid] = opacity_
|
||||
|
||||
comp_rgb += 1 - opacity[..., None]
|
||||
comp_rgb = comp_rgb.view(*rays_shape, 3)
|
||||
|
||||
return comp_rgb
|
||||
|
||||
def forward(
|
||||
self,
|
||||
decoder: torch.nn.Module,
|
||||
triplane: torch.Tensor,
|
||||
rays_o: torch.Tensor,
|
||||
rays_d: torch.Tensor,
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
if triplane.ndim == 4:
|
||||
comp_rgb = self._forward(decoder, triplane, rays_o, rays_d)
|
||||
else:
|
||||
comp_rgb = torch.stack(
|
||||
[
|
||||
self._forward(decoder, triplane[i], rays_o[i], rays_d[i])
|
||||
for i in range(triplane.shape[0])
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
|
||||
return comp_rgb
|
||||
|
||||
def train(self, mode=True):
|
||||
self.randomized = mode and self.cfg.randomized
|
||||
return super().train(mode=mode)
|
||||
|
||||
def eval(self):
|
||||
self.randomized = False
|
||||
return super().eval()
|
||||
@@ -0,0 +1,124 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from ..utils import BaseModule
|
||||
|
||||
|
||||
class TriplaneUpsampleNetwork(BaseModule):
|
||||
@dataclass
|
||||
class Config(BaseModule.Config):
|
||||
in_channels: int
|
||||
out_channels: int
|
||||
|
||||
cfg: Config
|
||||
|
||||
def configure(self) -> None:
|
||||
self.upsample = nn.ConvTranspose2d(
|
||||
self.cfg.in_channels, self.cfg.out_channels, kernel_size=2, stride=2
|
||||
)
|
||||
|
||||
def forward(self, triplanes: torch.Tensor) -> torch.Tensor:
|
||||
triplanes_up = rearrange(
|
||||
self.upsample(
|
||||
rearrange(triplanes, "B Np Ci Hp Wp -> (B Np) Ci Hp Wp", Np=3)
|
||||
),
|
||||
"(B Np) Co Hp Wp -> B Np Co Hp Wp",
|
||||
Np=3,
|
||||
)
|
||||
return triplanes_up
|
||||
|
||||
|
||||
class NeRFMLP(BaseModule):
|
||||
@dataclass
|
||||
class Config(BaseModule.Config):
|
||||
in_channels: int
|
||||
n_neurons: int
|
||||
n_hidden_layers: int
|
||||
activation: str = "relu"
|
||||
bias: bool = True
|
||||
weight_init: Optional[str] = "kaiming_uniform"
|
||||
bias_init: Optional[str] = None
|
||||
|
||||
cfg: Config
|
||||
|
||||
def configure(self) -> None:
|
||||
layers = [
|
||||
self.make_linear(
|
||||
self.cfg.in_channels,
|
||||
self.cfg.n_neurons,
|
||||
bias=self.cfg.bias,
|
||||
weight_init=self.cfg.weight_init,
|
||||
bias_init=self.cfg.bias_init,
|
||||
),
|
||||
self.make_activation(self.cfg.activation),
|
||||
]
|
||||
for i in range(self.cfg.n_hidden_layers - 1):
|
||||
layers += [
|
||||
self.make_linear(
|
||||
self.cfg.n_neurons,
|
||||
self.cfg.n_neurons,
|
||||
bias=self.cfg.bias,
|
||||
weight_init=self.cfg.weight_init,
|
||||
bias_init=self.cfg.bias_init,
|
||||
),
|
||||
self.make_activation(self.cfg.activation),
|
||||
]
|
||||
layers += [
|
||||
self.make_linear(
|
||||
self.cfg.n_neurons,
|
||||
4, # density 1 + features 3
|
||||
bias=self.cfg.bias,
|
||||
weight_init=self.cfg.weight_init,
|
||||
bias_init=self.cfg.bias_init,
|
||||
)
|
||||
]
|
||||
self.layers = nn.Sequential(*layers)
|
||||
|
||||
def make_linear(
|
||||
self,
|
||||
dim_in,
|
||||
dim_out,
|
||||
bias=True,
|
||||
weight_init=None,
|
||||
bias_init=None,
|
||||
):
|
||||
layer = nn.Linear(dim_in, dim_out, bias=bias)
|
||||
|
||||
if weight_init is None:
|
||||
pass
|
||||
elif weight_init == "kaiming_uniform":
|
||||
torch.nn.init.kaiming_uniform_(layer.weight, nonlinearity="relu")
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
if bias:
|
||||
if bias_init is None:
|
||||
pass
|
||||
elif bias_init == "zero":
|
||||
torch.nn.init.zeros_(layer.bias)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
return layer
|
||||
|
||||
def make_activation(self, activation):
|
||||
if activation == "relu":
|
||||
return nn.ReLU(inplace=True)
|
||||
elif activation == "silu":
|
||||
return nn.SiLU(inplace=True)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, x):
|
||||
inp_shape = x.shape[:-1]
|
||||
x = x.reshape(-1, x.shape[-1])
|
||||
|
||||
features = self.layers(x)
|
||||
features = features.reshape(*inp_shape, -1)
|
||||
out = {"density": features[..., 0:1], "features": features[..., 1:4]}
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,72 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from huggingface_hub import hf_hub_download
|
||||
from transformers.models.vit.modeling_vit import ViTModel
|
||||
|
||||
from ...utils import BaseModule
|
||||
import os
|
||||
import folder_paths
|
||||
model_path=os.path.join(folder_paths.models_dir,'triposr')
|
||||
|
||||
|
||||
class DINOSingleImageTokenizer(BaseModule):
|
||||
@dataclass
|
||||
class Config(BaseModule.Config):
|
||||
pretrained_model_name_or_path: str = "facebook/dino-vitb16"
|
||||
enable_gradient_checkpointing: bool = False
|
||||
|
||||
cfg: Config
|
||||
|
||||
def configure(self) -> None:
|
||||
print('#Loading ViTModel:',os.path.join(model_path,self.cfg.pretrained_model_name_or_path))
|
||||
self.model: ViTModel = ViTModel(
|
||||
ViTModel.config_class.from_pretrained(
|
||||
hf_hub_download(
|
||||
repo_id=self.cfg.pretrained_model_name_or_path,
|
||||
filename="config.json",
|
||||
local_dir=model_path,
|
||||
endpoint='https://hf-mirror.com'
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
if self.cfg.enable_gradient_checkpointing:
|
||||
self.model.encoder.gradient_checkpointing = True
|
||||
|
||||
self.register_buffer(
|
||||
"image_mean",
|
||||
torch.as_tensor([0.485, 0.456, 0.406]).reshape(1, 1, 3, 1, 1),
|
||||
persistent=False,
|
||||
)
|
||||
self.register_buffer(
|
||||
"image_std",
|
||||
torch.as_tensor([0.229, 0.224, 0.225]).reshape(1, 1, 3, 1, 1),
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
def forward(self, images: torch.FloatTensor, **kwargs) -> torch.FloatTensor:
|
||||
packed = False
|
||||
if images.ndim == 4:
|
||||
packed = True
|
||||
images = images.unsqueeze(1)
|
||||
|
||||
batch_size, n_input_views = images.shape[:2]
|
||||
images = (images - self.image_mean) / self.image_std
|
||||
out = self.model(
|
||||
rearrange(images, "B N C H W -> (B N) C H W"), interpolate_pos_encoding=True
|
||||
)
|
||||
local_features, global_features = out.last_hidden_state, out.pooler_output
|
||||
local_features = local_features.permute(0, 2, 1)
|
||||
local_features = rearrange(
|
||||
local_features, "(B N) Ct Nt -> B N Ct Nt", B=batch_size
|
||||
)
|
||||
if packed:
|
||||
local_features = local_features.squeeze(1)
|
||||
|
||||
return local_features
|
||||
|
||||
def detokenize(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,45 @@
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange, repeat
|
||||
|
||||
from ...utils import BaseModule
|
||||
|
||||
|
||||
class Triplane1DTokenizer(BaseModule):
|
||||
@dataclass
|
||||
class Config(BaseModule.Config):
|
||||
plane_size: int
|
||||
num_channels: int
|
||||
|
||||
cfg: Config
|
||||
|
||||
def configure(self) -> None:
|
||||
self.embeddings = nn.Parameter(
|
||||
torch.randn(
|
||||
(3, self.cfg.num_channels, self.cfg.plane_size, self.cfg.plane_size),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
* 1
|
||||
/ math.sqrt(self.cfg.num_channels)
|
||||
)
|
||||
|
||||
def forward(self, batch_size: int) -> torch.Tensor:
|
||||
return rearrange(
|
||||
repeat(self.embeddings, "Np Ct Hp Wp -> B Np Ct Hp Wp", B=batch_size),
|
||||
"B Np Ct Hp Wp -> B Ct (Np Hp Wp)",
|
||||
)
|
||||
|
||||
def detokenize(self, tokens: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, Ct, Nt = tokens.shape
|
||||
assert Nt == self.cfg.plane_size**2 * 3
|
||||
assert Ct == self.cfg.num_channels
|
||||
return rearrange(
|
||||
tokens,
|
||||
"B Ct (Np Hp Wp) -> B Np Ct Hp Wp",
|
||||
Np=3,
|
||||
Hp=self.cfg.plane_size,
|
||||
Wp=self.cfg.plane_size,
|
||||
)
|
||||
@@ -0,0 +1,653 @@
|
||||
# Copyright 2023 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
# --------
|
||||
#
|
||||
# Modified 2024 by the Tripo AI and Stability AI Team.
|
||||
#
|
||||
# Copyright (c) 2024 Tripo AI & Stability AI
|
||||
#
|
||||
# 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.
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
r"""
|
||||
A cross attention layer.
|
||||
|
||||
Parameters:
|
||||
query_dim (`int`):
|
||||
The number of channels in the query.
|
||||
cross_attention_dim (`int`, *optional*):
|
||||
The number of channels in the encoder_hidden_states. If not given, defaults to `query_dim`.
|
||||
heads (`int`, *optional*, defaults to 8):
|
||||
The number of heads to use for multi-head attention.
|
||||
dim_head (`int`, *optional*, defaults to 64):
|
||||
The number of channels in each head.
|
||||
dropout (`float`, *optional*, defaults to 0.0):
|
||||
The dropout probability to use.
|
||||
bias (`bool`, *optional*, defaults to False):
|
||||
Set to `True` for the query, key, and value linear layers to contain a bias parameter.
|
||||
upcast_attention (`bool`, *optional*, defaults to False):
|
||||
Set to `True` to upcast the attention computation to `float32`.
|
||||
upcast_softmax (`bool`, *optional*, defaults to False):
|
||||
Set to `True` to upcast the softmax computation to `float32`.
|
||||
cross_attention_norm (`str`, *optional*, defaults to `None`):
|
||||
The type of normalization to use for the cross attention. Can be `None`, `layer_norm`, or `group_norm`.
|
||||
cross_attention_norm_num_groups (`int`, *optional*, defaults to 32):
|
||||
The number of groups to use for the group norm in the cross attention.
|
||||
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
|
||||
The number of channels to use for the added key and value projections. If `None`, no projection is used.
|
||||
norm_num_groups (`int`, *optional*, defaults to `None`):
|
||||
The number of groups to use for the group norm in the attention.
|
||||
spatial_norm_dim (`int`, *optional*, defaults to `None`):
|
||||
The number of channels to use for the spatial normalization.
|
||||
out_bias (`bool`, *optional*, defaults to `True`):
|
||||
Set to `True` to use a bias in the output linear layer.
|
||||
scale_qk (`bool`, *optional*, defaults to `True`):
|
||||
Set to `True` to scale the query and key by `1 / sqrt(dim_head)`.
|
||||
only_cross_attention (`bool`, *optional*, defaults to `False`):
|
||||
Set to `True` to only use cross attention and not added_kv_proj_dim. Can only be set to `True` if
|
||||
`added_kv_proj_dim` is not `None`.
|
||||
eps (`float`, *optional*, defaults to 1e-5):
|
||||
An additional value added to the denominator in group normalization that is used for numerical stability.
|
||||
rescale_output_factor (`float`, *optional*, defaults to 1.0):
|
||||
A factor to rescale the output by dividing it with this value.
|
||||
residual_connection (`bool`, *optional*, defaults to `False`):
|
||||
Set to `True` to add the residual connection to the output.
|
||||
_from_deprecated_attn_block (`bool`, *optional*, defaults to `False`):
|
||||
Set to `True` if the attention block is loaded from a deprecated state dict.
|
||||
processor (`AttnProcessor`, *optional*, defaults to `None`):
|
||||
The attention processor to use. If `None`, defaults to `AttnProcessor2_0` if `torch 2.x` is used and
|
||||
`AttnProcessor` otherwise.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
heads: int = 8,
|
||||
dim_head: int = 64,
|
||||
dropout: float = 0.0,
|
||||
bias: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
upcast_softmax: bool = False,
|
||||
cross_attention_norm: Optional[str] = None,
|
||||
cross_attention_norm_num_groups: int = 32,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
norm_num_groups: Optional[int] = None,
|
||||
out_bias: bool = True,
|
||||
scale_qk: bool = True,
|
||||
only_cross_attention: bool = False,
|
||||
eps: float = 1e-5,
|
||||
rescale_output_factor: float = 1.0,
|
||||
residual_connection: bool = False,
|
||||
_from_deprecated_attn_block: bool = False,
|
||||
processor: Optional["AttnProcessor"] = None,
|
||||
out_dim: int = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
|
||||
self.query_dim = query_dim
|
||||
self.cross_attention_dim = (
|
||||
cross_attention_dim if cross_attention_dim is not None else query_dim
|
||||
)
|
||||
self.upcast_attention = upcast_attention
|
||||
self.upcast_softmax = upcast_softmax
|
||||
self.rescale_output_factor = rescale_output_factor
|
||||
self.residual_connection = residual_connection
|
||||
self.dropout = dropout
|
||||
self.fused_projections = False
|
||||
self.out_dim = out_dim if out_dim is not None else query_dim
|
||||
|
||||
# we make use of this private variable to know whether this class is loaded
|
||||
# with an deprecated state dict so that we can convert it on the fly
|
||||
self._from_deprecated_attn_block = _from_deprecated_attn_block
|
||||
|
||||
self.scale_qk = scale_qk
|
||||
self.scale = dim_head**-0.5 if self.scale_qk else 1.0
|
||||
|
||||
self.heads = out_dim // dim_head if out_dim is not None else heads
|
||||
# for slice_size > 0 the attention score computation
|
||||
# is split across the batch axis to save memory
|
||||
# You can set slice_size with `set_attention_slice`
|
||||
self.sliceable_head_dim = heads
|
||||
|
||||
self.added_kv_proj_dim = added_kv_proj_dim
|
||||
self.only_cross_attention = only_cross_attention
|
||||
|
||||
if self.added_kv_proj_dim is None and self.only_cross_attention:
|
||||
raise ValueError(
|
||||
"`only_cross_attention` can only be set to True if `added_kv_proj_dim` is not None. Make sure to set either `only_cross_attention=False` or define `added_kv_proj_dim`."
|
||||
)
|
||||
|
||||
if norm_num_groups is not None:
|
||||
self.group_norm = nn.GroupNorm(
|
||||
num_channels=query_dim, num_groups=norm_num_groups, eps=eps, affine=True
|
||||
)
|
||||
else:
|
||||
self.group_norm = None
|
||||
|
||||
self.spatial_norm = None
|
||||
|
||||
if cross_attention_norm is None:
|
||||
self.norm_cross = None
|
||||
elif cross_attention_norm == "layer_norm":
|
||||
self.norm_cross = nn.LayerNorm(self.cross_attention_dim)
|
||||
elif cross_attention_norm == "group_norm":
|
||||
if self.added_kv_proj_dim is not None:
|
||||
# The given `encoder_hidden_states` are initially of shape
|
||||
# (batch_size, seq_len, added_kv_proj_dim) before being projected
|
||||
# to (batch_size, seq_len, cross_attention_dim). The norm is applied
|
||||
# before the projection, so we need to use `added_kv_proj_dim` as
|
||||
# the number of channels for the group norm.
|
||||
norm_cross_num_channels = added_kv_proj_dim
|
||||
else:
|
||||
norm_cross_num_channels = self.cross_attention_dim
|
||||
|
||||
self.norm_cross = nn.GroupNorm(
|
||||
num_channels=norm_cross_num_channels,
|
||||
num_groups=cross_attention_norm_num_groups,
|
||||
eps=1e-5,
|
||||
affine=True,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"unknown cross_attention_norm: {cross_attention_norm}. Should be None, 'layer_norm' or 'group_norm'"
|
||||
)
|
||||
|
||||
linear_cls = nn.Linear
|
||||
|
||||
self.linear_cls = linear_cls
|
||||
self.to_q = linear_cls(query_dim, self.inner_dim, bias=bias)
|
||||
|
||||
if not self.only_cross_attention:
|
||||
# only relevant for the `AddedKVProcessor` classes
|
||||
self.to_k = linear_cls(self.cross_attention_dim, self.inner_dim, bias=bias)
|
||||
self.to_v = linear_cls(self.cross_attention_dim, self.inner_dim, bias=bias)
|
||||
else:
|
||||
self.to_k = None
|
||||
self.to_v = None
|
||||
|
||||
if self.added_kv_proj_dim is not None:
|
||||
self.add_k_proj = linear_cls(added_kv_proj_dim, self.inner_dim)
|
||||
self.add_v_proj = linear_cls(added_kv_proj_dim, self.inner_dim)
|
||||
|
||||
self.to_out = nn.ModuleList([])
|
||||
self.to_out.append(linear_cls(self.inner_dim, self.out_dim, bias=out_bias))
|
||||
self.to_out.append(nn.Dropout(dropout))
|
||||
|
||||
# set attention processor
|
||||
# We use the AttnProcessor2_0 by default when torch 2.x is used which uses
|
||||
# torch.nn.functional.scaled_dot_product_attention for native Flash/memory_efficient_attention
|
||||
# but only if it has the default `scale` argument. TODO remove scale_qk check when we move to torch 2.1
|
||||
if processor is None:
|
||||
processor = (
|
||||
AttnProcessor2_0()
|
||||
if hasattr(F, "scaled_dot_product_attention") and self.scale_qk
|
||||
else AttnProcessor()
|
||||
)
|
||||
self.set_processor(processor)
|
||||
|
||||
def set_processor(self, processor: "AttnProcessor") -> None:
|
||||
self.processor = processor
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: Optional[torch.FloatTensor] = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
**cross_attention_kwargs,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
The forward method of the `Attention` class.
|
||||
|
||||
Args:
|
||||
hidden_states (`torch.Tensor`):
|
||||
The hidden states of the query.
|
||||
encoder_hidden_states (`torch.Tensor`, *optional*):
|
||||
The hidden states of the encoder.
|
||||
attention_mask (`torch.Tensor`, *optional*):
|
||||
The attention mask to use. If `None`, no mask is applied.
|
||||
**cross_attention_kwargs:
|
||||
Additional keyword arguments to pass along to the cross attention.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`: The output of the attention layer.
|
||||
"""
|
||||
# The `Attention` class can call different attention processors / attention functions
|
||||
# here we simply pass along all tensors to the selected processor class
|
||||
# For standard processors that are defined here, `**cross_attention_kwargs` is empty
|
||||
return self.processor(
|
||||
self,
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
**cross_attention_kwargs,
|
||||
)
|
||||
|
||||
def batch_to_head_dim(self, tensor: torch.Tensor) -> torch.Tensor:
|
||||
r"""
|
||||
Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size // heads, seq_len, dim * heads]`. `heads`
|
||||
is the number of heads initialized while constructing the `Attention` class.
|
||||
|
||||
Args:
|
||||
tensor (`torch.Tensor`): The tensor to reshape.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`: The reshaped tensor.
|
||||
"""
|
||||
head_size = self.heads
|
||||
batch_size, seq_len, dim = tensor.shape
|
||||
tensor = tensor.reshape(batch_size // head_size, head_size, seq_len, dim)
|
||||
tensor = tensor.permute(0, 2, 1, 3).reshape(
|
||||
batch_size // head_size, seq_len, dim * head_size
|
||||
)
|
||||
return tensor
|
||||
|
||||
def head_to_batch_dim(self, tensor: torch.Tensor, out_dim: int = 3) -> torch.Tensor:
|
||||
r"""
|
||||
Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size, seq_len, heads, dim // heads]` `heads` is
|
||||
the number of heads initialized while constructing the `Attention` class.
|
||||
|
||||
Args:
|
||||
tensor (`torch.Tensor`): The tensor to reshape.
|
||||
out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor. If `3`, the tensor is
|
||||
reshaped to `[batch_size * heads, seq_len, dim // heads]`.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`: The reshaped tensor.
|
||||
"""
|
||||
head_size = self.heads
|
||||
batch_size, seq_len, dim = tensor.shape
|
||||
tensor = tensor.reshape(batch_size, seq_len, head_size, dim // head_size)
|
||||
tensor = tensor.permute(0, 2, 1, 3)
|
||||
|
||||
if out_dim == 3:
|
||||
tensor = tensor.reshape(batch_size * head_size, seq_len, dim // head_size)
|
||||
|
||||
return tensor
|
||||
|
||||
def get_attention_scores(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
attention_mask: torch.Tensor = None,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Compute the attention scores.
|
||||
|
||||
Args:
|
||||
query (`torch.Tensor`): The query tensor.
|
||||
key (`torch.Tensor`): The key tensor.
|
||||
attention_mask (`torch.Tensor`, *optional*): The attention mask to use. If `None`, no mask is applied.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`: The attention probabilities/scores.
|
||||
"""
|
||||
dtype = query.dtype
|
||||
if self.upcast_attention:
|
||||
query = query.float()
|
||||
key = key.float()
|
||||
|
||||
if attention_mask is None:
|
||||
baddbmm_input = torch.empty(
|
||||
query.shape[0],
|
||||
query.shape[1],
|
||||
key.shape[1],
|
||||
dtype=query.dtype,
|
||||
device=query.device,
|
||||
)
|
||||
beta = 0
|
||||
else:
|
||||
baddbmm_input = attention_mask
|
||||
beta = 1
|
||||
|
||||
attention_scores = torch.baddbmm(
|
||||
baddbmm_input,
|
||||
query,
|
||||
key.transpose(-1, -2),
|
||||
beta=beta,
|
||||
alpha=self.scale,
|
||||
)
|
||||
del baddbmm_input
|
||||
|
||||
if self.upcast_softmax:
|
||||
attention_scores = attention_scores.float()
|
||||
|
||||
attention_probs = attention_scores.softmax(dim=-1)
|
||||
del attention_scores
|
||||
|
||||
attention_probs = attention_probs.to(dtype)
|
||||
|
||||
return attention_probs
|
||||
|
||||
def prepare_attention_mask(
|
||||
self,
|
||||
attention_mask: torch.Tensor,
|
||||
target_length: int,
|
||||
batch_size: int,
|
||||
out_dim: int = 3,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Prepare the attention mask for the attention computation.
|
||||
|
||||
Args:
|
||||
attention_mask (`torch.Tensor`):
|
||||
The attention mask to prepare.
|
||||
target_length (`int`):
|
||||
The target length of the attention mask. This is the length of the attention mask after padding.
|
||||
batch_size (`int`):
|
||||
The batch size, which is used to repeat the attention mask.
|
||||
out_dim (`int`, *optional*, defaults to `3`):
|
||||
The output dimension of the attention mask. Can be either `3` or `4`.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`: The prepared attention mask.
|
||||
"""
|
||||
head_size = self.heads
|
||||
if attention_mask is None:
|
||||
return attention_mask
|
||||
|
||||
current_length: int = attention_mask.shape[-1]
|
||||
if current_length != target_length:
|
||||
if attention_mask.device.type == "mps":
|
||||
# HACK: MPS: Does not support padding by greater than dimension of input tensor.
|
||||
# Instead, we can manually construct the padding tensor.
|
||||
padding_shape = (
|
||||
attention_mask.shape[0],
|
||||
attention_mask.shape[1],
|
||||
target_length,
|
||||
)
|
||||
padding = torch.zeros(
|
||||
padding_shape,
|
||||
dtype=attention_mask.dtype,
|
||||
device=attention_mask.device,
|
||||
)
|
||||
attention_mask = torch.cat([attention_mask, padding], dim=2)
|
||||
else:
|
||||
# TODO: for pipelines such as stable-diffusion, padding cross-attn mask:
|
||||
# we want to instead pad by (0, remaining_length), where remaining_length is:
|
||||
# remaining_length: int = target_length - current_length
|
||||
# TODO: re-enable tests/models/test_models_unet_2d_condition.py#test_model_xattn_padding
|
||||
attention_mask = F.pad(attention_mask, (0, target_length), value=0.0)
|
||||
|
||||
if out_dim == 3:
|
||||
if attention_mask.shape[0] < batch_size * head_size:
|
||||
attention_mask = attention_mask.repeat_interleave(head_size, dim=0)
|
||||
elif out_dim == 4:
|
||||
attention_mask = attention_mask.unsqueeze(1)
|
||||
attention_mask = attention_mask.repeat_interleave(head_size, dim=1)
|
||||
|
||||
return attention_mask
|
||||
|
||||
def norm_encoder_hidden_states(
|
||||
self, encoder_hidden_states: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Normalize the encoder hidden states. Requires `self.norm_cross` to be specified when constructing the
|
||||
`Attention` class.
|
||||
|
||||
Args:
|
||||
encoder_hidden_states (`torch.Tensor`): Hidden states of the encoder.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`: The normalized encoder hidden states.
|
||||
"""
|
||||
assert (
|
||||
self.norm_cross is not None
|
||||
), "self.norm_cross must be defined to call self.norm_encoder_hidden_states"
|
||||
|
||||
if isinstance(self.norm_cross, nn.LayerNorm):
|
||||
encoder_hidden_states = self.norm_cross(encoder_hidden_states)
|
||||
elif isinstance(self.norm_cross, nn.GroupNorm):
|
||||
# Group norm norms along the channels dimension and expects
|
||||
# input to be in the shape of (N, C, *). In this case, we want
|
||||
# to norm along the hidden dimension, so we need to move
|
||||
# (batch_size, sequence_length, hidden_size) ->
|
||||
# (batch_size, hidden_size, sequence_length)
|
||||
encoder_hidden_states = encoder_hidden_states.transpose(1, 2)
|
||||
encoder_hidden_states = self.norm_cross(encoder_hidden_states)
|
||||
encoder_hidden_states = encoder_hidden_states.transpose(1, 2)
|
||||
else:
|
||||
assert False
|
||||
|
||||
return encoder_hidden_states
|
||||
|
||||
@torch.no_grad()
|
||||
def fuse_projections(self, fuse=True):
|
||||
is_cross_attention = self.cross_attention_dim != self.query_dim
|
||||
device = self.to_q.weight.data.device
|
||||
dtype = self.to_q.weight.data.dtype
|
||||
|
||||
if not is_cross_attention:
|
||||
# fetch weight matrices.
|
||||
concatenated_weights = torch.cat(
|
||||
[self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]
|
||||
)
|
||||
in_features = concatenated_weights.shape[1]
|
||||
out_features = concatenated_weights.shape[0]
|
||||
|
||||
# create a new single projection layer and copy over the weights.
|
||||
self.to_qkv = self.linear_cls(
|
||||
in_features, out_features, bias=False, device=device, dtype=dtype
|
||||
)
|
||||
self.to_qkv.weight.copy_(concatenated_weights)
|
||||
|
||||
else:
|
||||
concatenated_weights = torch.cat(
|
||||
[self.to_k.weight.data, self.to_v.weight.data]
|
||||
)
|
||||
in_features = concatenated_weights.shape[1]
|
||||
out_features = concatenated_weights.shape[0]
|
||||
|
||||
self.to_kv = self.linear_cls(
|
||||
in_features, out_features, bias=False, device=device, dtype=dtype
|
||||
)
|
||||
self.to_kv.weight.copy_(concatenated_weights)
|
||||
|
||||
self.fused_projections = fuse
|
||||
|
||||
|
||||
class AttnProcessor:
|
||||
r"""
|
||||
Default processor for performing attention-related computations.
|
||||
"""
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: Optional[torch.FloatTensor] = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
residual = hidden_states
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(
|
||||
batch_size, channel, height * width
|
||||
).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape
|
||||
if encoder_hidden_states is None
|
||||
else encoder_hidden_states.shape
|
||||
)
|
||||
attention_mask = attn.prepare_attention_mask(
|
||||
attention_mask, sequence_length, batch_size
|
||||
)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(
|
||||
1, 2
|
||||
)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(
|
||||
encoder_hidden_states
|
||||
)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
query = attn.head_to_batch_dim(query)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(
|
||||
batch_size, channel, height, width
|
||||
)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AttnProcessor2_0:
|
||||
r"""
|
||||
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError(
|
||||
"AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0."
|
||||
)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: Optional[torch.FloatTensor] = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
residual = hidden_states
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(
|
||||
batch_size, channel, height * width
|
||||
).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape
|
||||
if encoder_hidden_states is None
|
||||
else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = attn.prepare_attention_mask(
|
||||
attention_mask, sequence_length, batch_size
|
||||
)
|
||||
# scaled_dot_product_attention expects attention_mask shape to be
|
||||
# (batch, heads, source_length, target_length)
|
||||
attention_mask = attention_mask.view(
|
||||
batch_size, attn.heads, -1, attention_mask.shape[-1]
|
||||
)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(
|
||||
1, 2
|
||||
)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(
|
||||
encoder_hidden_states
|
||||
)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
# TODO: add support for attn.scale when we move to Torch 2.1
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(
|
||||
batch_size, -1, attn.heads * head_dim
|
||||
)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(
|
||||
batch_size, channel, height, width
|
||||
)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,334 @@
|
||||
# Copyright 2023 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
# --------
|
||||
#
|
||||
# Modified 2024 by the Tripo AI and Stability AI Team.
|
||||
#
|
||||
# Copyright (c) 2024 Tripo AI & Stability AI
|
||||
#
|
||||
# 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.
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from .attention import Attention
|
||||
|
||||
|
||||
class BasicTransformerBlock(nn.Module):
|
||||
r"""
|
||||
A basic Transformer block.
|
||||
|
||||
Parameters:
|
||||
dim (`int`): The number of channels in the input and output.
|
||||
num_attention_heads (`int`): The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`): The number of channels in each head.
|
||||
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
|
||||
cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
|
||||
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
|
||||
attention_bias (:
|
||||
obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
|
||||
only_cross_attention (`bool`, *optional*):
|
||||
Whether to use only cross-attention layers. In this case two cross attention layers are used.
|
||||
double_self_attention (`bool`, *optional*):
|
||||
Whether to use two self-attention layers. In this case no cross attention layers are used.
|
||||
upcast_attention (`bool`, *optional*):
|
||||
Whether to upcast the attention computation to float32. This is useful for mixed precision training.
|
||||
norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use learnable elementwise affine parameters for normalization.
|
||||
norm_type (`str`, *optional*, defaults to `"layer_norm"`):
|
||||
The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
|
||||
final_dropout (`bool` *optional*, defaults to False):
|
||||
Whether to apply a final dropout after the last feed-forward layer.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
dropout=0.0,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
activation_fn: str = "geglu",
|
||||
attention_bias: bool = False,
|
||||
only_cross_attention: bool = False,
|
||||
double_self_attention: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
norm_elementwise_affine: bool = True,
|
||||
norm_type: str = "layer_norm",
|
||||
final_dropout: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.only_cross_attention = only_cross_attention
|
||||
|
||||
assert norm_type == "layer_norm"
|
||||
|
||||
# Define 3 blocks. Each block has its own normalization layer.
|
||||
# 1. Self-Attn
|
||||
self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine)
|
||||
self.attn1 = Attention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
cross_attention_dim=cross_attention_dim if only_cross_attention else None,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
|
||||
# 2. Cross-Attn
|
||||
if cross_attention_dim is not None or double_self_attention:
|
||||
# We currently only use AdaLayerNormZero for self attention where there will only be one attention block.
|
||||
# I.e. the number of returned modulation chunks from AdaLayerZero would not make sense if returned during
|
||||
# the second cross attention block.
|
||||
self.norm2 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine)
|
||||
|
||||
self.attn2 = Attention(
|
||||
query_dim=dim,
|
||||
cross_attention_dim=(
|
||||
cross_attention_dim if not double_self_attention else None
|
||||
),
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
) # is self-attn if encoder_hidden_states is none
|
||||
else:
|
||||
self.norm2 = None
|
||||
self.attn2 = None
|
||||
|
||||
# 3. Feed-forward
|
||||
self.norm3 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine)
|
||||
self.ff = FeedForward(
|
||||
dim,
|
||||
dropout=dropout,
|
||||
activation_fn=activation_fn,
|
||||
final_dropout=final_dropout,
|
||||
)
|
||||
|
||||
# let chunk size default to None
|
||||
self._chunk_size = None
|
||||
self._chunk_dim = 0
|
||||
|
||||
def set_chunk_feed_forward(self, chunk_size: Optional[int], dim: int):
|
||||
# Sets chunk feed-forward
|
||||
self._chunk_size = chunk_size
|
||||
self._chunk_dim = dim
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
encoder_hidden_states: Optional[torch.FloatTensor] = None,
|
||||
encoder_attention_mask: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
# Notice that normalization is always applied before the real computation in the following blocks.
|
||||
# 0. Self-Attention
|
||||
norm_hidden_states = self.norm1(hidden_states)
|
||||
|
||||
attn_output = self.attn1(
|
||||
norm_hidden_states,
|
||||
encoder_hidden_states=(
|
||||
encoder_hidden_states if self.only_cross_attention else None
|
||||
),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
hidden_states = attn_output + hidden_states
|
||||
|
||||
# 3. Cross-Attention
|
||||
if self.attn2 is not None:
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
|
||||
attn_output = self.attn2(
|
||||
norm_hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=encoder_attention_mask,
|
||||
)
|
||||
hidden_states = attn_output + hidden_states
|
||||
|
||||
# 4. Feed-forward
|
||||
norm_hidden_states = self.norm3(hidden_states)
|
||||
|
||||
if self._chunk_size is not None:
|
||||
# "feed_forward_chunk_size" can be used to save memory
|
||||
if norm_hidden_states.shape[self._chunk_dim] % self._chunk_size != 0:
|
||||
raise ValueError(
|
||||
f"`hidden_states` dimension to be chunked: {norm_hidden_states.shape[self._chunk_dim]} has to be divisible by chunk size: {self._chunk_size}. Make sure to set an appropriate `chunk_size` when calling `unet.enable_forward_chunking`."
|
||||
)
|
||||
|
||||
num_chunks = norm_hidden_states.shape[self._chunk_dim] // self._chunk_size
|
||||
ff_output = torch.cat(
|
||||
[
|
||||
self.ff(hid_slice)
|
||||
for hid_slice in norm_hidden_states.chunk(
|
||||
num_chunks, dim=self._chunk_dim
|
||||
)
|
||||
],
|
||||
dim=self._chunk_dim,
|
||||
)
|
||||
else:
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
|
||||
hidden_states = ff_output + hidden_states
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
r"""
|
||||
A feed-forward layer.
|
||||
|
||||
Parameters:
|
||||
dim (`int`): The number of channels in the input.
|
||||
dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
|
||||
mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
|
||||
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
|
||||
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
|
||||
final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
dim_out: Optional[int] = None,
|
||||
mult: int = 4,
|
||||
dropout: float = 0.0,
|
||||
activation_fn: str = "geglu",
|
||||
final_dropout: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
inner_dim = int(dim * mult)
|
||||
dim_out = dim_out if dim_out is not None else dim
|
||||
linear_cls = nn.Linear
|
||||
|
||||
if activation_fn == "gelu":
|
||||
act_fn = GELU(dim, inner_dim)
|
||||
if activation_fn == "gelu-approximate":
|
||||
act_fn = GELU(dim, inner_dim, approximate="tanh")
|
||||
elif activation_fn == "geglu":
|
||||
act_fn = GEGLU(dim, inner_dim)
|
||||
elif activation_fn == "geglu-approximate":
|
||||
act_fn = ApproximateGELU(dim, inner_dim)
|
||||
|
||||
self.net = nn.ModuleList([])
|
||||
# project in
|
||||
self.net.append(act_fn)
|
||||
# project dropout
|
||||
self.net.append(nn.Dropout(dropout))
|
||||
# project out
|
||||
self.net.append(linear_cls(inner_dim, dim_out))
|
||||
# FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout
|
||||
if final_dropout:
|
||||
self.net.append(nn.Dropout(dropout))
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
for module in self.net:
|
||||
hidden_states = module(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class GELU(nn.Module):
|
||||
r"""
|
||||
GELU activation function with tanh approximation support with `approximate="tanh"`.
|
||||
|
||||
Parameters:
|
||||
dim_in (`int`): The number of channels in the input.
|
||||
dim_out (`int`): The number of channels in the output.
|
||||
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
|
||||
"""
|
||||
|
||||
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none"):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(dim_in, dim_out)
|
||||
self.approximate = approximate
|
||||
|
||||
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
||||
if gate.device.type != "mps":
|
||||
return F.gelu(gate, approximate=self.approximate)
|
||||
# mps: gelu is not implemented for float16
|
||||
return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(
|
||||
dtype=gate.dtype
|
||||
)
|
||||
|
||||
def forward(self, hidden_states):
|
||||
hidden_states = self.proj(hidden_states)
|
||||
hidden_states = self.gelu(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class GEGLU(nn.Module):
|
||||
r"""
|
||||
A variant of the gated linear unit activation function from https://arxiv.org/abs/2002.05202.
|
||||
|
||||
Parameters:
|
||||
dim_in (`int`): The number of channels in the input.
|
||||
dim_out (`int`): The number of channels in the output.
|
||||
"""
|
||||
|
||||
def __init__(self, dim_in: int, dim_out: int):
|
||||
super().__init__()
|
||||
linear_cls = nn.Linear
|
||||
|
||||
self.proj = linear_cls(dim_in, dim_out * 2)
|
||||
|
||||
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
||||
if gate.device.type != "mps":
|
||||
return F.gelu(gate)
|
||||
# mps: gelu is not implemented for float16
|
||||
return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype)
|
||||
|
||||
def forward(self, hidden_states, scale: float = 1.0):
|
||||
args = ()
|
||||
hidden_states, gate = self.proj(hidden_states, *args).chunk(2, dim=-1)
|
||||
return hidden_states * self.gelu(gate)
|
||||
|
||||
|
||||
class ApproximateGELU(nn.Module):
|
||||
r"""
|
||||
The approximate form of Gaussian Error Linear Unit (GELU). For more details, see section 2:
|
||||
https://arxiv.org/abs/1606.08415.
|
||||
|
||||
Parameters:
|
||||
dim_in (`int`): The number of channels in the input.
|
||||
dim_out (`int`): The number of channels in the output.
|
||||
"""
|
||||
|
||||
def __init__(self, dim_in: int, dim_out: int):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(dim_in, dim_out)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.proj(x)
|
||||
return x * torch.sigmoid(1.702 * x)
|
||||
@@ -0,0 +1,219 @@
|
||||
# Copyright 2023 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
# --------
|
||||
#
|
||||
# Modified 2024 by the Tripo AI and Stability AI Team.
|
||||
#
|
||||
# Copyright (c) 2024 Tripo AI & Stability AI
|
||||
#
|
||||
# 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.
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from ...utils import BaseModule
|
||||
from .basic_transformer_block import BasicTransformerBlock
|
||||
|
||||
|
||||
class Transformer1D(BaseModule):
|
||||
@dataclass
|
||||
class Config(BaseModule.Config):
|
||||
num_attention_heads: int = 16
|
||||
attention_head_dim: int = 88
|
||||
in_channels: Optional[int] = None
|
||||
out_channels: Optional[int] = None
|
||||
num_layers: int = 1
|
||||
dropout: float = 0.0
|
||||
norm_num_groups: int = 32
|
||||
cross_attention_dim: Optional[int] = None
|
||||
attention_bias: bool = False
|
||||
activation_fn: str = "geglu"
|
||||
only_cross_attention: bool = False
|
||||
double_self_attention: bool = False
|
||||
upcast_attention: bool = False
|
||||
norm_type: str = "layer_norm"
|
||||
norm_elementwise_affine: bool = True
|
||||
gradient_checkpointing: bool = False
|
||||
|
||||
cfg: Config
|
||||
|
||||
def configure(self) -> None:
|
||||
self.num_attention_heads = self.cfg.num_attention_heads
|
||||
self.attention_head_dim = self.cfg.attention_head_dim
|
||||
inner_dim = self.num_attention_heads * self.attention_head_dim
|
||||
|
||||
linear_cls = nn.Linear
|
||||
|
||||
# 2. Define input layers
|
||||
self.in_channels = self.cfg.in_channels
|
||||
|
||||
self.norm = torch.nn.GroupNorm(
|
||||
num_groups=self.cfg.norm_num_groups,
|
||||
num_channels=self.cfg.in_channels,
|
||||
eps=1e-6,
|
||||
affine=True,
|
||||
)
|
||||
self.proj_in = linear_cls(self.cfg.in_channels, inner_dim)
|
||||
|
||||
# 3. Define transformers blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
self.num_attention_heads,
|
||||
self.attention_head_dim,
|
||||
dropout=self.cfg.dropout,
|
||||
cross_attention_dim=self.cfg.cross_attention_dim,
|
||||
activation_fn=self.cfg.activation_fn,
|
||||
attention_bias=self.cfg.attention_bias,
|
||||
only_cross_attention=self.cfg.only_cross_attention,
|
||||
double_self_attention=self.cfg.double_self_attention,
|
||||
upcast_attention=self.cfg.upcast_attention,
|
||||
norm_type=self.cfg.norm_type,
|
||||
norm_elementwise_affine=self.cfg.norm_elementwise_affine,
|
||||
)
|
||||
for d in range(self.cfg.num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Define output layers
|
||||
self.out_channels = (
|
||||
self.cfg.in_channels
|
||||
if self.cfg.out_channels is None
|
||||
else self.cfg.out_channels
|
||||
)
|
||||
|
||||
self.proj_out = linear_cls(inner_dim, self.cfg.in_channels)
|
||||
|
||||
self.gradient_checkpointing = self.cfg.gradient_checkpointing
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
encoder_attention_mask: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""
|
||||
The [`Transformer1DModel`] forward method.
|
||||
|
||||
Args:
|
||||
hidden_states (`torch.LongTensor` of shape `(batch size, num latent pixels)` if discrete, `torch.FloatTensor` of shape `(batch size, channel, height, width)` if continuous):
|
||||
Input `hidden_states`.
|
||||
encoder_hidden_states ( `torch.FloatTensor` of shape `(batch size, sequence len, embed dims)`, *optional*):
|
||||
Conditional embeddings for cross attention layer. If not given, cross-attention defaults to
|
||||
self-attention.
|
||||
attention_mask ( `torch.Tensor`, *optional*):
|
||||
An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask
|
||||
is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large
|
||||
negative values to the attention scores corresponding to "discard" tokens.
|
||||
encoder_attention_mask ( `torch.Tensor`, *optional*):
|
||||
Cross-attention mask applied to `encoder_hidden_states`. Two formats supported:
|
||||
|
||||
* Mask `(batch, sequence_length)` True = keep, False = discard.
|
||||
* Bias `(batch, 1, sequence_length)` 0 = keep, -10000 = discard.
|
||||
|
||||
If `ndim == 2`: will be interpreted as a mask, then converted into a bias consistent with the format
|
||||
above. This bias will be added to the cross-attention scores.
|
||||
|
||||
Returns:
|
||||
torch.FloatTensor
|
||||
"""
|
||||
# ensure attention_mask is a bias, and give it a singleton query_tokens dimension.
|
||||
# we may have done this conversion already, e.g. if we came here via UNet2DConditionModel#forward.
|
||||
# we can tell by counting dims; if ndim == 2: it's a mask rather than a bias.
|
||||
# expects mask of shape:
|
||||
# [batch, key_tokens]
|
||||
# adds singleton query_tokens dimension:
|
||||
# [batch, 1, key_tokens]
|
||||
# this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes:
|
||||
# [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn)
|
||||
# [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn)
|
||||
if attention_mask is not None and attention_mask.ndim == 2:
|
||||
# assume that mask is expressed as:
|
||||
# (1 = keep, 0 = discard)
|
||||
# convert mask into a bias that can be added to attention scores:
|
||||
# (keep = +0, discard = -10000.0)
|
||||
attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0
|
||||
attention_mask = attention_mask.unsqueeze(1)
|
||||
|
||||
# convert encoder_attention_mask to a bias the same way we do for attention_mask
|
||||
if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2:
|
||||
encoder_attention_mask = (
|
||||
1 - encoder_attention_mask.to(hidden_states.dtype)
|
||||
) * -10000.0
|
||||
encoder_attention_mask = encoder_attention_mask.unsqueeze(1)
|
||||
|
||||
# 1. Input
|
||||
batch, _, seq_len = hidden_states.shape
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 1).reshape(
|
||||
batch, seq_len, inner_dim
|
||||
)
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
# 2. Blocks
|
||||
for block in self.transformer_blocks:
|
||||
if self.training and self.gradient_checkpointing:
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(
|
||||
block,
|
||||
hidden_states,
|
||||
attention_mask,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
|
||||
# 3. Output
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = (
|
||||
hidden_states.reshape(batch, seq_len, inner_dim)
|
||||
.permute(0, 2, 1)
|
||||
.contiguous()
|
||||
)
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
return output
|
||||
@@ -0,0 +1,218 @@
|
||||
import math
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Union
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import trimesh
|
||||
from einops import rearrange
|
||||
from huggingface_hub import hf_hub_download
|
||||
from omegaconf import OmegaConf
|
||||
from PIL import Image
|
||||
|
||||
from .models.isosurface import MarchingCubeHelper
|
||||
from .utils import (
|
||||
BaseModule,
|
||||
ImagePreprocessor,
|
||||
find_class,
|
||||
get_spherical_cameras,
|
||||
scale_tensor,
|
||||
)
|
||||
|
||||
|
||||
class TSR(BaseModule):
|
||||
@dataclass
|
||||
class Config(BaseModule.Config):
|
||||
cond_image_size: int
|
||||
|
||||
image_tokenizer_cls: str
|
||||
image_tokenizer: dict
|
||||
|
||||
tokenizer_cls: str
|
||||
tokenizer: dict
|
||||
|
||||
backbone_cls: str
|
||||
backbone: dict
|
||||
|
||||
post_processor_cls: str
|
||||
post_processor: dict
|
||||
|
||||
decoder_cls: str
|
||||
decoder: dict
|
||||
|
||||
renderer_cls: str
|
||||
renderer: dict
|
||||
|
||||
cfg: Config
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(
|
||||
cls, pretrained_model_name_or_path: str, config_name: str, weight_name: str
|
||||
):
|
||||
if os.path.isdir(pretrained_model_name_or_path):
|
||||
config_path = os.path.join(pretrained_model_name_or_path, config_name)
|
||||
weight_path = os.path.join(pretrained_model_name_or_path, weight_name)
|
||||
else:
|
||||
config_path = hf_hub_download(
|
||||
repo_id=pretrained_model_name_or_path, filename=config_name
|
||||
)
|
||||
weight_path = hf_hub_download(
|
||||
repo_id=pretrained_model_name_or_path, filename=weight_name
|
||||
)
|
||||
|
||||
cfg = OmegaConf.load(config_path)
|
||||
OmegaConf.resolve(cfg)
|
||||
model = cls(cfg)
|
||||
ckpt = torch.load(weight_path, map_location="cpu")
|
||||
model.load_state_dict(ckpt)
|
||||
return model
|
||||
|
||||
@classmethod
|
||||
def from_pretrained_custom(
|
||||
cls, weight_path: str, config_path: str
|
||||
):
|
||||
cfg = OmegaConf.load(config_path)
|
||||
OmegaConf.resolve(cfg)
|
||||
model = cls(cfg)
|
||||
ckpt = torch.load(weight_path, map_location="cpu")
|
||||
model.load_state_dict(ckpt)
|
||||
return model
|
||||
|
||||
def configure(self):
|
||||
self.image_tokenizer = find_class(self.cfg.image_tokenizer_cls)(
|
||||
self.cfg.image_tokenizer
|
||||
)
|
||||
self.tokenizer = find_class(self.cfg.tokenizer_cls)(self.cfg.tokenizer)
|
||||
self.backbone = find_class(self.cfg.backbone_cls)(self.cfg.backbone)
|
||||
self.post_processor = find_class(self.cfg.post_processor_cls)(
|
||||
self.cfg.post_processor
|
||||
)
|
||||
self.decoder = find_class(self.cfg.decoder_cls)(self.cfg.decoder)
|
||||
self.renderer = find_class(self.cfg.renderer_cls)(self.cfg.renderer)
|
||||
self.image_processor = ImagePreprocessor()
|
||||
self.isosurface_helper = None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image: Union[
|
||||
PIL.Image.Image,
|
||||
np.ndarray,
|
||||
torch.FloatTensor,
|
||||
List[PIL.Image.Image],
|
||||
List[np.ndarray],
|
||||
List[torch.FloatTensor],
|
||||
],
|
||||
device: str,
|
||||
) -> torch.FloatTensor:
|
||||
rgb_cond = self.image_processor(image, self.cfg.cond_image_size)[:, None].to(
|
||||
device
|
||||
)
|
||||
batch_size = rgb_cond.shape[0]
|
||||
|
||||
input_image_tokens: torch.Tensor = self.image_tokenizer(
|
||||
rearrange(rgb_cond, "B Nv H W C -> B Nv C H W", Nv=1),
|
||||
)
|
||||
|
||||
input_image_tokens = rearrange(
|
||||
input_image_tokens, "B Nv C Nt -> B (Nv Nt) C", Nv=1
|
||||
)
|
||||
|
||||
tokens: torch.Tensor = self.tokenizer(batch_size)
|
||||
|
||||
tokens = self.backbone(
|
||||
tokens,
|
||||
encoder_hidden_states=input_image_tokens,
|
||||
)
|
||||
|
||||
scene_codes = self.post_processor(self.tokenizer.detokenize(tokens))
|
||||
return scene_codes
|
||||
|
||||
def render(
|
||||
self,
|
||||
scene_codes,
|
||||
n_views: int,
|
||||
elevation_deg: float = 0.0,
|
||||
camera_distance: float = 1.9,
|
||||
fovy_deg: float = 40.0,
|
||||
height: int = 256,
|
||||
width: int = 256,
|
||||
return_type: str = "pil",
|
||||
):
|
||||
rays_o, rays_d = get_spherical_cameras(
|
||||
n_views, elevation_deg, camera_distance, fovy_deg, height, width
|
||||
)
|
||||
rays_o, rays_d = rays_o.to(scene_codes.device), rays_d.to(scene_codes.device)
|
||||
|
||||
def process_output(image: torch.FloatTensor):
|
||||
if return_type == "pt":
|
||||
return image
|
||||
elif return_type == "np":
|
||||
return image.detach().cpu().numpy()
|
||||
elif return_type == "pil":
|
||||
return Image.fromarray(
|
||||
(image.detach().cpu().numpy() * 255.0).astype(np.uint8)
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
images = []
|
||||
for scene_code in scene_codes:
|
||||
images_ = []
|
||||
for i in range(n_views):
|
||||
with torch.no_grad():
|
||||
image = self.renderer(
|
||||
self.decoder, scene_code, rays_o[i], rays_d[i]
|
||||
)
|
||||
images_.append(process_output(image))
|
||||
images.append(images_)
|
||||
|
||||
return images
|
||||
|
||||
def set_marching_cubes_resolution(self, resolution: int):
|
||||
if (
|
||||
self.isosurface_helper is not None
|
||||
and self.isosurface_helper.resolution == resolution
|
||||
):
|
||||
return
|
||||
self.isosurface_helper = MarchingCubeHelper(resolution)
|
||||
|
||||
def extract_mesh(self, scene_codes, resolution: int = 256, threshold: float = 25.0,callback=None):
|
||||
self.set_marching_cubes_resolution(resolution)
|
||||
meshes = []
|
||||
for scene_code in scene_codes:
|
||||
with torch.no_grad():
|
||||
density = self.renderer.query_triplane(
|
||||
self.decoder,
|
||||
scale_tensor(
|
||||
self.isosurface_helper.grid_vertices.to(scene_codes.device),
|
||||
self.isosurface_helper.points_range,
|
||||
(-self.renderer.cfg.radius, self.renderer.cfg.radius),
|
||||
),
|
||||
scene_code,
|
||||
)["density_act"]
|
||||
v_pos, t_pos_idx = self.isosurface_helper(-(density - threshold))
|
||||
v_pos = scale_tensor(
|
||||
v_pos,
|
||||
self.isosurface_helper.points_range,
|
||||
(-self.renderer.cfg.radius, self.renderer.cfg.radius),
|
||||
)
|
||||
with torch.no_grad():
|
||||
color = self.renderer.query_triplane(
|
||||
self.decoder,
|
||||
v_pos,
|
||||
scene_code,
|
||||
)["color"]
|
||||
mesh = trimesh.Trimesh(
|
||||
vertices=v_pos.cpu().numpy(),
|
||||
faces=t_pos_idx.cpu().numpy(),
|
||||
vertex_colors=color.cpu().numpy(),
|
||||
)
|
||||
meshes.append(mesh)
|
||||
|
||||
if callback:
|
||||
callback(len(meshes))
|
||||
|
||||
return meshes
|
||||
@@ -0,0 +1,475 @@
|
||||
import importlib
|
||||
import math
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
#import rembg
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import trimesh
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
#from PIL import Image
|
||||
|
||||
|
||||
def parse_structured(fields: Any, cfg: Optional[Union[dict, DictConfig]] = None) -> Any:
|
||||
scfg = OmegaConf.merge(OmegaConf.structured(fields), cfg)
|
||||
return scfg
|
||||
|
||||
|
||||
def find_class(cls_string):
|
||||
module_string = ".".join(cls_string.split(".")[:-1])
|
||||
cls_name = cls_string.split(".")[-1]
|
||||
module = importlib.import_module(module_string, package=None)
|
||||
cls = getattr(module, cls_name)
|
||||
return cls
|
||||
|
||||
|
||||
def get_intrinsic_from_fov(fov, H, W, bs=-1):
|
||||
focal_length = 0.5 * H / np.tan(0.5 * fov)
|
||||
intrinsic = np.identity(3, dtype=np.float32)
|
||||
intrinsic[0, 0] = focal_length
|
||||
intrinsic[1, 1] = focal_length
|
||||
intrinsic[0, 2] = W / 2.0
|
||||
intrinsic[1, 2] = H / 2.0
|
||||
|
||||
if bs > 0:
|
||||
intrinsic = intrinsic[None].repeat(bs, axis=0)
|
||||
|
||||
return torch.from_numpy(intrinsic)
|
||||
|
||||
|
||||
class BaseModule(nn.Module):
|
||||
@dataclass
|
||||
class Config:
|
||||
pass
|
||||
|
||||
cfg: Config # add this to every subclass of BaseModule to enable static type checking
|
||||
|
||||
def __init__(
|
||||
self, cfg: Optional[Union[dict, DictConfig]] = None, *args, **kwargs
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.cfg = parse_structured(self.Config, cfg)
|
||||
self.configure(*args, **kwargs)
|
||||
|
||||
def configure(self, *args, **kwargs) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class ImagePreprocessor:
|
||||
def convert_and_resize(
|
||||
self,
|
||||
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
|
||||
size: int,
|
||||
):
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
image = torch.from_numpy(np.array(image).astype(np.float32) / 255.0)
|
||||
elif isinstance(image, np.ndarray):
|
||||
if image.dtype == np.uint8:
|
||||
image = torch.from_numpy(image.astype(np.float32) / 255.0)
|
||||
else:
|
||||
image = torch.from_numpy(image)
|
||||
elif isinstance(image, torch.Tensor):
|
||||
pass
|
||||
|
||||
batched = image.ndim == 4
|
||||
|
||||
if not batched:
|
||||
image = image[None, ...]
|
||||
image = F.interpolate(
|
||||
image.permute(0, 3, 1, 2),
|
||||
(size, size),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
antialias=True,
|
||||
).permute(0, 2, 3, 1)
|
||||
if not batched:
|
||||
image = image[0]
|
||||
return image
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
image: Union[
|
||||
PIL.Image.Image,
|
||||
np.ndarray,
|
||||
torch.FloatTensor,
|
||||
List[PIL.Image.Image],
|
||||
List[np.ndarray],
|
||||
List[torch.FloatTensor],
|
||||
],
|
||||
size: int,
|
||||
) -> Any:
|
||||
if isinstance(image, (np.ndarray, torch.FloatTensor)) and image.ndim == 4:
|
||||
image = self.convert_and_resize(image, size)
|
||||
else:
|
||||
if not isinstance(image, list):
|
||||
image = [image]
|
||||
image = [self.convert_and_resize(im, size) for im in image]
|
||||
image = torch.stack(image, dim=0)
|
||||
return image
|
||||
|
||||
|
||||
def rays_intersect_bbox(
|
||||
rays_o: torch.Tensor,
|
||||
rays_d: torch.Tensor,
|
||||
radius: float,
|
||||
near: float = 0.0,
|
||||
valid_thresh: float = 0.01,
|
||||
):
|
||||
input_shape = rays_o.shape[:-1]
|
||||
rays_o, rays_d = rays_o.view(-1, 3), rays_d.view(-1, 3)
|
||||
rays_d_valid = torch.where(
|
||||
rays_d.abs() < 1e-6, torch.full_like(rays_d, 1e-6), rays_d
|
||||
)
|
||||
if type(radius) in [int, float]:
|
||||
radius = torch.FloatTensor(
|
||||
[[-radius, radius], [-radius, radius], [-radius, radius]]
|
||||
).to(rays_o.device)
|
||||
radius = (
|
||||
1.0 - 1.0e-3
|
||||
) * radius # tighten the radius to make sure the intersection point lies in the bounding box
|
||||
interx0 = (radius[..., 1] - rays_o) / rays_d_valid
|
||||
interx1 = (radius[..., 0] - rays_o) / rays_d_valid
|
||||
t_near = torch.minimum(interx0, interx1).amax(dim=-1).clamp_min(near)
|
||||
t_far = torch.maximum(interx0, interx1).amin(dim=-1)
|
||||
|
||||
# check wheter a ray intersects the bbox or not
|
||||
rays_valid = t_far - t_near > valid_thresh
|
||||
|
||||
t_near[torch.where(~rays_valid)] = 0.0
|
||||
t_far[torch.where(~rays_valid)] = 0.0
|
||||
|
||||
t_near = t_near.view(*input_shape, 1)
|
||||
t_far = t_far.view(*input_shape, 1)
|
||||
rays_valid = rays_valid.view(*input_shape)
|
||||
|
||||
return t_near, t_far, rays_valid
|
||||
|
||||
|
||||
def chunk_batch(func: Callable, chunk_size: int, *args, **kwargs) -> Any:
|
||||
if chunk_size <= 0:
|
||||
return func(*args, **kwargs)
|
||||
B = None
|
||||
for arg in list(args) + list(kwargs.values()):
|
||||
if isinstance(arg, torch.Tensor):
|
||||
B = arg.shape[0]
|
||||
break
|
||||
assert (
|
||||
B is not None
|
||||
), "No tensor found in args or kwargs, cannot determine batch size."
|
||||
out = defaultdict(list)
|
||||
out_type = None
|
||||
# max(1, B) to support B == 0
|
||||
for i in range(0, max(1, B), chunk_size):
|
||||
out_chunk = func(
|
||||
*[
|
||||
arg[i : i + chunk_size] if isinstance(arg, torch.Tensor) else arg
|
||||
for arg in args
|
||||
],
|
||||
**{
|
||||
k: arg[i : i + chunk_size] if isinstance(arg, torch.Tensor) else arg
|
||||
for k, arg in kwargs.items()
|
||||
},
|
||||
)
|
||||
if out_chunk is None:
|
||||
continue
|
||||
out_type = type(out_chunk)
|
||||
if isinstance(out_chunk, torch.Tensor):
|
||||
out_chunk = {0: out_chunk}
|
||||
elif isinstance(out_chunk, tuple) or isinstance(out_chunk, list):
|
||||
chunk_length = len(out_chunk)
|
||||
out_chunk = {i: chunk for i, chunk in enumerate(out_chunk)}
|
||||
elif isinstance(out_chunk, dict):
|
||||
pass
|
||||
else:
|
||||
print(
|
||||
f"Return value of func must be in type [torch.Tensor, list, tuple, dict], get {type(out_chunk)}."
|
||||
)
|
||||
exit(1)
|
||||
for k, v in out_chunk.items():
|
||||
v = v if torch.is_grad_enabled() else v.detach()
|
||||
out[k].append(v)
|
||||
|
||||
if out_type is None:
|
||||
return None
|
||||
|
||||
out_merged: Dict[Any, Optional[torch.Tensor]] = {}
|
||||
for k, v in out.items():
|
||||
if all([vv is None for vv in v]):
|
||||
# allow None in return value
|
||||
out_merged[k] = None
|
||||
elif all([isinstance(vv, torch.Tensor) for vv in v]):
|
||||
out_merged[k] = torch.cat(v, dim=0)
|
||||
else:
|
||||
raise TypeError(
|
||||
f"Unsupported types in return value of func: {[type(vv) for vv in v if not isinstance(vv, torch.Tensor)]}"
|
||||
)
|
||||
|
||||
if out_type is torch.Tensor:
|
||||
return out_merged[0]
|
||||
elif out_type in [tuple, list]:
|
||||
return out_type([out_merged[i] for i in range(chunk_length)])
|
||||
elif out_type is dict:
|
||||
return out_merged
|
||||
|
||||
|
||||
ValidScale = Union[Tuple[float, float], torch.FloatTensor]
|
||||
|
||||
|
||||
def scale_tensor(dat: torch.FloatTensor, inp_scale: ValidScale, tgt_scale: ValidScale):
|
||||
if inp_scale is None:
|
||||
inp_scale = (0, 1)
|
||||
if tgt_scale is None:
|
||||
tgt_scale = (0, 1)
|
||||
if isinstance(tgt_scale, torch.FloatTensor):
|
||||
assert dat.shape[-1] == tgt_scale.shape[-1]
|
||||
dat = (dat - inp_scale[0]) / (inp_scale[1] - inp_scale[0])
|
||||
dat = dat * (tgt_scale[1] - tgt_scale[0]) + tgt_scale[0]
|
||||
return dat
|
||||
|
||||
|
||||
def get_activation(name) -> Callable:
|
||||
if name is None:
|
||||
return lambda x: x
|
||||
name = name.lower()
|
||||
if name == "none":
|
||||
return lambda x: x
|
||||
elif name == "exp":
|
||||
return lambda x: torch.exp(x)
|
||||
elif name == "sigmoid":
|
||||
return lambda x: torch.sigmoid(x)
|
||||
elif name == "tanh":
|
||||
return lambda x: torch.tanh(x)
|
||||
elif name == "softplus":
|
||||
return lambda x: F.softplus(x)
|
||||
else:
|
||||
try:
|
||||
return getattr(F, name)
|
||||
except AttributeError:
|
||||
raise ValueError(f"Unknown activation function: {name}")
|
||||
|
||||
|
||||
def get_ray_directions(
|
||||
H: int,
|
||||
W: int,
|
||||
focal: Union[float, Tuple[float, float]],
|
||||
principal: Optional[Tuple[float, float]] = None,
|
||||
use_pixel_centers: bool = True,
|
||||
normalize: bool = True,
|
||||
) -> torch.FloatTensor:
|
||||
"""
|
||||
Get ray directions for all pixels in camera coordinate.
|
||||
Reference: https://www.scratchapixel.com/lessons/3d-basic-rendering/
|
||||
ray-tracing-generating-camera-rays/standard-coordinate-systems
|
||||
|
||||
Inputs:
|
||||
H, W, focal, principal, use_pixel_centers: image height, width, focal length, principal point and whether use pixel centers
|
||||
Outputs:
|
||||
directions: (H, W, 3), the direction of the rays in camera coordinate
|
||||
"""
|
||||
pixel_center = 0.5 if use_pixel_centers else 0
|
||||
|
||||
if isinstance(focal, float):
|
||||
fx, fy = focal, focal
|
||||
cx, cy = W / 2, H / 2
|
||||
else:
|
||||
fx, fy = focal
|
||||
assert principal is not None
|
||||
cx, cy = principal
|
||||
|
||||
i, j = torch.meshgrid(
|
||||
torch.arange(W, dtype=torch.float32) + pixel_center,
|
||||
torch.arange(H, dtype=torch.float32) + pixel_center,
|
||||
indexing="xy",
|
||||
)
|
||||
|
||||
directions = torch.stack([(i - cx) / fx, -(j - cy) / fy, -torch.ones_like(i)], -1)
|
||||
|
||||
if normalize:
|
||||
directions = F.normalize(directions, dim=-1)
|
||||
|
||||
return directions
|
||||
|
||||
|
||||
def get_rays(
|
||||
directions,
|
||||
c2w,
|
||||
keepdim=False,
|
||||
normalize=False,
|
||||
) -> Tuple[torch.FloatTensor, torch.FloatTensor]:
|
||||
# Rotate ray directions from camera coordinate to the world coordinate
|
||||
assert directions.shape[-1] == 3
|
||||
|
||||
if directions.ndim == 2: # (N_rays, 3)
|
||||
if c2w.ndim == 2: # (4, 4)
|
||||
c2w = c2w[None, :, :]
|
||||
assert c2w.ndim == 3 # (N_rays, 4, 4) or (1, 4, 4)
|
||||
rays_d = (directions[:, None, :] * c2w[:, :3, :3]).sum(-1) # (N_rays, 3)
|
||||
rays_o = c2w[:, :3, 3].expand(rays_d.shape)
|
||||
elif directions.ndim == 3: # (H, W, 3)
|
||||
assert c2w.ndim in [2, 3]
|
||||
if c2w.ndim == 2: # (4, 4)
|
||||
rays_d = (directions[:, :, None, :] * c2w[None, None, :3, :3]).sum(
|
||||
-1
|
||||
) # (H, W, 3)
|
||||
rays_o = c2w[None, None, :3, 3].expand(rays_d.shape)
|
||||
elif c2w.ndim == 3: # (B, 4, 4)
|
||||
rays_d = (directions[None, :, :, None, :] * c2w[:, None, None, :3, :3]).sum(
|
||||
-1
|
||||
) # (B, H, W, 3)
|
||||
rays_o = c2w[:, None, None, :3, 3].expand(rays_d.shape)
|
||||
elif directions.ndim == 4: # (B, H, W, 3)
|
||||
assert c2w.ndim == 3 # (B, 4, 4)
|
||||
rays_d = (directions[:, :, :, None, :] * c2w[:, None, None, :3, :3]).sum(
|
||||
-1
|
||||
) # (B, H, W, 3)
|
||||
rays_o = c2w[:, None, None, :3, 3].expand(rays_d.shape)
|
||||
|
||||
if normalize:
|
||||
rays_d = F.normalize(rays_d, dim=-1)
|
||||
if not keepdim:
|
||||
rays_o, rays_d = rays_o.reshape(-1, 3), rays_d.reshape(-1, 3)
|
||||
|
||||
return rays_o, rays_d
|
||||
|
||||
|
||||
def get_spherical_cameras(
|
||||
n_views: int,
|
||||
elevation_deg: float,
|
||||
camera_distance: float,
|
||||
fovy_deg: float,
|
||||
height: int,
|
||||
width: int,
|
||||
):
|
||||
azimuth_deg = torch.linspace(0, 360.0, n_views + 1)[:n_views]
|
||||
elevation_deg = torch.full_like(azimuth_deg, elevation_deg)
|
||||
camera_distances = torch.full_like(elevation_deg, camera_distance)
|
||||
|
||||
elevation = elevation_deg * math.pi / 180
|
||||
azimuth = azimuth_deg * math.pi / 180
|
||||
|
||||
# convert spherical coordinates to cartesian coordinates
|
||||
# right hand coordinate system, x back, y right, z up
|
||||
# elevation in (-90, 90), azimuth from +x to +y in (-180, 180)
|
||||
camera_positions = torch.stack(
|
||||
[
|
||||
camera_distances * torch.cos(elevation) * torch.cos(azimuth),
|
||||
camera_distances * torch.cos(elevation) * torch.sin(azimuth),
|
||||
camera_distances * torch.sin(elevation),
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
|
||||
# default scene center at origin
|
||||
center = torch.zeros_like(camera_positions)
|
||||
# default camera up direction as +z
|
||||
up = torch.as_tensor([0, 0, 1], dtype=torch.float32)[None, :].repeat(n_views, 1)
|
||||
|
||||
fovy = torch.full_like(elevation_deg, fovy_deg) * math.pi / 180
|
||||
|
||||
lookat = F.normalize(center - camera_positions, dim=-1)
|
||||
right = F.normalize(torch.cross(lookat, up), dim=-1)
|
||||
up = F.normalize(torch.cross(right, lookat), dim=-1)
|
||||
c2w3x4 = torch.cat(
|
||||
[torch.stack([right, up, -lookat], dim=-1), camera_positions[:, :, None]],
|
||||
dim=-1,
|
||||
)
|
||||
c2w = torch.cat([c2w3x4, torch.zeros_like(c2w3x4[:, :1])], dim=1)
|
||||
c2w[:, 3, 3] = 1.0
|
||||
|
||||
# get directions by dividing directions_unit_focal by focal length
|
||||
focal_length = 0.5 * height / torch.tan(0.5 * fovy)
|
||||
directions_unit_focal = get_ray_directions(
|
||||
H=height,
|
||||
W=width,
|
||||
focal=1.0,
|
||||
)
|
||||
directions = directions_unit_focal[None, :, :, :].repeat(n_views, 1, 1, 1)
|
||||
directions[:, :, :, :2] = (
|
||||
directions[:, :, :, :2] / focal_length[:, None, None, None]
|
||||
)
|
||||
# must use normalize=True to normalize directions here
|
||||
rays_o, rays_d = get_rays(directions, c2w, keepdim=True, normalize=True)
|
||||
|
||||
return rays_o, rays_d
|
||||
|
||||
|
||||
# def remove_background(
|
||||
# image: PIL.Image.Image,
|
||||
# rembg_session: Any = None,
|
||||
# force: bool = False,
|
||||
# **rembg_kwargs,
|
||||
# ) -> PIL.Image.Image:
|
||||
# do_remove = True
|
||||
# if image.mode == "RGBA" and image.getextrema()[3][0] < 255:
|
||||
# do_remove = False
|
||||
# do_remove = do_remove or force
|
||||
# if do_remove:
|
||||
# image = rembg.remove(image, session=rembg_session, **rembg_kwargs)
|
||||
# return image
|
||||
|
||||
|
||||
def resize_foreground(
|
||||
image: PIL.Image.Image,
|
||||
ratio: float,
|
||||
) -> PIL.Image.Image:
|
||||
image = np.array(image)
|
||||
assert image.shape[-1] == 4
|
||||
alpha = np.where(image[..., 3] > 0)
|
||||
y1, y2, x1, x2 = (
|
||||
alpha[0].min(),
|
||||
alpha[0].max(),
|
||||
alpha[1].min(),
|
||||
alpha[1].max(),
|
||||
)
|
||||
# crop the foreground
|
||||
fg = image[y1:y2, x1:x2]
|
||||
# pad to square
|
||||
size = max(fg.shape[0], fg.shape[1])
|
||||
ph0, pw0 = (size - fg.shape[0]) // 2, (size - fg.shape[1]) // 2
|
||||
ph1, pw1 = size - fg.shape[0] - ph0, size - fg.shape[1] - pw0
|
||||
new_image = np.pad(
|
||||
fg,
|
||||
((ph0, ph1), (pw0, pw1), (0, 0)),
|
||||
mode="constant",
|
||||
constant_values=((0, 0), (0, 0), (0, 0)),
|
||||
)
|
||||
|
||||
# compute padding according to the ratio
|
||||
new_size = int(new_image.shape[0] / ratio)
|
||||
# pad to size, double side
|
||||
ph0, pw0 = (new_size - size) // 2, (new_size - size) // 2
|
||||
ph1, pw1 = new_size - size - ph0, new_size - size - pw0
|
||||
new_image = np.pad(
|
||||
new_image,
|
||||
((ph0, ph1), (pw0, pw1), (0, 0)),
|
||||
mode="constant",
|
||||
constant_values=((0, 0), (0, 0), (0, 0)),
|
||||
)
|
||||
new_image = PIL.Image.fromarray(new_image)
|
||||
return new_image
|
||||
|
||||
|
||||
def save_video(
|
||||
frames: List[PIL.Image.Image],
|
||||
output_path: str,
|
||||
fps: int = 30,
|
||||
):
|
||||
# use imageio to save video
|
||||
frames = [np.array(frame) for frame in frames]
|
||||
writer = imageio.get_writer(output_path, fps=fps)
|
||||
for frame in frames:
|
||||
writer.append_data(frame)
|
||||
writer.close()
|
||||
|
||||
|
||||
def to_gradio_3d_orientation(mesh):
|
||||
mesh.apply_transform(trimesh.transformations.rotation_matrix(-np.pi/2, [1, 0, 0]))
|
||||
mesh.apply_scale([1, 1, -1])
|
||||
mesh.apply_transform(trimesh.transformations.rotation_matrix(np.pi/2, [0, 1, 0]))
|
||||
return mesh
|
||||
@@ -0,0 +1,10 @@
|
||||
{
|
||||
"main_pass":
|
||||
[
|
||||
"-n", "-c:v", "libsvtav1",
|
||||
"-pix_fmt", "yuv420p10le",
|
||||
"-crf", "23"
|
||||
],
|
||||
"extension": "webm",
|
||||
"environment": {"SVT_LOG": "1"}
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
{
|
||||
"main_pass":
|
||||
[
|
||||
"-n", "-c:v", "libx264",
|
||||
"-pix_fmt", "yuv420p",
|
||||
"-crf", "19"
|
||||
],
|
||||
"extension": "mp4"
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"main_pass":
|
||||
[
|
||||
"-n", "-c:v", "libx265",
|
||||
"-pix_fmt", "yuv420p10le",
|
||||
"-preset", "medium",
|
||||
"-crf", "22",
|
||||
"-x265-params", "log-level=quiet"
|
||||
],
|
||||
"extension": "mp4"
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
{
|
||||
"main_pass":
|
||||
[
|
||||
"-n",
|
||||
"-pix_fmt", "yuv420p",
|
||||
"-crf", "23"
|
||||
],
|
||||
"extension": "webm"
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
[project]
|
||||
name = "comfyui-mixlab-nodes"
|
||||
description = "3D, ScreenShareNode & FloatingVideoNode, SpeechRecognition & SpeechSynthesis, GPT, LoadImagesFromLocal, Layers, Other Nodes, ..."
|
||||
version = "0.28.3"
|
||||
license = "MIT"
|
||||
dependencies = ["numpy", "pyOpenSSL", "watchdog", "opencv-python-headless", "matplotlib", "openai", "simple-lama-inpainting", "clip-interrogator==0.6.0", "transformers>=4.36.0", "lark-parser", "imageio-ffmpeg", "rembg[gpu]", "omegaconf==2.3.0", "Pillow>=9.5.0", "einops==0.7.0", "trimesh>=4.0.5", "huggingface-hub", "scikit-image"]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/shadowcz007/comfyui-mixlab-nodes"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "shadow"
|
||||
DisplayName = "comfyui-mixlab-nodes"
|
||||
Icon = ""
|
||||
@@ -4,4 +4,15 @@ watchdog
|
||||
opencv-python-headless
|
||||
matplotlib
|
||||
openai
|
||||
simple-lama-inpainting
|
||||
simple-lama-inpainting
|
||||
clip-interrogator==0.6.0
|
||||
transformers>=4.36.0
|
||||
lark-parser
|
||||
imageio-ffmpeg
|
||||
rembg[gpu]
|
||||
omegaconf==2.3.0
|
||||
Pillow>=9.5.0
|
||||
einops==0.7.0
|
||||
trimesh>=4.0.5
|
||||
huggingface-hub
|
||||
scikit-image
|
||||
@@ -196,6 +196,9 @@ app.registerExtension({
|
||||
}
|
||||
if (bg) {
|
||||
data.bg_image = await parseImage(bg)
|
||||
if (!data.bg_image.match('data:image/')) {
|
||||
delete data.bg_image
|
||||
}
|
||||
}
|
||||
|
||||
if (material) {
|
||||
|
||||
@@ -2,6 +2,46 @@ import { app } from '../../../scripts/app.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
import { api } from '../../../scripts/api.js'
|
||||
|
||||
//本机安装的插件节点全集
|
||||
window._nodesAll = null
|
||||
|
||||
//获取当前系统的插件,节点清单
|
||||
function getObjectInfo () {
|
||||
return new Promise(async (resolve, reject) => {
|
||||
let url = getUrl()
|
||||
|
||||
try {
|
||||
const response = await fetch(`${url}/object_info`)
|
||||
const data = await response.json()
|
||||
resolve(data)
|
||||
} catch (error) {
|
||||
reject(error)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
const base64Df =
|
||||
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
|
||||
|
||||
const parseImageToBase64 = url => {
|
||||
return new Promise((res, rej) => {
|
||||
fetch(url)
|
||||
.then(response => response.blob())
|
||||
.then(blob => {
|
||||
const reader = new FileReader()
|
||||
reader.onloadend = () => {
|
||||
const base64data = reader.result
|
||||
res(base64data)
|
||||
// 在这里可以将base64数据用于进一步处理或显示图片
|
||||
}
|
||||
reader.readAsDataURL(blob)
|
||||
})
|
||||
.catch(error => {
|
||||
console.log('发生错误:', error)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
function get_position_style (ctx, widget_width, y, node_height) {
|
||||
const MARGIN = 12 // the margin around the html element
|
||||
|
||||
@@ -28,20 +68,21 @@ function get_position_style (ctx, widget_width, y, node_height) {
|
||||
// height: `${node_height * 0.3 - MARGIN * 2}px`,
|
||||
// background: '#EEEEEE',
|
||||
display: 'flex',
|
||||
flexDirection: 'row',
|
||||
flexDirection: 'column',
|
||||
// alignItems: 'center',
|
||||
justifyContent: 'flex-start'
|
||||
justifyContent: 'flex-start',
|
||||
zIndex: 9999999
|
||||
}
|
||||
}
|
||||
|
||||
async function drawImageToCanvas (imageUrl) {
|
||||
async function drawImageToCanvas (imageUrl, sFactor = 320) {
|
||||
var canvas = document.createElement('canvas')
|
||||
var ctx = canvas.getContext('2d')
|
||||
var img = new Image()
|
||||
|
||||
await new Promise((resolve, reject) => {
|
||||
img.onload = function () {
|
||||
var scaleFactor = 320 / img.width
|
||||
var scaleFactor = sFactor / img.width
|
||||
var canvasWidth = img.width * scaleFactor
|
||||
var canvasHeight = img.height * scaleFactor
|
||||
|
||||
@@ -66,18 +107,27 @@ async function drawImageToCanvas (imageUrl) {
|
||||
// 可以在这里执行其他操作,比如将Base64数据保存到服务器或显示在页面上
|
||||
}
|
||||
|
||||
function extractInputAndOutputData (jsonData, inputIds = [], outputIds = []) {
|
||||
const data = jsonData
|
||||
const input = []
|
||||
const output = []
|
||||
async function extractInputAndOutputData (
|
||||
jsonData,
|
||||
inputIds = [],
|
||||
outputIds = []
|
||||
) {
|
||||
// workflow
|
||||
// const workflow=jsonData.workflow;
|
||||
// const nodes=workflow.nodes;
|
||||
|
||||
const data = jsonData.output
|
||||
let input = []
|
||||
let output = []
|
||||
const seed = {}
|
||||
const seedTitle = {}
|
||||
|
||||
for (const id in data) {
|
||||
if (data.hasOwnProperty(id)) {
|
||||
let node = app.graph.getNodeById(id)
|
||||
if (inputIds.includes(id)) {
|
||||
// let node = app.graph.getNodeById(id)
|
||||
let options = []
|
||||
let options = {}
|
||||
// 模型
|
||||
try {
|
||||
if (node.type === 'CheckpointLoaderSimple') {
|
||||
@@ -99,6 +149,54 @@ function extractInputAndOutputData (jsonData, inputIds = [], outputIds = []) {
|
||||
if (node.type == 'PromptSlide') {
|
||||
// min max step
|
||||
options = node.widgets.filter(w => w.type === 'slider')[0].options
|
||||
// 备选的keywords清单
|
||||
try {
|
||||
let keywords = node.widgets.filter(w => w.name === 'upload')[0]
|
||||
.value
|
||||
keywords = JSON.parse(keywords)
|
||||
options.keywords = keywords
|
||||
} catch (error) {
|
||||
console.log(error)
|
||||
}
|
||||
}
|
||||
|
||||
if (node.type == 'ImagesPrompt_') {
|
||||
//图库
|
||||
// console.log('ImagesPrompt_', data[id])
|
||||
let image_base64 = data[id].inputs.image_base64
|
||||
let img_index = 0
|
||||
let imgsData = JSON.parse(data[id].inputs.upload)
|
||||
for (let index = 0; index < imgsData.length; index++) {
|
||||
const imgd = imgsData[index].imgurl
|
||||
imgsData[index].index = index
|
||||
//TODO缩放大小
|
||||
imgsData[index].imgurl = await parseImageToBase64(imgd)
|
||||
if (image_base64 == imgsData[index].imgurl) {
|
||||
img_index = index
|
||||
}
|
||||
}
|
||||
options.images = imgsData
|
||||
delete data[id].inputs.upload
|
||||
delete data[id].inputs.image_base64
|
||||
|
||||
data[id].inputs.imageIndex = img_index
|
||||
}
|
||||
|
||||
if (node.type == 'Color') {
|
||||
}
|
||||
|
||||
if (node.type === 'LoadImage') {
|
||||
// loadImage的mask支持
|
||||
let output = node.outputs.filter(ot => ot.type == 'MASK')[0]
|
||||
if (output.links) {
|
||||
// 有输出
|
||||
options.hasMask = true
|
||||
}
|
||||
// loadImage的默认图,转为base64
|
||||
let imgurl = app.graph.getNodeById(id).imgs[0].src + '&channel=rgb'
|
||||
|
||||
options.defaultImage = await drawImageToCanvas(imgurl, 512)
|
||||
console.log('#loadImage的默认图', options)
|
||||
}
|
||||
|
||||
input[inputIds.indexOf(id)] = {
|
||||
@@ -110,23 +208,51 @@ function extractInputAndOutputData (jsonData, inputIds = [], outputIds = []) {
|
||||
// input.push()
|
||||
}
|
||||
if (outputIds.includes(id)) {
|
||||
let options = {}
|
||||
//输出的默认图
|
||||
if (
|
||||
node.type === 'SaveImageAndMetadata_' &&
|
||||
app.graph.getNodeById(id).imgs
|
||||
) {
|
||||
// SaveImageAndMetadata_的默认图,转为base64
|
||||
let imgurl = app.graph.getNodeById(id).imgs[0].src
|
||||
|
||||
options.defaultImage = await drawImageToCanvas(imgurl, 512)
|
||||
console.log('#SaveImageAndMetadata_的默认图', options)
|
||||
}
|
||||
|
||||
// let node = app.graph.getNodeById(id)
|
||||
// output.push()
|
||||
output[outputIds.indexOf(id)] = { ...data[id], title: node.title, id }
|
||||
output[outputIds.indexOf(id)] = {
|
||||
...data[id],
|
||||
title: node.title,
|
||||
id,
|
||||
options
|
||||
}
|
||||
}
|
||||
|
||||
if (node.type === 'KSampler') {
|
||||
if (
|
||||
node.type === 'KSampler' ||
|
||||
node.type == 'SamplerCustom' ||
|
||||
node.type === 'ChinesePrompt_Mix' ||
|
||||
node.type === 'Seed_'
|
||||
) {
|
||||
// seed 的类型收集
|
||||
try {
|
||||
seed[id] = node.widgets.filter(
|
||||
w => w.name === 'seed'
|
||||
w => w.name === 'seed' || w.name == 'noise_seed'
|
||||
)[0].linkedWidgets[0].value
|
||||
seedTitle[id] = node.title
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return { input, output, seed }
|
||||
// 修复bug,当节点不存在时
|
||||
input = input.filter(i => i)
|
||||
output = output.filter(i => i)
|
||||
|
||||
return { input, output, seed, seedTitle }
|
||||
}
|
||||
|
||||
function getUrl () {
|
||||
@@ -136,6 +262,16 @@ function getUrl () {
|
||||
return url
|
||||
}
|
||||
|
||||
const getLocalData = key => {
|
||||
let data = {}
|
||||
try {
|
||||
data = JSON.parse(localStorage.getItem(key)) || {}
|
||||
} catch (error) {
|
||||
return {}
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
async function save_app (json) {
|
||||
let url = getUrl()
|
||||
|
||||
@@ -144,7 +280,8 @@ async function save_app (json) {
|
||||
body: JSON.stringify({
|
||||
data: json,
|
||||
task: 'save_app',
|
||||
filename: json.app.filename
|
||||
filename: json.app.filename,
|
||||
category: json.app.category
|
||||
})
|
||||
})
|
||||
return await res.json()
|
||||
@@ -166,11 +303,16 @@ function downloadJsonFile (jsonData, fileName = 'mix_app.json') {
|
||||
}, 0)
|
||||
}
|
||||
|
||||
async function save (json, download = false) {
|
||||
async function save (json, download = false, showInfo = true) {
|
||||
let nodesAll = window._nodesAll || (await getObjectInfo())
|
||||
|
||||
console.log('####SAVE', nodesAll, json[0])
|
||||
|
||||
const name = json[0],
|
||||
version = json[5],
|
||||
share_prefix = json[6], //用于分享的功能扩展
|
||||
link=json[7],//用于创建界面上的跳转链接
|
||||
link = json[7], //用于创建界面上的跳转链接
|
||||
category = json[8] || '', //用于分类
|
||||
description = json[4],
|
||||
inputIds = json[2].split('\n').filter(f => f),
|
||||
outputIds = json[3].split('\n').filter(f => f)
|
||||
@@ -186,12 +328,26 @@ async function save (json, download = false) {
|
||||
try {
|
||||
let data = await app.graphToPrompt()
|
||||
|
||||
const { input, output, seed } = extractInputAndOutputData(
|
||||
data.output,
|
||||
//从output数据里把工作流的节点,插件数据统计出来
|
||||
data.nodesMap = {}
|
||||
for (const id in data.output) {
|
||||
data.nodesMap[data.output[id].class_type] =
|
||||
nodesAll[data.output[id].class_type]
|
||||
}
|
||||
|
||||
let { input, output, seed, seedTitle } = await extractInputAndOutputData(
|
||||
data,
|
||||
inputIds,
|
||||
outputIds
|
||||
)
|
||||
|
||||
let authorAvatar =
|
||||
localStorage.getItem('_mixlab_author_avatar') || base64Df,
|
||||
authorName =
|
||||
localStorage.getItem('_mixlab_author_name') ||
|
||||
localStorage.getItem('Comfy.userName'),
|
||||
authorLink = localStorage.getItem('_mixlab_author_link') || ''
|
||||
|
||||
data.app = {
|
||||
name,
|
||||
description,
|
||||
@@ -199,9 +355,16 @@ async function save (json, download = false) {
|
||||
input,
|
||||
output,
|
||||
seed, //控制是fixed 还是random
|
||||
seedTitle,
|
||||
share_prefix,
|
||||
link,
|
||||
filename: `${name}_${version}.json`
|
||||
category,
|
||||
filename: `${name}_${version}.json`,
|
||||
author: {
|
||||
avatar: authorAvatar,
|
||||
name: authorName,
|
||||
link: authorLink
|
||||
}
|
||||
}
|
||||
|
||||
try {
|
||||
@@ -209,42 +372,80 @@ async function save (json, download = false) {
|
||||
} catch (error) {}
|
||||
// console.log(data.app)
|
||||
// let http_workflow = app.graph.serialize()
|
||||
|
||||
await save_app(data)
|
||||
if (download) {
|
||||
await save_app(data)
|
||||
await downloadJsonFile(data, data.app.filename)
|
||||
}
|
||||
|
||||
if (showInfo) {
|
||||
let open = window.confirm(
|
||||
`You can now access the standalone application on a new page!\n${getUrl()}/mixlab/app?filename=${encodeURIComponent(
|
||||
data.app.filename
|
||||
)}`
|
||||
)}&category=${encodeURIComponent(data.app.category)}`
|
||||
)
|
||||
if (open)
|
||||
window.open(
|
||||
`${getUrl()}/mixlab/app?filename=${encodeURIComponent(
|
||||
data.app.filename
|
||||
)}`
|
||||
)}&category=${encodeURIComponent(data.app.category)}`
|
||||
)
|
||||
} else {
|
||||
await save_app(data)
|
||||
|
||||
let open = window.confirm(
|
||||
`You can now access the standalone application on a new page!\n${getUrl()}/mixlab/app`
|
||||
)
|
||||
if (open) window.open(`${getUrl()}/mixlab/app`)
|
||||
}
|
||||
} catch (error) {
|
||||
console.log('###SpeechRecognition', error)
|
||||
console.log('###error', error)
|
||||
}
|
||||
}
|
||||
|
||||
function getInputsAndOutputs () {
|
||||
const inputs =
|
||||
`LoadImage LoadImagesToBatch ImagesPrompt_ VHS_LoadVideo CLIPTextEncode PromptSlide TextInput_ Color FloatSlider IntNumber CheckpointLoaderSimple LoraLoader`.split(
|
||||
' '
|
||||
),
|
||||
outputs =
|
||||
`SaveTripoSRMesh,PreviewImage,SaveImage,TransparentImage,ShowTextForGPT,VHS_VideoCombine,VideoCombine_Adv,Image Save,SaveImageAndMetadata_,ClipInterrogator`.split(
|
||||
','
|
||||
)
|
||||
|
||||
let inputsId = [],
|
||||
outputsId = []
|
||||
|
||||
for (let node of app.graph._nodes) {
|
||||
if (inputs.includes(node.type)) {
|
||||
inputsId.push(node.id)
|
||||
}
|
||||
|
||||
if (outputs.includes(node.type)) {
|
||||
outputsId.push(node.id)
|
||||
}
|
||||
}
|
||||
|
||||
return {
|
||||
input: inputsId,
|
||||
output: outputsId
|
||||
}
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.utils.AppInfo',
|
||||
init () {
|
||||
if (!window._nodesAll) {
|
||||
getObjectInfo().then(r => (window._nodesAll = r))
|
||||
}
|
||||
},
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'AppInfo') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
console.log('#orig_nodeCreated', this)
|
||||
// console.log('#orig_nodeCreated', this)
|
||||
|
||||
// 自动计算workflow里哪些节点支持
|
||||
let input_ids = this.widgets.filter(w => w.name == 'input_ids')[0],
|
||||
output_ids = this.widgets.filter(w => w.name == 'output_ids')[0]
|
||||
|
||||
const { input, output } = getInputsAndOutputs()
|
||||
input_ids.value = input.join('\n')
|
||||
output_ids.value = output.join('\n')
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'AppInfoRun',
|
||||
@@ -302,10 +503,164 @@ app.registerExtension({
|
||||
}
|
||||
})
|
||||
|
||||
document.body.appendChild(widget.div)
|
||||
widget.div.appendChild(btn)
|
||||
widget.div.appendChild(download)
|
||||
// author
|
||||
let author = document.createElement('div')
|
||||
// author.style=`display: flex`
|
||||
|
||||
let authorAvatar = document.createElement('img')
|
||||
authorAvatar.className = `${'comfy-multiline-input'}`
|
||||
authorAvatar.style = `outline: none;
|
||||
border: none;
|
||||
padding: 4px;
|
||||
width: 32px;
|
||||
cursor: pointer;
|
||||
height: 32px;`
|
||||
|
||||
if (localStorage.getItem('_mixlab_author_avatar')) {
|
||||
authorAvatar.src =
|
||||
localStorage.getItem('_mixlab_author_avatar') || base64Df
|
||||
}
|
||||
|
||||
let authorAvatarUpload = document.createElement('input')
|
||||
authorAvatarUpload.type = 'file'
|
||||
authorAvatarUpload.style = `display:none`
|
||||
|
||||
let authorAvatarInput = document.createElement('div')
|
||||
authorAvatarInput.style = `display: flex;justify-content: flex-start;
|
||||
align-items: center;`
|
||||
let authorAvatarInputLabel = document.createElement('p')
|
||||
authorAvatarInputLabel.innerText = 'Author Avatar'
|
||||
authorAvatarInputLabel.className = `${'comfy-multiline-input'}`
|
||||
authorAvatarInputLabel.style = `font-size:12px`
|
||||
|
||||
authorAvatar.addEventListener('click', e => {
|
||||
authorAvatarUpload.click()
|
||||
})
|
||||
|
||||
authorAvatarInputLabel.addEventListener('click', e => {
|
||||
authorAvatarUpload.click()
|
||||
})
|
||||
|
||||
authorAvatarUpload.addEventListener('change', event => {
|
||||
const file = event.target.files[0]
|
||||
const reader = new FileReader()
|
||||
|
||||
reader.onload = async e => {
|
||||
let im = new Image()
|
||||
im.src = e.target.result
|
||||
authorAvatar.src = e.target.result
|
||||
im.onload = () => {
|
||||
let c = document.createElement('canvas')
|
||||
let ctx = c.getContext('2d')
|
||||
c.width = 72
|
||||
c.height = 72
|
||||
ctx.drawImage(
|
||||
im,
|
||||
0,
|
||||
0,
|
||||
im.naturalWidth,
|
||||
im.naturalHeight,
|
||||
0,
|
||||
0,
|
||||
c.width,
|
||||
c.height
|
||||
)
|
||||
window._mixlab_author_avatar = c.toDataURL()
|
||||
localStorage.setItem(
|
||||
'_mixlab_author_avatar',
|
||||
window._mixlab_author_avatar
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// 以文本形式读取文件
|
||||
reader.readAsDataURL(file)
|
||||
})
|
||||
|
||||
author.appendChild(authorAvatarInput)
|
||||
authorAvatarInput.appendChild(authorAvatarInputLabel)
|
||||
authorAvatarInput.appendChild(authorAvatar)
|
||||
authorAvatarInput.appendChild(authorAvatarUpload)
|
||||
|
||||
let authorName = document.createElement('input')
|
||||
authorName.type = 'text'
|
||||
authorName.value =
|
||||
localStorage.getItem('_mixlab_author_name') ||
|
||||
localStorage.getItem('Comfy.userName')
|
||||
authorName.placeholder = 'author name'
|
||||
authorName.className = `${'comfy-multiline-input'}`
|
||||
authorName.style = `
|
||||
outline: none;
|
||||
border: none;
|
||||
padding: 4px;
|
||||
width: 100%;
|
||||
cursor: pointer;
|
||||
height: 32px;`
|
||||
|
||||
let authorNameInput = document.createElement('div')
|
||||
authorNameInput.style = `display: flex;justify-content: flex-start;
|
||||
align-items: center;`
|
||||
let authorNameInputLabel = document.createElement('p')
|
||||
authorNameInputLabel.innerText = 'Author Name'
|
||||
authorNameInputLabel.className = `${'comfy-multiline-input'}`
|
||||
authorNameInputLabel.style = `font-size:12px;width: 110px`
|
||||
|
||||
authorName.addEventListener('change', e => {
|
||||
window._mixlab_author_name = authorName.value.trim()
|
||||
localStorage.setItem(
|
||||
'_mixlab_author_name',
|
||||
window._mixlab_author_name
|
||||
)
|
||||
})
|
||||
|
||||
author.appendChild(authorNameInput)
|
||||
authorNameInput.appendChild(authorNameInputLabel)
|
||||
authorNameInput.appendChild(authorName)
|
||||
|
||||
// 社交链接
|
||||
let authorLink = document.createElement('input')
|
||||
authorLink.type = 'text'
|
||||
authorLink.value = localStorage.getItem('_mixlab_author_link') || ''
|
||||
authorLink.placeholder = 'author link'
|
||||
authorLink.className = `${'comfy-multiline-input'}`
|
||||
authorLink.style = `
|
||||
outline: none;
|
||||
border: none;
|
||||
padding: 4px;
|
||||
width: 100%;
|
||||
cursor: pointer;
|
||||
height: 32px;`
|
||||
|
||||
let authorLinkInput = document.createElement('div')
|
||||
authorLinkInput.style = `display: flex;justify-content: flex-start;
|
||||
align-items: center;`
|
||||
let authorLinkInputLabel = document.createElement('p')
|
||||
authorLinkInputLabel.innerText = 'Author Link'
|
||||
authorLinkInputLabel.className = `${'comfy-multiline-input'}`
|
||||
authorLinkInputLabel.style = `font-size:12px;width: 110px`
|
||||
|
||||
authorLink.addEventListener('change', e => {
|
||||
window._mixlab_author_link = authorLink.value.trim()
|
||||
localStorage.setItem(
|
||||
'_mixlab_author_link',
|
||||
window._mixlab_author_link
|
||||
)
|
||||
})
|
||||
|
||||
author.appendChild(authorLinkInput)
|
||||
authorLinkInput.appendChild(authorLinkInputLabel)
|
||||
authorLinkInput.appendChild(authorLink)
|
||||
|
||||
widget.div.appendChild(author)
|
||||
|
||||
let btns = document.createElement('div')
|
||||
|
||||
widget.div.appendChild(btns)
|
||||
|
||||
btns.appendChild(btn)
|
||||
btns.appendChild(download)
|
||||
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
@@ -315,14 +670,22 @@ app.registerExtension({
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
|
||||
window._mixlab_app_json = null
|
||||
}
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = async function (message) {
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
onExecuted?.apply(this, arguments)
|
||||
console.log(message.json)
|
||||
window._mixlab_app_json = message.json
|
||||
try {
|
||||
let a = this.widgets.filter(w => w.name === 'AppInfoRun')[0]
|
||||
if (a) {
|
||||
if (!a.value) a.value = 0
|
||||
a.value += 1
|
||||
}
|
||||
|
||||
const div = this.widgets.filter(w => w.div)[0].div
|
||||
Array.from(
|
||||
div.querySelectorAll('button'),
|
||||
@@ -331,5 +694,42 @@ app.registerExtension({
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
// console.log('#loadedGraphNode1111')
|
||||
window._mixlab_app_json = null //切换workflow需要清空
|
||||
if (node.type === 'AppInfo') {
|
||||
let auto_save = node.widgets.filter(w => w.name == 'auto_save')[0]
|
||||
if (auto_save) {
|
||||
if (!['enable', 'disable'].includes(auto_save.value)) {
|
||||
auto_save.value = 'enable'
|
||||
}
|
||||
}
|
||||
|
||||
// app.canvas.centerOnNode(node)
|
||||
// app.canvas.setZoom(0.45)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
api.addEventListener('execution_start', async ({ detail }) => {
|
||||
console.log('#execution_start', detail)
|
||||
window._mixlab_app_json = null
|
||||
})
|
||||
|
||||
api.addEventListener('executed', async ({ detail }) => {
|
||||
console.log('#executed', detail)
|
||||
// window._mixlab_app_json=null;
|
||||
const { output } = getInputsAndOutputs()
|
||||
if (output.includes(parseInt(detail.node))) {
|
||||
let appinfo = app.graph.findNodesByType('AppInfo')[0]
|
||||
if (appinfo) {
|
||||
let auto_save = appinfo.widgets.filter(w => w.name == 'auto_save')[0]
|
||||
if (auto_save?.value === 'enable') {
|
||||
// 自动保存
|
||||
console.log('auto_save')
|
||||
if (window._mixlab_app_json) save(window._mixlab_app_json, false, false)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
async function* completion (url, messages, controller) {
|
||||
let data = {
|
||||
model: 'gpt-3.5-turbo-16k',
|
||||
messages,
|
||||
temperature: 0.05,
|
||||
stream: true
|
||||
}
|
||||
// if (imageNode) {
|
||||
// data = { ...data, image_data: [imageNode] }
|
||||
// }
|
||||
|
||||
// let controller = new AbortController()
|
||||
|
||||
let response = await fetch(url, {
|
||||
method: 'POST',
|
||||
body: JSON.stringify(data),
|
||||
headers: {
|
||||
Connection: 'keep-alive',
|
||||
'Content-Type': 'application/json',
|
||||
Accept: 'text/event-stream'
|
||||
},
|
||||
signal: controller.signal
|
||||
})
|
||||
|
||||
const reader = response.body.getReader()
|
||||
const decoder = new TextDecoder()
|
||||
|
||||
let content = ''
|
||||
let leftover = '' // Buffer for partially read lines
|
||||
|
||||
try {
|
||||
let cont = true
|
||||
while (cont) {
|
||||
let result = await reader.read()
|
||||
if (result.done) {
|
||||
break
|
||||
}
|
||||
|
||||
// Add any leftover data to the current chunk of data
|
||||
const text = leftover + decoder.decode(result.value)
|
||||
|
||||
// Check if the last character is a line break
|
||||
const endsWithLineBreak = text.endsWith('\n')
|
||||
|
||||
// Split the text into lines
|
||||
let lines = text.split('\n')
|
||||
|
||||
// If the text doesn't end with a line break, then the last line is incomplete
|
||||
// Store it in leftover to be added to the next chunk of data
|
||||
if (!endsWithLineBreak) {
|
||||
leftover = lines.pop()
|
||||
} else {
|
||||
leftover = '' // Reset leftover if we have a line break at the end
|
||||
}
|
||||
|
||||
// Parse all sse events and add them to result
|
||||
const regex = /^(\S+):\s(.*)$/gm
|
||||
for (const line of lines) {
|
||||
const match = regex.exec(line)
|
||||
if (match) {
|
||||
result[match[1]] = match[2]
|
||||
// since we know this is llama.cpp, let's just decode the json in data
|
||||
if (result.data) {
|
||||
result.data = JSON.parse(result.data)
|
||||
// console.log('#result.data',result.data)
|
||||
|
||||
content += result.data.choices[0].delta?.content || ''
|
||||
|
||||
// yield
|
||||
yield result
|
||||
|
||||
// if we got a stop token from server, we will break here
|
||||
if (result.data.choices[0].finish_reason == 'stop') {
|
||||
if (result.data.generation_settings) {
|
||||
// generation_settings = result.data.generation_settings;
|
||||
}
|
||||
cont = false
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} catch (e) {
|
||||
console.error('llama error: ', e)
|
||||
throw e
|
||||
} finally {
|
||||
controller.abort()
|
||||
}
|
||||
|
||||
return content
|
||||
// return (await response.json()).content
|
||||
}
|
||||
|
||||
export async function completion_ (url, messages, controller, callback) {
|
||||
let request = await completion(url, messages, controller)
|
||||
for await (const chunk of request) {
|
||||
let content = chunk.data.choices[0].delta.content || ''
|
||||
if (chunk.data.choices[0].role == 'assistant') {
|
||||
//开始
|
||||
content = ''
|
||||
}
|
||||
|
||||
if (callback) callback(content)
|
||||
}
|
||||
}
|
||||
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
|
||||
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
|
||||
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
|
||||
|
||||
const version = 'v0.8.0'
|
||||
const version = 'v0.28.3'
|
||||
|
||||
fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
|
||||
.then(response => response.json())
|
||||
@@ -17,7 +17,13 @@ fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
|
||||
return
|
||||
if (latestVersion && latestVersion != version) {
|
||||
localStorage.setItem('_mixlab_nodes_vesion', latestVersion)
|
||||
app.ui.dialog.show(`<h4 style="font-size: 18px;">${repoName} <br>
|
||||
app.ui.dialog.show(`<a style="color: white;
|
||||
font-size: 18px;
|
||||
font-weight: 800;
|
||||
letter-spacing: 2px;
|
||||
}"
|
||||
href="https://discord.gg/cXs9vZSqeK">Welcome to Mixlab nodes discord</a>
|
||||
<h4 style="font-size: 18px;">${repoName} <br>
|
||||
Latest release version: ${latestVersion}</h4>
|
||||
<p>Please proceed to the official repository to download the latest version.</p>
|
||||
<a style="color: #2196F3;
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
// import { api } from '../../../scripts/api.js'
|
||||
import { ComfyWidgets } from '../../../scripts/widgets.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
function getRandomElements (arr, num) {
|
||||
var result = []
|
||||
var len = arr.length
|
||||
|
||||
for (var i = 0; i < num; i++) {
|
||||
var randomIndex = Math.floor(Math.random() * len)
|
||||
result.push(arr[randomIndex])
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
const createPrompt = (node, prompts, items, sample) => {
|
||||
const w = ComfyWidgets['STRING'](
|
||||
node,
|
||||
'text',
|
||||
['STRING', { multiline: true }],
|
||||
app
|
||||
).widget
|
||||
w.inputEl.readOnly = true
|
||||
w.inputEl.style.opacity = 0.6
|
||||
|
||||
w.value = typeof prompts === 'string' ? prompts : prompts.join('\n\n')
|
||||
|
||||
const w2 = ComfyWidgets['STRING'](
|
||||
node,
|
||||
'text',
|
||||
['STRING', { multiline: true }],
|
||||
app
|
||||
).widget
|
||||
w2.inputEl.readOnly = true
|
||||
w2.inputEl.style.opacity = 0.6
|
||||
|
||||
w2.value = typeof items === 'string' ? items : JSON.stringify(items, null, 2)
|
||||
|
||||
const w3 = ComfyWidgets['STRING'](
|
||||
node,
|
||||
'text',
|
||||
['STRING', { multiline: true }],
|
||||
app
|
||||
).widget
|
||||
w3.inputEl.readOnly = true
|
||||
w3.inputEl.style.opacity = 0.6
|
||||
w3.value = typeof sample === 'string' ? sample : sample.join('\n\n')
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.prompt.ClipInterrogator',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeData.name === 'ClipInterrogator') {
|
||||
function populate (prompts, items, random_samples) {
|
||||
if (this.widgets) {
|
||||
for (let i = 0; i < this.widgets.length; i++) {
|
||||
if (this.widgets[i].type !== 'combo') this.widgets[i].onRemove?.()
|
||||
}
|
||||
this.widgets.length = 2
|
||||
}
|
||||
|
||||
createPrompt(this, prompts, items, random_samples)
|
||||
|
||||
// console.log('ClipInterrogator', w, w2)
|
||||
requestAnimationFrame(() => {
|
||||
const sz = this.computeSize()
|
||||
if (sz[0] < this.size[0]) {
|
||||
sz[0] = this.size[0]
|
||||
}
|
||||
if (sz[1] < this.size[1]) {
|
||||
sz[1] = this.size[1]
|
||||
}
|
||||
this.onResize?.(sz)
|
||||
app.graph.setDirtyCanvas(true, false)
|
||||
})
|
||||
}
|
||||
|
||||
// When the node is executed we will be sent the input text, display this in the widget
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
onExecuted?.apply(this, arguments)
|
||||
// console.log('##', message)
|
||||
populate.call(
|
||||
this,
|
||||
message.prompt,
|
||||
message.analysis,
|
||||
message.random_samples
|
||||
)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
// Fires every time a node is constructed
|
||||
// You can modify widgets/add handlers/etc here
|
||||
|
||||
if (node.type === 'ClipInterrogator') {
|
||||
try {
|
||||
|
||||
let widgets_values = node.widgets_values
|
||||
console.log(widgets_values )
|
||||
try {
|
||||
if (widgets_values[2] && widgets_values[3] && widgets_values[4])
|
||||
createPrompt(
|
||||
node,
|
||||
widgets_values[2],
|
||||
widgets_values[3],
|
||||
widgets_values[4]
|
||||
)
|
||||
} catch (error) {
|
||||
console.log(error)
|
||||
}
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
})
|
||||
@@ -61,14 +61,14 @@ app.registerExtension({
|
||||
async getCustomWidgets (app) {
|
||||
return {
|
||||
KEY (node, inputName, inputData, app) {
|
||||
console.log('##inputData', inputData)
|
||||
// console.log('##inputData', inputData)
|
||||
const widget = {
|
||||
type: inputData[0], // the type, CHEESE
|
||||
name: inputName, // the name, slice
|
||||
size: [128, 32], // a default size
|
||||
draw (ctx, node, width, y) {},
|
||||
computeSize (...args) {
|
||||
return [128,32] // a method to compute the current size of the widget
|
||||
return [128, 32] // a method to compute the current size of the widget
|
||||
},
|
||||
async serializeValue (nodeId, widgetIndex) {
|
||||
let data = getLocalData('_mixlab_api_key')
|
||||
@@ -203,75 +203,82 @@ app.registerExtension({
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.GPT.ShowTextForGPT',
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (nodeData.name === "ShowTextForGPT") {
|
||||
function populate(text) {
|
||||
if (this.widgets) {
|
||||
|
||||
const pos = this.widgets.findIndex((w) => w.name === "text");
|
||||
if (pos !== -1) {
|
||||
for (let i = pos; i < this.widgets.length; i++) {
|
||||
this.widgets[i].onRemove?.();
|
||||
}
|
||||
this.widgets.length = pos;
|
||||
}
|
||||
}
|
||||
// console.log('ShowTextForGPT',text)
|
||||
for (let list of text) {
|
||||
const w = ComfyWidgets["STRING"](this, "text", ["STRING", { multiline: true }], app).widget;
|
||||
w.inputEl.readOnly = true;
|
||||
w.inputEl.style.opacity = 0.6;
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeData.name === 'ShowTextForGPT') {
|
||||
function populate (text) {
|
||||
text = text.filter(t => t && t?.trim())
|
||||
|
||||
try {
|
||||
let data=JSON.parse(list);
|
||||
data=Array.from(data,d=>{
|
||||
return {
|
||||
...d,
|
||||
content:decodeURIComponent(d.content)
|
||||
}
|
||||
})
|
||||
list=JSON.stringify(data,null,2)
|
||||
} catch (error) {
|
||||
// console.log(error)
|
||||
if (this.widgets) {
|
||||
// const pos = this.widgets.findIndex(w => w.name === 'text')
|
||||
for (let i = 0; i < this.widgets.length; i++) {
|
||||
if (this.widgets[i].name == 'show_text') this.widgets[i].onRemove?.()
|
||||
}
|
||||
this.widgets.length = 1
|
||||
}
|
||||
// console.log('ShowTextForGPT',text)
|
||||
for (let list of text) {
|
||||
if (list) {
|
||||
// console.log('#####', list)
|
||||
const w = ComfyWidgets['STRING'](
|
||||
this,
|
||||
'show_text',
|
||||
['STRING', { multiline: true }],
|
||||
app
|
||||
).widget
|
||||
w.inputEl.readOnly = true
|
||||
w.inputEl.style.opacity = 0.6
|
||||
|
||||
w.value =list;
|
||||
|
||||
}
|
||||
try {
|
||||
if (typeof list != 'string') {
|
||||
let data = JSON.parse(list)
|
||||
data = Array.from(data, d => {
|
||||
return {
|
||||
...d,
|
||||
content: decodeURIComponent(d.content)
|
||||
}
|
||||
})
|
||||
list = JSON.stringify(data, null, 2)
|
||||
}
|
||||
} catch (error) {
|
||||
console.log(error)
|
||||
}
|
||||
|
||||
w.value = list
|
||||
}
|
||||
}
|
||||
// console.log('ShowTextForGPT',this.widgets.length)
|
||||
requestAnimationFrame(() => {
|
||||
const sz = this.computeSize();
|
||||
if (sz[0] < this.size[0]) {
|
||||
sz[0] = this.size[0];
|
||||
}
|
||||
if (sz[1] < this.size[1]) {
|
||||
sz[1] = this.size[1];
|
||||
}
|
||||
this.onResize?.(sz);
|
||||
app.graph.setDirtyCanvas(true, false);
|
||||
});
|
||||
}
|
||||
requestAnimationFrame(() => {
|
||||
if (this) {
|
||||
const sz = this.computeSize()
|
||||
if (sz[0] < this.size[0]) {
|
||||
sz[0] = this.size[0]
|
||||
}
|
||||
if (sz[1] < this.size[1]) {
|
||||
sz[1] = this.size[1]
|
||||
}
|
||||
this.onResize?.(sz)
|
||||
app.graph.setDirtyCanvas(true, false)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// When the node is executed we will be sent the input text, display this in the widget
|
||||
const onExecuted = nodeType.prototype.onExecuted;
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
onExecuted?.apply(this, arguments);
|
||||
populate.call(this, message.text);
|
||||
};
|
||||
// When the node is executed we will be sent the input text, display this in the widget
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
onExecuted?.apply(this, arguments)
|
||||
// console.log('##onExecuted', this, message)
|
||||
if (message.text) populate.call(this, message.text)
|
||||
}
|
||||
|
||||
const onConfigure = nodeType.prototype.onConfigure;
|
||||
nodeType.prototype.onConfigure = function () {
|
||||
onConfigure?.apply(this, arguments);
|
||||
if (this.widgets_values?.length) {
|
||||
|
||||
populate.call(this, this.widgets_values);
|
||||
}
|
||||
};
|
||||
const onConfigure = nodeType.prototype.onConfigure
|
||||
nodeType.prototype.onConfigure = function () {
|
||||
onConfigure?.apply(this, arguments)
|
||||
if (this.widgets_values?.length) {
|
||||
populate.call(this, this.widgets_values)
|
||||
}
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
|
||||
}
|
||||
|
||||
|
||||
},
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
@@ -1,7 +1,39 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
import { api } from '../../../scripts/api.js'
|
||||
import { ComfyWidgets } from '../../../scripts/widgets.js'
|
||||
// import { ComfyWidgets } from '../../../scripts/widgets.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
import { applyTextReplacements } from '../../../scripts/utils.js'
|
||||
|
||||
function loadImageToCanvas (base64Image) {
|
||||
var img = new Image()
|
||||
var canvas = document.createElement('canvas')
|
||||
var ctx = canvas.getContext('2d')
|
||||
return new Promise((res, rej) => {
|
||||
img.onload = function () {
|
||||
// 等比例缩放图片
|
||||
var width = img.width
|
||||
var height = img.height
|
||||
var max_width = 1024
|
||||
if (width > max_width) {
|
||||
height *= max_width / width
|
||||
width = max_width
|
||||
}
|
||||
|
||||
// 设置canvas尺寸
|
||||
canvas.width = width
|
||||
canvas.height = height
|
||||
|
||||
// 在canvas上绘制图片
|
||||
ctx.drawImage(img, 0, 0, width, height)
|
||||
|
||||
// 将canvas转换为base64图片数据
|
||||
var canvasData = canvas.toDataURL()
|
||||
res(canvasData) // canvas转换后的base64图片数据
|
||||
}
|
||||
|
||||
img.src = base64Image
|
||||
})
|
||||
}
|
||||
|
||||
async function uploadImage (blob, fileType = '.svg', filename) {
|
||||
// const blob = await (await fetch(src)).blob();
|
||||
@@ -28,6 +60,9 @@ async function uploadImage (blob, fileType = '.svg', filename) {
|
||||
return src
|
||||
}
|
||||
|
||||
const base64Df =
|
||||
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
|
||||
|
||||
function base64ToBlobFromURL (base64URL, contentType) {
|
||||
return fetch(base64URL).then(response => response.blob())
|
||||
}
|
||||
@@ -108,7 +143,7 @@ function createImage (url) {
|
||||
})
|
||||
}
|
||||
|
||||
const parseImage = url => {
|
||||
const parseImageToBase64 = url => {
|
||||
return new Promise((res, rej) => {
|
||||
fetch(url)
|
||||
.then(response => response.blob())
|
||||
@@ -406,9 +441,7 @@ app.registerExtension({
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
// Fires every time a node is constructed
|
||||
@@ -442,3 +475,518 @@ app.registerExtension({
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const createSelect = (imgDiv, select, opts, targetWidget, textWidget) => {
|
||||
select.style.display = 'block'
|
||||
let html = ''
|
||||
let isMatch = false
|
||||
for (const opt of opts) {
|
||||
html += `<option value='${opt.keyword}' ${opt.selected ? 'selected' : ''}>${
|
||||
opt.keyword
|
||||
}</option>`
|
||||
if (opt.selected) {
|
||||
isMatch = true
|
||||
imgDiv.src = opt.imgurl
|
||||
// targetWidget.value = opt.keyword
|
||||
}
|
||||
}
|
||||
select.innerHTML = html
|
||||
if (!isMatch) {
|
||||
// targetWidget.value = opts[0].keyword
|
||||
imgDiv.src = opts[0].imgurl
|
||||
}
|
||||
|
||||
// 添加change事件监听器
|
||||
select.addEventListener('change', async function () {
|
||||
// 获取选中的选项的值
|
||||
var selectedOption = select.options[select.selectedIndex].value
|
||||
let t = opts.filter(opt => opt.keyword === selectedOption)[0]
|
||||
|
||||
targetWidget.value = await parseImageToBase64(t.imgurl)
|
||||
imgDiv.src = targetWidget.value
|
||||
textWidget.value = t.keyword
|
||||
})
|
||||
// console.log(select)
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.prompt.ImagesPrompt_',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'ImagesPrompt_') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const image_prompt = this.widgets.filter(
|
||||
w => w.name == 'image_base64'
|
||||
)[0]
|
||||
const image_text = this.widgets.filter(w => w.name == 'text')[0]
|
||||
|
||||
const node = this
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'upload',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1])
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
// console.log('image_prompt',image_prompt)
|
||||
const img = new Image()
|
||||
img.src = image_prompt?.value || base64Df
|
||||
widget.div.appendChild(img)
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Upload Images JSON'
|
||||
|
||||
btn.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid;height: 30px;min-width: 122px;
|
||||
`
|
||||
|
||||
const select = document.createElement('select')
|
||||
select.style = `display:none;cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid;height: 30px;min-width: 100px;
|
||||
`
|
||||
widget.select = select
|
||||
|
||||
// const btn=document.createElement('button');
|
||||
// btn.innerText='Upload'
|
||||
btn.addEventListener('click', () => {
|
||||
let inp = document.createElement('input')
|
||||
inp.type = 'file'
|
||||
inp.accept = '.json'
|
||||
inp.click()
|
||||
inp.addEventListener('change', event => {
|
||||
// 获取选择的文件
|
||||
// [{title,imageUrl}]
|
||||
const file = event.target.files[0]
|
||||
this.title = file.name.split('.')[0]
|
||||
|
||||
// console.log(file.name.split('.')[0])
|
||||
// 创建文件读取器
|
||||
const reader = new FileReader()
|
||||
|
||||
// 定义读取完成事件的回调函数
|
||||
reader.onload = async event => {
|
||||
// 读取完成后的文本内容
|
||||
const json = JSON.parse(event.target.result)
|
||||
console.log(node, json)
|
||||
|
||||
widget.value = JSON.stringify(json)
|
||||
|
||||
let img = widget.div.querySelector('img')
|
||||
|
||||
createSelect(img, select, json, image_prompt, image_text)
|
||||
|
||||
image_prompt.value = await parseImageToBase64(json[0].imgurl)
|
||||
image_text.value = json[0].keyword
|
||||
|
||||
if (img) {
|
||||
img.src = image_prompt.value
|
||||
}
|
||||
|
||||
inp.remove()
|
||||
}
|
||||
|
||||
// 以文本方式读取文件
|
||||
reader.readAsText(file)
|
||||
})
|
||||
})
|
||||
|
||||
widget.div.appendChild(btn)
|
||||
widget.div.appendChild(select)
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
if (this.onResize) {
|
||||
this.onResize(this.size)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'ImagesPrompt_') {
|
||||
try {
|
||||
let prompt = node.widgets.filter(w => w.name === 'image_base64')[0]
|
||||
let text = node.widgets.filter(w => w.name === 'text')[0]
|
||||
let uploadWidget = node.widgets.filter(w => w.name == 'upload')[0]
|
||||
// console.log('##prompt',prompt.value)
|
||||
let img = uploadWidget.div.querySelector('img')
|
||||
let json = JSON.parse(uploadWidget.value)
|
||||
|
||||
for (let index = 0; index < json.length; index++) {
|
||||
const j = json[index]
|
||||
let base64 = await parseImageToBase64(j.imgurl)
|
||||
if (base64 === prompt.value) {
|
||||
json[index].selected = true
|
||||
}
|
||||
}
|
||||
|
||||
if (json && json[0]) {
|
||||
uploadWidget.select.style.display = 'block'
|
||||
createSelect(img, uploadWidget.select, json, prompt, text)
|
||||
}
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const createInputImageForBatch = (base64, widget) => {
|
||||
let im = new Image()
|
||||
im.src = base64
|
||||
im.style = `width: 88px;`
|
||||
|
||||
im.addEventListener('click', e => {
|
||||
let newValue = []
|
||||
let items = widget.value?.base64 || []
|
||||
for (const v of items) {
|
||||
if (v != base64) newValue.push(v)
|
||||
}
|
||||
widget.value.base64 = newValue
|
||||
im.remove()
|
||||
})
|
||||
|
||||
return im
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.Comfy.LoadImagesToBatch',
|
||||
async getCustomWidgets (app) {
|
||||
return {
|
||||
IMAGEBASE64 (node, inputName, inputData, app) {
|
||||
// console.log('##node', node)
|
||||
const widget = {
|
||||
value: {
|
||||
base64: []
|
||||
}, // 不能[x,x,x]
|
||||
type: inputData[0], // the type
|
||||
name: inputName, // the name, slice
|
||||
size: [128, 32], // a default size
|
||||
draw (ctx, node, width, y) {},
|
||||
computeSize (...args) {
|
||||
return [128, 32] // a method to compute the current size of the widget
|
||||
}
|
||||
// serializeValue (nodeId, widgetIndex) {
|
||||
// return widget.value
|
||||
// },
|
||||
}
|
||||
// widget.something = something; // maybe adds stuff to it
|
||||
node.addCustomWidget(widget) // adds it to the node
|
||||
return widget // and returns it.
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'LoadImagesToBatch') {
|
||||
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
let imagesWidget = this.widgets.filter(w => w.name == 'images')[0]
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'image_base64',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, 44, node.size[1])
|
||||
)
|
||||
},
|
||||
serialize: false
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
document.body.appendChild(widget.div)
|
||||
|
||||
let imagePreview = document.createElement('div')
|
||||
let imagesDiv = document.createElement('div') //显示图片
|
||||
imagesDiv.className = 'images_preview'
|
||||
imagesDiv.style = `width: calc(100% - 14px);
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
padding: 7px; justify-content: space-between;
|
||||
align-items: center;`
|
||||
|
||||
let inputImage = document.createElement('input')
|
||||
inputImage.type = 'file'
|
||||
inputImage.style.display = 'none'
|
||||
inputImage.addEventListener('change', e => {
|
||||
e.preventDefault()
|
||||
const file = e.target.files[0]
|
||||
const reader = new FileReader()
|
||||
reader.onload = async event => {
|
||||
let base64 = event.target.result
|
||||
//压缩图片,控制1024以内
|
||||
base64 = await loadImageToCanvas(base64)
|
||||
// console.log(base64)
|
||||
if (!imagesWidget.value) imagesWidget.value = { base64: [] }
|
||||
imagesWidget.value.base64.push(base64)
|
||||
let im = createInputImageForBatch(base64, imagesWidget)
|
||||
imagesDiv.appendChild(im)
|
||||
}
|
||||
reader.readAsDataURL(file)
|
||||
})
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Upload Image'
|
||||
|
||||
btn.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid;height: 30px;min-width: 122px;
|
||||
`
|
||||
|
||||
btn.addEventListener('click', e => {
|
||||
e.preventDefault()
|
||||
inputImage.click()
|
||||
})
|
||||
|
||||
widget.div.appendChild(imagePreview)
|
||||
imagePreview.appendChild(imagesDiv)
|
||||
imagePreview.appendChild(btn)
|
||||
imagePreview.appendChild(inputImage)
|
||||
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
// document.addEventListener('wheel', handleMouseWheel)
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
inputImage.remove()
|
||||
widget.div.remove()
|
||||
try {
|
||||
// document.removeEventListener('wheel', handleMouseWheel)
|
||||
} catch (error) {
|
||||
console.log(error)
|
||||
}
|
||||
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
}
|
||||
|
||||
if (nodeData.name === 'SaveImageAndMetadata_') {
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
// /web/extensions/core/saveImageExtraOutput.js
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
const r = onNodeCreated
|
||||
? onNodeCreated.apply(this, arguments)
|
||||
: undefined
|
||||
const widget = this.widgets.find(w => w.name === 'filename_prefix')
|
||||
widget.serializeValue = () => {
|
||||
return applyTextReplacements(app, widget.value)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
onExecuted?.apply(this, arguments)
|
||||
console.log('##onExecuted', this, message)
|
||||
//TODO 是否 保存base64
|
||||
if (message.base64) {
|
||||
if (Array.isArray(message.base64)) {
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'LoadImagesToBatch') {
|
||||
// await sleep(0)
|
||||
let imagesWidget = node.widgets.filter(w => w.name === 'images')[0]
|
||||
let imagePreview = node.widgets.filter(w => w.name == 'image_base64')[0]
|
||||
|
||||
let pre = imagePreview.div.querySelector('.images_preview')
|
||||
for (const d of imagesWidget.value?.base64 || []) {
|
||||
let im = createInputImageForBatch(d, imagesWidget)
|
||||
pre.appendChild(im)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
// 如何引入css
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.output.ComparingTwoFrames_',
|
||||
init () {
|
||||
$el('link', {
|
||||
rel: 'stylesheet',
|
||||
href: '/extensions/comfyui-mixlab-nodes/lib/juxtapose.css',
|
||||
parent: document.head
|
||||
})
|
||||
|
||||
$el('style', {
|
||||
textContent: `
|
||||
.juxtapose-name{
|
||||
display: none!important;
|
||||
}
|
||||
`,
|
||||
parent: document.body
|
||||
})
|
||||
},
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'ComparingTwoFrames_') {
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
const r = onNodeCreated
|
||||
? onNodeCreated.apply(this, arguments)
|
||||
: undefined
|
||||
|
||||
this.size = [400, this.size[1]]
|
||||
console.log('##onNodeCreated', this)
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'preview',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, 400, 44, node.size[1])
|
||||
)
|
||||
},
|
||||
serialize: false
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
|
||||
return r
|
||||
|
||||
}
|
||||
|
||||
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
onExecuted?.apply(this, arguments)
|
||||
console.log('##onExecuted', this, message)
|
||||
|
||||
this.widgets[0].div.id = 'mix_comparingtowframes_' + this.id
|
||||
|
||||
let after_image = message.after_images[0]
|
||||
let before_image = message.before_images[0]
|
||||
|
||||
after_image = `${window.location.protocol}//${
|
||||
window.location.hostname
|
||||
}:${window.location.port}/view?filename=${encodeURIComponent(
|
||||
after_image.filename
|
||||
)}&type=${after_image.type}&subfolder=${encodeURIComponent(
|
||||
after_image.subfolder
|
||||
)}&t=${+new Date()}`
|
||||
|
||||
before_image = `${window.location.protocol}//${
|
||||
window.location.hostname
|
||||
}:${window.location.port}/view?filename=${encodeURIComponent(
|
||||
before_image.filename
|
||||
)}&type=${before_image.type}&subfolder=${encodeURIComponent(
|
||||
before_image.subfolder
|
||||
)}&t=${+new Date()}`
|
||||
|
||||
this.widgets[0].div.innerHTML = ''
|
||||
|
||||
let slider = new juxtapose.JXSlider(
|
||||
'#mix_comparingtowframes_' + this.id,
|
||||
[
|
||||
{
|
||||
src: before_image,
|
||||
label: 'Before'
|
||||
},
|
||||
{
|
||||
src: after_image,
|
||||
label: 'After'
|
||||
}
|
||||
],
|
||||
{
|
||||
animate: true,
|
||||
showLabels: true,
|
||||
showCredits: false,
|
||||
startingPosition: '50%',
|
||||
makeResponsive: false
|
||||
}
|
||||
)
|
||||
|
||||
this.widgets_values = [
|
||||
{
|
||||
src: before_image,
|
||||
label: 'Before'
|
||||
},
|
||||
{
|
||||
src: after_image,
|
||||
label: 'After'
|
||||
}
|
||||
]
|
||||
this.size=[this.size[0],300]
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
// console.log('##loadedGraphNode', node)
|
||||
if (node.type === 'ComparingTwoFrames_') {
|
||||
// node.widgets[0].div.id = 'mix_comparingtowframes_' + node.id
|
||||
// if (node.widgets_values && node.widgets_values[0]) {
|
||||
// node.widgets[0].div.innerHTML = ''
|
||||
|
||||
// let slider = new juxtapose.JXSlider(
|
||||
// '#mix_comparingtowframes_' + node.id,
|
||||
// node.widgets_values,
|
||||
// {
|
||||
// animate: true,
|
||||
// showLabels: true,
|
||||
// showCredits: false,
|
||||
// startingPosition: '50%',
|
||||
// makeResponsive: false
|
||||
// }
|
||||
// )
|
||||
// }
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
@@ -1,8 +1,70 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
// import { api } from '../../../scripts/api.js'
|
||||
import { api } from '../../../scripts/api.js'
|
||||
import { ComfyWidgets } from '../../../scripts/widgets.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
function downloadJsonFile (jsonData, fileName = 'grid.json') {
|
||||
const dataString = JSON.stringify(jsonData)
|
||||
const blob = new Blob([dataString], { type: 'application/json' })
|
||||
const url = URL.createObjectURL(blob)
|
||||
|
||||
const link = document.createElement('a')
|
||||
link.href = url
|
||||
link.download = fileName
|
||||
link.click()
|
||||
|
||||
// 释放URL对象
|
||||
setTimeout(() => {
|
||||
URL.revokeObjectURL(url)
|
||||
}, 0)
|
||||
}
|
||||
|
||||
function createSelectWithOptions (options) {
|
||||
const select = document.createElement('select')
|
||||
|
||||
options.forEach(option => {
|
||||
const optionElement = document.createElement('option')
|
||||
optionElement.text = option
|
||||
optionElement.value = option
|
||||
select.appendChild(optionElement)
|
||||
})
|
||||
|
||||
select.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
height: 30px;
|
||||
min-width: 122px;
|
||||
position: absolute;
|
||||
top: 24px;
|
||||
left: 88px;
|
||||
z-index: 999999999999999;
|
||||
`
|
||||
|
||||
return select
|
||||
}
|
||||
|
||||
function drawCanvasWithText (w, h, tag, color = 'rgba(255,255,255,0.4)') {
|
||||
const canvas = document.createElement('canvas')
|
||||
const ctx = canvas.getContext('2d')
|
||||
|
||||
// 设置画布大小
|
||||
canvas.width = w
|
||||
canvas.height = h
|
||||
|
||||
// 绘制白色背景
|
||||
ctx.fillStyle = color
|
||||
ctx.fillRect(0, 0, canvas.width, canvas.height)
|
||||
|
||||
// 绘制文字
|
||||
ctx.fillStyle = '#000000'
|
||||
ctx.font = '20px Arial'
|
||||
ctx.fillText(tag, 50, 50)
|
||||
|
||||
// 导出为Base64
|
||||
const base64 = canvas.toDataURL()
|
||||
|
||||
return base64
|
||||
}
|
||||
|
||||
function get_position_style (ctx, widget_width, y, node_height) {
|
||||
const MARGIN = 4 // the margin around the html element
|
||||
|
||||
@@ -156,6 +218,29 @@ const parseSvg = async svgContent => {
|
||||
return { data, image: base64, svgElement }
|
||||
}
|
||||
|
||||
function findImages (nodeId) {
|
||||
// 检查当前节点是否有 imgs 字段
|
||||
const n = app.graph.getNodeById(nodeId)
|
||||
if (n.imgs) {
|
||||
return n.imgs
|
||||
}
|
||||
|
||||
// 检查当前节点的 inputs 是否有 image 字段
|
||||
if (n.inputs) {
|
||||
for (let i = 0; i < n.inputs.length; i++) {
|
||||
if (n.inputs[i].name === 'image' || n.inputs[i].name === 'images') {
|
||||
// 获取新的 nodeId,并递归调用 findImages 函数
|
||||
var linkId = n.inputs[i]?.link
|
||||
var origin_id = app.graph.links[linkId].origin_id
|
||||
return findImages(origin_id)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 如果没有找到 imgs 字段或者 image 字段,则返回 null
|
||||
return null
|
||||
}
|
||||
|
||||
async function setArea (cw, ch, topBase64, base64, data, fn) {
|
||||
let displayHeight = Math.round(window.screen.availHeight * 0.8)
|
||||
let div = document.createElement('div')
|
||||
@@ -327,6 +412,196 @@ async function setArea (cw, ch, topBase64, base64, data, fn) {
|
||||
}
|
||||
}
|
||||
|
||||
async function setAreaTags (cw, ch, grids, fn) {
|
||||
let base64 = drawCanvasWithText(cw, ch, '', 'white')
|
||||
let displayHeight = Math.round(window.screen.availHeight * 0.8)
|
||||
let div = document.createElement('div')
|
||||
div.innerHTML = `
|
||||
<div id='ml_overlay' style='position: absolute;top:0;background: #251f1fc4;
|
||||
height: 100vh;
|
||||
z-index:999999;
|
||||
width: 100%;'>
|
||||
<img id='ml_video' style='position: absolute;
|
||||
height: ${displayHeight}px;user-select: none;
|
||||
-webkit-user-drag: none;
|
||||
outline: 2px solid #eaeaea;
|
||||
box-shadow: 8px 9px 17px #575757;' />
|
||||
${Array.from(grids, g => {
|
||||
const { label: tag, grid } = g
|
||||
const [dx, dy, dw, dh] = grid
|
||||
const base64Data = drawCanvasWithText(dw, dh, tag)
|
||||
|
||||
let x = 0,
|
||||
y = 0,
|
||||
width = (cw * displayHeight) / ch,
|
||||
height = displayHeight
|
||||
|
||||
let imgWidth = cw
|
||||
let imgHeight = ch
|
||||
|
||||
if (dw > 0 && dh > 0) {
|
||||
// 相同尺寸窗口,恢复选区
|
||||
x = (width * dx) / imgWidth
|
||||
y = (height * dy) / imgHeight
|
||||
width = (width * dw) / imgWidth
|
||||
height = (height * dh) / imgHeight
|
||||
}
|
||||
|
||||
return `<div class='ml_selection'
|
||||
data-tag="${tag}"
|
||||
style='position:absolute;
|
||||
border: 2px dashed red;
|
||||
pointer-events: none;
|
||||
background-image: url("${base64Data}");
|
||||
background-repeat: no-repeat;
|
||||
background-size: cover;
|
||||
left:${x}px;
|
||||
top:${y}px;
|
||||
width:${width}px;
|
||||
height:${height}px;
|
||||
'></div>`
|
||||
})}
|
||||
<div class="mx_close"> X </div>
|
||||
</div>`
|
||||
// document.body.querySelector('#ml_overlay')
|
||||
document.body.appendChild(div)
|
||||
|
||||
const tags = Array.from(grids, g => g.label)
|
||||
let select = createSelectWithOptions(tags)
|
||||
document.body.appendChild(select)
|
||||
|
||||
let img = div.querySelector('#ml_video')
|
||||
// let overlay = div.querySelector('#ml_overlay')
|
||||
let selections = [...div.querySelectorAll('.ml_selection')]
|
||||
|
||||
let selection = selections.filter(
|
||||
s => s.getAttribute('data-tag') === select.value
|
||||
)[0]
|
||||
|
||||
select.addEventListener('change', e => {
|
||||
selection = selections.filter(
|
||||
s => s.getAttribute('data-tag') === select.value
|
||||
)[0]
|
||||
})
|
||||
|
||||
// console.log(select.value,selection)
|
||||
let close = div.querySelector('.mx_close')
|
||||
let startX, startY, endX, endY
|
||||
let start = false
|
||||
let setDone = false
|
||||
// Set video source
|
||||
img.src = base64
|
||||
// canvas.toDataURL();
|
||||
close.style = `cursor: pointer;
|
||||
position: fixed;
|
||||
left: 12px;
|
||||
top: 12px;
|
||||
z-index: 99999999;
|
||||
background: black;
|
||||
width: 44px;
|
||||
height: 44px;
|
||||
text-align: center;
|
||||
line-height: 44px;`
|
||||
|
||||
// Add mouse events
|
||||
img.addEventListener('mousedown', startSelection)
|
||||
img.addEventListener('mousemove', updateSelection)
|
||||
img.addEventListener('mouseup', endSelection)
|
||||
|
||||
const removeDiv = () => {
|
||||
div.remove()
|
||||
select?.remove()
|
||||
close.removeEventListener('click', removeDiv)
|
||||
img.removeEventListener('mousedown', startSelection)
|
||||
img.removeEventListener('mousemove', updateSelection)
|
||||
img.removeEventListener('mouseup', endSelection)
|
||||
img.removeEventListener('mousedown', setDoneCheck)
|
||||
}
|
||||
close.addEventListener('click', removeDiv)
|
||||
|
||||
const setDoneCheck = event => {
|
||||
console.log(setDone)
|
||||
if (setDone) {
|
||||
img.addEventListener('mousedown', startSelection)
|
||||
img.addEventListener('mousemove', updateSelection)
|
||||
img.addEventListener('mouseup', endSelection)
|
||||
setDone = false
|
||||
start = false
|
||||
startX = event.clientX
|
||||
startY = event.clientY
|
||||
}
|
||||
}
|
||||
img.addEventListener('mousedown', setDoneCheck)
|
||||
|
||||
function remove () {
|
||||
img.removeEventListener('mousedown', startSelection)
|
||||
img.removeEventListener('mousemove', updateSelection)
|
||||
img.removeEventListener('mouseup', endSelection)
|
||||
setDone = true
|
||||
// select?.remove()
|
||||
}
|
||||
|
||||
function startSelection (event) {
|
||||
if (start == false) {
|
||||
startX = event.clientX
|
||||
startY = event.clientY
|
||||
updateSelection(event)
|
||||
start = true
|
||||
} else {
|
||||
}
|
||||
}
|
||||
|
||||
function updateSelection (event) {
|
||||
endX = event.clientX
|
||||
endY = event.clientY
|
||||
|
||||
// Calculate width, height, and coordinates
|
||||
let width = Math.abs(endX - startX)
|
||||
let height = Math.abs(endY - startY)
|
||||
let left = Math.min(startX, endX)
|
||||
let top = Math.min(startY, endY)
|
||||
|
||||
// Set selection style
|
||||
selection.style.left = left + 'px'
|
||||
selection.style.top = top + 'px'
|
||||
selection.style.width = width + 'px'
|
||||
selection.style.height = height + 'px'
|
||||
}
|
||||
|
||||
function endSelection (event) {
|
||||
endX = event.clientX
|
||||
endY = event.clientY
|
||||
|
||||
// 获取img元素的真实宽度和高度
|
||||
let imgWidth = img.naturalWidth
|
||||
let imgHeight = img.naturalHeight
|
||||
|
||||
// 换算起始坐标
|
||||
let realStartX = (startX / img.offsetWidth) * imgWidth
|
||||
let realStartY = (startY / img.offsetHeight) * imgHeight
|
||||
|
||||
// 换算起始坐标
|
||||
let realEndX = (endX / img.offsetWidth) * imgWidth
|
||||
let realEndY = (endY / img.offsetHeight) * imgHeight
|
||||
|
||||
startX = realStartX
|
||||
startY = realStartY
|
||||
endX = realEndX
|
||||
endY = realEndY
|
||||
// Calculate width, height, and coordinates
|
||||
let width = Math.round(Math.abs(endX - startX))
|
||||
let height = Math.round(Math.abs(endY - startY))
|
||||
let left = Math.round(Math.min(startX, endX))
|
||||
let top = Math.round(Math.min(startY, endY))
|
||||
|
||||
if (width <= 0 && height <= 0) return remove()
|
||||
|
||||
if (!!fn) fn(select.value, left, top, width, height)
|
||||
|
||||
remove()
|
||||
}
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.layer.ShowLayer',
|
||||
async getCustomWidgets (app) {
|
||||
@@ -571,15 +846,19 @@ app.registerExtension({
|
||||
}
|
||||
}
|
||||
try {
|
||||
console.log('this.inputs', this.inputs)
|
||||
let topLinkId = this.inputs[0].link
|
||||
let topNodeId = app.graph.links[topLinkId].origin_id
|
||||
let topIm = app.graph.getNodeById(topNodeId).imgs[0]
|
||||
console.log('this.inputs', this.id)
|
||||
let imgs = findImages(this.id)
|
||||
|
||||
// let topLinkId = this.inputs[0].link
|
||||
// let topNodeId = app.graph.links[topLinkId].origin_id
|
||||
let topIm = imgs[0]
|
||||
|
||||
let linkId = this.inputs[3].link
|
||||
let nodeId = app.graph.links[linkId].origin_id
|
||||
// console.log(linkId,this.inputs)
|
||||
let im = app.graph.getNodeById(nodeId).imgs[0]
|
||||
let imgs2 = findImages(nodeId)
|
||||
let im = imgs2[0]
|
||||
console.log(topIm, im)
|
||||
// let src = im.src
|
||||
setArea(
|
||||
im.naturalWidth,
|
||||
@@ -589,7 +868,9 @@ app.registerExtension({
|
||||
data,
|
||||
updateValue
|
||||
)
|
||||
} catch (error) {}
|
||||
} catch (error) {
|
||||
console.log(error)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -609,3 +890,306 @@ app.registerExtension({
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.layer.GridInput',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'GridInput') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const grids_widget = this.widgets.filter(w => w.name == 'grids')[0]
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'upload',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1]),
|
||||
{
|
||||
justifyContent: 'flex-start'
|
||||
}
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
const addBtn = document.createElement('button')
|
||||
addBtn.innerText = 'Add Box'
|
||||
addBtn.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid;height: 30px;min-width: 122px;
|
||||
`
|
||||
|
||||
const vbtn = document.createElement('button')
|
||||
vbtn.innerText = 'Set Box'
|
||||
vbtn.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid;height: 30px;min-width: 122px;
|
||||
`
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Upload JSON'
|
||||
|
||||
btn.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid;height: 30px;min-width: 122px;
|
||||
`
|
||||
|
||||
addBtn.addEventListener('click', () => {
|
||||
const { width, height, grids } = JSON.parse(grids_widget.value)
|
||||
grids.push({
|
||||
label: 'background',
|
||||
grid: [12, 12, width - 24, height - 24]
|
||||
})
|
||||
grids_widget.value = JSON.stringify(
|
||||
{
|
||||
width,
|
||||
height,
|
||||
grids
|
||||
},
|
||||
null,
|
||||
2
|
||||
)
|
||||
})
|
||||
|
||||
vbtn.addEventListener('click', () => {
|
||||
const { width, height, grids } = JSON.parse(grids_widget.value)
|
||||
|
||||
setAreaTags(width, height, grids, (tag, x, y, w, h) => {
|
||||
grids_widget.value = JSON.stringify(
|
||||
{
|
||||
width,
|
||||
height,
|
||||
grids: Array.from(grids, g => {
|
||||
if (g.label === tag) {
|
||||
g.grid = [x, y, w, h]
|
||||
}
|
||||
return g
|
||||
})
|
||||
},
|
||||
null,
|
||||
2
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
btn.addEventListener('click', () => {
|
||||
let inp = document.createElement('input')
|
||||
inp.type = 'file'
|
||||
inp.accept = '.json'
|
||||
inp.click()
|
||||
inp.addEventListener('change', event => {
|
||||
// 获取选择的文件
|
||||
const file = event.target.files[0]
|
||||
this.title = file.name.split('.')[0]
|
||||
|
||||
// console.log(file.name.split('.')[0])
|
||||
// 创建文件读取器
|
||||
const reader = new FileReader()
|
||||
|
||||
// 定义读取完成事件的回调函数
|
||||
reader.onload = event => {
|
||||
// 读取完成后的文本内容
|
||||
const fileContent = JSON.parse(event.target.result)
|
||||
const grids = fileContent
|
||||
grids_widget.value = JSON.stringify(grids, null, 2)
|
||||
// widget.value = grids
|
||||
|
||||
inp.remove()
|
||||
}
|
||||
|
||||
// 以文本方式读取文件
|
||||
reader.readAsText(file)
|
||||
})
|
||||
})
|
||||
|
||||
widget.div.appendChild(addBtn)
|
||||
widget.div.appendChild(vbtn)
|
||||
widget.div.appendChild(btn)
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
const r = onExecuted?.apply?.(this, arguments)
|
||||
|
||||
let json = message.json
|
||||
if (json) {
|
||||
json = {
|
||||
width: json[0],
|
||||
height: json[1],
|
||||
grids: json[2]
|
||||
}
|
||||
grids_widget.value = JSON.stringify(json, null, 2)
|
||||
// widget.value = json
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
if (this.onResize) {
|
||||
this.onResize(this.size)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'GridInput') {
|
||||
try {
|
||||
const grids_widget = node.widgets.filter(w => w.name == 'grids')[0]
|
||||
const { width, height, grids } = JSON.parse(grids_widget.value)
|
||||
console.log('#GridInput', node, grids)
|
||||
|
||||
const div = node.widgets.filter(w => w.name == 'upload')[0]
|
||||
div.div.querySelector('select').innerHTML = Array.from(
|
||||
grids,
|
||||
g => `<option value="${g.label}">${g.label}</option>`
|
||||
).join('')
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.layer.GridDisplayAndSave',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'GridDisplayAndSave') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const grids_widget = this.widgets.filter(w => w.name == 'grids')[0]
|
||||
console.log('GridDisplayAndSave', grids_widget)
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'save_json',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1]),
|
||||
{
|
||||
justifyContent: 'flex-start',
|
||||
flexDirection: 'column'
|
||||
}
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Save JSON'
|
||||
|
||||
btn.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid;height: 30px;min-width: 122px;
|
||||
max-width: 122px;
|
||||
`
|
||||
|
||||
btn.addEventListener('click', () => {
|
||||
if (window._mixlab_grid)
|
||||
downloadJsonFile(
|
||||
window._mixlab_grid,
|
||||
this.widgets.filter(w => w.name == 'filename_prefix')[0]?.value +
|
||||
'_grid.json'
|
||||
)
|
||||
})
|
||||
|
||||
widget.div.appendChild(btn)
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
const r = onExecuted?.apply?.(this, arguments)
|
||||
let save_json = this.widgets.filter(d => d.name == 'save_json')[0]
|
||||
let div = save_json?.div
|
||||
// console.log('Test',message)
|
||||
|
||||
let image = message.image[0]
|
||||
let json = message.json
|
||||
if (image) {
|
||||
const { filename, subfolder, type } = image
|
||||
|
||||
if (!div.querySelector('img')) {
|
||||
let im = new Image()
|
||||
div.appendChild(im)
|
||||
im.style.width = '100%'
|
||||
}
|
||||
div.querySelector('img').src = api.apiURL(
|
||||
`/view?filename=${encodeURIComponent(
|
||||
filename
|
||||
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
|
||||
)
|
||||
|
||||
window._mixlab_grid = {
|
||||
width: json[0],
|
||||
height: json[1],
|
||||
grids: json[2]
|
||||
}
|
||||
// console.log(src)
|
||||
}
|
||||
|
||||
this.onResize?.(this.size)
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
if (this.onResize) {
|
||||
this.onResize(this.size)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'GridDisplayAndSave') {
|
||||
try {
|
||||
let grids_widget = node.widgets.filter(w => w.name === 'grids')[0]
|
||||
// let ks = getLocalData(`_mixlab_PromptSlide`)
|
||||
let uploadWidget = node.widgets.filter(w => w.name == 'upload')[0]
|
||||
// console.log('##widget', uploadWidget.value)
|
||||
let grids = JSON.parse(uploadWidget.value)
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
@@ -1331,7 +1331,13 @@ app.registerExtension({
|
||||
let w = 360,
|
||||
s = widget.preview.videoWidth / widget.preview.videoHeight,
|
||||
h = w / s || w
|
||||
console.log(h)
|
||||
// console.log(h)
|
||||
|
||||
if (!window.documentPictureInPicture) {
|
||||
window.alert(
|
||||
'This feature is available only in secure contexts (HTTPS), in some or all supporting browsers. https://developer.mozilla.org/en-US/docs/Web/API/Document_Picture-in-Picture_API'
|
||||
)
|
||||
}
|
||||
|
||||
const pipWindow = await documentPictureInPicture.requestWindow({
|
||||
width: w,
|
||||
@@ -1800,15 +1806,19 @@ const updateUI = node => {
|
||||
pw.inputEl.title = `Total of ${prompts.length} prompts`
|
||||
} else {
|
||||
// 动态添加
|
||||
console.log('ComfyWidgets',ComfyWidgets.STRING(
|
||||
node,
|
||||
'prompts',
|
||||
['STRING', { multiline: true }]
|
||||
))
|
||||
// console.log('ComfyWidgets',ComfyWidgets.STRING(
|
||||
// node,
|
||||
// 'prompts',
|
||||
// ['STRING', { multiline: true }]
|
||||
// ))
|
||||
|
||||
// ComfyWidgets.STRING(this, "", ["", {default:this.properties.text, multiline: true}], app)
|
||||
|
||||
const w = ComfyWidgets.STRING(
|
||||
node,
|
||||
'prompts',
|
||||
['STRING', { multiline: true }]
|
||||
['STRING', { multiline: true }],
|
||||
app
|
||||
).widget
|
||||
w.inputEl.readOnly = true
|
||||
w.inputEl.style.opacity = 0.6
|
||||
@@ -2089,13 +2099,13 @@ const node = {
|
||||
name: 'RandomPrompt',
|
||||
async init (app) {
|
||||
// Any initial setup to run as soon as the page loads
|
||||
console.log('[logging]', 'extension init')
|
||||
// console.log('[logging]', 'extension init')
|
||||
|
||||
if (window.location.href.match('/?')) {
|
||||
const { workflow } = getURLParameters(window.location.href)
|
||||
if (workflow)
|
||||
get_my_workflow().then(data => {
|
||||
console.log('#get_my_workflow', data)
|
||||
// console.log('#get_my_workflow', data)
|
||||
let my_workflow = data.filter(
|
||||
d => d.filename == 'my_workflow.json'
|
||||
)[0]
|
||||
@@ -2131,10 +2141,15 @@ const node = {
|
||||
// }
|
||||
},
|
||||
loadedGraphNode (node, app) {
|
||||
// Fires for each node when loading/dragging/etc a workflow json or png
|
||||
// If you break something in the backend and want to patch workflows in the frontend
|
||||
// This is the place to do this
|
||||
// console.log("[logging]", "loaded graph node: ", exportGraph(node.graph));
|
||||
if (node.type === 'RandomPrompt') {
|
||||
try {
|
||||
let max_count = node.widgets.filter(w => w.name === 'max_count')[0]
|
||||
max_count.value = node.widgets_values[0]
|
||||
// console.log('RandomPrompt',max_count,node.widgets_values[0])
|
||||
} catch (error) {
|
||||
console.log(error)
|
||||
}
|
||||
}
|
||||
},
|
||||
async nodeCreated (node) {
|
||||
if (node.type === 'RandomPrompt') {
|
||||
@@ -2227,7 +2242,7 @@ const node = {
|
||||
const r = onExecuted?.apply?.(this, arguments)
|
||||
|
||||
let prompts = message.prompts
|
||||
console.log('executed', message)
|
||||
// console.log('executed', message)
|
||||
// console.log('#RandomPrompt', this.widgets)
|
||||
const pw = this.widgets.filter(w => w.name === 'prompts')[0]
|
||||
|
||||
@@ -2238,7 +2253,7 @@ const node = {
|
||||
} else {
|
||||
// 动态添加
|
||||
const w = ComfyWidgets.STRING(
|
||||
node,
|
||||
this,
|
||||
'prompts',
|
||||
['STRING', { multiline: true }],
|
||||
app
|
||||
|
||||
@@ -0,0 +1,565 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
import { api } from '../../../scripts/api.js'
|
||||
import { ComfyWidgets } from '../../../scripts/widgets.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
import PhotoSwipeLightbox from '/extensions/comfyui-mixlab-nodes/lib/photoswipe-lightbox.esm.min.js'
|
||||
function loadCSS (url) {
|
||||
var link = document.createElement('link')
|
||||
link.rel = 'stylesheet'
|
||||
link.type = 'text/css'
|
||||
link.href = url
|
||||
document.getElementsByTagName('head')[0].appendChild(link)
|
||||
|
||||
// Create a style element
|
||||
const style = document.createElement('style')
|
||||
// Define the CSS rule for scrollbar width
|
||||
const cssRule = `.pswp__custom-caption {
|
||||
background: rgb(20 27 70);
|
||||
font-size: 16px;
|
||||
color: #fff;
|
||||
width: calc(100% - 32px);
|
||||
max-width: 980px;
|
||||
padding: 2px 8px;
|
||||
border-radius: 4px;
|
||||
position: absolute;
|
||||
left: 50%;
|
||||
bottom: 16px;
|
||||
transform: translateX(-50%);
|
||||
}
|
||||
.pswp__custom-caption a {
|
||||
color: #fff;
|
||||
text-decoration: underline;
|
||||
}
|
||||
.hidden-caption-content {
|
||||
display: none;
|
||||
}`
|
||||
// Add the CSS rule to the style element
|
||||
style.appendChild(document.createTextNode(cssRule))
|
||||
|
||||
// Append the style element to the document head
|
||||
document.head.appendChild(style)
|
||||
}
|
||||
loadCSS('/extensions/comfyui-mixlab-nodes/lib/photoswipe.min.css')
|
||||
|
||||
function initLightBox () {
|
||||
const lightbox = new PhotoSwipeLightbox({
|
||||
gallery: '.prompt_image_output',
|
||||
children: 'a',
|
||||
pswpModule: () =>
|
||||
import('/extensions/comfyui-mixlab-nodes/lib/photoswipe.esm.min.js')
|
||||
})
|
||||
|
||||
lightbox.on('uiRegister', function () {
|
||||
lightbox.pswp.ui.registerElement({
|
||||
name: 'custom-caption',
|
||||
order: 9,
|
||||
isButton: false,
|
||||
appendTo: 'root',
|
||||
html: 'Caption text',
|
||||
onInit: (el, pswp) => {
|
||||
lightbox.pswp.on('change', () => {
|
||||
const currSlideElement = lightbox.pswp.currSlide.data.element
|
||||
let captionHTML = ''
|
||||
if (currSlideElement) {
|
||||
const hiddenCaption = currSlideElement.querySelector(
|
||||
'.hidden-caption-content'
|
||||
)
|
||||
if (hiddenCaption) {
|
||||
// get caption from element with class hidden-caption-content
|
||||
captionHTML = hiddenCaption.innerHTML
|
||||
} else {
|
||||
// get caption from alt attribute
|
||||
captionHTML = currSlideElement
|
||||
.querySelector('img')
|
||||
.getAttribute('alt')
|
||||
}
|
||||
}
|
||||
el.innerHTML = captionHTML || ''
|
||||
})
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
lightbox.init()
|
||||
}
|
||||
|
||||
function get_position_style (ctx, widget_width, y, node_height) {
|
||||
const MARGIN = 4 // the margin around the html element
|
||||
|
||||
/* Create a transform that deals with all the scrolling and zooming */
|
||||
const elRect = ctx.canvas.getBoundingClientRect()
|
||||
const transform = new DOMMatrix()
|
||||
.scaleSelf(
|
||||
elRect.width / ctx.canvas.width,
|
||||
elRect.height / ctx.canvas.height
|
||||
)
|
||||
.multiplySelf(ctx.getTransform())
|
||||
.translateSelf(MARGIN, MARGIN + y)
|
||||
|
||||
return {
|
||||
transformOrigin: '0 0',
|
||||
transform: transform,
|
||||
left: `0`,
|
||||
top: `0`,
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
maxWidth: `${widget_width - MARGIN * 2}px`,
|
||||
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
|
||||
width: `${widget_width - MARGIN * 2 - 24}px`,
|
||||
// height: `${node_height * 0.3 - MARGIN * 2}px`,
|
||||
// background: '#EEEEEE',
|
||||
paddingLeft: '12px',
|
||||
display: 'flex',
|
||||
flexDirection: 'row',
|
||||
// alignItems: 'center',
|
||||
justifyContent: 'space-between'
|
||||
}
|
||||
}
|
||||
function createImage (url) {
|
||||
let im = new Image()
|
||||
return new Promise((res, rej) => {
|
||||
im.onload = () => res(im)
|
||||
im.src = url
|
||||
})
|
||||
}
|
||||
|
||||
async function fetchImage (url) {
|
||||
try {
|
||||
const response = await fetch(url)
|
||||
const blob = await response.blob()
|
||||
|
||||
return blob
|
||||
} catch (error) {
|
||||
console.error('出现错误:', error)
|
||||
}
|
||||
}
|
||||
|
||||
const getLocalData = key => {
|
||||
let data = {}
|
||||
try {
|
||||
data = JSON.parse(localStorage.getItem(key)) || {}
|
||||
} catch (error) {
|
||||
return {}
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
const setLocalDataOfWin = (key, value) => {
|
||||
localStorage.setItem(key, JSON.stringify(value))
|
||||
// window[key] = value
|
||||
}
|
||||
|
||||
const createSelect = (select, opts, targetWidget) => {
|
||||
select.style.display = 'block'
|
||||
let html = ''
|
||||
let isMatch = false
|
||||
for (const opt of opts) {
|
||||
html += `<option value='${opt}' ${
|
||||
targetWidget.value === opt ? 'selected' : ''
|
||||
}>${opt}</option>`
|
||||
if (targetWidget.value === opt) isMatch = true
|
||||
}
|
||||
select.innerHTML = html
|
||||
if (!isMatch) targetWidget.value = opts[0]
|
||||
// 添加change事件监听器
|
||||
select.addEventListener('change', function () {
|
||||
// 获取选中的选项的值
|
||||
var selectedOption = select.options[select.selectedIndex].value
|
||||
targetWidget.value = selectedOption
|
||||
// console.log(widget,selectedOption)
|
||||
})
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.prompt.RandomPrompt',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'RandomPrompt') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const mutable_prompt = this.widgets.filter(
|
||||
w => w.name == 'mutable_prompt'
|
||||
)[0]
|
||||
// console.log('PromptSlide nodeData', prompt_keyword)
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'upload',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1])
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Upload Keywords'
|
||||
|
||||
btn.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid; height: 30px;min-width: 122px;
|
||||
`
|
||||
|
||||
// const btn=document.createElement('button');
|
||||
// btn.innerText='Upload'
|
||||
btn.addEventListener('click', () => {
|
||||
let inp = document.createElement('input')
|
||||
inp.type = 'file'
|
||||
inp.accept = '.txt'
|
||||
inp.click()
|
||||
inp.addEventListener('change', event => {
|
||||
// 获取选择的文件
|
||||
const file = event.target.files[0]
|
||||
this.title = file.name.split('.')[0]
|
||||
|
||||
// console.log(file.name.split('.')[0])
|
||||
// 创建文件读取器
|
||||
const reader = new FileReader()
|
||||
|
||||
// 定义读取完成事件的回调函数
|
||||
reader.onload = event => {
|
||||
// 读取完成后的文本内容
|
||||
const fileContent = event.target.result.split('\n')
|
||||
const keywords = Array.from(fileContent, f => f.trim()).filter(
|
||||
f => f
|
||||
)
|
||||
// 打印文件内容
|
||||
// console.log(keywords)
|
||||
|
||||
mutable_prompt.value = keywords.join('\n')
|
||||
|
||||
inp.remove()
|
||||
}
|
||||
|
||||
// 以文本方式读取文件
|
||||
reader.readAsText(file)
|
||||
})
|
||||
})
|
||||
|
||||
widget.div.appendChild(btn)
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
if (this.onResize) {
|
||||
this.onResize(this.size)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'RandomPrompt') {
|
||||
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.prompt.PromptSlide',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'PromptSlide') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const prompt_keyword = this.widgets.filter(
|
||||
w => w.name == 'prompt_keyword'
|
||||
)[0]
|
||||
// console.log('PromptSlide nodeData', prompt_keyword)
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'upload',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1])
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Upload Keywords'
|
||||
|
||||
btn.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid; height: 30px;min-width: 122px;
|
||||
`
|
||||
|
||||
const select = document.createElement('select')
|
||||
select.style = `display:none;cursor: pointer;
|
||||
font-weight: 300;
|
||||
margin: 2px;
|
||||
color: var(--descrip-text);
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid; height: 30px;min-width: 100px;
|
||||
`
|
||||
widget.select = select
|
||||
|
||||
// const btn=document.createElement('button');
|
||||
// btn.innerText='Upload'
|
||||
btn.addEventListener('click', () => {
|
||||
let inp = document.createElement('input')
|
||||
inp.type = 'file'
|
||||
inp.accept = '.txt'
|
||||
inp.click()
|
||||
inp.addEventListener('change', event => {
|
||||
// 获取选择的文件
|
||||
const file = event.target.files[0]
|
||||
this.title = file.name.split('.')[0]
|
||||
|
||||
// console.log(file.name.split('.')[0])
|
||||
// 创建文件读取器
|
||||
const reader = new FileReader()
|
||||
|
||||
// 定义读取完成事件的回调函数
|
||||
reader.onload = event => {
|
||||
// 读取完成后的文本内容
|
||||
const fileContent = event.target.result.split('\n')
|
||||
const keywords = Array.from(fileContent, f => f.trim()).filter(
|
||||
f => f
|
||||
)
|
||||
// 打印文件内容
|
||||
// console.log(keywords)
|
||||
|
||||
widget.value = JSON.stringify(keywords)
|
||||
|
||||
// let ks = getLocalData(`_mixlab_PromptSlide`)
|
||||
// ks[this.id] = keywords
|
||||
// setLocalDataOfWin(`_mixlab_PromptSlide`, ks)
|
||||
|
||||
createSelect(select, keywords, prompt_keyword)
|
||||
|
||||
inp.remove()
|
||||
}
|
||||
|
||||
// 以文本方式读取文件
|
||||
reader.readAsText(file)
|
||||
})
|
||||
})
|
||||
|
||||
widget.div.appendChild(btn)
|
||||
widget.div.appendChild(select)
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
if (this.onResize) {
|
||||
this.onResize(this.size)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'PromptSlide') {
|
||||
try {
|
||||
let prompt = node.widgets.filter(w => w.name === 'prompt_keyword')[0]
|
||||
// let ks = getLocalData(`_mixlab_PromptSlide`)
|
||||
let uploadWidget = node.widgets.filter(w => w.name == 'upload')[0]
|
||||
// console.log('##widget', uploadWidget.value)
|
||||
let keywords = JSON.parse(uploadWidget.value)
|
||||
// console.log('keywords',keywords)
|
||||
let widget = node.widgets.filter(w => w.select)[0]
|
||||
if (keywords && keywords[0]) {
|
||||
widget.select.style.display = 'block'
|
||||
createSelect(widget.select, keywords, prompt)
|
||||
}
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const _createResult = async (node, widget, message) => {
|
||||
widget.div.innerHTML = ``
|
||||
|
||||
const width = node.size[0] * 0.5 - 12
|
||||
|
||||
let height_add = 0
|
||||
|
||||
for (let index = 0; index < message._images.length; index++) {
|
||||
const imgs = message._images[index]
|
||||
|
||||
for (const img of imgs) {
|
||||
let url = api.apiURL(
|
||||
`/view?filename=${encodeURIComponent(img.filename)}&type=${
|
||||
img.type
|
||||
}&subfolder=${
|
||||
img.subfolder
|
||||
}${app.getPreviewFormatParam()}${app.getRandParam()}`
|
||||
)
|
||||
|
||||
let image = await createImage(url)
|
||||
|
||||
// 创建card
|
||||
let div = document.createElement('div')
|
||||
div.className = 'card'
|
||||
div.draggable = true
|
||||
|
||||
div.ondragend = async event => {
|
||||
console.log('拖动停止')
|
||||
let url = div.querySelector('img').src
|
||||
|
||||
let blob = await fetchImage(url)
|
||||
|
||||
let imageNode = null
|
||||
// No image node selected: add a new one
|
||||
if (!imageNode) {
|
||||
const newNode = LiteGraph.createNode('LoadImage')
|
||||
newNode.pos = [...app.canvas.graph_mouse]
|
||||
imageNode = app.graph.add(newNode)
|
||||
app.graph.change()
|
||||
}
|
||||
|
||||
// const blob = item.getAsFile();
|
||||
imageNode.pasteFile(blob)
|
||||
}
|
||||
|
||||
div.setAttribute('data-scale', image.naturalHeight / image.naturalWidth)
|
||||
|
||||
let h = (image.naturalHeight * width) / image.naturalWidth
|
||||
if (index % 2 === 0) height_add += h
|
||||
div.style = `width: ${width}px;height:${h}px;position: relative;margin: 4px;`
|
||||
|
||||
div.innerHTML = `<a href="${url}"
|
||||
data-pswp-width="${image.naturalWidth}"
|
||||
data-pswp-height="${image.naturalHeight}"
|
||||
target="_blank">
|
||||
<img src="${url}" style='width: 100%' alt="${message.prompts[index]}"/>
|
||||
</a>
|
||||
<p style="position: absolute;
|
||||
bottom: 0;
|
||||
left: 0;
|
||||
opacity: 0.6;
|
||||
background-color: var(--comfy-input-bg);
|
||||
color: var(--descrip-text);
|
||||
margin: 0;
|
||||
font-size: 12px;
|
||||
padding: 5px;
|
||||
text-align: left;">${message.prompts[index]}</p>`
|
||||
widget.div.appendChild(div)
|
||||
}
|
||||
}
|
||||
|
||||
node.size[1] = 98 + height_add
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.prompt.PromptImage',
|
||||
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'PromptImage') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
console.log('#orig_nodeCreated', this)
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'result',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(this.div.style, {
|
||||
...get_position_style(ctx, widget_width, y, node.size[1]),
|
||||
flexWrap: 'wrap',
|
||||
justifyContent: 'space-between',
|
||||
// outline: '1px solid red',
|
||||
paddingLeft: '0px',
|
||||
width: widget_width + 'px'
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
widget.div.className = 'prompt_image_output'
|
||||
|
||||
document.body.appendChild(widget.div)
|
||||
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
initLightBox()
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
const onResize = this.onResize
|
||||
this.onResize = function () {
|
||||
// 缩放发生
|
||||
// console.log('##缩放发生', this.size)
|
||||
let w = this.size[0] * 0.5 - 12
|
||||
Array.from(widget.div.querySelectorAll('.card'), card => {
|
||||
card.style.width = `${w}px`
|
||||
card.style.height = `${
|
||||
w * parseFloat(card.getAttribute('data-scale'))
|
||||
}px`
|
||||
})
|
||||
return onResize?.apply(this, arguments)
|
||||
}
|
||||
|
||||
// this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = async function (message) {
|
||||
onExecuted?.apply(this, arguments)
|
||||
console.log('#PromptImage', message.prompts, message._images)
|
||||
// window._mixlab_app_json = message.json
|
||||
try {
|
||||
let widget = this.widgets.filter(w => w.name === 'result')[0]
|
||||
widget.value = message
|
||||
_createResult(this, widget, { ...message })
|
||||
} catch (error) {
|
||||
console.log(error)
|
||||
}
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'PromptImage') {
|
||||
// await sleep(0)
|
||||
let widget = node.widgets.filter(w => w.name === 'result')[0]
|
||||
console.log('widget.value', widget.value)
|
||||
|
||||
initLightBox()
|
||||
|
||||
let cards = widget.div.querySelectorAll('.card')
|
||||
if (cards.length == 0) node.size = [280, 120]
|
||||
if(widget.value) _createResult(node, widget, widget.value)
|
||||
}
|
||||
}
|
||||
})
|
||||
@@ -0,0 +1,203 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
import { api } from '../../../scripts/api.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
|
||||
function get_position_style (ctx, widget_width, y, node_height) {
|
||||
const MARGIN = 14 // the margin around the html element
|
||||
|
||||
/* Create a transform that deals with all the scrolling and zooming */
|
||||
const elRect = ctx.canvas.getBoundingClientRect()
|
||||
const transform = new DOMMatrix()
|
||||
.scaleSelf(
|
||||
elRect.width / ctx.canvas.width,
|
||||
elRect.height / ctx.canvas.height
|
||||
)
|
||||
.multiplySelf(ctx.getTransform())
|
||||
.translateSelf(MARGIN, MARGIN + y)
|
||||
|
||||
return {
|
||||
transformOrigin: '0 0',
|
||||
transform: transform,
|
||||
left: `0`,
|
||||
top: `0`,
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
maxWidth: `${widget_width - MARGIN * 2}px`,
|
||||
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
|
||||
width: `${widget_width - MARGIN * 2}px`,
|
||||
// height: `${node_height * 0.3 - MARGIN * 2}px`,
|
||||
// background: '#EEEEEE',
|
||||
// outline: '1px solid red',
|
||||
display: 'flex',
|
||||
flexDirection: 'column',
|
||||
// alignItems: 'center',
|
||||
justifyContent: 'space-around'
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.3D.SaveTripoSRMesh',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'SaveTripoSRMesh') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'preview',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, 88, node.size[1])
|
||||
)
|
||||
}
|
||||
// value: [],
|
||||
// async serializeValue (nodeId, widgetIndex) {
|
||||
// return widget.value
|
||||
// }
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
widget.div.style.width = `120px`
|
||||
|
||||
document.body.appendChild(widget.div)
|
||||
|
||||
// preview.style = `margin-top: 12px;display: flex;
|
||||
// justify-content: center;
|
||||
// align-items: center;background-repeat: no-repeat;background-size: contain;`
|
||||
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onResize = this.onResize
|
||||
this.onResize = () => {
|
||||
widget.div.style.width = `${this.size[0]}px`
|
||||
widget.div.style.height = `${this.size[1] - 112}px`
|
||||
let mvs = widget.div.querySelectorAll('model-viewer')
|
||||
for (const m of mvs) {
|
||||
m.style.height = `${Math.round(
|
||||
(this.size[1] - 112) / mvs.length
|
||||
)}px`
|
||||
// console.log(m.style.height)
|
||||
}
|
||||
// console.log('resize', this.size)
|
||||
return onResize?.apply(this, arguments)
|
||||
}
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
if (this.onResize) {
|
||||
this.onResize(this.size)
|
||||
}
|
||||
// this.isVirtualNode = true
|
||||
this.serialize_widgets = false //需要保存参数
|
||||
}
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
const r = onExecuted?.apply?.(this, arguments)
|
||||
|
||||
let widget = this.widgets.filter(d => d.name == 'preview')[0]
|
||||
console.log('Test', widget, message)
|
||||
|
||||
let meshes = message.mesh
|
||||
widget.div.innerHTML = ''
|
||||
|
||||
for (const mesh of meshes) {
|
||||
if (mesh) {
|
||||
const { filename, subfolder, type } = mesh
|
||||
const fileURL = api.apiURL(
|
||||
`/view?filename=${encodeURIComponent(
|
||||
filename
|
||||
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
|
||||
)
|
||||
|
||||
let modelViewer = document.createElement('div')
|
||||
modelViewer.innerHTML = `<model-viewer src="${fileURL}"
|
||||
min-field-of-view="0deg" max-field-of-view="180deg"
|
||||
shadow-intensity="1"
|
||||
camera-controls
|
||||
touch-action="pan-y"
|
||||
style="width:100%;margin:4px;min-height:88px"
|
||||
>
|
||||
|
||||
<div class="controls">
|
||||
|
||||
<div><button class="export" style="
|
||||
background-color: var(--comfy-input-bg);
|
||||
border-radius: 8px;
|
||||
border-color: var(--border-color);
|
||||
border-style: solid;
|
||||
color: var(--descrip-text);cursor: pointer;">Export GLB</button></div>
|
||||
|
||||
</div></model-viewer>`
|
||||
widget.div.appendChild(modelViewer)
|
||||
let modelViewerVariants= modelViewer
|
||||
.querySelector('model-viewer');
|
||||
|
||||
modelViewer
|
||||
.querySelector('.export')
|
||||
.addEventListener('click', async e => {
|
||||
e.preventDefault()
|
||||
const glTF = await modelViewerVariants.exportScene()
|
||||
const file = new File([glTF], filename)
|
||||
const link = document.createElement('a')
|
||||
link.download = file.name
|
||||
link.href = URL.createObjectURL(file)
|
||||
link.click()
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// widget.value = [meshes]
|
||||
|
||||
this.onResize?.(this.size)
|
||||
|
||||
return r
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
const sleep = (t = 1000) => {
|
||||
return new Promise((res, rej) => {
|
||||
setTimeout(() => res(1), t)
|
||||
})
|
||||
}
|
||||
// if (node.type === 'SaveTripoSRMesh') {
|
||||
// await sleep(0)
|
||||
// let widget = node.widgets.filter(w => w.name === 'preview')[0]
|
||||
// widget.div.innerHTML = ''
|
||||
|
||||
// for (const mesh of widget.value) {
|
||||
// if (mesh) {
|
||||
// const { filename, subfolder, type } = mesh
|
||||
// const fileURL = api.apiURL(
|
||||
// `/view?filename=${encodeURIComponent(
|
||||
// filename
|
||||
// )}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
|
||||
// )
|
||||
|
||||
// let modelViewer = document.createElement('div')
|
||||
// modelViewer.innerHTML = `<model-viewer src="${fileURL}"
|
||||
// min-field-of-view="0deg" max-field-of-view="180deg"
|
||||
// shadow-intensity="1"
|
||||
// camera-controls
|
||||
// touch-action="pan-y">
|
||||
|
||||
// <div class="controls">
|
||||
|
||||
// <div><button class="export">Export GLB</button></div>
|
||||
|
||||
// </div></model-viewer>`
|
||||
// widget.div.appendChild(modelViewer)
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
}
|
||||
})
|
||||
@@ -0,0 +1,329 @@
|
||||
const smart_connect_config_input = [
|
||||
{
|
||||
node_type: 'CLIPTextEncode',
|
||||
node_widget_name: 'text',
|
||||
inputNodeName: 'RandomPrompt',
|
||||
inputNode_output_name: 'STRING'
|
||||
},
|
||||
{
|
||||
node_type: 'CLIPTextEncode',
|
||||
node_widget_name: 'text',
|
||||
inputNodeName: 'EmbeddingPrompt',
|
||||
inputNode_output_name: 'STRING'
|
||||
},
|
||||
{
|
||||
node_type: 'CLIPTextEncode',
|
||||
node_widget_name: 'text',
|
||||
inputNodeName: 'ChinesePrompt_Mix',
|
||||
inputNode_output_name: 'prompt'
|
||||
},
|
||||
{
|
||||
node_type: 'CheckpointLoaderSimple',
|
||||
node_widget_name: 'ckpt_name',
|
||||
inputNodeName: 'CkptNames_',
|
||||
inputNode_output_name: 'ckpt_names'
|
||||
},
|
||||
{
|
||||
node_type: 'KSampler',
|
||||
node_widget_name: 'sampler_name',
|
||||
inputNodeName: 'SamplerNames_',
|
||||
inputNode_output_name: 'sampler_names'
|
||||
},
|
||||
{
|
||||
node_type: 'LoraLoaderModelOnly',
|
||||
node_widget_name: 'lora_name',
|
||||
inputNodeName: 'LoraNames_',
|
||||
inputNode_output_name: 'lora_names'
|
||||
},
|
||||
{
|
||||
node_type: 'LoadLoRA',
|
||||
node_widget_name: 'lora_name',
|
||||
inputNodeName: 'LoraNames_',
|
||||
inputNode_output_name: 'lora_names'
|
||||
},
|
||||
{
|
||||
node_type: 'Moondream',
|
||||
node_widget_name: 'image',
|
||||
inputNodeName: 'LoadImage',
|
||||
inputNode_output_name: 'IMAGE'
|
||||
},
|
||||
{
|
||||
node_type: 'TripoSRSampler_',
|
||||
node_widget_name: 'image',
|
||||
inputNodeName: 'LoadImagesToBatch',
|
||||
inputNode_output_name: 'IMAGE'
|
||||
},
|
||||
{
|
||||
node_type: 'TripoSRSampler_',
|
||||
node_widget_name: 'mask',
|
||||
inputNodeName: 'RembgNode_Mix',
|
||||
inputNode_output_name: 'masks'
|
||||
}
|
||||
]
|
||||
|
||||
const smart_connect_config_output = [
|
||||
{
|
||||
node_type: 'LoadImage',
|
||||
node_output_name: 'IMAGE',
|
||||
outputNodeName: 'ClipInterrogator',
|
||||
outputNode_input_name: 'image'
|
||||
},
|
||||
{
|
||||
node_type: 'VAEDecode',
|
||||
node_output_name: 'IMAGE',
|
||||
outputNodeName: 'PromptImage',
|
||||
outputNode_input_name: 'images'
|
||||
},
|
||||
{
|
||||
node_type: 'VAEDecode',
|
||||
node_output_name: 'IMAGE',
|
||||
outputNodeName: 'PreviewImage',
|
||||
outputNode_input_name: 'images'
|
||||
},
|
||||
{
|
||||
node_type: 'VAEDecode',
|
||||
node_output_name: 'IMAGE',
|
||||
outputNodeName: 'SaveImage',
|
||||
outputNode_input_name: 'images'
|
||||
},
|
||||
{
|
||||
node_type: 'VAEDecode',
|
||||
node_output_name: 'IMAGE',
|
||||
outputNodeName: 'AppInfo',
|
||||
outputNode_input_name: 'IMAGE'
|
||||
},
|
||||
{
|
||||
node_type: 'VAEDecode',
|
||||
node_output_name: 'IMAGE',
|
||||
outputNodeName: 'SaveImageAndMetadata_',
|
||||
outputNode_input_name: 'images'
|
||||
},
|
||||
{
|
||||
node_type: 'Moondream',
|
||||
node_output_name: 'STRING',
|
||||
outputNodeName: 'ShowTextForGPT',
|
||||
outputNode_input_name: 'text'
|
||||
}
|
||||
]
|
||||
|
||||
// import {
|
||||
// convertToInput,
|
||||
// getConfig,
|
||||
// isConvertableWidget
|
||||
// } from '../../../extensions/core/widgetInputs.js'
|
||||
|
||||
const CONVERTED_TYPE = 'converted-widget'
|
||||
const GET_CONFIG = Symbol()
|
||||
|
||||
function getConfig (widgetName) {
|
||||
const { nodeData } = this.constructor
|
||||
return (
|
||||
nodeData?.input?.required[widgetName] ??
|
||||
nodeData?.input?.optional?.[widgetName]
|
||||
)
|
||||
}
|
||||
|
||||
function hideWidget (node, widget, suffix = '') {
|
||||
widget.origType = widget.type
|
||||
widget.origComputeSize = widget.computeSize
|
||||
widget.origSerializeValue = widget.serializeValue
|
||||
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
|
||||
widget.type = CONVERTED_TYPE + suffix
|
||||
widget.serializeValue = () => {
|
||||
// Prevent serializing the widget if we have no input linked
|
||||
if (!node.inputs) {
|
||||
return undefined
|
||||
}
|
||||
let node_input = node.inputs.find(i => i.widget?.name === widget.name)
|
||||
|
||||
if (!node_input || !node_input.link) {
|
||||
return undefined
|
||||
}
|
||||
return widget.origSerializeValue
|
||||
? widget.origSerializeValue()
|
||||
: widget.value
|
||||
}
|
||||
|
||||
// Hide any linked widgets, e.g. seed+seedControl
|
||||
if (widget.linkedWidgets) {
|
||||
for (const w of widget.linkedWidgets) {
|
||||
hideWidget(node, w, ':' + widget.name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function convertToInput (node, widget, config) {
|
||||
hideWidget(node, widget)
|
||||
|
||||
const type = config[0]
|
||||
|
||||
// Add input and store widget config for creating on primitive node
|
||||
const sz = node.size
|
||||
node.addInput(widget.name, type, {
|
||||
widget: { name: widget.name, [GET_CONFIG]: () => config }
|
||||
})
|
||||
|
||||
for (const widget of node.widgets) {
|
||||
widget.last_y += LiteGraph.NODE_SLOT_HEIGHT
|
||||
}
|
||||
|
||||
// Restore original size but grow if needed
|
||||
node.setSize([Math.max(sz[0], node.size[0]), Math.max(sz[1], node.size[1])])
|
||||
}
|
||||
|
||||
export function smart_init () {
|
||||
LGraphCanvas.prototype._createNodeForInput = function (
|
||||
node,
|
||||
widget,
|
||||
inputNodeName,
|
||||
inputNode_slot
|
||||
) {
|
||||
// console.log(node.pos)
|
||||
|
||||
// var widget = node.widgets.filter(w => w.name === node_widget_name)[0]
|
||||
if (widget) {
|
||||
// 如果有存在的,没有连线输出的,自动连,不新建
|
||||
let input_node = null
|
||||
|
||||
Array.from(app.graph.findNodesByType(inputNodeName), n => {
|
||||
var links = n.outputs.filter(o => o.name === inputNode_slot)[0].links
|
||||
// console.log(links)
|
||||
if (!links || links?.length === 0) input_node = n
|
||||
})
|
||||
// 新建
|
||||
if (!input_node) {
|
||||
input_node = LiteGraph.createNode(inputNodeName)
|
||||
input_node.pos = [node.pos[0] - node.size[0] - 24, node.pos[1] - 48]
|
||||
app.canvas.graph.add(input_node, false)
|
||||
} else {
|
||||
input_node.pos = [node.pos[0] - node.size[0] - 24, node.pos[1] - 48]
|
||||
}
|
||||
|
||||
const config = getConfig.call(node, widget.name) ?? [
|
||||
widget.type,
|
||||
widget.options || {}
|
||||
]
|
||||
let node_slotType = config[0]
|
||||
// 如果input没有,则创建
|
||||
if (
|
||||
!node.inputs?.filter(inp => inp.name === widget.name)[0] ||
|
||||
!node.inputs
|
||||
)
|
||||
convertToInput(node, widget, config)
|
||||
input_node.connectByType(inputNode_slot, node, node_slotType)
|
||||
}
|
||||
}
|
||||
|
||||
LGraphCanvas.prototype._createNodeForOutput = function (
|
||||
node,
|
||||
widget,
|
||||
outputNodeName,
|
||||
outputNode_slot
|
||||
) {
|
||||
if (widget) {
|
||||
let output_node
|
||||
Array.from(app.graph.findNodesByType(outputNodeName), n => {
|
||||
var links = n.inputs.filter(o => o.name === outputNode_slot)[0].links
|
||||
// console.log(links)
|
||||
if (!links || links?.length === 0) output_node = n
|
||||
})
|
||||
console.log('output_node', output_node, widget.name)
|
||||
|
||||
if (!output_node) {
|
||||
// 新建
|
||||
output_node = LiteGraph.createNode(outputNodeName)
|
||||
output_node.pos = [node.pos[0] + node.size[0] + 24, node.pos[1] - 48]
|
||||
app.canvas.graph.add(output_node, false)
|
||||
} else {
|
||||
output_node.pos = [node.pos[0] + node.size[0] + 24, node.pos[1] - 48]
|
||||
}
|
||||
|
||||
const config = getConfig.call(node, widget.name) ?? [
|
||||
widget.type,
|
||||
widget.options || {}
|
||||
]
|
||||
let node_slotType = config[0]
|
||||
console.log(node_slotType, output_node, outputNode_slot)
|
||||
let type = output_node.inputs.filter(
|
||||
inp => inp.name == outputNode_slot
|
||||
)[0].type
|
||||
node.connectByType(node_slotType, output_node, type)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export function addSmartMenu (options, node) {
|
||||
let sopts = []
|
||||
|
||||
for (const sc of smart_connect_config_input) {
|
||||
// 有智能推荐,则出现
|
||||
if (node.type === sc.node_type) {
|
||||
// console.log('smart',node)
|
||||
// 则出现 randomPrompt
|
||||
// CLIPTextEncode 的widget ,name== 'text'
|
||||
let node_widget_name = sc.node_widget_name
|
||||
let widget = node.widgets.filter(w => w.name === node_widget_name)[0]
|
||||
if (!widget) {
|
||||
// 控件没有,则查找inputs
|
||||
widget = node.inputs.filter(w => w.name === node_widget_name)[0]
|
||||
}
|
||||
|
||||
let isLinkNull = true
|
||||
// 如果input里已经有,但是link为空
|
||||
if (node.inputs?.filter(inp => inp.name === node_widget_name)[0]) {
|
||||
isLinkNull =
|
||||
node.inputs.filter(inp => inp.name === node_widget_name)[0].link ===
|
||||
null
|
||||
}
|
||||
|
||||
if (widget && isLinkNull) {
|
||||
sopts.push({
|
||||
content: sc.inputNodeName.split('_')[0] + '➡️',
|
||||
callback: () => {
|
||||
LGraphCanvas.prototype._createNodeForInput(
|
||||
node, //当前node
|
||||
widget, //当前node里需要自动连线的widget
|
||||
sc.inputNodeName, //作为input的node type
|
||||
sc.inputNode_output_name // 作为input的node的outputs的name. the input slot type of the target node
|
||||
)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for (const sc of smart_connect_config_output) {
|
||||
if (node.type === sc.node_type) {
|
||||
let node_output_name = sc.node_output_name
|
||||
const widget = node.outputs.filter(w => w.name === node_output_name)[0]
|
||||
|
||||
let isLinkNull = true
|
||||
// 如果output里 link为空
|
||||
if (node.outputs?.filter(inp => inp.name === node_output_name)[0]) {
|
||||
isLinkNull =
|
||||
node.outputs.filter(inp => inp.name === node_output_name)[0].links
|
||||
?.length === 0
|
||||
if (!node.outputs.filter(inp => inp.name === node_output_name)[0].links)
|
||||
isLinkNull = true
|
||||
}
|
||||
|
||||
if (widget && isLinkNull) {
|
||||
sopts.push({
|
||||
content: '➡️' + sc.outputNodeName.split('_')[0],
|
||||
callback: () => {
|
||||
LGraphCanvas.prototype._createNodeForOutput(
|
||||
node, //当前node
|
||||
widget, //当前node里需要自动连线的widget
|
||||
sc.outputNodeName, //作为input的node type
|
||||
sc.outputNode_input_name // 作为input的node的outputs的name. the input slot type of the target node
|
||||
)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (sopts.length > 0) options = [...sopts, null, ...options]
|
||||
|
||||
return options
|
||||
}
|
||||