Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
67afdc8204 | ||
|
|
8a0f2fc412 | ||
|
|
22145befb3 | ||
|
|
8730ffd140 | ||
|
|
32931f09a7 | ||
|
|
86873e7bda | ||
|
|
375f3b77e0 | ||
|
|
450b1ce4ce | ||
|
|
457b3a81e8 | ||
|
|
271685698b | ||
|
|
859af7e7b4 | ||
|
|
8b522f121d | ||
|
|
5b469409bc | ||
|
|
80e1261b1c | ||
|
|
005c57839c | ||
|
|
cf15032ab6 | ||
|
|
ca93381de8 | ||
|
|
58e077a743 | ||
|
|
4de1ab3b66 | ||
|
|
595e0738a9 | ||
|
|
7535cd0dfd | ||
|
|
960862223b | ||
|
|
54d080bf6a | ||
|
|
625efbfa2f | ||
|
|
5618a748c1 | ||
|
|
130c1b5796 | ||
|
|
3cf9ab4e63 | ||
|
|
ff5e3a34fc | ||
|
|
ec4ca6717f | ||
|
|
b82bb48948 | ||
|
|
d08eedabd3 | ||
|
|
8ba21d0b44 | ||
|
|
337a03bb19 | ||
|
|
aef19b8772 | ||
|
|
d60b61d575 | ||
|
|
a3f051f0c3 | ||
|
|
8ca6ace667 | ||
|
|
81c510c06e | ||
|
|
7601371923 | ||
|
|
b11c634872 | ||
|
|
7c470c67d6 | ||
|
|
b5865efd16 | ||
|
|
5ec3b5ef86 | ||
|
|
b5e31ef12a | ||
|
|
21b3c15040 | ||
|
|
5dfcbcf51d | ||
|
|
070001b36b | ||
|
|
6b4c89adc4 | ||
|
|
32ad26f0e1 | ||
|
|
d9c2072a2d | ||
|
|
e94405e610 | ||
|
|
03d5a4cf12 | ||
|
|
ad43ed3154 | ||
|
|
5cc1f8535a | ||
|
|
9f42ead9db | ||
|
|
23d9c365bd | ||
|
|
7a17ad010d | ||
|
|
3b38a5ae60 | ||
|
|
3b224fbccd | ||
|
|
00c69fc816 | ||
|
|
d39b5e13e2 | ||
|
|
e900c2ca9c | ||
|
|
9947b8be70 | ||
|
|
bc4cec287e | ||
|
|
84a4348bc5 | ||
|
|
29cca677c5 | ||
|
|
76b5896f08 | ||
|
|
be70a8671a | ||
|
|
0f6cc6958a | ||
|
|
499ed4eafc | ||
|
|
125b3a4905 | ||
|
|
a1b402b4a9 | ||
|
|
755723ce64 | ||
|
|
737feef400 | ||
|
|
8ecc929cd4 | ||
|
|
58e9d594cd | ||
|
|
b6deb5f515 | ||
|
|
1148fbe5fe | ||
|
|
5026d36489 | ||
|
|
fe0cad5981 | ||
|
|
3f37ba9491 | ||
|
|
a72a4dc40f | ||
|
|
11794f7d71 | ||
|
|
74226d8b2a | ||
|
|
f550ce83a8 | ||
|
|
34882bca10 | ||
|
|
43b94be806 | ||
|
|
0349d81694 | ||
|
|
93254a4c07 | ||
|
|
8485447325 | ||
|
|
717092a3ce | ||
|
|
14a1121860 | ||
|
|
8c1eec2858 | ||
|
|
b6bb4a3055 | ||
|
|
6873492872 | ||
|
|
2d71b3e647 | ||
|
|
e7320ec0c4 | ||
|
|
560be6aee7 | ||
|
|
54614079ca | ||
|
|
b0cd0bcb5b | ||
|
|
e46f8a45d0 | ||
|
|
1616dd6602 | ||
|
|
282eedfea6 | ||
|
|
17b163e234 | ||
|
|
501d97bb5c | ||
|
|
de92038f88 | ||
|
|
530333d72d | ||
|
|
041f49540c | ||
|
|
71c7865d2d | ||
|
|
c7fbf05970 | ||
|
|
2986a01469 | ||
|
|
fa7c5d8b4d | ||
|
|
1d8db7510b | ||
|
|
7ef0612ce7 | ||
|
|
7ff4790493 | ||
|
|
640ef31625 | ||
|
|
e4ac947d96 | ||
|
|
d287e28e5c | ||
|
|
f33c17f762 | ||
|
|
e07b8cc7bf | ||
|
|
6abe07bb79 | ||
|
|
7fbd03bda7 | ||
|
|
9cc2ac02da | ||
|
|
2f2a3035a2 | ||
|
|
5d8f0a3b0a | ||
|
|
8aadd72494 | ||
|
|
fceec754a4 | ||
|
|
419b7c985c | ||
|
|
ce62fc73da | ||
|
|
4f31641da3 | ||
|
|
deec62ab76 | ||
|
|
5c8cdb58c7 | ||
|
|
d820842e39 | ||
|
|
2f78a523b3 | ||
|
|
b2a8666423 | ||
|
|
2c02a471d0 | ||
|
|
ea521e0303 | ||
|
|
9b5daac023 | ||
|
|
0de83f88dc | ||
|
|
342ce8ccad | ||
|
|
b7881d84b1 | ||
|
|
0f5ad38384 | ||
|
|
a3f487c822 | ||
|
|
d0f496adc1 | ||
|
|
665861ff35 | ||
|
|
61568e021c | ||
|
|
a2edc37d89 | ||
|
|
1c4cb43f7b | ||
|
|
66143b0e20 | ||
|
|
f0da5e25c9 | ||
|
|
ebf25b585f | ||
|
|
fb8968d438 | ||
|
|
15cfeedf7b | ||
|
|
50ae13a993 | ||
|
|
aedf917067 | ||
|
|
69ac5e52a0 | ||
|
|
368f7e508d | ||
|
|
44f0676323 | ||
|
|
98273b37f2 | ||
|
|
eff718c13f | ||
|
|
615a2abcfe | ||
|
|
2b4b38ce03 | ||
|
|
dbbd2ffef3 | ||
|
|
b1a875b151 | ||
|
|
6a39ea1188 | ||
|
|
69aac075e8 | ||
|
|
e1dc9250b9 | ||
|
|
8b9c577f55 | ||
|
|
6d8c266b04 | ||
|
|
9292f22862 | ||
|
|
10e9629ca3 | ||
|
|
4f694195a2 | ||
|
|
ff6c0f0e39 | ||
|
|
a6e8783605 | ||
|
|
3e84b8cd77 | ||
|
|
9e70cc0090 | ||
|
|
7dddd2d6e5 | ||
|
|
63a1ca5ec6 | ||
|
|
0104f7f6a9 | ||
|
|
6b1f5cbf69 | ||
|
|
f888e3d75d | ||
|
|
16631d21d9 | ||
|
|
ccb4ba08fc | ||
|
|
0daf114fe8 | ||
|
|
4e9c9c897c | ||
|
|
31fde1ae34 | ||
|
|
aadbb0b389 | ||
|
|
52a8e7faf3 | ||
|
|
3893873085 | ||
|
|
037080ac39 | ||
|
|
4738313b64 | ||
|
|
e842c3bd06 | ||
|
|
fa73da5a00 | ||
|
|
ffe26e8571 | ||
|
|
123917da9a | ||
|
|
d4fb74df19 | ||
|
|
3175716585 | ||
|
|
bc19ed63fc | ||
|
|
6cfb0585da | ||
|
|
daf10e96f8 | ||
|
|
94882b7da7 | ||
|
|
0b64d4c297 | ||
|
|
7bacb16c89 | ||
|
|
45d5c08bbb | ||
|
|
7866b053a3 | ||
|
|
68c96e0a2e | ||
|
|
862cde4bcd | ||
|
|
ca1fa507d0 | ||
|
|
b05af806d7 | ||
|
|
1bf3b2d7a4 | ||
|
|
756f60a01a | ||
|
|
06ed9f33a3 | ||
|
|
9ad997ccab | ||
|
|
bd149ca8de | ||
|
|
01ab8f4ac2 | ||
|
|
48d06c4485 | ||
|
|
05b9182196 | ||
|
|
991a62fc51 | ||
|
|
46f0126339 | ||
|
|
f961596092 | ||
|
|
a157b55835 | ||
|
|
2b160cc789 | ||
|
|
4d9f791cf7 | ||
|
|
3515268de5 | ||
|
|
a80845f641 | ||
|
|
39fa6ef37a | ||
|
|
7a65c2f5d7 | ||
|
|
65937a75eb | ||
|
|
17e022a7aa | ||
|
|
138fb519e7 | ||
|
|
fec464b015 | ||
|
|
9b9c1b3cc7 | ||
|
|
17379c5156 | ||
|
|
4eb433281c | ||
|
|
e17a81d335 | ||
|
|
68a286ae4a | ||
|
|
60fb13e068 | ||
|
|
f52dd53ed0 | ||
|
|
2aae0affd2 | ||
|
|
765462549c | ||
|
|
bf21bbfd93 | ||
|
|
3700d010ba | ||
|
|
286e6ba336 | ||
|
|
3a8fcbbcb9 | ||
|
|
570ea601ce | ||
|
|
be8306b17a | ||
|
|
5ec744927b | ||
|
|
5e0cc2ea71 | ||
|
|
a1125b20bc | ||
|
|
db14b955a5 | ||
|
|
6a61c1cf89 | ||
|
|
c974a60749 | ||
|
|
a844119335 | ||
|
|
c22e434e38 | ||
|
|
2c2751a762 | ||
|
|
4f283c3b55 | ||
|
|
d68f0804ed | ||
|
|
9edc20e810 | ||
|
|
615289f00e | ||
|
|
d0f269807e | ||
|
|
f5efee7f23 | ||
|
|
b26fcef6d7 | ||
|
|
54a7c22296 | ||
|
|
0ab09df6c4 | ||
|
|
aa57e309ba | ||
|
|
d56cbf572d | ||
|
|
d416ad21f0 | ||
|
|
a46d80b6be | ||
|
|
da57b55c03 | ||
|
|
ebcad2bb54 | ||
|
|
694673bc1c | ||
|
|
2461869aae | ||
|
|
bb9dc79325 | ||
|
|
3939e9d525 | ||
|
|
b36b68a648 | ||
|
|
5d0ad29657 | ||
|
|
25a4420b4f | ||
|
|
bed6ab1df1 | ||
|
|
ff8ba6b209 | ||
|
|
b0e892b083 | ||
|
|
523205b6b4 | ||
|
|
83bbe7b7f7 | ||
|
|
97519d816c | ||
|
|
483b858abe | ||
|
|
f28a3f3ed1 | ||
|
|
e94ece1b1d | ||
|
|
178e9402c9 | ||
|
|
ee25139e53 | ||
|
|
9c1806f71d | ||
|
|
20e360036f | ||
|
|
9d6e210921 | ||
|
|
8bc0caa057 | ||
|
|
07d9b1a225 | ||
|
|
7fb85eb987 | ||
|
|
acfdd7713c | ||
|
|
3c1ea86bc6 | ||
|
|
b8d31fde80 | ||
|
|
631f2f80c9 | ||
|
|
e76a8e634c | ||
|
|
2166920cf0 | ||
|
|
cf32e868d6 | ||
|
|
3c37489c0a | ||
|
|
976dffed60 | ||
|
|
876210a197 | ||
|
|
b869fee891 | ||
|
|
1c82506ab9 | ||
|
|
edb0e409df | ||
|
|
a32f850225 | ||
|
|
918bd85865 | ||
|
|
1be8fa596c | ||
|
|
46dd9f16fd | ||
|
|
3fe0b9ba40 | ||
|
|
5011099081 | ||
|
|
b75a247435 | ||
|
|
cc4997cd94 | ||
|
|
2a4f89dab0 | ||
|
|
e4b331cd93 | ||
|
|
9666ef733b | ||
|
|
b008fa162f | ||
|
|
b498cbd5f8 | ||
|
|
b44511b78d | ||
|
|
9d10c9f5a6 | ||
|
|
df2b4edc65 | ||
|
|
628499ad1c | ||
|
|
ede22dfa27 | ||
|
|
5aaaaffa2e | ||
|
|
aa30d9c495 | ||
|
|
5b7980facd | ||
|
|
c51d1fdea2 | ||
|
|
727f8b87e1 | ||
|
|
60784bf262 | ||
|
|
4cc0273a4c | ||
|
|
d5ec95ec0f | ||
|
|
523be189e3 | ||
|
|
ba7bd1f542 | ||
|
|
7d2f16595a | ||
|
|
82fd658894 | ||
|
|
a52f10255a | ||
|
|
da5b3de3eb | ||
|
|
122d00f8fc | ||
|
|
41715bb263 | ||
|
|
079b65332e | ||
|
|
88cf2a6688 | ||
|
|
cdbcb7f033 | ||
|
|
5c278e8d56 | ||
|
|
4e8daffcd2 | ||
|
|
944051f210 | ||
|
|
0922da0b66 | ||
|
|
e754b97b99 | ||
|
|
80ede24bb5 | ||
|
|
2f63c5c385 | ||
|
|
0adf673854 | ||
|
|
576c52746c | ||
|
|
7410d7c865 | ||
|
|
435558b778 | ||
|
|
701cb45770 | ||
|
|
6f67a49251 | ||
|
|
a9d985c666 | ||
|
|
f543f13668 | ||
|
|
42f6e81bef | ||
|
|
51ee274e40 | ||
|
|
609ccce401 | ||
|
|
c8331f8656 | ||
|
|
dea68c212b | ||
|
|
93580e635a | ||
|
|
b71e6c0900 | ||
|
|
c7bab3cc98 | ||
|
|
503f4a756b | ||
|
|
f641bc15de | ||
|
|
39abd72526 | ||
|
|
d2bf013dcb | ||
|
|
ad516a08b7 | ||
|
|
be4c62b923 | ||
|
|
8119bfd962 | ||
|
|
11054532d2 | ||
|
|
5e0caf6f6f | ||
|
|
7a2f9f95fc | ||
|
|
d5a72214a8 | ||
|
|
7f0766231a | ||
|
|
b5e41f2108 | ||
|
|
fe08a2270b | ||
|
|
775db58e91 | ||
|
|
f5d7f7f575 | ||
|
|
67a650c570 | ||
|
|
847a9c6c7d | ||
|
|
e0f45f51a6 | ||
|
|
5b1eb92c75 | ||
|
|
1657342edd | ||
|
|
c17a0ee889 | ||
|
|
d8f3aaf713 | ||
|
|
5cb59dd1d5 | ||
|
|
de2c27d7d1 | ||
|
|
df5fb224fb | ||
|
|
e9b18dbe48 | ||
|
|
9bbef76417 | ||
|
|
06ed579310 | ||
|
|
02695fe0df | ||
|
|
9b5b2399e1 | ||
|
|
4167733d39 | ||
|
|
6964e11f3d | ||
|
|
e9439f0dc0 | ||
|
|
f1ee79a9ae | ||
|
|
d0118ca742 | ||
|
|
d31c9076c3 | ||
|
|
ac25ebad3c | ||
|
|
9b05d46ff2 | ||
|
|
e514bc1d8a | ||
|
|
ffb0bf5de9 | ||
|
|
1897c3acfd | ||
|
|
7b1dc8ce62 | ||
|
|
d7ee354fe4 | ||
|
|
d4a443607f | ||
|
|
df918829dd | ||
|
|
ba701d1d59 | ||
|
|
12002acd93 | ||
|
|
5d1721d0c3 | ||
|
|
0d4e1ede0f | ||
|
|
a808e30a23 | ||
|
|
913722813e | ||
|
|
49374e012e | ||
|
|
2b69b4f33a | ||
|
|
8ba8c214a3 | ||
|
|
a060322a88 | ||
|
|
52508f0f35 | ||
|
|
a6c4158af3 | ||
|
|
d4a8f0d415 | ||
|
|
a8ca28ebca | ||
|
|
4608de5cbf | ||
|
|
8c14ffdafe | ||
|
|
5ab388d690 | ||
|
|
a78ba7dd35 | ||
|
|
81d97ca81f | ||
|
|
51dcc04be4 | ||
|
|
9bf1e808b2 | ||
|
|
a0e195d1c1 | ||
|
|
6bce70780f | ||
|
|
c1a85e3aa2 | ||
|
|
763cff49e4 | ||
|
|
daaa44eddf | ||
|
|
7f25ce6cd2 | ||
|
|
5d68d617f9 | ||
|
|
400ffebb49 | ||
|
|
99e05344bf | ||
|
|
6a1b5b8d69 | ||
|
|
24ee3ffea5 | ||
|
|
500eae510d | ||
|
|
2e79ec2504 | ||
|
|
33c7e9513f | ||
|
|
ef64db05e3 | ||
|
|
f1ab83f429 | ||
|
|
9e85e3a25c | ||
|
|
6cab7a18e3 | ||
|
|
6d6004ce0d | ||
|
|
625b295a9c | ||
|
|
4145471d24 | ||
|
|
a5e12ff375 | ||
|
|
5f60119ab9 | ||
|
|
bd0fee0bf3 | ||
|
|
2269e27952 | ||
|
|
1807037e54 | ||
|
|
fcf6dca1b9 | ||
|
|
e8a12b0e8d | ||
|
|
4670f22f2a | ||
|
|
686aef4409 | ||
|
|
91f2ebcf74 | ||
|
|
795efcb25f | ||
|
|
df32c09852 | ||
|
|
6a763b32c4 | ||
|
|
ad8351691d | ||
|
|
6be9170701 | ||
|
|
466234d517 | ||
|
|
1a82ad74fb | ||
|
|
2a855046c6 | ||
|
|
d27d18e3f0 | ||
|
|
c209b1ce6d | ||
|
|
567dfe9db5 | ||
|
|
119f706779 | ||
|
|
cb3ac02a0f | ||
|
|
1f8b19fe5e | ||
|
|
a5a118c98d | ||
|
|
d569a14597 | ||
|
|
e8007d575f | ||
|
|
cbcb01b770 | ||
|
|
a01a2ecc0d | ||
|
|
fa0e8c34e8 | ||
|
|
96d747acef | ||
|
|
598a54a2dc | ||
|
|
8a272004dd | ||
|
|
c68258304c | ||
|
|
56c8b64bd1 | ||
|
|
0b9a76454d | ||
|
|
ad3f695c18 | ||
|
|
195a2b514b | ||
|
|
85d3e6619f | ||
|
|
ee09de3f16 | ||
|
|
1eb1e1a5c1 | ||
|
|
f5219ab516 | ||
|
|
282a121379 | ||
|
|
8baedc78fa | ||
|
|
84f8cc92d4 |
+1
-1
@@ -1,3 +1,3 @@
|
|||||||
# These are supported funding model platforms
|
# These are supported funding model platforms
|
||||||
|
|
||||||
custom: ["https://afdian.net/a/yolain"]
|
custom: ["https://space.bilibili.com/1840885116"]
|
||||||
@@ -7,15 +7,36 @@ on:
|
|||||||
paths:
|
paths:
|
||||||
- "pyproject.toml"
|
- "pyproject.toml"
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
issues: write
|
||||||
|
contents: write
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
publish-node:
|
publish-node:
|
||||||
name: Publish Custom Node to registry
|
name: Publish Custom Node to registry
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
|
if: ${{ github.repository_owner == 'yolain' }}
|
||||||
steps:
|
steps:
|
||||||
- name: Check out code
|
- name: Check out code
|
||||||
uses: actions/checkout@v4
|
uses: actions/checkout@v4
|
||||||
|
- name: Extract version from pyproject.toml
|
||||||
|
id: version
|
||||||
|
run: |
|
||||||
|
VERSION=$(grep -E '^\s*version\s*=' pyproject.toml | head -1 | sed -E 's/.*version\s*=\s*"([^"]+)".*/\1/')
|
||||||
|
if [ -z "$VERSION" ]; then
|
||||||
|
echo "ERROR: Could not extract version from pyproject.toml" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
echo "version=$VERSION" >> $GITHUB_OUTPUT
|
||||||
|
echo "Extracted version: $VERSION"
|
||||||
|
|
||||||
|
- name: Create GitHub Release
|
||||||
|
uses: softprops/action-gh-release@v2
|
||||||
|
with:
|
||||||
|
tag_name: v${{ steps.version.outputs.version }}
|
||||||
|
generate_release_notes: true
|
||||||
- name: Publish Custom Node
|
- name: Publish Custom Node
|
||||||
uses: Comfy-Org/publish-node-action@main
|
uses: Comfy-Org/publish-node-action@v1
|
||||||
with:
|
with:
|
||||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||||
|
|||||||
+9
-1
@@ -7,9 +7,17 @@ wildcards/**
|
|||||||
styles/**
|
styles/**
|
||||||
workflow/**
|
workflow/**
|
||||||
autocomplete/**
|
autocomplete/**
|
||||||
|
web_beta/**
|
||||||
|
web_version/dev/**
|
||||||
docs/**
|
docs/**
|
||||||
.vscode/
|
.vscode/
|
||||||
|
.vs/
|
||||||
.idea/
|
.idea/
|
||||||
|
.claude/**
|
||||||
mmb-preset.custom.txt
|
mmb-preset.custom.txt
|
||||||
config.yaml
|
config.yaml
|
||||||
node.tar.gz
|
node.tar.gz
|
||||||
|
.codex
|
||||||
|
|
||||||
|
.cursorrules
|
||||||
|
tools/ComfyUI-Easy-Use.json
|
||||||
|
|||||||
@@ -0,0 +1,4 @@
|
|||||||
|
[submodule "ComfyUI-Easy-Use-Frontend"]
|
||||||
|
path = ComfyUI-Easy-Use-Frontend
|
||||||
|
url = https://github.com/yolain/ComfyUI-Easy-Use-Frontend.git
|
||||||
|
branch = main
|
||||||
Submodule
+1
Submodule ComfyUI-Easy-Use-Frontend added at 656ae09121
+603
@@ -0,0 +1,603 @@
|
|||||||
|

|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<a href="https://space.bilibili.com/1840885116">视频介绍</a> |
|
||||||
|
<a href="https://docs.easyuse.yolain.com">文档</a> |
|
||||||
|
<a href="https://github.com/yolain/ComfyUI-Yolain-Workflows">工作流合集</a> |
|
||||||
|
<a href="#%EF%B8%8F-donation">捐助</a>
|
||||||
|
<br><br>
|
||||||
|
<a href="./README.md"><img src="https://img.shields.io/badge/🇬🇧English-e9e9e9"></a>
|
||||||
|
<a href="./README.ZH_CN.md"><img src="https://img.shields.io/badge/🇨🇳中文简体-0b8cf5"></a>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
**ComfyUI-Easy-Use** 是一个化繁为简的节点整合包, 在 [tinyterraNodes](https://github.com/TinyTerra/ComfyUI_tinyterraNodes) 的基础上进行延展,并针对了诸多主流的节点包做了整合与优化,以达到更快更方便使用ComfyUI的目的,在保证自由度的同时还原了本属于Stable Diffusion的极致畅快出图体验。
|
||||||
|
|
||||||
|
## 👨🏻🎨 特色介绍
|
||||||
|
|
||||||
|
- 沿用了 [tinyterraNodes](https://github.com/TinyTerra/ComfyUI_tinyterraNodes) 的思路,大大减少了折腾工作流的时间成本。
|
||||||
|
- UI界面美化,首次安装的用户,如需使用UI主题,请在 Settings -> Color Palette 中自行切换主题并**刷新页面**即可
|
||||||
|
- 增加了预采样参数配置的节点,可与采样节点分离,更方便预览。
|
||||||
|
- 支持通配符与Lora的提示词节点,如需使用Lora Block Weight用法,需先保证自定义节点包中安装了 [ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack)
|
||||||
|
- 可多选的风格化提示词选择器,默认是Fooocus的样式json,可自定义json放在styles底下,samples文件夹里可放预览图(名称和name一致,图片文件名如有空格需转为下划线'_')
|
||||||
|
- 加载器可开启A1111提示词风格模式,可重现与webui生成近乎相同的图像
|
||||||
|
- 可使用`easy latentNoisy`或`easy preSamplingNoiseIn`节点实现对潜空间的噪声注入
|
||||||
|
- 简化 SD1.x、SD2.x、SDXL、SVD、Zero123等流程
|
||||||
|
- 简化 Stable Cascade [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#1-13-stable-cascade)
|
||||||
|
- 简化 Layer Diffuse [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#2-3-layerdiffusion)
|
||||||
|
- 简化 InstantID [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#2-2-instantid), 需先保证自定义节点包中安装了 [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID)
|
||||||
|
- 简化 IPAdapter, 需先保证自定义节点包中安装最新版v2的 [ComfyUI_IPAdapter_plus](https://github.com/cubiq/ComfyUI_IPAdapter_plus)
|
||||||
|
- 扩展 XYplot 的可用性
|
||||||
|
- 整合了Fooocus Inpaint功能
|
||||||
|
- 整合了常用的逻辑计算、转换类型、展示所有类型等
|
||||||
|
- 支持节点上checkpoint、lora模型子目录分类及预览图 (请在设置中开启上下文菜单嵌套子目录)
|
||||||
|
- 支持BriaAI的RMBG-1.4模型的背景去除节点,[技术参考](https://huggingface.co/briaai/RMBG-1.4)
|
||||||
|
- 支持 强制清理comfyUI模型显存占用
|
||||||
|
- 支持Stable Diffusion 3 多账号API节点
|
||||||
|
- 支持IC-Light的应用 [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#2-5-ic-light) | [代码整合来源](https://github.com/huchenlei/ComfyUI-IC-Light) | [技术参考](https://github.com/lllyasviel/IC-Light)
|
||||||
|
- 中文提示词自动识别,使用[opus-mt-zh-en模型](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en)
|
||||||
|
- 支持 sd3 模型
|
||||||
|
- 支持 kolors 模型
|
||||||
|
- 支持 flux 模型
|
||||||
|
- 支持 惰性条件判断(ifElse)和 for循环
|
||||||
|
- 支持 Anima 与 Krea2 diffusion 模型,可通过 `easy diffusionModelLoader` 加载(需显式选择文本编码器与 VAE),并使用 `easy XYInputs: DiffusionModel` 进行 XY 对比
|
||||||
|
|
||||||
|
## 👨🏻🔧 安装
|
||||||
|
|
||||||
|
1. 将存储库克隆到 **custom_nodes** 目录并安装依赖
|
||||||
|
```shell
|
||||||
|
#1. git下载
|
||||||
|
git clone https://github.com/yolain/ComfyUI-Easy-Use
|
||||||
|
#2. 安装依赖
|
||||||
|
双击install.bat安装依赖
|
||||||
|
```
|
||||||
|
|
||||||
|
## 📜 更新日志
|
||||||
|
|
||||||
|
**v1.4.1**
|
||||||
|
|
||||||
|
- 修复 `easy saveText` 将文本输出限制在输出目录 #1032
|
||||||
|
|
||||||
|
**v1.4.0**
|
||||||
|
|
||||||
|
- 添加 `easy tableEditor` 节点 - 用于编辑和显示表格数据的节点
|
||||||
|
- 修复 `easy showAnything` 在最新版 ComfyUI 前端无法工作的问题
|
||||||
|
- 修复 `easy multiAnglePrompt` 设置保存失败的问题
|
||||||
|
- 修复 `easy detailer` 在子图中无法工作的问题
|
||||||
|
- 修复新版 ComfyUI 前端中"刷新节点"功能连接线丢失的问题
|
||||||
|
- 使用原生 `VAEDecodeTiled` 进行分块解码(支持 Qwen Image VAE)
|
||||||
|
- 修复 `easy forLoopStart` - 允许 `total=0` 以防止不必要的循环执行
|
||||||
|
- 修复 `easy preSampling` - `samplerCustomSettings.ip2p` 中 `vae`/`pixels` 现在是可选的
|
||||||
|
- 修复 `easy pixart` ControlNet 包装器 - 使用 `pe_interpolation` 替代已移除的 `lewei_scale`
|
||||||
|
- 修复 `easy promptConcat` - 当输入为列表时的 TypeError 问题
|
||||||
|
- 使用专用 RNG 进行全局种子生成
|
||||||
|
- 修复 `easy imageDetailTransfer` - 多帧蒙版在通道广播时崩溃
|
||||||
|
- 修复 Windows 环境下 PrimeVue 对话框遮罩未清除的问题
|
||||||
|
- 修复 `LockedMeta` 对象 TypeError(`object of type 'LockedMeta' has no len()`)
|
||||||
|
- 增强 `easy simpleMath` `evaluate_formula` 以处理列表输入
|
||||||
|
- 修复 `loraStack`/`controlnetStack` - 禁用时不再清除上游堆栈
|
||||||
|
- 修复 XYPlot 在 `Seeds++ Batch` 和元组 X/Y 输入时崩溃的问题
|
||||||
|
- 回滚循环节点到 v1 版本
|
||||||
|
- 修复 `easy NodesMap` - 避免递归组引用
|
||||||
|
- 修复 `easy CleanVRAM` 清理顺序并正确清除 Easy-Use 缓存
|
||||||
|
- 指定文件读取的 UTF-8 编码
|
||||||
|
- 移除不必要的 print 语句
|
||||||
|
|
||||||
|
**v1.3.6**
|
||||||
|
|
||||||
|
- 恢复 `easy showAnything` 对于列表类型的支持(但一些情况下展示庞大数据时仍会导致ComfyUI崩溃)
|
||||||
|
- 修复自定义小部件以支持子图和 Nodes 2.0 #942
|
||||||
|
- 添加 `easy multiAngle` 节点
|
||||||
|
- 将 `prompt.py` 转换为 V3 Schema
|
||||||
|
- 修复 `easy humanSegmentation` 错误
|
||||||
|
- 添加 `easy stringJoinLines`、`easy stringToIntList`、`easy simpleMath`
|
||||||
|
- 修复 `easy ifElse` 和 `easy anythingIndexSwitch` 在某些环境下失败的问题
|
||||||
|
|
||||||
|
**v1.3.5**
|
||||||
|
|
||||||
|
- 修复`isNone`
|
||||||
|
- 将`preview_rescale`添加到`easy imageChooser`
|
||||||
|
- 修复小部件隐藏#910
|
||||||
|
- 将 max 参数添加到 `wildcardsPromptMatrix` 偏移量 #909
|
||||||
|
- 修复子图节点上的标题框样式
|
||||||
|
- 在 `easypromptLine` 上添加 `remove_empty_lines`
|
||||||
|
|
||||||
|
**v1.3.4**
|
||||||
|
|
||||||
|
- 修复 `easy seedList` 最大值 #879
|
||||||
|
- 为xyplot添加controlnet input #877
|
||||||
|
- 为 `easy indexAnything` 支持 `反向索引`
|
||||||
|
|
||||||
|
**v1.3.3**
|
||||||
|
|
||||||
|
- 删除CSS类名称`gird-cols-1` #859
|
||||||
|
- 修复锁定种子在 `easy promptAwait` 中不起作用
|
||||||
|
- 重命名节点图
|
||||||
|
- 修复`easy ImageChooser`输出错误类型 #845
|
||||||
|
|
||||||
|
**v1.3.2**
|
||||||
|
|
||||||
|
- 改造 `easy imageChooser` 节点以兼容 frontend>=v1.24.2, 解决方案参考自 [Comfyui_LG_Tools](https://github.com/LAOGOU-666/Comfyui_LG_Tools)
|
||||||
|
- 改造 `easy stylesSelector` 节点, 你可在 [other styles files](https://github.com/yolain/EasyUse-Styles-Templates) 下载到 `styles` 文件夹下
|
||||||
|
- 改造 `easy humanSegmentation` 节点
|
||||||
|
- 修复 `easy makeImageForICLora` 节点.
|
||||||
|
- 添加 `easy joycaption3API` 节点
|
||||||
|
- 添加 `easy promptAwait` 节点
|
||||||
|
|
||||||
|
**v1.3.1**
|
||||||
|
|
||||||
|
- 重写 drawNodeWidget 修复组节点预览的问题.
|
||||||
|
- 更新了一些 XYPlot 的功能 by [mekinney](https://github.com/mekinney)
|
||||||
|
- 添加 `easy seedList` 节点 (它对循环节点有用)
|
||||||
|
|
||||||
|
**v1.3.0**
|
||||||
|
|
||||||
|
- 将循环节点设置为最大输入和输出数量为20
|
||||||
|
- 添加 `uniform width` 方式到 `easy makeImageForICLora`
|
||||||
|
- 增加 `wildcardsPromptMatrix` 通配符提示词矩阵,由 [Rosmeowtis](https://github.com/Rosmeowtis) 贡献
|
||||||
|
|
||||||
|
**v1.2.9**
|
||||||
|
|
||||||
|
- 修复 Imagechooser 会导致工作流处理取消
|
||||||
|
- 修复 brushnet tensor(640) 错误
|
||||||
|
- 修复v1.6.0前端之后无法隐藏小部件的bug
|
||||||
|
- 修复图像选择器无法选择图像
|
||||||
|
- 修复ContextMenu Monkey修补以影响自定义脚本(PYSSSS)节点
|
||||||
|
|
||||||
|
**v1.2.8**
|
||||||
|
|
||||||
|
- 修复了一些BUG (😹)
|
||||||
|
- 增加了多语言目录
|
||||||
|
|
||||||
|
**v1.2.7**
|
||||||
|
|
||||||
|
- 优化管理节点组显示
|
||||||
|
- 在 `easy imageRemBg` 上添加 `ben2`
|
||||||
|
- 添加 joyCaption2 API版节点( https://github.com/siliconflow/BizyAir )
|
||||||
|
- 使用一种新的方式在 loader 中显示模型缩略图(支持 diffusion_models、lors、checkpoints)
|
||||||
|
|
||||||
|
**v1.2.6**
|
||||||
|
|
||||||
|
- 修复了在缺少自定义节点时缺少 “红色框框” 样式的问题。
|
||||||
|
- 在一些简单的加载器中,将 `clip_skip` 的默认值从 `-1` 调整为 `-2`。
|
||||||
|
- 修复因设置节点中缺少相连接的自定义节点而导致弄乱画布的问题
|
||||||
|
- 修复 'easy imageChooser' 不能循环使用的问题。
|
||||||
|
|
||||||
|
**v1.2.5**
|
||||||
|
|
||||||
|
- 在 `easy preSamplingCustom` 和 `easy preSamplingAdvanced` 上增加 `enable (GPU=A1111)` 噪波生成模式选择项
|
||||||
|
- 增加 `easy makeImageForICLora`
|
||||||
|
- 在 `easy ipadapterApply` 添加 `REGULAR - FLUX and SD3.5 only (high strength)` 预置项以支持 InstantX Flux ipadapter
|
||||||
|
- 修复brushnet 无法在 `--fast` 模式下使用
|
||||||
|
- 支持briaai RMBG-2.0
|
||||||
|
- 支持mochi模型
|
||||||
|
- 实现在循环主体中重复使用终端节点输出(例如预览图像和显示任何内容等输出节点...)
|
||||||
|
|
||||||
|
**v1.2.4**
|
||||||
|
|
||||||
|
- 增加 `easy imageSplitTiles` and `easy imageTilesFromBatch` - 图像分块
|
||||||
|
- 支持 `model_override`,`vae_override`,`clip_override` 可以在 `easy fullLoader` 中单独输入
|
||||||
|
- 增加 `easy saveImageLazy`
|
||||||
|
- 增加 `easy loadImageForLoop`
|
||||||
|
- 增加 `easy isFileExist`
|
||||||
|
- 增加 `easy saveText`
|
||||||
|
|
||||||
|
**v1.2.3**
|
||||||
|
|
||||||
|
- `easy showAnything` 和 `easy cleanGPUUsed` 增加输出插槽
|
||||||
|
- 添加新的人体分割在 `easy humanSegmentation` 节点上 - 代码从 [ComfyUI_Human_Parts](https://github.com/metal3d/ComfyUI_Human_Parts) 整合
|
||||||
|
- 当你在 `easy preSamplingCustom` 节点上选择basicGuider,CFG>0 且当前模型为Flux时,将使用FluxGuidance
|
||||||
|
- 增加 `easy loraStackApply` and `easy controlnetStackApply`
|
||||||
|
|
||||||
|
**v1.2.2**
|
||||||
|
|
||||||
|
- 增加 `easy batchAny`
|
||||||
|
- 增加 `easy anythingIndexSwitch`
|
||||||
|
- 增加 `easy forLoopStart` 和 `easy forLoopEnd`
|
||||||
|
- 增加 `easy ifElse`
|
||||||
|
- 增加 v2 版本新前端代码
|
||||||
|
- 增加 `easy fluxLoader`
|
||||||
|
- 增加 `controlnetApply` 相关节点对sd3和hunyuanDiT的支持
|
||||||
|
- 修复 当使用fooocus inpaint后,再使用Lora模型无法生效的问题
|
||||||
|
|
||||||
|
**v1.2.1**
|
||||||
|
|
||||||
|
- 增加 `easy ipadapterApplyFaceIDKolors`
|
||||||
|
- `easy ipadapterApply` 和 `easy ipadapterApplyADV` 增加 **PLUS (kolors genernal)** 和 **FACEID PLUS KOLORS** 预置项
|
||||||
|
- `easy imageRemBg` 增加 **inspyrenet** 选项
|
||||||
|
- 增加 `easy controlnetLoader++`
|
||||||
|
- 去除 `easy positive` `easy negative` 等prompt节点的自动将中文翻译功能,自动翻译仅在 `easy a1111Loader` 等不支持中文TE的加载器中生效
|
||||||
|
- 增加 `easy kolorsLoader` - 可灵加载器,参考了 [MinusZoneAI](https://github.com/MinusZoneAI/ComfyUI-Kolors-MZ) 和 [kijai](https://github.com/kijai/ComfyUI-KwaiKolorsWrapper) 的代码。
|
||||||
|
|
||||||
|
**v1.2.0**
|
||||||
|
|
||||||
|
- 增加 `easy pulIDApply` 和 `easy pulIDApplyADV`
|
||||||
|
- 增加 `easy hunyuanDiTLoader` 和 `easy pixArtLoader`
|
||||||
|
- 当新菜单的位置在上或者下时增加上 crystools 的显示,推荐开两个就好(如果后续crystools有更新UI适配我可能会删除掉)
|
||||||
|
- 增加 **easy sliderControl** - 滑块控制节点,当前可用于控制ipadapterMS的参数 (双击滑块可重置为默认值)
|
||||||
|
- 增加 **layer_weights** 属性在 `easy ipadapterApplyADV` 节点
|
||||||
|
|
||||||
|
**v1.1.9**
|
||||||
|
|
||||||
|
- 增加 新的调度器 **gitsScheduler**
|
||||||
|
- 增加 `easy imageBatchToImageList` 和 `easy imageListToImageBatch` (修复Impact版的一点小问题)
|
||||||
|
- 递归模型子目录嵌套
|
||||||
|
- 支持 sd3 模型
|
||||||
|
- 增加 `easy applyInpaint` - 局部重绘全模式节点 (相比与之前的kSamplerInpating节点逻辑会更合理些)
|
||||||
|
|
||||||
|
**v1.1.8**
|
||||||
|
|
||||||
|
- 增加中文提示词自动翻译,使用[opus-mt-zh-en模型](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en), 默认已对wildcard、lora正则处理, 其他需要保留的中文,可使用`@你的提示词@`包裹 (若依赖安装完成后报错, 请重启),测算大约会占0.3GB显存
|
||||||
|
- 增加 `easy controlnetStack` - controlnet堆
|
||||||
|
- 增加 `easy applyBrushNet` - [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4brushnet_1.1.8.json)
|
||||||
|
- 增加 `easy applyPowerPaint` - [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4powerpaint_outpaint_1.1.8.json)
|
||||||
|
|
||||||
|
**v1.1.7**
|
||||||
|
|
||||||
|
- 修复 一些模型(如controlnet模型等)未成功写入缓存,导致修改前置节点束参数(如提示词)需要二次载入模型的问题
|
||||||
|
- 增加 `easy prompt` - 主体和光影预置项,后期可能会调整
|
||||||
|
- 增加 `easy icLightApply` - 重绘光影, 从[ComfyUI-IC-Light](https://github.com/huchenlei/ComfyUI-IC-Light)优化
|
||||||
|
- 增加 `easy imageSplitGrid` - 图像网格拆分
|
||||||
|
- `easy kSamplerInpainting` 的 **additional** 属性增加差异扩散和brushnet等相关选项
|
||||||
|
- 增加 brushnet模型加载的支持 - [ComfyUI-BrushNet](https://github.com/nullquant/ComfyUI-BrushNet)
|
||||||
|
- 增加 `easy applyFooocusInpaint` - Fooocus内补节点 替代原有的 FooocusInpaintLoader
|
||||||
|
- 移除 `easy fooocusInpaintLoader` - 容易bug,不再使用
|
||||||
|
- 修改 easy kSampler等采样器中并联的model 不再替换输出中pipe里的model
|
||||||
|
|
||||||
|
**v1.1.6**
|
||||||
|
|
||||||
|
- 增加步调齐整适配 - 在所有的预采样和全采样器节点中的 调度器(schedulder) 增加了 **alignYourSteps** 选项
|
||||||
|
- `easy kSampler` 和 `easy fullkSampler` 的 **image_output** 增加 **Preview&Choose**选项
|
||||||
|
- 增加 `easy styleAlignedBatchAlign` - 风格对齐 [style_aligned_comfy](https://github.com/brianfitzgerald/style_aligned_comfy)
|
||||||
|
- 增加 `easy ckptNames`
|
||||||
|
- 增加 `easy controlnetNames`
|
||||||
|
- 增加 `easy imagesSplitimage` - 批次图像拆分单张
|
||||||
|
- 增加 `easy imageCount` - 图像数量
|
||||||
|
- 增加 `easy textSwitch` - 文字切换
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.1.5</b></summary>
|
||||||
|
|
||||||
|
- 重写 `easy cleanGPUUsed` - 可强制清理comfyUI的模型显存占用
|
||||||
|
- 增加 `easy humanSegmentation` - 多类分割、人像分割
|
||||||
|
- 增加 `easy imageColorMatch`
|
||||||
|
- 增加 `easy ipadapterApplyRegional`
|
||||||
|
- 增加 `easy ipadapterApplyFromParams`
|
||||||
|
- 增加 `easy imageInterrogator` - 图像反推
|
||||||
|
- 增加 `easy stableDiffusion3API` - 简易的Stable Diffusion 3 多账号API节点
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.1.4</b></summary>
|
||||||
|
|
||||||
|
- 增加 `easy imageChooser` - 从[cg-image-picker](https://github.com/chrisgoringe/cg-image-picker)简化的图片选择器
|
||||||
|
- 增加 `easy preSamplingCustom` - 自定义预采样,可支持cosXL-edit
|
||||||
|
- 增加 `easy ipadapterStyleComposition`
|
||||||
|
- 增加 在Loaders上右键菜单可查看 checkpoints、lora 信息
|
||||||
|
- 修复 `easy preSamplingNoiseIn`、`easy latentNoisy`、`east Unsampler` 以兼容ComfyUI Revision>=2098 [0542088e] 以上版本
|
||||||
|
- 修复 FooocusInpaint修改ModelPatcher计算权重引发的问题,理应在生成model后重置ModelPatcher为默认值
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.1.3</b></summary>
|
||||||
|
|
||||||
|
- `easy ipadapterApply` 增加 **COMPOSITION** 预置项
|
||||||
|
- 增加 对[ResAdapter](https://huggingface.co/jiaxiangc/res-adapter) lora模型 的加载支持
|
||||||
|
- 增加 `easy promptLine`
|
||||||
|
- 增加 `easy promptReplace`
|
||||||
|
- 增加 `easy promptConcat`
|
||||||
|
- `easy wildcards` 增加 **multiline_mode**属性
|
||||||
|
- 增加 当节点需要下载模型时,若huggingface连接超时,会切换至镜像地址下载模型
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.1.2</b></summary>
|
||||||
|
|
||||||
|
- 改写 EasyUse 相关节点的部分插槽推荐节点
|
||||||
|
- 增加 **启用上下文菜单自动嵌套子目录** 设置项,默认为启用状态,可分类子目录及checkpoints、loras预览图
|
||||||
|
- 增加 `easy sv3dLoader`
|
||||||
|
- 增加 `easy dynamiCrafterLoader`
|
||||||
|
- 增加 `easy ipadapterApply`
|
||||||
|
- 增加 `easy ipadapterApplyADV`
|
||||||
|
- 增加 `easy ipadapterApplyEncoder`
|
||||||
|
- 增加 `easy ipadapterApplyEmbeds`
|
||||||
|
- 增加 `easy preMaskDetailerFix`
|
||||||
|
- `easy kSamplerInpainting` 增加 **additional** 属性,可设置成 Differential Diffusion 或 Only InpaintModelConditioning
|
||||||
|
- 修复 `easy stylesSelector` 当未选择样式时,原有提示词发生了变化
|
||||||
|
- 修复 `easy pipeEdit` 提示词输入lora时报错
|
||||||
|
- 修复 layerDiffuse xyplot相关bug
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.1.1</b></summary>
|
||||||
|
|
||||||
|
- 修复首次添加含seed的节点且当前模式为control_before_generate时,seed为0的问题
|
||||||
|
- `easy preSamplingAdvanced` 增加 **return_with_leftover_noise**
|
||||||
|
- 修复 `easy stylesSelector` 当选择自定义样式文件时运行队列报错
|
||||||
|
- `easy preSamplingLayerDiffusion` 增加 mask 可选传入参数
|
||||||
|
- 将所有 **seed_num** 调整回 **seed**
|
||||||
|
- 修补官方BUG: 当control_mode为before 在首次加载页面时未修改节点中widget名称为 control_before_generate
|
||||||
|
- 去除强制**control_before_generate**设定
|
||||||
|
- 增加 `easy imageRemBg` - 默认为BriaAI的RMBG-1.4模型, 移除背景效果更加,速度更快
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.1.0</b></summary>
|
||||||
|
|
||||||
|
- 增加 `easy imageSplitList` - 拆分每 N 张图像
|
||||||
|
- 增加 `easy preSamplingDiffusionADDTL` - 可配置前景、背景、blended的additional_prompt等
|
||||||
|
- 增加 `easy preSamplingNoiseIn` 可替代需要前置的`easy latentNoisy`节点 实现效果更好的噪声注入
|
||||||
|
- `easy pipeEdit` 增加 条件拼接模式选择,可选择替换、合并、联结、平均、设置条件时间
|
||||||
|
- 增加 `easy pipeEdit` - 可编辑Pipe的节点(包含可重新输入提示词)
|
||||||
|
- 增加 `easy preSamplingLayerDiffusion` 与 `easy kSamplerLayerDiffusion` (连接 `easy kSampler` 也能通)
|
||||||
|
- 增加 在 加载器、预采样、采样器、Controlnet等节点上右键可快速替换同类型节点的便捷菜单
|
||||||
|
- 增加 `easy instantIDApplyADV` 可连入 positive 与 negative
|
||||||
|
- 修复 `easy wildcards` 读取lora未填写完整路径时未自动检索导致加载lora失败的问题
|
||||||
|
- 修复 `easy instantIDApply` mask 未传入正确值
|
||||||
|
- 修复 在 非a1111提示词风格下 BREAK 不生效的问题
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.9</b></summary>
|
||||||
|
|
||||||
|
- 修复未安装 ComfyUI-Impack-Pack 和 ComfyUI_InstantID 时报错
|
||||||
|
- 修复 `easy pipeIn` - pipe设为可不必选
|
||||||
|
- 增加 `easy instantIDApply` - 需要先安装 [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID), 工作流参考[示例](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#2-2-instantid)
|
||||||
|
- 修复 `easy detailerFix` 未添加到保存图片格式化扩展名可用节点列表
|
||||||
|
- 修复 `easy XYInputs: PromptSR` 在替换负面提示词时报错
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.8</b></summary>
|
||||||
|
|
||||||
|
- `easy cascadeLoader` stage_c 与 stage_b 支持checkpoint模型 (需要下载[checkpoints](https://huggingface.co/stabilityai/stable-cascade/tree/main/comfyui_checkpoints))
|
||||||
|
- `easy styleSelector` 搜索框修改为不区分大小写匹配
|
||||||
|
- `easy fullLoader` 增加 **positive**、**negative**、**latent** 输出项
|
||||||
|
- 修复 SDXLClipModel 在 ComfyUI 修订版本号 2016[c2cb8e88] 及以上的报错(判断了版本号可兼容老版本)
|
||||||
|
- 修复 `easy detailerFix` 批次大小大于1时生成出错
|
||||||
|
- 修复`easy preSampling`等 latent传入后无法根据批次索引生成的问题
|
||||||
|
- 修复 `easy svdLoader` 报错
|
||||||
|
- 优化代码,减少了诸多冗余,提升运行速度
|
||||||
|
- 去除中文翻译对照文本
|
||||||
|
|
||||||
|
(翻译对照已由 [AIGODLIKE-COMFYUI-TRANSLATION](https://github.com/AIGODLIKE/AIGODLIKE-ComfyUI-Translation) 统一维护啦!
|
||||||
|
首次下载或者版本较早的朋友请更新 AIGODLIKE-COMFYUI-TRANSLATION 和本节点包至最新版本。)
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.7</b></summary>
|
||||||
|
|
||||||
|
- 增加 `easy cascadeLoader` - stable cascade 加载器
|
||||||
|
- 增加 `easy preSamplingCascade` - stabled cascade stage_c 预采样参数
|
||||||
|
- 增加 `easy fullCascadeKSampler` - stable cascade stage_c 完整版采样器
|
||||||
|
- 增加 `easy cascadeKSampler` - stable cascade stage-c ksampler simple
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.6</b></summary>
|
||||||
|
|
||||||
|
- 增加 `easy XYInputs: Checkpoint`
|
||||||
|
- 增加 `easy XYInputs: Lora`
|
||||||
|
- `easy seed` 增加固定种子值时可手动切换随机种
|
||||||
|
- 修复 `easy fullLoader`等加载器切换lora时自动调整节点大小的问题
|
||||||
|
- 去除原有ttn的图片保存逻辑并适配ComfyUI默认的图片保存格式化扩展
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.5</b></summary>
|
||||||
|
|
||||||
|
- 增加 `easy isSDXL`
|
||||||
|
- `easy svdLoader` 增加提示词控制, 可配合open_clip模型进行使用
|
||||||
|
- `easy wildcards` 增加 **populated_text** 可输出通配填充后文本
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.4</b></summary>
|
||||||
|
|
||||||
|
- 增加 `easy showLoaderSettingsNames` 可显示与输出加载器部件中的 模型与VAE名称
|
||||||
|
- 增加 `easy promptList` - 提示词列表
|
||||||
|
- 增加 `easy fooocusInpaintLoader` - Fooocus内补节点(仅支持XL模型的流程)
|
||||||
|
- 增加 **Logic** 逻辑类节点 - 包含类型、计算、判断和转换类型等
|
||||||
|
- 增加 `easy imageSave` - 带日期转换和宽高格式化的图像保存节点
|
||||||
|
- 增加 `easy joinImageBatch` - 合并图像批次
|
||||||
|
- `easy showAnything` 增加支持转换其他类型(如:tensor类型的条件、图像等)
|
||||||
|
- `easy kSamplerInpainting` 增加 **patch** 传入值,配合Fooocus内补节点使用
|
||||||
|
- `easy imageSave` 增加 **only_preivew**
|
||||||
|
|
||||||
|
- 修复 xyplot在pillow>9.5中报错
|
||||||
|
- 修复 `easy wildcards` 在使用PS扩展插件运行时报错
|
||||||
|
- 修复 `easy latentCompositeMaskedWithCond`
|
||||||
|
- 修复 `easy XYInputs: ControlNet` 报错
|
||||||
|
- 修复 `easy loraStack` **toggle** 为 disabled 时报错
|
||||||
|
|
||||||
|
- 修改首次安装节点包不再自动替换主题,需手动调整并刷新页面
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.3</b></summary>
|
||||||
|
|
||||||
|
- 增加 `easy stylesSelector` 风格化提示词选择器
|
||||||
|
- 增加队列进度条设置项,默认为未启用状态
|
||||||
|
- `easy controlnetLoader` 和 `easy controlnetLoaderADV` 增加参数 **scale_soft_weights**
|
||||||
|
|
||||||
|
|
||||||
|
- 修复 `easy XYInputs: Sampler/Scheduler` 报错
|
||||||
|
- 修复 右侧菜单 点击按钮时老是跑位的问题
|
||||||
|
- 修复 styles 路径在其他环境报错
|
||||||
|
- 修复 `easy comfyLoader` 读取错误
|
||||||
|
- 修复 xyPlot 在连接 zero123 时报错
|
||||||
|
- 修复加载器中提示词为组件时报错
|
||||||
|
- 修复 `easy getNode` 和 `easy setNode` 加载时标题未更改
|
||||||
|
- 修复所有采样器中存储图片使用子目录前缀不生效的问题
|
||||||
|
|
||||||
|
|
||||||
|
- 调整UI主题
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.2</b></summary>
|
||||||
|
|
||||||
|
- 增加 **autocomplete** 文件夹,如果您安装了 [ComfyUI-Custom-Scripts](https://github.com/pythongosssss/ComfyUI-Custom-Scripts), 将在启动时合并该文件夹下的所有txt文件并覆盖到pyssss包里的autocomplete.txt文件。
|
||||||
|
- 增加 `easy XYPlotAdvanced` 和 `easy XYInputs` 等相关节点
|
||||||
|
- 增加 **Alt+1到9** 快捷键,可快速粘贴 Node templates 的节点预设 (对应 1到9 顺序)
|
||||||
|
|
||||||
|
- 修复 `easy imageInsetCrop` 测量值为百分比时步进为1
|
||||||
|
- 修复 开启 `a1111_prompt_style` 时XY图表无法使用的问题
|
||||||
|
- 右键菜单中增加了一个 `📜Groups Map(EasyUse)`
|
||||||
|
|
||||||
|
- 修复在Comfy新版本中UI加载失败
|
||||||
|
- 修复 `easy pipeToBasicPipe` 报错
|
||||||
|
- 修改 `easy fullLoader` 和 `easy a1111Loader` 中的 **a1111_prompt_style** 默认值为 False
|
||||||
|
- `easy XYInputs ModelMergeBlocks` 支持csv文件导入数值
|
||||||
|
|
||||||
|
- 替换了XY图生成时的字体文件
|
||||||
|
|
||||||
|
- 移除 `easy imageRemBg`
|
||||||
|
- 移除包中的介绍图和工作流文件,减少包体积
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.1</b></summary>
|
||||||
|
|
||||||
|
- 新增 `easy seed` - 简易随机种
|
||||||
|
- `easy preDetailerFix` 新增了 `optional_image` 传入图像可选,如未传默认取值为pipe里的图像
|
||||||
|
- 新增 `easy kSamplerInpainting` 用于内补潜空间的采样器
|
||||||
|
- 新增 `easy pipeToBasicPipe` 用于转换到Impact的某些节点上
|
||||||
|
|
||||||
|
- 修复 `easy comfyLoader` 报错
|
||||||
|
- 修复所有包含输出图片尺寸的节点取值方式无法批处理的问题
|
||||||
|
- 修复 `width` 和 `height` 无法在 `easy svdLoader` 自定义的报错问题
|
||||||
|
- 修复所有采样器预览图片的地址链接 (解决在 MACOS 系统中图片无法在采样器中预览的问题)
|
||||||
|
- 修复 `vae_name` 在 `easy fullLoader` 和 `easy a1111Loader` 和 `easy comfyLoader` 中选择但未替换原始vae问题
|
||||||
|
- 修复 `easy fullkSampler` 除pipe外其他输出值的报错
|
||||||
|
- 修复 `easy hiresFix` 输入连接pipe和image、vae同时存在时报错
|
||||||
|
- 修复 `easy fullLoader` 中 `model_override` 连接后未执行
|
||||||
|
- 修复 因新增`easy seed` 导致action错误
|
||||||
|
- 修复 `easy xyplot` 的字体文件路径读取错误
|
||||||
|
- 修复 convert 到 `easy seed` 随机种无法固定的问题
|
||||||
|
- 修复 `easy pipeIn` 值传入的报错问题
|
||||||
|
- 修复 `easy zero123Loader` 和 `easy svdLoader` 读取模型时将模型加入到缓存中
|
||||||
|
- 修复 `easy kSampler` `easy kSamplerTiled` `easy detailerFix` 的 `image_output` 默认值为 Preview
|
||||||
|
- `easy fullLoader` 和 `easy a1111Loader` 新增了 `a1111_prompt_style` 参数可以重现和webui生成相同的图像,当前您需要安装 [ComfyUI_smZNodes](https://github.com/shiimizu/ComfyUI_smZNodes) 才能使用此功能
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v1.0.0</b></summary>
|
||||||
|
|
||||||
|
- 新增`easy positive` - 简易正面提示词文本
|
||||||
|
- 新增`easy negative` - 简易负面提示词文本
|
||||||
|
- 新增`easy wildcards` - 支持通配符和Lora选择的提示词文本
|
||||||
|
- 新增`easy portraitMaster` - 肖像大师v2.2
|
||||||
|
- 新增`easy loraStack` - Lora堆
|
||||||
|
- 新增`easy fullLoader` - 完整版的加载器
|
||||||
|
- 新增`easy zero123Loader` - 简易zero123加载器
|
||||||
|
- 新增`easy svdLoader` - 简易svd加载器
|
||||||
|
- 新增`easy fullkSampler` - 完整版的采样器(无分离)
|
||||||
|
- 新增`easy hiresFix` - 支持Pipe的高清修复
|
||||||
|
- 新增`easy predetailerFix` `easy DetailerFix` - 支持Pipe的细节修复
|
||||||
|
- 新增`easy ultralyticsDetectorPipe` `easy samLoaderPipe` - 检测加载器(细节修复的输入项)
|
||||||
|
- 新增`easy pipein` `easy pipeout` - Pipe的输入与输出
|
||||||
|
- 新增`easy xyPlot` - 简易的xyplot (后续会更新更多可控参数)
|
||||||
|
- 新增`easy imageRemoveBG` - 图像去除背景
|
||||||
|
- 新增`easy imagePixelPerfect` - 图像完美像素
|
||||||
|
- 新增`easy poseEditor` - 姿势编辑器
|
||||||
|
- 新增UI主题(黑曜石)- 默认自动加载UI, 也可在设置中自行更替
|
||||||
|
|
||||||
|
- 修复 `easy globalSeed` 不生效问题
|
||||||
|
- 修复所有的`seed_num` 因 [cg-use-everywhere](https://github.com/chrisgoringe/cg-use-everywhere) 实时更新图表导致值错乱的问题
|
||||||
|
- 修复`easy imageSize` `easy imageSizeBySide` `easy imageSizeByLongerSide` 可作为终节点
|
||||||
|
- 修复 `seed_num` (随机种子值) 在历史记录中读取无法一致的Bug
|
||||||
|
</details>
|
||||||
|
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>v0.5</b></summary>
|
||||||
|
|
||||||
|
- 新增 `easy controlnetLoaderADV` 节点
|
||||||
|
- 新增 `easy imageSizeBySide` 节点,可选输出为长边或短边
|
||||||
|
- 新增 `easy LLLiteLoader` 节点,如果您预先安装过 kohya-ss/ControlNet-LLLite-ComfyUI 包,请将 models 里的模型文件移动至 ComfyUI\models\controlnet\ (即comfy默认的controlnet路径里,请勿修改模型的文件名,不然会读取不到)。
|
||||||
|
- 新增 `easy imageSize` 和 `easy imageSizeByLongerSize` 输出的尺寸显示。
|
||||||
|
- 新增 `easy showSpentTime` 节点用于展示图片推理花费时间与VAE解码花费时间。
|
||||||
|
- `easy controlnetLoaderADV` 和 `easy controlnetLoader` 新增 `control_net` 可选传入参数
|
||||||
|
- `easy preSampling` 和 `easy preSamplingAdvanced` 新增 `image_to_latent` 可选传入参数
|
||||||
|
- `easy a1111Loader` 和 `easy comfyLoader` 新增 `batch_size` 传入参数
|
||||||
|
|
||||||
|
- 修改 `easy controlnetLoader` 到 loader 分类底下。
|
||||||
|
</details>
|
||||||
|
|
||||||
|
## 整合参考到的相关节点包
|
||||||
|
|
||||||
|
声明: 非常尊重这些原作者们的付出,开源不易,我仅仅只是做了一些整合与优化。
|
||||||
|
|
||||||
|
| 节点名 (搜索名) | 相关的库 | 库相关的节点 |
|
||||||
|
|:-------------------------------|:----------------------------------------------------------------------------|:------------------------|
|
||||||
|
| easy setNode | [ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) | diffus3.SetNode |
|
||||||
|
| easy getNode | [ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) | diffus3.GetNode |
|
||||||
|
| easy bookmark | [rgthree-comfy](https://github.com/rgthree/rgthree-comfy) | Bookmark 🔖 |
|
||||||
|
| easy portraitMarker | [comfyui-portrait-master](https://github.com/florestefano1975/comfyui-portrait-master) | Portrait Master |
|
||||||
|
| easy LLLiteLoader | [ControlNet-LLLite-ComfyUI](https://github.com/kohya-ss/ControlNet-LLLite-ComfyUI) | LLLiteLoader |
|
||||||
|
| easy globalSeed | [ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) | Global Seed (Inspire) |
|
||||||
|
| easy preSamplingDynamicCFG | [sd-dynamic-thresholding](https://github.com/mcmonkeyprojects/sd-dynamic-thresholding) | DynamicThresholdingFull |
|
||||||
|
| dynamicThresholdingFull | [sd-dynamic-thresholding](https://github.com/mcmonkeyprojects/sd-dynamic-thresholding) | DynamicThresholdingFull |
|
||||||
|
| easy imageInsetCrop | [rgthree-comfy](https://github.com/rgthree/rgthree-comfy) | ImageInsetCrop |
|
||||||
|
| easy poseEditor | [ComfyUI_Custom_Nodes_AlekPet](https://github.com/AlekPet/ComfyUI_Custom_Nodes_AlekPet) | poseNode |
|
||||||
|
| easy if | [ComfyUI-Logic](https://github.com/theUpsider/ComfyUI-Logic) | IfExecute |
|
||||||
|
| easy preSamplingLayerDiffusion | [ComfyUI-layerdiffusion](https://github.com/huchenlei/ComfyUI-layerdiffusion) | LayeredDiffusionApply等 |
|
||||||
|
| easy dynamiCrafterLoader | [ComfyUI-layerdiffusion](https://github.com/ExponentialML/ComfyUI_Native_DynamiCrafter) | Apply Dynamicrafter |
|
||||||
|
| easy imageChooser | [cg-image-picker](https://github.com/chrisgoringe/cg-image-picker) | Preview Chooser |
|
||||||
|
| easy styleAlignedBatchAlign | [style_aligned_comfy](https://github.com/chrisgoringe/cg-image-picker) | styleAlignedBatchAlign |
|
||||||
|
| easy icLightApply | [ComfyUI-IC-Light](https://github.com/huchenlei/ComfyUI-IC-Light) | ICLightApply等 |
|
||||||
|
| easy kolorsLoader | [ComfyUI-Kolors-MZ](https://github.com/MinusZoneAI/ComfyUI-Kolors-MZ) | kolorsLoader |
|
||||||
|
|
||||||
|
## Credits
|
||||||
|
|
||||||
|
[ComfyUI](https://github.com/comfyanonymous/ComfyUI) - 功能强大且模块化的Stable Diffusion GUI
|
||||||
|
|
||||||
|
[ComfyUI-ComfyUI-Manager](https://github.com/ltdrdata/ComfyUI-Manager) - ComfyUI管理器
|
||||||
|
|
||||||
|
[tinyterraNodes](https://github.com/TinyTerra/ComfyUI_tinyterraNodes) - 管道节点(节点束)让用户减少了不必要的连接
|
||||||
|
|
||||||
|
[ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) - diffus3的获取与设置点让用户可以分离工作流构成
|
||||||
|
|
||||||
|
[ComfyUI-Impact-Pack](https://github.com/ltdrdata/ComfyUI-Impact-Pack) - 常规整合包1
|
||||||
|
|
||||||
|
[ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) - 常规整合包2
|
||||||
|
|
||||||
|
[ComfyUI-Logic](https://github.com/theUpsider/ComfyUI-Logic) - ComfyUI逻辑运算
|
||||||
|
|
||||||
|
[ComfyUI-ResAdapter](https://github.com/jiaxiangc/ComfyUI-ResAdapter) - 让模型生成不受训练分辨率限制
|
||||||
|
|
||||||
|
[ComfyUI_IPAdapter_plus](https://github.com/cubiq/ComfyUI_IPAdapter_plus) - 风格迁移
|
||||||
|
|
||||||
|
[ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID) - 人脸迁移
|
||||||
|
|
||||||
|
[ComfyUI_PuLID](https://github.com/cubiq/PuLID_ComfyUI) - 人脸迁移
|
||||||
|
|
||||||
|
[ComfyUI-Custom-Scripts](https://github.com/pythongosssss/ComfyUI-Custom-Scripts) - pyssss 小蛇🐍脚本
|
||||||
|
|
||||||
|
[cg-image-picker](https://github.com/chrisgoringe/cg-image-picker) - 图片选择器
|
||||||
|
|
||||||
|
[ComfyUI-BrushNet](https://github.com/nullquant/ComfyUI-BrushNet) - BrushNet 内补节点
|
||||||
|
|
||||||
|
[ComfyUI_ExtraModels](https://github.com/city96/ComfyUI_ExtraModels) - DiT架构相关节点(Pixart、混元DiT等)
|
||||||
|
|
||||||
|
## 免责声明
|
||||||
|
|
||||||
|
本开源项目及其内容按 “原样 ”提供,不作任何明示或暗示的保证,包括但不限于适销性、特定用途适用性和非侵权保证。在任何情况下,作者或其他版权所有者均不对因本软件或本软件的使用或其他交易而产生、引起或与之相关的任何索赔、损害或其他责任承担责任,无论是合同诉讼、侵权诉讼还是其他诉讼。
|
||||||
|
|
||||||
|
用户应自行负责确保在使用本软件或发布由本软件生成的内容时,遵守所在司法管辖区的所有适用法律和法规。作者和版权所有者不对用户在其各自所在地违反法律或法规的行为负责。
|
||||||
|
|
||||||
|
## ☕️ 投喂
|
||||||
|
|
||||||
|
**Comfyui-Easy-Use** 是一个 GPL 许可的开源项目。为了项目取得更好、可持续的发展,我希望能够获得更多的支持。 如果我的自定义节点为您的一天增添了价值,请考虑喝杯咖啡来进一步补充能量! 💖感谢您的支持,每一杯咖啡都是我创作的动力!
|
||||||
|
|
||||||
|
- [BiliBili充电](https://space.bilibili.com/1840885116)
|
||||||
|
- [Wechat/Alipay](https://github.com/user-attachments/assets/803469bd-ed6a-4fab-932d-50e5088a2d03)
|
||||||
|
|
||||||
|
感谢您的捐助,我将用这些费用来租用 GPU 或购买其他 GPT 服务,以便更好地调试和完善 ComfyUI-Easy-Use 功能
|
||||||
|
|
||||||
|
## 🌟大富大贵的人儿
|
||||||
|
|
||||||
|
我对那些慷慨的赐予一颗星的人表示感谢。非常感谢您的支持!
|
||||||
|
|
||||||
|
[](https://github.com/yolain/ComfyUI-Easy-Use/stargazers)
|
||||||
-402
@@ -1,402 +0,0 @@
|
|||||||
<p align="right">
|
|
||||||
<a href="./README.md">中文</a> | <strong>English</strong>
|
|
||||||
</p>
|
|
||||||
|
|
||||||
<div align="center">
|
|
||||||
|
|
||||||
# ComfyUI Easy Use
|
|
||||||
</div>
|
|
||||||
|
|
||||||
**ComfyUI-Easy-Use** is a simplified node integration package, which is extended on the basis of [tinyterraNodes](https://github.com/TinyTerra/ComfyUI_tinyterraNodes), and has been integrated and optimized for many mainstream node packages to achieve the purpose of faster and more convenient use of ComfyUI. While ensuring the degree of freedom, it restores the ultimate smooth image production experience that belongs to Stable Diffusion.
|
|
||||||
|
|
||||||
[](https://github.com/yolain/ComfyUI-Yolain-Workflows)
|
|
||||||
|
|
||||||
## Introduce
|
|
||||||
|
|
||||||
- Inspire by [tinyterraNodes](https://github.com/TinyTerra/ComfyUI_tinyterraNodes), which greatly reduces the time cost of tossing workflows。
|
|
||||||
- UI interface beautification, the first time you install the user, if you need to use the UI theme, please switch the theme in Settings -> Color Palette and refresh page.
|
|
||||||
- Added a node for pre-sampling parameter configuration, which can be separated from the sampling node for easier previewing
|
|
||||||
- Wildcards and lora's are supported, for Lora Block Weight usage, ensure that the custom node package has the [ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack)
|
|
||||||
- Multi-selectable styled cue word selector, default is Fooocus style json, custom json can be placed under styles, samples folder can be placed in the preview image (name and name consistent, image file name such as spaces need to be converted to underscores '_')
|
|
||||||
- The loader enables the A1111 prompt mode, which reproduces nearly identical images to those generated by webui, and needs to be installed [ComfyUI_smZNodes](https://github.com/shiimizu/ComfyUI_smZNodes) first.
|
|
||||||
- Noise injection into the latent space can be achieved using the `easy latentNoisy` or `easy preSamplingNoiseIn` node
|
|
||||||
- Simplified processes for SD1.x, SD2.x, SDXL, SVD, Zero123, etc. [Example](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#StableDiffusion)
|
|
||||||
- Simplified Stable Cascade [Example](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#StableCascade)
|
|
||||||
- Simplified Layer Diffuse [Example](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#LayerDiffusion),The first time you use it you may need to run `pip install -r requirements.txt` to install the required dependencies.
|
|
||||||
- Simplified InstantID [Example](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#InstantID), You need to make sure that the custom node package has the [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID)
|
|
||||||
- Extending the usability of XYplot
|
|
||||||
- Fooocus Inpaint integration
|
|
||||||
- Integration of common logical calculations, conversion of types, display of all types, etc.
|
|
||||||
- Background removal nodes for the RMBG-1.4 model supporting BriaAI, [BriaAI Guide](https://huggingface.co/briaai/RMBG-1.4)
|
|
||||||
- Forcibly cleared the memory usage of the comfy UI model are supported
|
|
||||||
- Stable Diffusion 3 multi-account API nodes are supported
|
|
||||||
- Support Stable Diffusion 3 model
|
|
||||||
|
|
||||||
## Installation
|
|
||||||
Clone the repo into the **custom_nodes** directory and install the requirements:
|
|
||||||
```shell
|
|
||||||
#1. Clone the repo
|
|
||||||
git clone https://github.com/yolain/ComfyUI-Easy-Use
|
|
||||||
#2. Install the requirements
|
|
||||||
Double-click install.bat to install the required dependencies
|
|
||||||
```
|
|
||||||
|
|
||||||
## Changelog
|
|
||||||
|
|
||||||
**v1.2.0**
|
|
||||||
|
|
||||||
- Added `easy pulIDApply` and `easy pulIDApplyADV`
|
|
||||||
- Added `easy huanyuanDiTLoader` and `easy pixArtLoader`
|
|
||||||
- Added **easy sliderControl** - Slider control node, which can currently be used to control the parameters of ipadapterMS (double-click the slider to reset to default)
|
|
||||||
- Added **layer_weights** in `easy ipadapterApplyADV`
|
|
||||||
|
|
||||||
**v1.1.9**
|
|
||||||
|
|
||||||
- Added **gitsScheduler**
|
|
||||||
- Added `easy imageBatchToImageList` and `easy imageListToImageBatch`
|
|
||||||
- Recursive subcategories nested for models
|
|
||||||
- Support for Stable Diffusion 3 model
|
|
||||||
- Added `easy applyInpaint` - All inpainting mode in this node
|
|
||||||
|
|
||||||
**v1.1.8**
|
|
||||||
|
|
||||||
- Added `easy controlnetStack`
|
|
||||||
- Added `easy applyBrushNet` - [Workflow Example](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4brushnet_1.1.8.json)
|
|
||||||
- Added `easy applyPowerPaint` - [Workflow Example](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4powerpaint_outpaint_1.1.8.json)
|
|
||||||
|
|
||||||
**v1.1.7**
|
|
||||||
|
|
||||||
- Added `easy prompt` - Subject and light presets, maybe adjusted later
|
|
||||||
- Added `easy icLightApply` - Light and shadow migration, Code based on [ComfyUI-IC-Light](https://github.com/huchenlei/ComfyUI-IC-Light)
|
|
||||||
- Added `easy imageSplitGrid`
|
|
||||||
- `easy kSamplerInpainting` added options such as different diffusion and brushnet in **additional** widget
|
|
||||||
- Support for brushnet model loading - [ComfyUI-BrushNet](https://github.com/nullquant/ComfyUI-BrushNet)
|
|
||||||
- Added `easy applyFooocusInpaint` - Replace FooocusInpaintLoader
|
|
||||||
- Removed `easy fooocusInpaintLoader`
|
|
||||||
|
|
||||||
**v1.1.6**
|
|
||||||
|
|
||||||
- Added **alignYourSteps** to **schedulder** widget in all `easy preSampling` and `easy fullkSampler`
|
|
||||||
- Added **Preview&Choose** to **image_output** widget in `easy kSampler` & `easy fullkSampler`
|
|
||||||
- Added `easy styleAlignedBatchAlign` - Credit of [style_aligned_comfy](https://github.com/brianfitzgerald/style_aligned_comfy)
|
|
||||||
- Added `easy ckptNames`
|
|
||||||
- Added `easy controlnetNames`
|
|
||||||
- Added `easy imagesSplitimage` - Batch images split into single images
|
|
||||||
- Added `easy imageCount` - Get Image Count
|
|
||||||
- Added `easy textSwitch` - Text Switch
|
|
||||||
|
|
||||||
**v1.1.5**
|
|
||||||
|
|
||||||
- Rewrite `easy cleanGPUUsed` - the memory usage of the comfyUI can to be cleared
|
|
||||||
- Added `easy humanSegmentation` - Human Part Segmentation
|
|
||||||
- Added `easy imageColorMatch`
|
|
||||||
- Added `easy ipadapterApplyRegional`
|
|
||||||
- Added `easy ipadapterApplyFromParams`
|
|
||||||
- Added `easy imageInterrogator` - Image To Prompt
|
|
||||||
- Added `easy stableDiffusion3API` - Easy Stable Diffusion 3 Multiple accounts API Node
|
|
||||||
|
|
||||||
**v1.1.4**
|
|
||||||
|
|
||||||
- Added `easy preSamplingCustom` - Custom-PreSampling, can be supported cosXL-edit
|
|
||||||
- Added `easy ipadapterStyleComposition`
|
|
||||||
- Added the right-click menu to view checkpoints and lora information in all Loaders
|
|
||||||
- Fixed `easy preSamplingNoiseIn`、`easy latentNoisy`、`east Unsampler` compatible with ComfyUI Revision>=2098 [0542088e] or later
|
|
||||||
|
|
||||||
|
|
||||||
**v1.1.3**
|
|
||||||
|
|
||||||
- `easy ipadapterApply` Added **COMPOSITION** preset
|
|
||||||
- Supported [ResAdapter](https://huggingface.co/jiaxiangc/res-adapter) when load ResAdapter lora
|
|
||||||
- Added `easy promptLine`
|
|
||||||
- Added `easy promptReplace`
|
|
||||||
- Added `easy promptConcat`
|
|
||||||
- `easy wildcards` Added **multiline_mode**
|
|
||||||
|
|
||||||
**v1.1.2**
|
|
||||||
|
|
||||||
- Optimized some of the recommended nodes for slots related to EasyUse
|
|
||||||
- Added **Enable ContextMenu Auto Nest Subdirectories** The setting item is enabled by default, and it can be classified into subdirectories, checkpoints and loras previews
|
|
||||||
- Added `easy sv3dLoader`
|
|
||||||
- Added `easy dynamiCrafterLoader`
|
|
||||||
- Added `easy ipadapterApply`
|
|
||||||
- Added `easy ipadapterApplyADV`
|
|
||||||
- Added `easy ipadapterApplyEncoder`
|
|
||||||
- Added `easy ipadapterApplyEmbeds`
|
|
||||||
- Added `easy preMaskDetailerFix`
|
|
||||||
- Fixed `easy stylesSelector` is change the prompt when not select the style
|
|
||||||
- Fixed `easy pipeEdit` error when add lora to prompt
|
|
||||||
- Fixed layerDiffuse xyplot bug
|
|
||||||
- `easy kSamplerInpainting` add *additional* widget,you can choose 'Differential Diffusion' or 'Only InpaintModelConditioning'
|
|
||||||
|
|
||||||
**v1.1.1**
|
|
||||||
|
|
||||||
- The issue that the seed is 0 when a node with a seed control is added and **control before generate** is fixed for the first time run queue prompt.
|
|
||||||
- `easy preSamplingAdvanced` Added **return_with_leftover_noise**
|
|
||||||
- Fixed `easy stylesSelector` error when choose the custom file
|
|
||||||
- `easy preSamplingLayerDiffusion` Added optional input parameter for mask
|
|
||||||
- Renamed all nodes widget name named seed_num to seed
|
|
||||||
- Remove forced **control_before_generate** settings。 If you want to use control_before_generate, change widget_value_control_mode to before in system settings
|
|
||||||
- Added `easy imageRemBg` - The default is BriaAI's RMBG-1.4 model, which removes the background effect more and faster
|
|
||||||
|
|
||||||
**v1.1.0**
|
|
||||||
|
|
||||||
- Added `easy imageSplitList` - to split every N images
|
|
||||||
- Added `easy preSamplingDiffusionADDTL` - It can modify foreground、background or blended additional prompt
|
|
||||||
- Added `easy preSamplingNoiseIn` It can replace the `easy latentNoisy` node that needs to be fronted to achieve better noise injection
|
|
||||||
- `easy pipeEdit` Added conditioning splicing mode selection, you can choose to replace, concat, combine, average, and set timestep range
|
|
||||||
- Added `easy pipeEdit` - nodes that can edit pipes (including re-enterable prompts)
|
|
||||||
- Added `easy preSamplingLayerDiffusion` and `easy kSamplerLayerDiffusion`
|
|
||||||
- Added a convenient menu to right-click on nodes such as Loader, Presampler, Sampler, Controlnet, etc. to quickly replace nodes of the same type
|
|
||||||
- Added `easy instantIDApplyADV` can link positive and negative
|
|
||||||
- Fixed layerDiffusion error when batch size greater than 1
|
|
||||||
- Fixed `easy wildcards` When LoRa is not filled in completely, LoRa is not automatically retrieved, resulting in failure to load LoRa
|
|
||||||
- Fixed the issue that 'BREAK' non-initiation when didn't use a1111 prompt style
|
|
||||||
- Fixed `easy instantIDApply` mask not input right
|
|
||||||
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.9</b></summary>
|
|
||||||
|
|
||||||
- Fixed the error when ComfyUI-Impack-Pack and ComfyUI_InstantID were not installed
|
|
||||||
- Fixed `easy pipeIn`
|
|
||||||
- Added `easy instantIDApply` - you need installed [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID) fisrt, Workflow[Example](https://github.com/yolain/ComfyUI-Easy-Use/blob/main/README.en.md#InstantID)
|
|
||||||
- Fixed `easy detailerFix` not added to the list of nodes available for saving images formatting extensions
|
|
||||||
- Fixed `easy XYInputs: PromptSR` errors are reported when replacing negative prompts
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.8</b></summary>
|
|
||||||
|
|
||||||
- `easy cascadeLoader` stage_c and stage_b support the checkpoint model (Download [checkpoints](https://huggingface.co/stabilityai/stable-cascade/tree/main/comfyui_checkpoints) models)
|
|
||||||
- `easy styleSelector` The search box is modified to be case-insensitive
|
|
||||||
- `easy fullLoader` **positive**、**negative**、**latent** added to the output items
|
|
||||||
- Fixed the issue that 'easy preSampling' and other similar node, latent could not be generated based on the batch index after passing in
|
|
||||||
- Fixed `easy svdLoader` error when the positive or negative is empty
|
|
||||||
- Fixed the error of SDXLClipModel in ComfyUI revision 2016[c2cb8e88] and above (the revision number was judged to be compatible with the old revision)
|
|
||||||
- Fixed `easy detailerFix` generation error when batch size is greater than 1
|
|
||||||
- Optimize the code, reduce a lot of redundant code and improve the running speed
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.7</b></summary>
|
|
||||||
|
|
||||||
- Added `easy cascadeLoader` - stable cascade Loader
|
|
||||||
- Added `easy preSamplingCascade` - stable cascade preSampling Settings
|
|
||||||
- Added `easy fullCascadeKSampler` - stable cascade stage-c ksampler full
|
|
||||||
- Added `easy cascadeKSampler` - stable cascade stage-c ksampler simple
|
|
||||||
-
|
|
||||||
- Optimize the image to image[Example](https://github.com/yolain/ComfyUI-Easy-Use/blob/main/README.en.md#image-to-image)
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.6</b></summary>
|
|
||||||
|
|
||||||
- Added `easy XYInputs: Checkpoint`
|
|
||||||
- Added `easy XYInputs: Lora`
|
|
||||||
- `easy seed` can manually switch the random seed when increasing the fixed seed value
|
|
||||||
- Fixed `easy fullLoader` and all loaders to automatically adjust the node size when switching LoRa
|
|
||||||
- Removed the original ttn image saving logic and adapted to the default image saving format extension of ComfyUI
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.5</b></summary>
|
|
||||||
|
|
||||||
- Added `easy isSDXL`
|
|
||||||
- Added prompt word control on `easy svdLoader`, which can be used with open_clip model
|
|
||||||
- Added **populated_text** on `easy wildcards`, wildcard populated text can be output
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.4</b></summary>
|
|
||||||
|
|
||||||
- `easy showAnything` added support for converting other types (e.g., tensor conditions, images, etc.)
|
|
||||||
- Added `easy showLoaderSettingsNames` can display the model and VAE name in the output loader assembly
|
|
||||||
- Added `easy promptList`
|
|
||||||
- Added `easy fooocusInpaintLoader` (only the process of SDXLModel is supported)
|
|
||||||
- Added **Logic** nodes
|
|
||||||
- Added `easy imageSave` - Image saving node with date conversion and aspect and height formatting
|
|
||||||
- Added `easy joinImageBatch`
|
|
||||||
- `easy kSamplerInpainting` Added the **patch** input value to be used with the FooocusInpaintLoader node
|
|
||||||
|
|
||||||
- Fixed xyplot error when with Pillow>9.5
|
|
||||||
- Fixed `easy wildcards` An error is reported when running with the PS extension
|
|
||||||
- Fixed `easy XYInputs: ControlNet` Error
|
|
||||||
- Fixed `easy loraStack` error when **toggle** is disabled
|
|
||||||
|
|
||||||
|
|
||||||
- Changing the first-time install node package no longer automatically replaces the theme, you need to manually adjust and refresh the page
|
|
||||||
- `easy imageSave` added **only_preivew**
|
|
||||||
- Adjust the `easy latentCompositeMaskedWithCond` node
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.3</b></summary>
|
|
||||||
|
|
||||||
- Added `easy stylesSelector`
|
|
||||||
- Added **scale_soft_weights** in `easy controlnetLoader` and `easy controlnetLoaderADV`
|
|
||||||
- Added the queue progress bar setting item, which is not enabled by default
|
|
||||||
|
|
||||||
|
|
||||||
- Fixed `easy XYInputs: Sampler/Scheduler` Error
|
|
||||||
- Fixed the right menu has a problem when clicking the button
|
|
||||||
- Fixed `easy comfyLoader` error
|
|
||||||
- Fixed xyPlot error when connecting to zero123
|
|
||||||
- Fixed the error message in the loader when the prompt word was component
|
|
||||||
- Fixed `easy getNode` and `easy setNode` the title does not change when loading
|
|
||||||
- Fixed all samplers using subdirectories to store images
|
|
||||||
|
|
||||||
|
|
||||||
- Adjust the UI theme, divided into two sets of styles: the official default background and the dark black background, which can be switched in the color palette in the settings
|
|
||||||
- Modify the styles path to be compatible with other environments
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.2</b></summary>
|
|
||||||
|
|
||||||
- Added `easy XYPlotAdvanced` and some nodes about `easy XYInputs`
|
|
||||||
- Added **Alt+1-Alt+9** Shortcut keys to quickly paste node presets for Node templates (corresponding to 1~9 sequences)
|
|
||||||
- Added a `📜Groups Map(EasyUse)` to the context menu.
|
|
||||||
- An `autocomplete` folder has been added, If you have [ComfyUI-Custom-Scripts](https://github.com/pythongosssss/ComfyUI-Custom-Scripts) installed, the txt files in that folder will be merged and overwritten to the autocomplete .txt file of the pyssss package at startup.
|
|
||||||
|
|
||||||
|
|
||||||
- Fixed XYPlot is not working when `a1111_prompt_style` is True
|
|
||||||
- Fixed UI loading failure in the new version of ComfyUI
|
|
||||||
- `easy XYInputs ModelMergeBlocks` Values can be imported from CSV files
|
|
||||||
- Fixed `easy pipeToBasicPipe` Bug
|
|
||||||
|
|
||||||
|
|
||||||
- Removed `easy imageRemBg`
|
|
||||||
- Remove the introductory diagram and workflow files from the package to reduce the package size
|
|
||||||
- Replaced the font file used in the generation of XY diagrams
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.1</b></summary>
|
|
||||||
|
|
||||||
- Fixed `easy comfyLoader` error
|
|
||||||
- Fixed All nodes that contain the value of the image size
|
|
||||||
- Added `easy kSamplerInpainting`
|
|
||||||
- Added `easy pipeToBasicPipe`
|
|
||||||
- Fixed `width` and `height` can not customize in `easy svdLoader`
|
|
||||||
- Fixed all preview image path (Previously, it was not possible to preview the image on the Mac system)
|
|
||||||
- Fixed `vae_name` is not working in `easy fullLoader` and `easy a1111Loader` and `easy comfyLoader`
|
|
||||||
- Fixed `easy fullkSampler` outputs error
|
|
||||||
- Fixed `model_override` is not working in `easy fullLoader`
|
|
||||||
- Fixed `easy hiresFix` error
|
|
||||||
- Fixed `easy xyplot` font file path error
|
|
||||||
- Fixed seed that cannot be fixed when you convert `seed_num` to `easy seed`
|
|
||||||
- Fixed `easy pipeIn` inputs bug
|
|
||||||
- `easy preDetailerFix` have added a new parameter `optional_image`
|
|
||||||
- Fixed `easy zero123Loader` and `easy svdLoader` model into cache.
|
|
||||||
- Added `easy seed`
|
|
||||||
- Fixed `image_output` default value is "Preview"
|
|
||||||
- `easy fullLoader` and `easy a1111Loader` have added a new parameter `a1111_prompt_style`,that can reproduce the same image generated from stable-diffusion-webui on comfyui, but you need to install [ComfyUI_smZNodes](https://github.com/shiimizu/ComfyUI_smZNodes) to use this feature in the current version
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>v1.0.0</b></summary>
|
|
||||||
|
|
||||||
- Added `easy positive` - simple positive prompt text
|
|
||||||
- Added `easy negative` - simple negative prompt text
|
|
||||||
- Added `easy wildcards` - support for wildcards and hint text selected by Lora
|
|
||||||
- Added `easy portraitMaster` - PortraitMaster v2.2
|
|
||||||
- Added `easy loraStack` - Lora stack
|
|
||||||
- Added `easy fullLoader` - full version of the loader
|
|
||||||
- Added `easy zero123Loader` - simple zero123 loader
|
|
||||||
- Added `easy svdLoader` - easy svd loader
|
|
||||||
- Added `easy fullkSampler` - full version of the sampler (no separation)
|
|
||||||
- Added `easy hiresFix` - support for HD repair of Pipe
|
|
||||||
- Added `easy predetailerFix` and `easy DetailerFix` - support for Pipe detail fixing
|
|
||||||
- Added `easy ultralyticsDetectorPipe` and `easy samLoaderPipe` - Detect loader (detail fixed input)
|
|
||||||
- Added `easy pipein` `easy pipeout` - Pipe input and output
|
|
||||||
- Added `easy xyPlot` - simple xyplot (more controllable parameters will be updated in the future)
|
|
||||||
- Added `easy imageRemoveBG` - image to remove background
|
|
||||||
- Added `easy imagePixelPerfect` - image pixel perfect
|
|
||||||
- Added `easy poseEditor` - Pose editor
|
|
||||||
- New UI Theme (Obsidian) - Auto-load UI by default, which can also be changed in the settings
|
|
||||||
|
|
||||||
- Fixed `easy globalSeed` is not working
|
|
||||||
- Fixed an issue where all `seed_num` values were out of order due to [cg-use-everywhere](https://github.com/chrisgoringe/cg-use-everywhere) updating the chart in real time
|
|
||||||
- Fixed `easy imageSize`, `easy imageSizeBySide`, `easy imageSizeByLongerSide` as end nodes
|
|
||||||
- Fixed the bug that `seed_num` (random seed value) could not be read consistently in history
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>Updated at 12/14/2023</b></summary>
|
|
||||||
|
|
||||||
- `easy a1111Loader` and `easy comfyLoader` added `batch_size` of required input parameters
|
|
||||||
- Added the `easy controlnetLoaderADV` node
|
|
||||||
- `easy controlnetLoaderADV` and `easy controlnetLoader` added `control_net ` of optional input parameters
|
|
||||||
- `easy preSampling` and `easy preSamplingAdvanced` added `image_to_latent` optional input parameters
|
|
||||||
- Added the `easy imageSizeBySide` node, which can be output as a long side or a short side
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>Updated at 12/13/2023</b></summary>
|
|
||||||
|
|
||||||
- Added the `easy LLLiteLoader` node, if you have pre-installed the kohya-ss/ControlNet-LLLite-ComfyUI package, please move the model files in the models to `ComfyUI\models\controlnet\` (i.e. in the default controlnet path of comfy, please do not change the file name of the model, otherwise it will not be read).
|
|
||||||
- Modify `easy controlnetLoader` to the bottom of the loader category.
|
|
||||||
- Added size display for `easy imageSize` and `easy imageSizeByLongerSize` outputs.
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><b>Updated at 12/11/2023</b></summary>
|
|
||||||
- Added the `showSpentTime` node to display the time spent on image diffusion and the time spent on VAE decoding images
|
|
||||||
</details>
|
|
||||||
|
|
||||||
## The relevant node package involved
|
|
||||||
|
|
||||||
Disclaimer: Opened source was not easy. I have a lot of respect for the contributions of these original authors. I just did some integration and optimization.
|
|
||||||
|
|
||||||
| Nodes Name(Search Name) | Related libraries | Library-related node |
|
|
||||||
|:-------------------------------|:----------------------------------------------------------------------------|:-------------------------|
|
|
||||||
| easy setNode | [ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) | diffus3.SetNode |
|
|
||||||
| easy getNode | [ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) | diffus3.GetNode |
|
|
||||||
| easy bookmark | [rgthree-comfy](https://github.com/rgthree/rgthree-comfy) | Bookmark 🔖 |
|
|
||||||
| easy portraitMarker | [comfyui-portrait-master](https://github.com/florestefano1975/comfyui-portrait-master) | Portrait Master |
|
|
||||||
| easy LLLiteLoader | [ControlNet-LLLite-ComfyUI](https://github.com/kohya-ss/ControlNet-LLLite-ComfyUI) | LLLiteLoader |
|
|
||||||
| easy globalSeed | [ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) | Global Seed (Inspire) |
|
|
||||||
| easy preSamplingDynamicCFG | [sd-dynamic-thresholding](https://github.com/mcmonkeyprojects/sd-dynamic-thresholding) | DynamicThresholdingFull |
|
|
||||||
| dynamicThresholdingFull | [sd-dynamic-thresholding](https://github.com/mcmonkeyprojects/sd-dynamic-thresholding) | DynamicThresholdingFull |
|
|
||||||
| easy imageInsetCrop | [rgthree-comfy](https://github.com/rgthree/rgthree-comfy) | ImageInsetCrop |
|
|
||||||
| easy poseEditor | [ComfyUI_Custom_Nodes_AlekPet](https://github.com/AlekPet/ComfyUI_Custom_Nodes_AlekPet) | poseNode |
|
|
||||||
| easy preSamplingLayerDiffusion | [ComfyUI-layerdiffusion](https://github.com/huchenlei/ComfyUI-layerdiffusion) | LayeredDiffusionApply... |
|
|
||||||
| easy dynamiCrafterLoader | [ComfyUI-layerdiffusion](https://github.com/ExponentialML/ComfyUI_Native_DynamiCrafter) | Apply Dynamicrafter |
|
|
||||||
| easy imageChooser | [cg-image-picker](https://github.com/chrisgoringe/cg-image-picker) | Preview Chooser |
|
|
||||||
| easy styleAlignedBatchAlign | [style_aligned_comfy](https://github.com/chrisgoringe/cg-image-picker) | styleAlignedBatchAlign |
|
|
||||||
|
|
||||||
|
|
||||||
## Credits
|
|
||||||
|
|
||||||
[ComfyUI](https://github.com/comfyanonymous/ComfyUI) - Powerful and modular Stable Diffusion GUI
|
|
||||||
|
|
||||||
[ComfyUI-ComfyUI-Manager](https://github.com/ltdrdata/ComfyUI-Manager) - ComfyUI Manager
|
|
||||||
|
|
||||||
[tinyterraNodes](https://github.com/TinyTerra/ComfyUI_tinyterraNodes) - Pipe nodes (node bundles) allow users to reduce unnecessary connections
|
|
||||||
|
|
||||||
[ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) - Diffus3 gets and sets points that allow the user to detach the composition of the workflow
|
|
||||||
|
|
||||||
[ComfyUI-Impact-Pack](https://github.com/ltdrdata/ComfyUI-Impact-Pack) - General modpack 1
|
|
||||||
|
|
||||||
[ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) - General Modpack 2
|
|
||||||
|
|
||||||
[ComfyUI-ResAdapter](https://github.com/jiaxiangc/ComfyUI-ResAdapter) - Make model generation independent of training resolution
|
|
||||||
|
|
||||||
[ComfyUI_IPAdapter_plus](https://github.com/cubiq/ComfyUI_IPAdapter_plus) - Style migration
|
|
||||||
|
|
||||||
[ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID) - Face migration
|
|
||||||
|
|
||||||
[ComfyUI_PuLID](https://github.com/cubiq/PuLID_ComfyUI) - Face migration
|
|
||||||
|
|
||||||
[ComfyUI-Custom-Scripts](https://github.com/pythongosssss/ComfyUI-Custom-Scripts) - pyssss🐍
|
|
||||||
|
|
||||||
[cg-image-picker](https://github.com/chrisgoringe/cg-image-picker) - Image Preview Chooser
|
|
||||||
|
|
||||||
[ComfyUI_ExtraModels](https://github.com/city96/ComfyUI_ExtraModels) - DiT custom nodes
|
|
||||||
|
|
||||||
|
|
||||||
## 🌟Stargazers
|
|
||||||
|
|
||||||
My gratitude extends to the generous souls who bestow a star. Your support is much appreciated!
|
|
||||||
|
|
||||||
[](https://github.com/yolain/ComfyUI-Easy-Use/stargazers)
|
|
||||||
+82
-35
@@ -1,36 +1,50 @@
|
|||||||
__version__ = "1.2.0"
|
__version__ = "1.4.1"
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
import json
|
||||||
import os
|
import os
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import importlib
|
import importlib
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
node_list = [
|
|
||||||
"server",
|
|
||||||
"api",
|
|
||||||
"easyNodes",
|
|
||||||
"image",
|
|
||||||
"logic"
|
|
||||||
]
|
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {}
|
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
|
||||||
|
|
||||||
for module_name in node_list:
|
|
||||||
imported_module = importlib.import_module(".py.{}".format(module_name), __name__)
|
|
||||||
NODE_CLASS_MAPPINGS = {**NODE_CLASS_MAPPINGS, **imported_module.NODE_CLASS_MAPPINGS}
|
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {**NODE_DISPLAY_NAME_MAPPINGS, **imported_module.NODE_DISPLAY_NAME_MAPPINGS}
|
|
||||||
|
|
||||||
cwd_path = os.path.dirname(os.path.realpath(__file__))
|
cwd_path = os.path.dirname(os.path.realpath(__file__))
|
||||||
comfy_path = folder_paths.base_path
|
comfy_path = folder_paths.base_path
|
||||||
|
|
||||||
#Wildcards读取
|
NODE_CLASS_MAPPINGS = {}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||||
|
|
||||||
|
try:
|
||||||
|
import comfy.supported_models as _supported_models
|
||||||
|
_HAS_DIFFUSION_XY_SUPPORT = (
|
||||||
|
hasattr(_supported_models, "Anima")
|
||||||
|
and hasattr(_supported_models, "Krea2")
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
_HAS_DIFFUSION_XY_SUPPORT = False
|
||||||
|
|
||||||
|
if not _HAS_DIFFUSION_XY_SUPPORT:
|
||||||
|
print("[ComfyUI-Easy-Use] Anima/Krea2 XY nodes need comfy.supported_models.Anima and Krea2")
|
||||||
|
|
||||||
|
importlib.import_module('.py.routes', __name__)
|
||||||
|
importlib.import_module('.py.server', __name__)
|
||||||
|
nodes_list = ["util", "seed", "prompt", "loaders", "adapter", "inpaint", "preSampling", "samplers", "fix", "pipe", "xyplot", "image", "logic", "api", "deprecated"]
|
||||||
|
for module_name in nodes_list:
|
||||||
|
imported_module = importlib.import_module(".py.nodes.{}".format(module_name), __name__)
|
||||||
|
NODE_CLASS_MAPPINGS = {**NODE_CLASS_MAPPINGS, **imported_module.NODE_CLASS_MAPPINGS}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {**NODE_DISPLAY_NAME_MAPPINGS, **imported_module.NODE_DISPLAY_NAME_MAPPINGS}
|
||||||
|
|
||||||
|
#Wildcards
|
||||||
from .py.libs.wildcards import read_wildcard_dict
|
from .py.libs.wildcards import read_wildcard_dict
|
||||||
wildcards_path = os.path.join(os.path.dirname(__file__), "wildcards")
|
wildcards_path = os.path.join(os.path.dirname(__file__), "wildcards")
|
||||||
if os.path.exists(wildcards_path):
|
if not os.path.exists(wildcards_path):
|
||||||
read_wildcard_dict(wildcards_path)
|
|
||||||
else:
|
|
||||||
os.mkdir(wildcards_path)
|
os.mkdir(wildcards_path)
|
||||||
|
|
||||||
|
# Add custom wildcards example
|
||||||
|
example_path = os.path.join(wildcards_path, "example.txt")
|
||||||
|
if not os.path.exists(example_path):
|
||||||
|
with open(example_path, 'w') as f:
|
||||||
|
text = "blue\nred\nyellow\ngreen\nbrown\npink\npurple\norange\nblack\nwhite"
|
||||||
|
f.write(text)
|
||||||
|
read_wildcard_dict(wildcards_path)
|
||||||
|
|
||||||
#Styles
|
#Styles
|
||||||
styles_path = os.path.join(os.path.dirname(__file__), "styles")
|
styles_path = os.path.join(os.path.dirname(__file__), "styles")
|
||||||
@@ -42,19 +56,52 @@ else:
|
|||||||
os.mkdir(styles_path)
|
os.mkdir(styles_path)
|
||||||
os.mkdir(samples_path)
|
os.mkdir(samples_path)
|
||||||
|
|
||||||
# Model thumbnails
|
# Add custom styles example
|
||||||
from .py.libs.add_resources import add_static_resource
|
example_path = os.path.join(styles_path, "your_styles.json.example")
|
||||||
from .py.libs.model import easyModelManager
|
if not os.path.exists(example_path):
|
||||||
model_config = easyModelManager().models_config
|
import json
|
||||||
for model in model_config:
|
data = [
|
||||||
paths = folder_paths.get_folder_paths(model)
|
{
|
||||||
for path in paths:
|
"name": "Example Style",
|
||||||
if not Path(path).exists():
|
"name_cn": "示例样式",
|
||||||
continue
|
"prompt": "(masterpiece), (best quality), (ultra-detailed), {prompt} ",
|
||||||
add_static_resource(path, path, limit=True)
|
"negative_prompt": "text, watermark, logo"
|
||||||
|
},
|
||||||
|
]
|
||||||
|
# Write to file
|
||||||
|
with open(example_path, 'w', encoding='utf-8') as f:
|
||||||
|
json.dump(data, f, indent=4, ensure_ascii=False)
|
||||||
|
|
||||||
|
|
||||||
|
web_default_version = 'v2'
|
||||||
|
# web directory
|
||||||
|
config_path = os.path.join(cwd_path, "config.yaml")
|
||||||
|
if os.path.isfile(config_path):
|
||||||
|
with open(config_path, 'r') as f:
|
||||||
|
data = yaml.load(f, Loader=yaml.FullLoader)
|
||||||
|
if data and "WEB_VERSION" in data:
|
||||||
|
directory = f"web_version/{data['WEB_VERSION']}"
|
||||||
|
with open(config_path, 'w') as f:
|
||||||
|
yaml.dump(data, f)
|
||||||
|
elif web_default_version != 'v1':
|
||||||
|
if not data:
|
||||||
|
data = {'WEB_VERSION': web_default_version}
|
||||||
|
elif 'WEB_VERSION' not in data:
|
||||||
|
data = {**data, 'WEB_VERSION': web_default_version}
|
||||||
|
with open(config_path, 'w') as f:
|
||||||
|
yaml.dump(data, f)
|
||||||
|
directory = f"web_version/{web_default_version}"
|
||||||
|
else:
|
||||||
|
directory = f"web_version/v1"
|
||||||
|
if not os.path.exists(os.path.join(cwd_path, directory)):
|
||||||
|
print(f"web root {data['WEB_VERSION']} not found, using default")
|
||||||
|
directory = f"web_version/{web_default_version}"
|
||||||
|
WEB_DIRECTORY = directory
|
||||||
|
else:
|
||||||
|
directory = f"web_version/{web_default_version}"
|
||||||
|
WEB_DIRECTORY = directory
|
||||||
|
|
||||||
WEB_DIRECTORY = "./web"
|
|
||||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', "WEB_DIRECTORY"]
|
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', "WEB_DIRECTORY"]
|
||||||
|
|
||||||
|
print(f'\033[34m[ComfyUI-Easy-Use] server: \033[0mv{__version__} \033[92mLoaded\033[0m')
|
||||||
print(f'\033[34mComfy-Easy-Use v{__version__}: \033[92mLoaded\033[0m')
|
print(f'\033[34m[ComfyUI-Easy-Use] web root: \033[0m{os.path.join(cwd_path, directory)} \033[92mLoaded\033[0m')
|
||||||
|
|||||||
+5
-1
@@ -1,6 +1,7 @@
|
|||||||
@echo off
|
@echo off
|
||||||
|
|
||||||
set "requirements_txt=%~dp0\requirements.txt"
|
set "requirements_txt=%~dp0\requirements.txt"
|
||||||
|
set "requirements_repair_txt=%~dp0\repair_dependency_list.txt"
|
||||||
set "python_exec=..\..\..\python_embeded\python.exe"
|
set "python_exec=..\..\..\python_embeded\python.exe"
|
||||||
set "aki_python_exec=..\..\python\python.exe"
|
set "aki_python_exec=..\..\python\python.exe"
|
||||||
|
|
||||||
@@ -12,7 +13,10 @@ if exist "%python_exec%" (
|
|||||||
)^
|
)^
|
||||||
else if exist "%aki_python_exec%" (
|
else if exist "%aki_python_exec%" (
|
||||||
echo Installing with ComfyUI Aki
|
echo Installing with ComfyUI Aki
|
||||||
"%python_exec%" -s -m pip install -r "%requirements_txt%"
|
"%aki_python_exec%" -s -m pip install -r "%requirements_txt%"
|
||||||
|
for /f "delims=" %%i in (%requirements_repair_txt%) do (
|
||||||
|
%aki_python_exec% -s -m pip install -i https://pypi.tuna.tsinghua.edu.cn/simple "%%i"
|
||||||
|
)
|
||||||
)^
|
)^
|
||||||
else (
|
else (
|
||||||
echo Installing with system Python
|
echo Installing with system Python
|
||||||
|
|||||||
+24
@@ -0,0 +1,24 @@
|
|||||||
|
#!/bin/bash
|
||||||
|
|
||||||
|
requirements_txt="$(dirname "$0")/requirements.txt"
|
||||||
|
requirements_repair_txt="$(dirname "$0")/repair_dependency_list.txt"
|
||||||
|
python_exec="../../../python_embeded/python.exe"
|
||||||
|
aki_python_exec="../../python/python.exe"
|
||||||
|
|
||||||
|
echo "Installing EasyUse Requirements..."
|
||||||
|
|
||||||
|
if [ -f "$python_exec" ]; then
|
||||||
|
echo "Installing with ComfyUI Portable"
|
||||||
|
"$python_exec" -s -m pip install -r "$requirements_txt"
|
||||||
|
elif [ -f "$aki_python_exec" ]; then
|
||||||
|
echo "Installing with ComfyUI Aki"
|
||||||
|
"$aki_python_exec" -s -m pip install -r "$requirements_txt"
|
||||||
|
while IFS= read -r line; do
|
||||||
|
"$aki_python_exec" -s -m pip install -i https://pypi.tuna.tsinghua.edu.cn/simple "$line"
|
||||||
|
done < "$requirements_repair_txt"
|
||||||
|
else
|
||||||
|
echo "Installing with system Python"
|
||||||
|
pip install -r "$requirements_txt"
|
||||||
|
fi
|
||||||
|
|
||||||
|
read -p "Press any key to continue..."
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
{
|
||||||
|
"settingsCategories": {
|
||||||
|
"Hotkeys": "Hotkeys",
|
||||||
|
"Nodes": "Nodes",
|
||||||
|
"NodesMap": "NodesMap",
|
||||||
|
"StylesSelector": "StylesSelector"
|
||||||
|
},
|
||||||
|
"nodeCategories": {
|
||||||
|
"Util": "Util",
|
||||||
|
"Seed": "Seed",
|
||||||
|
"Prompt": "Prompt",
|
||||||
|
"Loaders": "Loaders",
|
||||||
|
"Adapter": "Adapter",
|
||||||
|
"Inpaint": "Inpaint",
|
||||||
|
"PreSampling": "PreSampling",
|
||||||
|
"Sampler": "Sampler",
|
||||||
|
"Fix": "Fix",
|
||||||
|
"Pipe": "Pipe",
|
||||||
|
"XY Inputs": "XY Inputs",
|
||||||
|
"Image": "Image",
|
||||||
|
"Segmentation": "Segmentation",
|
||||||
|
"\uD83D\uDEAB Deprecated": "\uD83D\uDEAB Deprecated",
|
||||||
|
"Type": "Type",
|
||||||
|
"Math": "Math",
|
||||||
|
"Switch": "Switch",
|
||||||
|
"Index Switch": "Index Switch",
|
||||||
|
"While Loop": "While Loop",
|
||||||
|
"For Loop": "For Loop",
|
||||||
|
"LoadImage": "Load Image"
|
||||||
|
}
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,67 @@
|
|||||||
|
{
|
||||||
|
"EasyUse_Hotkeys_AddGroup": {
|
||||||
|
"name": "Enable Shift+g to add the selected nodes to a group",
|
||||||
|
"tooltip": "From v1.2.39, you can use Ctrl+g instead"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_cleanVRAMUsed": {
|
||||||
|
"name": "Enable Shift+r to unload model and node cache"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_toggleNodesMap": {
|
||||||
|
"name": "Enable Shift+m to toggle nodes map"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_AlignSelectedNodes": {
|
||||||
|
"name": "Enable Shift+Up/Down/Left/Right and Shift+Ctrl+Alt+Left/Right to align selected nodes",
|
||||||
|
"tooltip": "Shift+Up/Down/Left/Right can align selected nodes, Shift+Ctrl+Alt+Left/Right can distribute nodes horizontally/vertically"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NormalizeSelectedNodes": {
|
||||||
|
"name": "Enable Shift+Ctrl+Left/Right to normalize selected nodes",
|
||||||
|
"tooltip": "Enable Shift+Ctrl+Left to normalize width and Shift+Ctrl+Right to normalize height"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NodesTemplate": {
|
||||||
|
"name": "Enable Alt+1~9 to paste node templates into the workflow"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_JumpNearestNodes": {
|
||||||
|
"name": "Enable Up/Down/Left/Right to jump to the nearest node"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_SubDirectories": {
|
||||||
|
"name": "Enable automatic nesting of subdirectories in the context menu"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_ModelsThumbnails": {
|
||||||
|
"name": "Enable model preview thumbnails"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_NodesSort": {
|
||||||
|
"name": "Enable A~Z sorting of new nodes in the context menu"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_QuickOptions": {
|
||||||
|
"name": "Use three quick buttons in the context menu",
|
||||||
|
"options": {
|
||||||
|
"At the forefront": "At the forefront",
|
||||||
|
"At the end": "At the end",
|
||||||
|
"Disable": "Disable"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_Runtime": {
|
||||||
|
"name": "Enable node runtime display"
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_ChainGetSet": {
|
||||||
|
"name": "Enable chaining of get and set points with the parent node"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Sorting": {
|
||||||
|
"name": "Manage nodes group sorting mode",
|
||||||
|
"tooltip": "Automatically sort by default. If set to manual, groups can be drag and dropped and the order will be saved.",
|
||||||
|
"options": {
|
||||||
|
"Auto sorting": "Auto sorting",
|
||||||
|
"Manual drag&drop sorting": "Manual drag&drop sorting"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayNodeID": {
|
||||||
|
"name": "Enable node ID display"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayGroupOnly": {
|
||||||
|
"name": "Show groups only"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Enable": {
|
||||||
|
"name": "Enable Group Map",
|
||||||
|
"tooltip": "You need to refresh the page to update successfully"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
{
|
||||||
|
"settingsCategories": {
|
||||||
|
"Hotkeys": "Raccourcis",
|
||||||
|
"Nodes": "Nœuds",
|
||||||
|
"NodesMap": "Carte des nœuds"
|
||||||
|
},
|
||||||
|
"nodeCategories": {
|
||||||
|
"Util": "Utilitaire",
|
||||||
|
"Seed": "Graine",
|
||||||
|
"Prompt": "Prompt",
|
||||||
|
"Loaders": "Chargeurs",
|
||||||
|
"Adapter": "Adaptateur",
|
||||||
|
"Inpaint": "Retouche",
|
||||||
|
"PreSampling": "Pré-échantillonnage",
|
||||||
|
"Sampler": "Échantillonneur",
|
||||||
|
"Fix": "Correction",
|
||||||
|
"Pipe": "Pipeline",
|
||||||
|
"XY Inputs": "Entrées XY",
|
||||||
|
"Image": "Image",
|
||||||
|
"Segmentation": "Segmentation",
|
||||||
|
"\uD83D\uDEAB Deprecated": "\uD83D\uDEAB Obsolète",
|
||||||
|
"Type": "Type",
|
||||||
|
"Math": "Mathématiques",
|
||||||
|
"Switch": "Interrupteur",
|
||||||
|
"Index Switch": "Interrupteur d'index",
|
||||||
|
"While Loop": "Boucle While",
|
||||||
|
"For Loop": "Boucle For",
|
||||||
|
"LoadImage": "Charger l'image"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
{
|
||||||
|
"EasyUse_Hotkeys_AddGroup": {
|
||||||
|
"name": "Activer Shift+g pour ajouter les nœuds sélectionnés à un groupe",
|
||||||
|
"tooltip": "Depuis la v1.2.39, vous pouvez utiliser Ctrl+g à la place"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_cleanVRAMUsed": {
|
||||||
|
"name": "Activer Shift+r pour décharger le cache du modèle et des nœuds"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_toggleNodesMap": {
|
||||||
|
"name": "Activer Shift+m pour basculer la carte des nœuds"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_AlignSelectedNodes": {
|
||||||
|
"name": "Activer Shift+Up/Down/Left/Right et Shift+Ctrl+Alt+Left/Right pour aligner les nœuds sélectionnés",
|
||||||
|
"tooltip": "Shift+Up/Down/Left/Right peut aligner les nœuds sélectionnés, Shift+Ctrl+Alt+Left/Right peut les répartir horizontalement/verticalement"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NormalizeSelectedNodes": {
|
||||||
|
"name": "Activer Shift+Ctrl+Left/Right pour normaliser les nœuds sélectionnés",
|
||||||
|
"tooltip": "Activer Shift+Ctrl+Left pour normaliser la largeur et Shift+Ctrl+Right pour normaliser la hauteur"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NodesTemplate": {
|
||||||
|
"name": "Activer Alt+1~9 pour coller les modèles de nœuds dans le workflow"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_JumpNearestNodes": {
|
||||||
|
"name": "Activer Up/Down/Left/Right pour passer au nœud le plus proche"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_SubDirectories": {
|
||||||
|
"name": "Activer l'imbrication automatique des sous-répertoires dans le menu contextuel"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_ModelsThumbnails": {
|
||||||
|
"name": "Activer les vignettes d'aperçu du modèle"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_NodesSort": {
|
||||||
|
"name": "Activer le tri A~Z des nouveaux nœuds dans le menu contextuel"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_QuickOptions": {
|
||||||
|
"name": "Utiliser trois boutons rapides dans le menu contextuel",
|
||||||
|
"options": {
|
||||||
|
"At the forefront": "À l'avant-plan",
|
||||||
|
"At the end": "À la fin",
|
||||||
|
"Disable": "Désactiver"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_Runtime": {
|
||||||
|
"name": "Activer l'affichage du temps d'exécution des nœuds"
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_ChainGetSet": {
|
||||||
|
"name": "Activer le chaînage des points get et set avec le nœud parent"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Sorting": {
|
||||||
|
"name": "Gérer le mode de tri des groupes de nœuds",
|
||||||
|
"tooltip": "Tri automatique par défaut. Si défini sur manuel, les groupes peuvent être glissés-déposés et l'ordre sera sauvegardé.",
|
||||||
|
"options": {
|
||||||
|
"Auto sorting": "Tri automatique",
|
||||||
|
"Manual drag&drop sorting": "Tri manuel par glisser-déposer"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayNodeID": {
|
||||||
|
"name": "Activer l'affichage de l'ID du nœud"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayGroupOnly": {
|
||||||
|
"name": "Afficher uniquement les groupes"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Enable": {
|
||||||
|
"name": "Activer la carte des groupes",
|
||||||
|
"tooltip": "Vous devez actualiser la page pour mettre à jour"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
{
|
||||||
|
"settingsCategories": {
|
||||||
|
"Hotkeys": "ショートカットキー",
|
||||||
|
"Nodes": "ノード",
|
||||||
|
"NodesMap": "ノードマップ"
|
||||||
|
},
|
||||||
|
"nodeCategories": {
|
||||||
|
"Util": "ユーティリティ",
|
||||||
|
"Seed": "シード",
|
||||||
|
"Prompt": "プロンプト",
|
||||||
|
"Loaders": "ローダー",
|
||||||
|
"Adapter": "アダプター",
|
||||||
|
"Inpaint": "インペイント",
|
||||||
|
"PreSampling": "プリサンプリング",
|
||||||
|
"Sampler": "サンプラー",
|
||||||
|
"Fix": "フィックス",
|
||||||
|
"Pipe": "パイプ",
|
||||||
|
"XY Inputs": "XY入力",
|
||||||
|
"Image": "画像",
|
||||||
|
"Segmentation": "セグメンテーション",
|
||||||
|
"\uD83D\uDEAB Deprecated": "🚫 非推奨",
|
||||||
|
"Type": "タイプ",
|
||||||
|
"Math": "数学",
|
||||||
|
"Switch": "スイッチ",
|
||||||
|
"Index Switch": "インデックススイッチ",
|
||||||
|
"While Loop": "Whileループ",
|
||||||
|
"For Loop": "Forループ",
|
||||||
|
"LoadImage": "画像読み込み"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
{
|
||||||
|
"EasyUse_Hotkeys_AddGroup": {
|
||||||
|
"name": "Shift+gを使用して選択したノードをグループに追加する",
|
||||||
|
"tooltip": "v1.2.39以降、Ctrl+gが使用できます"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_cleanVRAMUsed": {
|
||||||
|
"name": "Shift+rを使用してモデルおよびノードキャッシュをアンロードする"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_toggleNodesMap": {
|
||||||
|
"name": "Shift+mを使用してノードマップを表示/非表示にします"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_AlignSelectedNodes": {
|
||||||
|
"name": "Shift+上/下/左/右およびShift+Ctrl+Alt+左/右を使用して選択したノードを整列する",
|
||||||
|
"tooltip": "Shift+上/下/左/右で選択したノードを整列し、Shift+Ctrl+Alt+左/右で水平方向/垂直方向に分布させる"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NormalizeSelectedNodes": {
|
||||||
|
"name": "Shift+Ctrl+左/右を使用して選択したノードのサイズを正規化する",
|
||||||
|
"tooltip": "Shift+Ctrl+左で幅を、Shift+Ctrl+右で高さを正規化する"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NodesTemplate": {
|
||||||
|
"name": "Alt+1~9を使用してワークフローにノードテンプレートを貼り付ける"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_JumpNearestNodes": {
|
||||||
|
"name": "上/下/左/右を使用して最も近いノードにジャンプする"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_SubDirectories": {
|
||||||
|
"name": "コンテキストメニューでサブディレクトリを自動でネストする"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_ModelsThumbnails": {
|
||||||
|
"name": "モデルプレビューサムネイルを有効にする"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_NodesSort": {
|
||||||
|
"name": "コンテキストメニューで新規ノードをA~Z順に並べ替える"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_QuickOptions": {
|
||||||
|
"name": "コンテキストメニューで3つのクイックボタンを使用する",
|
||||||
|
"options": {
|
||||||
|
"At the forefront": "最前面に",
|
||||||
|
"At the end": "最後に",
|
||||||
|
"Disable": "無効"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_Runtime": {
|
||||||
|
"name": "ノードの実行時間表示を有効にする"
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_ChainGetSet": {
|
||||||
|
"name": "親ノードと取得/設定ポイントを連結することを有効にする"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Sorting": {
|
||||||
|
"name": "ノードグループの並べ替えモードを管理する",
|
||||||
|
"tooltip": "デフォルトで自動的に並べ替えます。マニュアルに設定した場合、グループをドラッグアンドドロップで並べ替え、順序が保存されます。",
|
||||||
|
"options": {
|
||||||
|
"Auto sorting": "自動並べ替え",
|
||||||
|
"Manual drag&drop sorting": "手動ドラッグアンドドロップによる並べ替え"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayNodeID": {
|
||||||
|
"name": "ノードIDの表示を有効にする"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayGroupOnly": {
|
||||||
|
"name": "グループのみ表示する"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Enable": {
|
||||||
|
"name": "グループマップを有効にする",
|
||||||
|
"tooltip": "ページを更新する必要があります"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
{
|
||||||
|
"settingsCategories": {
|
||||||
|
"Hotkeys": "단축키",
|
||||||
|
"Nodes": "노드",
|
||||||
|
"NodesMap": "노드 맵"
|
||||||
|
},
|
||||||
|
"nodeCategories": {
|
||||||
|
"Util": "유틸",
|
||||||
|
"Seed": "시드",
|
||||||
|
"Prompt": "프롬프트",
|
||||||
|
"Loaders": "로더",
|
||||||
|
"Adapter": "어댑터",
|
||||||
|
"Inpaint": "인페인트",
|
||||||
|
"PreSampling": "사전 샘플링",
|
||||||
|
"Sampler": "샘플러",
|
||||||
|
"Fix": "픽스",
|
||||||
|
"Pipe": "파이프",
|
||||||
|
"XY Inputs": "XY 입력",
|
||||||
|
"Image": "이미지",
|
||||||
|
"Segmentation": "분할",
|
||||||
|
"\uD83D\uDEAB Deprecated": "\uD83D\uDEAB 사용 중단",
|
||||||
|
"Type": "유형",
|
||||||
|
"Math": "수학",
|
||||||
|
"Switch": "스위치",
|
||||||
|
"Index Switch": "인덱스 스위치",
|
||||||
|
"While Loop": "while 루프",
|
||||||
|
"For Loop": "for 루프",
|
||||||
|
"LoadImage": "이미지 로드"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
{
|
||||||
|
"EasyUse_Hotkeys_AddGroup": {
|
||||||
|
"name": "Shift+g 를 사용하여 선택된 노드를 그룹에 추가합니다",
|
||||||
|
"tooltip": "v1.2.39부터는 Ctrl+g 를 사용할 수 있습니다"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_cleanVRAMUsed": {
|
||||||
|
"name": "Shift+r 를 사용하여 모델 및 노드 캐시를 언로드합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_toggleNodesMap": {
|
||||||
|
"name": "Shift+m 를 사용하여 노드 맵을 전환합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_AlignSelectedNodes": {
|
||||||
|
"name": "Shift+Up/Down/Left/Right 와 Shift+Ctrl+Alt+Left/Right 를 사용하여 선택된 노드를 정렬합니다",
|
||||||
|
"tooltip": "Shift+Up/Down/Left/Right 는 선택된 노드를 정렬하며, Shift+Ctrl+Alt+Left/Right 는 노드를 수평/수직으로 분배합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NormalizeSelectedNodes": {
|
||||||
|
"name": "Shift+Ctrl+Left/Right 를 사용하여 선택된 노드를 정규화합니다",
|
||||||
|
"tooltip": "Shift+Ctrl+Left 는 너비를, Shift+Ctrl+Right 는 높이를 정규화합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NodesTemplate": {
|
||||||
|
"name": "Alt+1~9 를 사용하여 워크플로우에 노드 템플릿을 붙여넣습니다"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_JumpNearestNodes": {
|
||||||
|
"name": "Up/Down/Left/Right 를 사용하여 가장 가까운 노드로 이동합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_SubDirectories": {
|
||||||
|
"name": "컨텍스트 메뉴에서 자동으로 하위 디렉토리를 중첩합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_ModelsThumbnails": {
|
||||||
|
"name": "모델 미리보기 썸네일을 활성화합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_NodesSort": {
|
||||||
|
"name": "컨텍스트 메뉴에서 새로운 노드를 A~Z 순으로 정렬합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_QuickOptions": {
|
||||||
|
"name": "컨텍스트 메뉴에 3개의 빠른 옵션 버튼을 사용합니다",
|
||||||
|
"options": {
|
||||||
|
"At the forefront": "앞쪽에",
|
||||||
|
"At the end": "뒤쪽에",
|
||||||
|
"Disable": "비활성화"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_Runtime": {
|
||||||
|
"name": "노드 실행 시간 표시를 활성화합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_ChainGetSet": {
|
||||||
|
"name": "부모 노드와 연결된 get/ set 포인트 체이닝을 활성화합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Sorting": {
|
||||||
|
"name": "노드 그룹 정렬 모드를 관리합니다",
|
||||||
|
"tooltip": "기본값은 자동 정렬입니다. 수동으로 설정하면 그룹을 드래그 앤 드롭할 수 있으며 순서가 저장됩니다.",
|
||||||
|
"options": {
|
||||||
|
"Auto sorting": "자동 정렬",
|
||||||
|
"Manual drag&drop sorting": "수동 드래그 앤 드롭 정렬"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayNodeID": {
|
||||||
|
"name": "노드 ID 표시를 활성화합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayGroupOnly": {
|
||||||
|
"name": "그룹만 표시합니다"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Enable": {
|
||||||
|
"name": "그룹 맵을 활성화합니다",
|
||||||
|
"tooltip": "업데이트를 위해 페이지를 새로고침해야 합니다"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
{
|
||||||
|
"settingsCategories": {
|
||||||
|
"Hotkeys": "Горячие клавиши",
|
||||||
|
"Nodes": "Узлы",
|
||||||
|
"NodesMap": "Карта узлов"
|
||||||
|
},
|
||||||
|
"nodeCategories": {
|
||||||
|
"Util": "Утилиты",
|
||||||
|
"Seed": "Сид",
|
||||||
|
"Prompt": "Подсказка",
|
||||||
|
"Loaders": "Загрузчики",
|
||||||
|
"Adapter": "Адаптер",
|
||||||
|
"Inpaint": "Ретушь",
|
||||||
|
"PreSampling": "Предвыборка",
|
||||||
|
"Sampler": "Сэмплер",
|
||||||
|
"Fix": "Исправление",
|
||||||
|
"Pipe": "Конвейер",
|
||||||
|
"XY Inputs": "Ввод XY",
|
||||||
|
"Image": "Изображение",
|
||||||
|
"Segmentation": "Сегментация",
|
||||||
|
"\uD83D\uDEAB Deprecated": "\uD83D\uDEAB Устарело",
|
||||||
|
"Type": "Тип",
|
||||||
|
"Math": "Математика",
|
||||||
|
"Switch": "Переключатель",
|
||||||
|
"Index Switch": "Переключатель индексов",
|
||||||
|
"While Loop": "Цикл while",
|
||||||
|
"For Loop": "Цикл for",
|
||||||
|
"LoadImage": "Загрузка изображения"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
{
|
||||||
|
"EasyUse_Hotkeys_AddGroup": {
|
||||||
|
"name": "Включить Shift+g для добавления выделенных узлов в группу",
|
||||||
|
"tooltip": "Начиная с версии v1.2.39, можно использовать Ctrl+g"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_cleanVRAMUsed": {
|
||||||
|
"name": "Включить Shift+r для выгрузки модели и кэша узлов"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_toggleNodesMap": {
|
||||||
|
"name": "Включить Shift+m для переключения карты узлов"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_AlignSelectedNodes": {
|
||||||
|
"name": "Включить Shift+Стрелки для выравнивания выделенных узлов и Shift+Ctrl+Alt+Стрелки для распределения узлов по горизонтали/вертикали",
|
||||||
|
"tooltip": "Shift+Стрелки выравнивают выделенные узлы, Shift+Ctrl+Alt+Стрелки распределяют узлы по горизонтали/вертикали"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NormalizeSelectedNodes": {
|
||||||
|
"name": "Включить Shift+Ctrl+Стрелки для нормализации выделенных узлов",
|
||||||
|
"tooltip": "Включить Shift+Ctrl+Лево для нормализации ширины и Shift+Ctrl+Право для нормализации высоты"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NodesTemplate": {
|
||||||
|
"name": "Включить Alt+1~9 для вставки шаблонов узлов в рабочий процесс"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_JumpNearestNodes": {
|
||||||
|
"name": "Включить Стрелки для перехода к ближайшему узлу"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_SubDirectories": {
|
||||||
|
"name": "Включить автоматическое вложение подкаталогов в контекстном меню"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_ModelsThumbnails": {
|
||||||
|
"name": "Включить превью миниатюр моделей"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_NodesSort": {
|
||||||
|
"name": "Включить A~Z сортировку новых узлов в контекстном меню"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_QuickOptions": {
|
||||||
|
"name": "Использовать три быстрых кнопки в контекстном меню",
|
||||||
|
"options": {
|
||||||
|
"At the forefront": "В начале",
|
||||||
|
"At the end": "В конце",
|
||||||
|
"Disable": "Отключено"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_Runtime": {
|
||||||
|
"name": "Включить отображение времени выполнения узлов"
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_ChainGetSet": {
|
||||||
|
"name": "Включить связывание точек получения и установки с родительским узлом"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Sorting": {
|
||||||
|
"name": "Управление режимом сортировки групп узлов",
|
||||||
|
"tooltip": "По умолчанию автоматическая сортировка. При ручном режиме группы можно перемещать методом перетаскивания, и порядок будет сохранён.",
|
||||||
|
"options": {
|
||||||
|
"Auto sorting": "Автоматическая сортировка",
|
||||||
|
"Manual drag&drop sorting": "Ручная сортировка перетаскиванием"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayNodeID": {
|
||||||
|
"name": "Включить отображение ID узлов"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayGroupOnly": {
|
||||||
|
"name": "Показывать только группы"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Enable": {
|
||||||
|
"name": "Включить карту групп",
|
||||||
|
"tooltip": "Необходимо обновить страницу для успешного обновления"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,33 @@
|
|||||||
|
{
|
||||||
|
"settingsCategories": {
|
||||||
|
"Hotkeys": "快捷键",
|
||||||
|
"Nodes": "节点相关",
|
||||||
|
"NodesMap": "管理节点组",
|
||||||
|
"StylesSelector": "样式选择器",
|
||||||
|
"MultiAngle": "摄影机多角度提示词"
|
||||||
|
},
|
||||||
|
"nodeCategories": {
|
||||||
|
"Util": "工具",
|
||||||
|
"Seed": "随机种",
|
||||||
|
"Prompt": "提示词",
|
||||||
|
"Loaders": "模型加载器",
|
||||||
|
"Adapter": "模型适配器",
|
||||||
|
"Inpaint": "内补重绘",
|
||||||
|
"PreSampling": "预采样参数",
|
||||||
|
"Sampler": "采样器",
|
||||||
|
"Fix": "修复相关",
|
||||||
|
"Pipe": "节点束",
|
||||||
|
"XY Inputs": "XY图表输入项",
|
||||||
|
"Image": "图像",
|
||||||
|
"Segmentation": "分割",
|
||||||
|
"Logic": "逻辑",
|
||||||
|
"\uD83D\uDEAB Deprecated": "\uD83D\uDEAB 已弃用",
|
||||||
|
"Type": "类型",
|
||||||
|
"Math": "数学计算",
|
||||||
|
"Switch": "开关",
|
||||||
|
"Index Switch": "索引开关",
|
||||||
|
"While Loop": "While循环",
|
||||||
|
"For Loop": "For循环",
|
||||||
|
"LoadImage": "加载图像"
|
||||||
|
}
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,86 @@
|
|||||||
|
{
|
||||||
|
"EasyUse_Hotkeys_AddGroup": {
|
||||||
|
"name": "启用 Shift+g 键将选中的节点添加一个组",
|
||||||
|
"tooltip": "从v1.2.39开始,可以使用Ctrl+g代替"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_cleanVRAMUsed": {
|
||||||
|
"name": "启用 Shift+r 键卸载模型和节点缓存"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_toggleNodesMap": {
|
||||||
|
"name": "启用 Shift+m 键显隐管理节点组"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_AlignSelectedNodes": {
|
||||||
|
"name": "启用 Shift+上/下/左/右 和 Shift+Ctrl+Alt+左/右 键对齐选中的节点",
|
||||||
|
"tooltip": "Shift+上/下/左/右 可以对齐选中的节点, Shift+Ctrl+Alt+左/右 可以水平/垂直分布节点"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NormalizeSelectedNodes": {
|
||||||
|
"name": "启用 Shift+Ctrl+左/右 键规范化选中的节点",
|
||||||
|
"tooltip": "启用 Shift+Ctrl+左 键规范化宽度和 Shift+Ctrl+右 键规范化高度"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_NodesTemplate": {
|
||||||
|
"name": "启用 Alt+1~9 从节点模板粘贴到工作流中"
|
||||||
|
},
|
||||||
|
"EasyUse_Hotkeys_JumpNearestNodes": {
|
||||||
|
"name": "启用 上/下/左/右 键跳转到最近的前后节点"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_SubDirectories": {
|
||||||
|
"name": "启用上下文菜单自动嵌套子目录"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_ModelsThumbnails": {
|
||||||
|
"name": "启动模型预览图显示"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_NodesSort": {
|
||||||
|
"name": "启用右键菜单中新建节点A~Z排序"
|
||||||
|
},
|
||||||
|
"EasyUse_ContextMenu_QuickOptions": {
|
||||||
|
"name": "在右键菜单中使用三个快捷按钮",
|
||||||
|
"options": {
|
||||||
|
"At the forefront": "在最前面",
|
||||||
|
"At the end": "在最后面",
|
||||||
|
"Disable": "禁用"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_Runtime": {
|
||||||
|
"name": "启动节点运行时间显示"
|
||||||
|
},
|
||||||
|
"EasyUse_Nodes_ChainGetSet": {
|
||||||
|
"name": "启用将获取点和设置点与父节点链在一起"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Sorting": {
|
||||||
|
"name": "管理节点组排序模式",
|
||||||
|
"tooltip": "默认自动排序,如果设置为手动,组可以拖放并保存排序结果。",
|
||||||
|
"options": {
|
||||||
|
"Auto sorting": "自动排序",
|
||||||
|
"Manual drag&drop sorting": "手动拖拽排序"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayNodeID": {
|
||||||
|
"name": "启用节点ID显示"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_DisplayGroupOnly": {
|
||||||
|
"name": "仅显示组"
|
||||||
|
},
|
||||||
|
"EasyUse_NodesMap_Enable": {
|
||||||
|
"name": "启用管理节点组",
|
||||||
|
"tooltip": "您需要刷新页面以成功更新"
|
||||||
|
},
|
||||||
|
"EasyUse_StylesSelector_DisplayType": {
|
||||||
|
"name": "样式选择器显示类型",
|
||||||
|
"tooltip": "样式选择器显示类型,如果设置为“网格”,则显示为网格,如果设置为“列表”,则显示为列表",
|
||||||
|
"options": {
|
||||||
|
"Grid": "网格",
|
||||||
|
"List": "列表"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"EasyUse_MultiAngle_InvertRotate": {
|
||||||
|
"name": "启用反转旋转模式",
|
||||||
|
"tooltip": "在多角度节点中启用反转旋转模式,使旋转方向与大多数3D软件一致"
|
||||||
|
},
|
||||||
|
"EasyUse_MultiAngle_HollowMode": {
|
||||||
|
"name": "启用多角度镂空展示模式",
|
||||||
|
"tooltip": "在多角度节点中启用镂空展示模式,可以更直观地查看相机角度"
|
||||||
|
},
|
||||||
|
"EasyUse_MultiAngle_AddAnglePrompt": {
|
||||||
|
"name": "启用添加多角度提示词"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -31,6 +31,4 @@ add_folder_path_and_extensions("mediapipe", [os.path.join(model_path, "mediapipe
|
|||||||
add_folder_path_and_extensions("inpaint", [os.path.join(model_path, "inpaint")], folder_paths.supported_pt_extensions)
|
add_folder_path_and_extensions("inpaint", [os.path.join(model_path, "inpaint")], folder_paths.supported_pt_extensions)
|
||||||
add_folder_path_and_extensions("prompt_generator", [os.path.join(model_path, "prompt_generator")], folder_paths.supported_pt_extensions)
|
add_folder_path_and_extensions("prompt_generator", [os.path.join(model_path, "prompt_generator")], folder_paths.supported_pt_extensions)
|
||||||
add_folder_path_and_extensions("t5", [os.path.join(model_path, "t5")], folder_paths.supported_pt_extensions)
|
add_folder_path_and_extensions("t5", [os.path.join(model_path, "t5")], folder_paths.supported_pt_extensions)
|
||||||
|
add_folder_path_and_extensions("llm", [os.path.join(model_path, "LLM")], folder_paths.supported_pt_extensions)
|
||||||
add_folder_path_and_extensions("checkpoints_thumb", [os.path.join(model_path, "checkpoints")], image_suffixs)
|
|
||||||
add_folder_path_and_extensions("loras_thumb", [os.path.join(model_path, "loras")], image_suffixs)
|
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
from .libs.loader import easyLoader
|
||||||
|
from .libs.sampler import easySampler
|
||||||
|
|
||||||
|
sampler = easySampler()
|
||||||
|
easyCache = easyLoader()
|
||||||
|
|
||||||
|
|||||||
@@ -1,298 +0,0 @@
|
|||||||
import os
|
|
||||||
import hashlib
|
|
||||||
import sys
|
|
||||||
import json
|
|
||||||
import shutil
|
|
||||||
import folder_paths
|
|
||||||
from folder_paths import get_directory_by_type
|
|
||||||
from server import PromptServer
|
|
||||||
from .config import RESOURCES_DIR, FOOOCUS_STYLES_DIR, FOOOCUS_STYLES_SAMPLES
|
|
||||||
from .logic import ConvertAnything
|
|
||||||
from .libs.model import easyModelManager
|
|
||||||
from .libs.utils import getMetadata, cleanGPUUsedForce, get_local_filepath
|
|
||||||
from .libs.cache import remove_cache
|
|
||||||
from .libs.translate import has_chinese, zh_to_en
|
|
||||||
|
|
||||||
try:
|
|
||||||
import aiohttp
|
|
||||||
from aiohttp import web
|
|
||||||
except ImportError:
|
|
||||||
print("Module 'aiohttp' not installed. Please install it via:")
|
|
||||||
print("pip install aiohttp")
|
|
||||||
sys.exit()
|
|
||||||
|
|
||||||
@PromptServer.instance.routes.post("/easyuse/cleangpu")
|
|
||||||
def cleanGPU(request):
|
|
||||||
try:
|
|
||||||
cleanGPUUsedForce()
|
|
||||||
remove_cache('*')
|
|
||||||
return web.Response(status=200)
|
|
||||||
except Exception as e:
|
|
||||||
return web.Response(status=500)
|
|
||||||
pass
|
|
||||||
|
|
||||||
@PromptServer.instance.routes.post("/easyuse/translate")
|
|
||||||
async def translate(request):
|
|
||||||
post = await request.post()
|
|
||||||
text = post.get("text")
|
|
||||||
if has_chinese(text):
|
|
||||||
return web.json_response({"text": zh_to_en([text])[0]})
|
|
||||||
else:
|
|
||||||
return web.json_response({"text": text})
|
|
||||||
|
|
||||||
@PromptServer.instance.routes.get("/easyuse/reboot")
|
|
||||||
def reboot(request):
|
|
||||||
try:
|
|
||||||
sys.stdout.close_log()
|
|
||||||
except Exception as e:
|
|
||||||
pass
|
|
||||||
|
|
||||||
return os.execv(sys.executable, [sys.executable] + sys.argv)
|
|
||||||
|
|
||||||
# parse csv
|
|
||||||
@PromptServer.instance.routes.post("/easyuse/upload/csv")
|
|
||||||
async def parse_csv(request):
|
|
||||||
post = await request.post()
|
|
||||||
csv = post.get("csv")
|
|
||||||
if csv and csv.file:
|
|
||||||
file = csv.file
|
|
||||||
text = ''
|
|
||||||
for line in file.readlines():
|
|
||||||
line = str(line.strip())
|
|
||||||
line = line.replace("'", "").replace("b",'')
|
|
||||||
text += line + '; \n'
|
|
||||||
return web.json_response(text)
|
|
||||||
|
|
||||||
#get style list
|
|
||||||
@PromptServer.instance.routes.get("/easyuse/prompt/styles")
|
|
||||||
async def getStylesList(request):
|
|
||||||
if "name" in request.rel_url.query:
|
|
||||||
name = request.rel_url.query["name"]
|
|
||||||
if name == 'fooocus_styles':
|
|
||||||
file = os.path.join(RESOURCES_DIR, name+'.json')
|
|
||||||
cn_file = os.path.join(RESOURCES_DIR, name + '_cn.json')
|
|
||||||
else:
|
|
||||||
file = os.path.join(FOOOCUS_STYLES_DIR, name+'.json')
|
|
||||||
cn_file = os.path.join(FOOOCUS_STYLES_DIR, name + '_cn.json')
|
|
||||||
cn_data = None
|
|
||||||
if os.path.isfile(cn_file):
|
|
||||||
f = open(cn_file, 'r', encoding='utf-8')
|
|
||||||
cn_data = json.load(f)
|
|
||||||
f.close()
|
|
||||||
if os.path.isfile(file):
|
|
||||||
f = open(file, 'r', encoding='utf-8')
|
|
||||||
data = json.load(f)
|
|
||||||
f.close()
|
|
||||||
if data:
|
|
||||||
ndata = []
|
|
||||||
for d in data:
|
|
||||||
nd = {}
|
|
||||||
name = d['name'].replace('-', ' ')
|
|
||||||
words = name.split(' ')
|
|
||||||
key = ' '.join(
|
|
||||||
word.upper() if word.lower() in ['mre', 'sai', '3d'] else word.capitalize() for word in
|
|
||||||
words)
|
|
||||||
img_name = '_'.join(words).lower()
|
|
||||||
if "name_cn" in d:
|
|
||||||
nd['name_cn'] = d['name_cn']
|
|
||||||
elif cn_data:
|
|
||||||
nd['name_cn'] = cn_data[key] if key in cn_data else key
|
|
||||||
nd["name"] = d['name']
|
|
||||||
nd['imgName'] = img_name
|
|
||||||
ndata.append(nd)
|
|
||||||
return web.json_response(ndata)
|
|
||||||
return web.Response(status=400)
|
|
||||||
|
|
||||||
# get style preview image
|
|
||||||
@PromptServer.instance.routes.get("/easyuse/prompt/styles/image")
|
|
||||||
async def getStylesImage(request):
|
|
||||||
styles_name = request.rel_url.query["styles_name"] if "styles_name" in request.rel_url.query else None
|
|
||||||
if "name" in request.rel_url.query:
|
|
||||||
name = request.rel_url.query["name"]
|
|
||||||
if os.path.exists(os.path.join(FOOOCUS_STYLES_DIR, 'samples')):
|
|
||||||
file = os.path.join(FOOOCUS_STYLES_DIR, 'samples', name + '.jpg')
|
|
||||||
if os.path.isfile(file):
|
|
||||||
return web.FileResponse(file)
|
|
||||||
elif styles_name == 'fooocus_styles':
|
|
||||||
return web.Response(text=FOOOCUS_STYLES_SAMPLES + name + '.jpg')
|
|
||||||
elif styles_name == 'fooocus_styles':
|
|
||||||
return web.Response(text=FOOOCUS_STYLES_SAMPLES + name + '.jpg')
|
|
||||||
return web.Response(status=400)
|
|
||||||
|
|
||||||
# convert type
|
|
||||||
@PromptServer.instance.routes.post("/easyuse/convert")
|
|
||||||
async def convertType(request):
|
|
||||||
post = await request.post()
|
|
||||||
type = post.get('type')
|
|
||||||
if type:
|
|
||||||
ConvertAnything.RETURN_TYPES = (type.upper(),)
|
|
||||||
ConvertAnything.RETURN_NAMES = (type,)
|
|
||||||
return web.Response(status=200)
|
|
||||||
else:
|
|
||||||
return web.Response(status=400)
|
|
||||||
|
|
||||||
# get models lists
|
|
||||||
@PromptServer.instance.routes.get("/easyuse/models/list")
|
|
||||||
async def getModelsList(request):
|
|
||||||
if "type" in request.rel_url.query:
|
|
||||||
type = request.rel_url.query["type"]
|
|
||||||
if type not in ['checkpoints', 'loras']:
|
|
||||||
return web.Response(status=400)
|
|
||||||
manager = easyModelManager()
|
|
||||||
return web.json_response(manager.get_model_lists(type))
|
|
||||||
else:
|
|
||||||
return web.Response(status=400)
|
|
||||||
|
|
||||||
# get models thumbnails
|
|
||||||
@PromptServer.instance.routes.get("/easyuse/models/thumbnail")
|
|
||||||
async def getModelsThumbnail(request):
|
|
||||||
checkpoints = folder_paths.get_filename_list("checkpoints_thumb")
|
|
||||||
loras = folder_paths.get_filename_list("loras_thumb")
|
|
||||||
checkpoints_full = []
|
|
||||||
loras_full = []
|
|
||||||
if len(checkpoints) + len(loras) >= 500:
|
|
||||||
return web.Response(status=400)
|
|
||||||
for index, i in enumerate(checkpoints):
|
|
||||||
full_path = folder_paths.get_full_path('checkpoints_thumb', str(i))
|
|
||||||
if full_path:
|
|
||||||
checkpoints_full.append(full_path)
|
|
||||||
for index, i in enumerate(loras):
|
|
||||||
full_path = folder_paths.get_full_path('loras_thumb', str(i))
|
|
||||||
if full_path:
|
|
||||||
loras_full.append(full_path)
|
|
||||||
return web.json_response(checkpoints_full + loras_full)
|
|
||||||
|
|
||||||
@PromptServer.instance.routes.post("/easyuse/metadata/notes/{name}")
|
|
||||||
async def save_notes(request):
|
|
||||||
name = request.match_info["name"]
|
|
||||||
pos = name.index("/")
|
|
||||||
type = name[0:pos]
|
|
||||||
name = name[pos+1:]
|
|
||||||
|
|
||||||
file_path = None
|
|
||||||
if type == "embeddings" or type == "loras":
|
|
||||||
name = name.lower()
|
|
||||||
files = folder_paths.get_filename_list(type)
|
|
||||||
for f in files:
|
|
||||||
lower_f = f.lower()
|
|
||||||
if lower_f == name:
|
|
||||||
file_path = folder_paths.get_full_path(type, f)
|
|
||||||
else:
|
|
||||||
n = os.path.splitext(f)[0].lower()
|
|
||||||
if n == name:
|
|
||||||
file_path = folder_paths.get_full_path(type, f)
|
|
||||||
|
|
||||||
if file_path is not None:
|
|
||||||
break
|
|
||||||
else:
|
|
||||||
file_path = folder_paths.get_full_path(
|
|
||||||
type, name)
|
|
||||||
if not file_path:
|
|
||||||
return web.Response(status=404)
|
|
||||||
|
|
||||||
file_no_ext = os.path.splitext(file_path)[0]
|
|
||||||
info_file = file_no_ext + ".txt"
|
|
||||||
with open(info_file, "w") as f:
|
|
||||||
f.write(await request.text())
|
|
||||||
|
|
||||||
return web.Response(status=200)
|
|
||||||
|
|
||||||
@PromptServer.instance.routes.get("/easyuse/metadata/{name}")
|
|
||||||
async def load_metadata(request):
|
|
||||||
name = request.match_info["name"]
|
|
||||||
pos = name.index("/")
|
|
||||||
type = name[0:pos]
|
|
||||||
name = name[pos+1:]
|
|
||||||
|
|
||||||
file_path = None
|
|
||||||
if type == "embeddings":
|
|
||||||
name = name.lower()
|
|
||||||
files = folder_paths.get_filename_list(type)
|
|
||||||
for f in files:
|
|
||||||
lower_f = f.lower()
|
|
||||||
if lower_f == name:
|
|
||||||
file_path = folder_paths.get_full_path(type, f)
|
|
||||||
else:
|
|
||||||
n = os.path.splitext(f)[0].lower()
|
|
||||||
if n == name:
|
|
||||||
file_path = folder_paths.get_full_path(type, f)
|
|
||||||
|
|
||||||
if file_path is not None:
|
|
||||||
break
|
|
||||||
else:
|
|
||||||
file_path = folder_paths.get_full_path(type, name)
|
|
||||||
if not file_path:
|
|
||||||
return web.Response(status=404)
|
|
||||||
|
|
||||||
try:
|
|
||||||
header = getMetadata(file_path)
|
|
||||||
header_json = json.loads(header)
|
|
||||||
meta = header_json["__metadata__"] if "__metadata__" in header_json else None
|
|
||||||
except:
|
|
||||||
meta = None
|
|
||||||
|
|
||||||
if meta is None:
|
|
||||||
meta = {}
|
|
||||||
|
|
||||||
file_no_ext = os.path.splitext(file_path)[0]
|
|
||||||
|
|
||||||
info_file = file_no_ext + ".txt"
|
|
||||||
if os.path.isfile(info_file):
|
|
||||||
with open(info_file, "r") as f:
|
|
||||||
meta["easyuse.notes"] = f.read()
|
|
||||||
|
|
||||||
hash_file = file_no_ext + ".sha256"
|
|
||||||
if os.path.isfile(hash_file):
|
|
||||||
with open(hash_file, "rt") as f:
|
|
||||||
meta["easyuse.sha256"] = f.read()
|
|
||||||
else:
|
|
||||||
with open(file_path, "rb") as f:
|
|
||||||
meta["easyuse.sha256"] = hashlib.sha256(f.read()).hexdigest()
|
|
||||||
with open(hash_file, "wt") as f:
|
|
||||||
f.write(meta["easyuse.sha256"])
|
|
||||||
|
|
||||||
return web.json_response(meta)
|
|
||||||
|
|
||||||
@PromptServer.instance.routes.post("/easyuse/save/{name}")
|
|
||||||
async def save_preview(request):
|
|
||||||
name = request.match_info["name"]
|
|
||||||
pos = name.index("/")
|
|
||||||
type = name[0:pos]
|
|
||||||
name = name[pos+1:]
|
|
||||||
|
|
||||||
body = await request.json()
|
|
||||||
|
|
||||||
dir = get_directory_by_type(body.get("type", "output"))
|
|
||||||
subfolder = body.get("subfolder", "")
|
|
||||||
full_output_folder = os.path.join(dir, os.path.normpath(subfolder))
|
|
||||||
|
|
||||||
if os.path.commonpath((dir, os.path.abspath(full_output_folder))) != dir:
|
|
||||||
return web.Response(status=400)
|
|
||||||
|
|
||||||
filepath = os.path.join(full_output_folder, body.get("filename", ""))
|
|
||||||
image_path = folder_paths.get_full_path(type, name)
|
|
||||||
image_path = os.path.splitext(
|
|
||||||
image_path)[0] + os.path.splitext(filepath)[1]
|
|
||||||
|
|
||||||
shutil.copyfile(filepath, image_path)
|
|
||||||
|
|
||||||
return web.json_response({
|
|
||||||
"image": type + "/" + os.path.basename(image_path)
|
|
||||||
})
|
|
||||||
|
|
||||||
@PromptServer.instance.routes.post("/easyuse/model/download")
|
|
||||||
async def download_model(request):
|
|
||||||
post = await request.post()
|
|
||||||
url = post.get("url")
|
|
||||||
local_dir = post.get("local_dir")
|
|
||||||
if local_dir not in ['checkpoints', 'loras', 'controlnet', 'onnx', 'instantid', 'ipadapter', 'dynamicrafter_models', 'mediapipe', 'rembg', 'layer_model']:
|
|
||||||
return web.Response(status=400)
|
|
||||||
local_path = os.path.join(folder_paths.models_dir, local_dir)
|
|
||||||
try:
|
|
||||||
get_local_filepath(url, local_path)
|
|
||||||
return web.Response(status=200)
|
|
||||||
except:
|
|
||||||
return web.Response(status=500)
|
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {}
|
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
|
||||||
+105
-16
@@ -190,6 +190,12 @@ REMBG_DIR = os.path.join(folder_paths.models_dir, "rembg")
|
|||||||
REMBG_MODELS = {
|
REMBG_MODELS = {
|
||||||
"RMBG-1.4": {
|
"RMBG-1.4": {
|
||||||
"model_url": "https://huggingface.co/briaai/RMBG-1.4/resolve/main/model.pth"
|
"model_url": "https://huggingface.co/briaai/RMBG-1.4/resolve/main/model.pth"
|
||||||
|
},
|
||||||
|
"RMBG-2.0": {
|
||||||
|
"model_url": "briaai/RMBG-2.0"
|
||||||
|
},
|
||||||
|
"BEN2": {
|
||||||
|
"model_url": "https://huggingface.co/PramaLLC/BEN2/resolve/main/BEN2_Base.pth"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -197,7 +203,7 @@ REMBG_MODELS = {
|
|||||||
IPADAPTER_DIR = os.path.join(folder_paths.models_dir, "ipadapter")
|
IPADAPTER_DIR = os.path.join(folder_paths.models_dir, "ipadapter")
|
||||||
IPADAPTER_MODELS = {
|
IPADAPTER_MODELS = {
|
||||||
"LIGHT - SD1.5 only (low strength)": {
|
"LIGHT - SD1.5 only (low strength)": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter_sd15_light_v11.bin"
|
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter_sd15_light_v11.bin"
|
||||||
},
|
},
|
||||||
"sdxl": {
|
"sdxl": {
|
||||||
@@ -205,7 +211,7 @@ IPADAPTER_MODELS = {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"STANDARD (medium strength)": {
|
"STANDARD (medium strength)": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter_sd15.safetensors"
|
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter_sd15.safetensors"
|
||||||
},
|
},
|
||||||
"sdxl": {
|
"sdxl": {
|
||||||
@@ -213,7 +219,7 @@ IPADAPTER_MODELS = {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"VIT-G (medium strength)": {
|
"VIT-G (medium strength)": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter_sd15_vit-G.safetensors"
|
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter_sd15_vit-G.safetensors"
|
||||||
},
|
},
|
||||||
"sdxl": {
|
"sdxl": {
|
||||||
@@ -221,15 +227,33 @@ IPADAPTER_MODELS = {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"PLUS (high strength)": {
|
"PLUS (high strength)": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter-plus_sd15.safetensors"
|
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter-plus_sd15.safetensors"
|
||||||
},
|
},
|
||||||
"sdxl": {
|
"sdxl": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/sdxl_models/ip-adapter-plus_sdxl_vit-h.safetensors"
|
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/sdxl_models/ip-adapter-plus_sdxl_vit-h.safetensors"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
"PLUS (kolors genernal)": {
|
||||||
|
"sd1": {
|
||||||
|
"model_url": ""
|
||||||
|
},
|
||||||
|
"sdxl": {
|
||||||
|
"model_url":"https://huggingface.co/Kwai-Kolors/Kolors-IP-Adapter-Plus/resolve/main/ip_adapter_plus_general.bin"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"REGULAR - FLUX and SD3.5 only (high strength)": {
|
||||||
|
"flux": {
|
||||||
|
"model_url": "https://huggingface.co/InstantX/FLUX.1-dev-IP-Adapter/resolve/main/ip-adapter.bin",
|
||||||
|
"model_file_name": "ip-adapter_flux_1_dev.bin",
|
||||||
|
},
|
||||||
|
"sd3": {
|
||||||
|
"model_url": "https://huggingface.co/InstantX/SD3.5-Large-IP-Adapter/resolve/main/ip-adapter.bin",
|
||||||
|
"model_file_name": "ip-adapter_sd35.bin",
|
||||||
|
},
|
||||||
|
},
|
||||||
"PLUS FACE (portraits)": {
|
"PLUS FACE (portraits)": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter-plus-face_sd15.safetensors"
|
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter-plus-face_sd15.safetensors"
|
||||||
},
|
},
|
||||||
"sdxl": {
|
"sdxl": {
|
||||||
@@ -237,7 +261,7 @@ IPADAPTER_MODELS = {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"FULL FACE - SD1.5 only (portraits stronger)": {
|
"FULL FACE - SD1.5 only (portraits stronger)": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter-full-face_sd15.safetensors"
|
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter-full-face_sd15.safetensors"
|
||||||
},
|
},
|
||||||
"sdxl": {
|
"sdxl": {
|
||||||
@@ -245,7 +269,7 @@ IPADAPTER_MODELS = {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"FACEID": {
|
"FACEID": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sd15.bin",
|
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sd15.bin",
|
||||||
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sd15_lora.safetensors"
|
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sd15_lora.safetensors"
|
||||||
},
|
},
|
||||||
@@ -255,7 +279,7 @@ IPADAPTER_MODELS = {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"FACEID PLUS - SD1.5 only": {
|
"FACEID PLUS - SD1.5 only": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plus_sd15.bin",
|
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plus_sd15.bin",
|
||||||
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plus_sd15_lora.safetensors"
|
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plus_sd15_lora.safetensors"
|
||||||
},
|
},
|
||||||
@@ -265,7 +289,7 @@ IPADAPTER_MODELS = {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"FACEID PLUS V2": {
|
"FACEID PLUS V2": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sd15.bin",
|
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sd15.bin",
|
||||||
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sd15_lora.safetensors"
|
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sd15_lora.safetensors"
|
||||||
},
|
},
|
||||||
@@ -274,24 +298,32 @@ IPADAPTER_MODELS = {
|
|||||||
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sdxl_lora.safetensors"
|
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sdxl_lora.safetensors"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
"FACEID PLUS KOLORS":{
|
||||||
|
"sd1":{
|
||||||
|
|
||||||
|
},
|
||||||
|
"sdxl":{
|
||||||
|
"model_url":"https://huggingface.co/Kwai-Kolors/Kolors-IP-Adapter-FaceID-Plus/resolve/main/ipa-faceid-plus.bin"
|
||||||
|
}
|
||||||
|
},
|
||||||
"FACEID PORTRAIT (style transfer)": {
|
"FACEID PORTRAIT (style transfer)": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait-v11_sd15.bin",
|
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait-v11_sd15.bin",
|
||||||
},
|
},
|
||||||
"sdxl": {
|
"sdxl": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait_sdxl.bin",
|
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait_sdxl.bin",
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"FACEID PORTRAIT UNNORM - SDXL only (strong)":{
|
"FACEID PORTRAIT UNNORM - SDXL only (strong)": {
|
||||||
"SD15":{
|
"sd1": {
|
||||||
"model_url":""
|
"model_url":""
|
||||||
},
|
},
|
||||||
"SDXL":{
|
"sdxl": {
|
||||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait_sdxl_unnorm.bin",
|
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait_sdxl_unnorm.bin",
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"COMPOSITION": {
|
"COMPOSITION": {
|
||||||
"sd15": {
|
"sd1": {
|
||||||
"model_url": "https://huggingface.co/ostris/ip-composition-adapter/resolve/main/ip_plus_composition_sd15.safetensors"
|
"model_url": "https://huggingface.co/ostris/ip-composition-adapter/resolve/main/ip_plus_composition_sd15.safetensors"
|
||||||
},
|
},
|
||||||
"sdxl": {
|
"sdxl": {
|
||||||
@@ -299,6 +331,17 @@ IPADAPTER_MODELS = {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
IPADAPTER_CLIPVISION_MODELS = {
|
||||||
|
"clip-vit-large-patch14-336":{
|
||||||
|
"model_url": "https://huggingface.co/openai/clip-vit-large-patch14-336/resolve/main/pytorch_model.bin"
|
||||||
|
},
|
||||||
|
"clip-vit-h-14-laion2B-s32B-b79K":{
|
||||||
|
"model_url": "https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/resolve/main/open_clip_model.safetensors"
|
||||||
|
},
|
||||||
|
"sigclip_vision_patch14_384":{
|
||||||
|
"model_url": "https://huggingface.co/Comfy-Org/sigclip_vision_384/resolve/main/sigclip_vision_patch14_384.safetensors"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
# dynamiCrafter
|
# dynamiCrafter
|
||||||
DYNAMICRAFTER_DIR = os.path.join(folder_paths.models_dir, "dynamicrafter_models")
|
DYNAMICRAFTER_DIR = os.path.join(folder_paths.models_dir, "dynamicrafter_models")
|
||||||
@@ -307,7 +350,7 @@ DYNAMICRAFTER_MODELS = {
|
|||||||
"model_url": "https://huggingface.co/ExponentialML/DynamiCrafterUNet/resolve/main/dynamicrafter_unet_512.safetensors",
|
"model_url": "https://huggingface.co/ExponentialML/DynamiCrafterUNet/resolve/main/dynamicrafter_unet_512.safetensors",
|
||||||
"vae_url": "https://huggingface.co/stabilityai/sd-vae-ft-mse-original/resolve/main/vae-ft-mse-840000-ema-pruned.safetensors",
|
"vae_url": "https://huggingface.co/stabilityai/sd-vae-ft-mse-original/resolve/main/vae-ft-mse-840000-ema-pruned.safetensors",
|
||||||
"clip_url": "https://huggingface.co/stabilityai/stable-diffusion-2-1/resolve/main/text_encoder/model.safetensors",
|
"clip_url": "https://huggingface.co/stabilityai/stable-diffusion-2-1/resolve/main/text_encoder/model.safetensors",
|
||||||
"clip_vision_url": "https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/resolve/main/open_clip_pytorch_model.safetensors",
|
"clip_vision_url": "https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/resolve/main/open_clip_model.safetensors",
|
||||||
},
|
},
|
||||||
"dynamicrafter_unet_512_interp (2.98GB)": {
|
"dynamicrafter_unet_512_interp (2.98GB)": {
|
||||||
"model_url": "https://huggingface.co/ExponentialML/DynamiCrafterUNet/resolve/main/dynamicrafter_unet_512_interp.safetensors"
|
"model_url": "https://huggingface.co/ExponentialML/DynamiCrafterUNet/resolve/main/dynamicrafter_unet_512_interp.safetensors"
|
||||||
@@ -325,6 +368,18 @@ HUMANPARSING_MODELS = {
|
|||||||
"parsing_lip": {
|
"parsing_lip": {
|
||||||
"model_url": "https://huggingface.co/levihsu/OOTDiffusion/resolve/main/checkpoints/humanparsing/parsing_lip.onnx",
|
"model_url": "https://huggingface.co/levihsu/OOTDiffusion/resolve/main/checkpoints/humanparsing/parsing_lip.onnx",
|
||||||
},
|
},
|
||||||
|
"human-parts":{
|
||||||
|
"model_url":"https://huggingface.co/Metal3d/deeplabv3p-resnet50-human/resolve/main/deeplabv3p-resnet50-human.onnx",
|
||||||
|
},
|
||||||
|
"segformer_b3_clothes":{
|
||||||
|
"model_name": "sayeed99/segformer_b3_clothes",
|
||||||
|
},
|
||||||
|
"segformer_b3_fashion":{
|
||||||
|
"model_name": "sayeed99/segformer-b3-fashion",
|
||||||
|
},
|
||||||
|
"face_parsing":{
|
||||||
|
"model_name": "jonathandinu/face-parsing"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#mediapipe
|
#mediapipe
|
||||||
@@ -333,4 +388,38 @@ MEDIAPIPE_MODELS = {
|
|||||||
"selfie_multiclass_256x256": {
|
"selfie_multiclass_256x256": {
|
||||||
"model_url": "https://huggingface.co/yolain/selfie_multiclass_256x256/resolve/main/selfie_multiclass_256x256.tflite"
|
"model_url": "https://huggingface.co/yolain/selfie_multiclass_256x256/resolve/main/selfie_multiclass_256x256.tflite"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
#prompt template
|
||||||
|
PROMPT_TEMPLATE = {
|
||||||
|
"prefix": ["Detailed photo of", "Amateur photo of", "Flicker 2008 photo of", "Fantastic artwork of",
|
||||||
|
"Vintage photograph of", "Unreal 5 render of", "Surrealist painting of",
|
||||||
|
"Professional advertising design of"],
|
||||||
|
"subject": ["a man", "a woman", "a young man", "a young woman", "a handsome man", "a beautiful woman", "a monster", "a toy", "a product", "a buddha", "a dog", "a cat"],
|
||||||
|
"action": ["looking at viewer", "looking away", "looking up", "looking down", "looking back", "open mouth", "half-closed mouth", "closed mouth", "open eyes", "half-closed eyes", "closed eyes", "wink", "standing", "sitting", "lying", "walking", "running", "adjusting hair", "waving", "hand on hip", "crossed arms", "smile", "sad", "angry", "sleepy", "tired", "expressionless"],
|
||||||
|
"clothes": ["underwear", "clothed", "casual", "dress", "swimsuit", "uniform", "bikini", "one-piece swimsuit", "shirt", "blouse", "sweater", "hoodie", "jeans", "pants", "shorts", "skirt", "vest", "coat", "trenchoat", "jacket", "short dress", "long dress", "off-shoulder", "backless", "hairbow", "hair ribbon", "hair tie", "hairband", "cap", "beanie", "bucket hat", "sun hat", "straw hat", "rice hat", "witch hat", "crown", "chain necklace", "tooth necklace", "choker", "pendant", "bracelet", "watch", "ring", "earring", "anklet", "belt", "scarf", "gloves", "mittens", "socks", "stockings", "tights", "leggings", "boots", "sneakers", "heels", "sandals", "flip-flops", "slippers", "loafers", "mules", "oxfords", "brogues", "derbies", "monk shoes", "chelsea boots", "combat boots", "riding boots", "rain boots", "wedge heels", "platform heels", "stilettos", "block heels", "kitten heels", "moccasins", "espadrilles", "pumps", "flats", "ballet flats", "mary janes", "slingbacks", "peep-toe", "mule sandals", "gladiator sandals", "thong sandals", "slide sandals", "espadrille sandals", "wedge sandals", "platform sandals", "ankle boots", "knee-high boots", "over-the-knee boots", "thigh-high boots", "wellington boots", "chukka boots", "desert boots", "chelsea boots", "hiking boots", "work boots", "snow boots", "rain boots", "riding boots", "cowboy boots", "combat boots", "biker boots", "duck boots", "military boots", "western boots", "ankle strap heels", "block heels", "chunky heels", "cone heels", "kitten heels", "platform heels", "pumps", "slingback heels", "stiletto heels", "wedge heels", "mules", "slingbacks", "slides", "thong sandals", "gladiator sandals", "espadrilles", "wedge sandals", "platform sandals", "ankle boots", "knee-high boots", "over-the-knee boots", "thigh-high boots", "wellington boots", "chukka boots", "desert boots", "chelsea boots", "hiking boots", "work boots", "snow boots", "rain boots", "riding boots", "cowboy boots", "combat boots", "biker boots", "duck boots", "military boots", "western boots", "ankle strap heels", "block heels" ],
|
||||||
|
"environment": ["sunshine from window", "neon night, city", "sunset over sea", "golden time", "sci-fi RGB glowing, cyberpunk", "natural lighting", "warm atmosphere, at home, bedroom", "magic lit", "evil, gothic, in a cave", "light and shadow", "shadow from window", "soft studio lighting", "home atmosphere, cozy bedroom illumination", "neon, Wong Kar-wai, warm", "moonlight through curtains", "stormy sky lighting", "underwater glow, deep sea", "foggy forest at dawn", "golden hour in a meadow", "rainbow reflections, neon", "cozy candlelight", "apocalyptic, smoky atmosphere", "red glow, emergency lights", "mystical glow, enchanted forest", "campfire light", "harsh, industrial lighting", "sunrise in the mountains", "evening glow in the desert", "moonlight in a dark alley", "golden glow at a fairground", "midnight in the forest", "purple and pink hues at twilight", "foggy morning, muted light", "candle-lit room, rustic vibe", "fluorescent office lighting", "lightning flash in storm", "night, cozy warm light from fireplace", "ethereal glow, magical forest", "dusky evening on a beach", "afternoon light filtering through trees", "blue neon light, urban street", "red and blue police lights in rain", "aurora borealis glow, arctic landscape", "sunrise through foggy mountains", "golden hour on a city skyline", "mysterious twilight, heavy mist", "early morning rays, forest clearing", "colorful lantern light at festival", "soft glow through stained glass", "harsh spotlight in dark room", "mellow evening glow on a lake", "crystal reflections in a cave", "vibrant autumn lighting in a forest", "gentle snowfall at dusk", "hazy light of a winter morning", "soft, diffused foggy glow", "underwater luminescence", "rain-soaked reflections in city lights", "golden sunlight streaming through trees", "fireflies lighting up a summer night", "glowing embers from a forge", "dim candlelight in a gothic castle", "midnight sky with bright starlight", "warm sunset in a rural village", "flickering light in a haunted house", "desert sunset with mirage-like glow", "golden beams piercing through storm clouds"],
|
||||||
|
"background": ["cars and people", "a cozy bed and a lamp", "a forest clearing with mist", "a bustling marketplace", "a quiet beach at dusk", "an old, cobblestone street", "a futuristic cityscape", "a tranquil lake with mountains", "a mysterious cave entrance", "bookshelves and plants in the background", "an ancient temple in ruins", "tall skyscrapers and neon signs", "a starry sky over a desert", "a bustling café", "rolling hills and farmland", "a modern living room with a fireplace", "an abandoned warehouse", "a picturesque mountain range", "a starry night sky", "the interior of a futuristic spaceship", "the cluttered workshop of an inventor", "the glowing embers of a bonfire", "a misty lake surrounded by trees", "an ornate palace hall", "a busy street market", "a vast desert landscape", "a peaceful library corner", "bustling train station", "a mystical, enchanted forest", "an underwater reef with colorful fish", "a quiet rural village", "a sandy beach with palm trees", "a vibrant coral reef, teeming with life", "snow-capped mountains in distance", "a stormy ocean, waves crashing", "a rustic barn in open fields", "a futuristic lab with glowing screens", "a dark, abandoned castle", "the ruins of an ancient civilization", "a bustling urban street in rain", "an elegant grand ballroom", "a sprawling field of wildflowers", "a dense jungle with sunlight filtering through", "a dimly lit, vintage bar", "an ice cave with sparkling crystals", "a serene riverbank at sunset", "a narrow alley with graffiti walls", "a peaceful zen garden with koi pond", "a high-tech control room", "a quiet mountain village at dawn", "a lighthouse on a rocky coast", "a rainy street with flickering lights", "a frozen lake with ice formations", "an abandoned theme park", "a small fishing village on a pier", "rolling sand dunes in a desert", "a dense forest with towering redwoods", "a snowy cabin in the mountains", "a mystical cave with bioluminescent plants", "a castle courtyard under moonlight", "a bustling open-air night market", "an old train station with steam", "a tranquil waterfall surrounded by trees", "a vineyard in the countryside", "a quaint medieval village", "a bustling harbor with boats", "a high-tech futuristic mall", "a lush tropical rainforest"],
|
||||||
|
"nsfw": ["nude", "breast", "small breast", "middle breast", "large breast", "nipples", "clothes lift", "pussy juice trail", "pussy juice puddle", "small testicles", "medium testicles", "large testicles", "disembodied penis", "cum on body", "cum inside", "cum outside", "fingering", "handjob", "fellatio", "licking penis", "paizuri", "doggystyle", "cowgirl", "reversed cowgirl", "piledriver", "suspended congress", "full nelson",],
|
||||||
|
}
|
||||||
|
|
||||||
|
NEW_SCHEDULERS = ['align_your_steps', 'gits']
|
||||||
|
|
||||||
|
DIFFUSION_MODEL_XY_DEFAULTS = {
|
||||||
|
"anima": {
|
||||||
|
"clip_name": "qwen_3_06b_base.safetensors",
|
||||||
|
"clip_type": "anima",
|
||||||
|
"vae_name": "qwen_image_vae.safetensors",
|
||||||
|
},
|
||||||
|
"krea2": {
|
||||||
|
"clip_name": "Huihui-Qwen3-VL-4B-Instruct-abliterated-fp8_scaled.safetensors",
|
||||||
|
"clip_type": "krea2",
|
||||||
|
"vae_name": "qwen_image_vae.safetensors",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
DIFFUSION_MODEL_CLIP_TYPES = {
|
||||||
|
"anima": "anima",
|
||||||
|
"krea2": "krea2",
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,74 +0,0 @@
|
|||||||
TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT
|
|
||||||
Tencent Hunyuan Release Date: 2024/5/14
|
|
||||||
By clicking to agree or by using, reproducing, modifying, distributing, performing or displaying any portion or element of the Tencent Hunyuan Works, including via any Hosted Service, You will be deemed to have recognized and accepted the content of this Agreement, which is effective immediately.
|
|
||||||
1. DEFINITIONS.
|
|
||||||
a. “Acceptable Use Policy” shall mean the policy made available by Tencent as set forth in the Exhibit A.
|
|
||||||
b. “Agreement” shall mean the terms and conditions for use, reproduction, distribution, modification, performance and displaying of the Hunyuan Works or any portion or element thereof set forth herein.
|
|
||||||
c. “Documentation” shall mean the specifications, manuals and documentation for Tencent Hunyuan made publicly available by Tencent.
|
|
||||||
d. “Hosted Service” shall mean a hosted service offered via an application programming interface (API), web access, or any other electronic or remote means.
|
|
||||||
e. “Licensee,” “You” or “Your” shall mean a natural person or legal entity exercising the rights granted by this Agreement and/or using the Tencent Hunyuan Works for any purpose and in any field of use.
|
|
||||||
f. “Materials” shall mean, collectively, Tencent’s proprietary Tencent Hunyuan and Documentation (and any portion thereof) as made available by Tencent under this Agreement.
|
|
||||||
g. “Model Derivatives” shall mean all: (i) modifications to Tencent Hunyuan or any Model Derivative of Tencent Hunyuan; (ii) works based on Tencent Hunyuan or any Model Derivative of Tencent Hunyuan; or (iii) any other machine learning model which is created by transfer of patterns of the weights, parameters, operations, or Output of Tencent Hunyuan or any Model Derivative of Tencent Hunyuan, to that model in order to cause that model to perform similarly to Tencent Hunyuan or a Model Derivative of Tencent Hunyuan, including distillation methods, methods that use intermediate data representations, or methods based on the generation of synthetic data Outputs by Tencent Hunyuan or a Model Derivative of Tencent Hunyuan for training that model. For clarity, Outputs by themselves are not deemed Model Derivatives.
|
|
||||||
h. “Output” shall mean the information and/or content output of Tencent Hunyuan or a Model Derivative that results from operating or otherwise using Tencent Hunyuan or a Model Derivative, including via a Hosted Service.
|
|
||||||
i. “Tencent,” “We” or “Us” shall mean THL A29 Limited.
|
|
||||||
j. “Tencent Hunyuan” shall mean the large language models, image/video/audio/3D generation models, and multimodal large language models and their software and algorithms, including trained model weights, parameters (including optimizer states), machine-learning model code, inference-enabling code, training-enabling code, fine-tuning enabling code and other elements of the foregoing made publicly available by Us at https://huggingface.co/Tencent-Hunyuan/HunyuanDiT and https://github.com/Tencent/HunyuanDiT .
|
|
||||||
k. “Tencent Hunyuan Works” shall mean: (i) the Materials; (ii) Model Derivatives; and (iii) all derivative works thereof.
|
|
||||||
l. “Third Party” or “Third Parties” shall mean individuals or legal entities that are not under common control with Us or You.
|
|
||||||
m. “including” shall mean including but not limited to.
|
|
||||||
2. GRANT OF RIGHTS.
|
|
||||||
We grant You a non-exclusive, worldwide, non-transferable and royalty-free limited license under Tencent’s intellectual property or other rights owned by Us embodied in or utilized by the Materials to use, reproduce, distribute, create derivative works of (including Model Derivatives), and make modifications to the Materials, only in accordance with the terms of this Agreement and the Acceptable Use Policy, and You must not violate (or encourage or permit anyone else to violate) any term of this Agreement or the Acceptable Use Policy.
|
|
||||||
3. DISTRIBUTION.
|
|
||||||
You may, subject to Your compliance with this Agreement, distribute or make available to Third Parties the Tencent Hunyuan Works, provided that You meet all of the following conditions:
|
|
||||||
a. You must provide all such Third Party recipients of the Tencent Hunyuan Works or products or services using them a copy of this Agreement;
|
|
||||||
b. You must cause any modified files to carry prominent notices stating that You changed the files;
|
|
||||||
c. You are encouraged to: (i) publish at least one technology introduction blogpost or one public statement expressing Your experience of using the Tencent Hunyuan Works; and (ii) mark the products or services developed by using the Tencent Hunyuan Works to indicate that the product/service is “Powered by Tencent Hunyuan”; and
|
|
||||||
d. All distributions to Third Parties (other than through a Hosted Service) must be accompanied by a “Notice” text file that contains the following notice: “Tencent Hunyuan is licensed under the Tencent Hunyuan Community License Agreement, Copyright © 2024 Tencent. All Rights Reserved. The trademark rights of “Tencent Hunyuan” are owned by Tencent or its affiliate.”
|
|
||||||
You may add Your own copyright statement to Your modifications and, except as set forth in this Section and in Section 5, may provide additional or different license terms and conditions for use, reproduction, or distribution of Your modifications, or for any such Model Derivatives as a whole, provided Your use, reproduction, modification, distribution, performance and display of the work otherwise complies with the terms and conditions of this Agreement. If You receive Tencent Hunyuan Works from a Licensee as part of an integrated end user product, then this Section 3 of this Agreement will not apply to You.
|
|
||||||
4. ADDITIONAL COMMERCIAL TERMS.
|
|
||||||
If, on the Tencent Hunyuan version release date, the monthly active users of all products or services made available by or for Licensee is greater than 100 million monthly active users in the preceding calendar month, You must request a license from Tencent, which Tencent may grant to You in its sole discretion, and You are not authorized to exercise any of the rights under this Agreement unless or until Tencent otherwise expressly grants You such rights.
|
|
||||||
5. RULES OF USE.
|
|
||||||
a. Your use of the Tencent Hunyuan Works must comply with applicable laws and regulations (including trade compliance laws and regulations) and adhere to the Acceptable Use Policy for the Tencent Hunyuan Works, which is hereby incorporated by reference into this Agreement. You must include the use restrictions referenced in these Sections 5(a) and 5(b) as an enforceable provision in any agreement (e.g., license agreement, terms of use, etc.) governing the use and/or distribution of Tencent Hunyuan Works and You must provide notice to subsequent users to whom You distribute that Tencent Hunyuan Works are subject to the use restrictions in these Sections 5(a) and 5(b).
|
|
||||||
b. You must not use the Tencent Hunyuan Works or any Output or results of the Tencent Hunyuan Works to improve any other large language model (other than Tencent Hunyuan or Model Derivatives thereof).
|
|
||||||
6. INTELLECTUAL PROPERTY.
|
|
||||||
a. Subject to Tencent’s ownership of Tencent Hunyuan Works made by or for Tencent and intellectual property rights therein, conditioned upon Your compliance with the terms and conditions of this Agreement, as between You and Tencent, You will be the owner of any derivative works and modifications of the Materials and any Model Derivatives that are made by or for You.
|
|
||||||
b. No trademark licenses are granted under this Agreement, and in connection with the Tencent Hunyuan Works, Licensee may not use any name or mark owned by or associated with Tencent or any of its affiliates, except as required for reasonable and customary use in describing and distributing the Tencent Hunyuan Works. Tencent hereby grants You a license to use “Tencent Hunyuan” (the “Mark”) solely as required to comply with the provisions of Section 3(c), provided that You comply with any applicable laws related to trademark protection. All goodwill arising out of Your use of the Mark will inure to the benefit of Tencent.
|
|
||||||
c. If You commence a lawsuit or other proceedings (including a cross-claim or counterclaim in a lawsuit) against Us or any person or entity alleging that the Materials or any Output, or any portion of any of the foregoing, infringe any intellectual property or other right owned or licensable by You, then all licenses granted to You under this Agreement shall terminate as of the date such lawsuit or other proceeding is filed. You will defend, indemnify and hold harmless Us from and against any claim by any Third Party arising out of or related to Your or the Third Party’s use or distribution of the Tencent Hunyuan Works.
|
|
||||||
d. Tencent claims no rights in Outputs You generate. You and Your users are solely responsible for Outputs and their subsequent uses.
|
|
||||||
7. DISCLAIMERS OF WARRANTY AND LIMITATIONS OF LIABILITY.
|
|
||||||
a. We are not obligated to support, update, provide training for, or develop any further version of the Tencent Hunyuan Works or to grant any license thereto.
|
|
||||||
b. UNLESS AND ONLY TO THE EXTENT REQUIRED BY APPLICABLE LAW, THE TENCENT HUNYUAN WORKS AND ANY OUTPUT AND RESULTS THEREFROM ARE PROVIDED “AS IS” WITHOUT ANY EXPRESS OR IMPLIED WARRANTIES OF ANY KIND INCLUDING ANY WARRANTIES OF TITLE, MERCHANTABILITY, NONINFRINGEMENT, COURSE OF DEALING, USAGE OF TRADE, OR FITNESS FOR A PARTICULAR PURPOSE. YOU ARE SOLELY RESPONSIBLE FOR DETERMINING THE APPROPRIATENESS OF USING, REPRODUCING, MODIFYING, PERFORMING, DISPLAYING OR DISTRIBUTING ANY OF THE TENCENT HUNYUAN WORKS OR OUTPUTS AND ASSUME ANY AND ALL RISKS ASSOCIATED WITH YOUR OR A THIRD PARTY’S USE OR DISTRIBUTION OF ANY OF THE TENCENT HUNYUAN WORKS OR OUTPUTS AND YOUR EXERCISE OF RIGHTS AND PERMISSIONS UNDER THIS AGREEMENT.
|
|
||||||
c. TO THE FULLEST EXTENT PERMITTED BY APPLICABLE LAW, IN NO EVENT SHALL TENCENT OR ITS AFFILIATES BE LIABLE UNDER ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, TORT, NEGLIGENCE, PRODUCTS LIABILITY, OR OTHERWISE, FOR ANY DAMAGES, INCLUDING ANY DIRECT, INDIRECT, SPECIAL, INCIDENTAL, EXEMPLARY, CONSEQUENTIAL OR PUNITIVE DAMAGES, OR LOST PROFITS OF ANY KIND ARISING FROM THIS AGREEMENT OR RELATED TO ANY OF THE TENCENT HUNYUAN WORKS OR OUTPUTS, EVEN IF TENCENT OR ITS AFFILIATES HAVE BEEN ADVISED OF THE POSSIBILITY OF ANY OF THE FOREGOING.
|
|
||||||
8. SURVIVAL AND TERMINATION.
|
|
||||||
a. The term of this Agreement shall commence upon Your acceptance of this Agreement or access to the Materials and will continue in full force and effect until terminated in accordance with the terms and conditions herein.
|
|
||||||
b. We may terminate this Agreement if You breach any of the terms or conditions of this Agreement. Upon termination of this Agreement, You must promptly delete and cease use of the Tencent Hunyuan Works. Sections 6(a), 6(c), 7 and 9 shall survive the termination of this Agreement.
|
|
||||||
9. GOVERNING LAW AND JURISDICTION.
|
|
||||||
a. This Agreement and any dispute arising out of or relating to it will be governed by the laws of the Hong Kong Special Administrative Region of the People’s Republic of China, without regard to conflict of law principles, and the UN Convention on Contracts for the International Sale of Goods does not apply to this Agreement.
|
|
||||||
b. Exclusive jurisdiction and venue for any dispute arising out of or relating to this Agreement will be a court of competent jurisdiction in the Hong Kong Special Administrative Region of the People’s Republic of China, and Tencent and Licensee consent to the exclusive jurisdiction of such court with respect to any such dispute.
|
|
||||||
|
|
||||||
|
|
||||||
EXHIBIT A
|
|
||||||
ACCEPTABLE USE POLICY
|
|
||||||
|
|
||||||
Tencent reserves the right to update this Acceptable Use Policy from time to time.
|
|
||||||
Last modified: 2024/5/14
|
|
||||||
|
|
||||||
Tencent endeavors to promote safe and fair use of its tools and features, including Tencent Hunyuan. You agree not to use Tencent Hunyuan or Model Derivatives:
|
|
||||||
1. In any way that violates any applicable national, federal, state, local, international or any other law or regulation;
|
|
||||||
2. To harm Yourself or others;
|
|
||||||
3. To repurpose or distribute output from Tencent Hunyuan or any Model Derivatives to harm Yourself or others;
|
|
||||||
4. To override or circumvent the safety guardrails and safeguards We have put in place;
|
|
||||||
5. For the purpose of exploiting, harming or attempting to exploit or harm minors in any way;
|
|
||||||
6. To generate or disseminate verifiably false information and/or content with the purpose of harming others or influencing elections;
|
|
||||||
7. To generate or facilitate false online engagement, including fake reviews and other means of fake online engagement;
|
|
||||||
8. To intentionally defame, disparage or otherwise harass others;
|
|
||||||
9. To generate and/or disseminate malware (including ransomware) or any other content to be used for the purpose of harming electronic systems;
|
|
||||||
10. To generate or disseminate personal identifiable information with the purpose of harming others;
|
|
||||||
11. To generate or disseminate information (including images, code, posts, articles), and place the information in any public context (including –through the use of bot generated tweets), without expressly and conspicuously identifying that the information and/or content is machine generated;
|
|
||||||
12. To impersonate another individual without consent, authorization, or legal right;
|
|
||||||
13. To make high-stakes automated decisions in domains that affect an individual’s safety, rights or wellbeing (e.g., law enforcement, migration, medicine/health, management of critical infrastructure, safety components of products, essential services, credit, employment, housing, education, social scoring, or insurance);
|
|
||||||
14. In a manner that violates or disrespects the social ethics and moral standards of other countries or regions;
|
|
||||||
15. To perform, facilitate, threaten, incite, plan, promote or encourage violent extremism or terrorism;
|
|
||||||
16. For any use intended to discriminate against or harm individuals or groups based on protected characteristics or categories, online or offline social behavior or known or predicted personal or personality characteristics;
|
|
||||||
17. To intentionally exploit any of the vulnerabilities of a specific group of persons based on their age, social, physical or mental characteristics, in order to materially distort the behavior of a person pertaining to that group in a manner that causes or is likely to cause that person or another person physical or psychological harm;
|
|
||||||
18. For military purposes;
|
|
||||||
19. To engage in the unauthorized or unlicensed practice of any profession including, but not limited to, financial, legal, medical/health, or other professional practices.
|
|
||||||
@@ -1,46 +0,0 @@
|
|||||||
"""List of all HYDiT model types / settings"""
|
|
||||||
sampling_settings = {
|
|
||||||
"beta_schedule" : "linear",
|
|
||||||
"linear_start" : 0.00085,
|
|
||||||
"linear_end" : 0.03,
|
|
||||||
"timesteps" : 1000,
|
|
||||||
}
|
|
||||||
|
|
||||||
from argparse import Namespace
|
|
||||||
hydit_args = Namespace(**{ # normally from argparse
|
|
||||||
"infer_mode": "torch",
|
|
||||||
"norm": "layer",
|
|
||||||
"learn_sigma": True,
|
|
||||||
"text_states_dim": 1024,
|
|
||||||
"text_states_dim_t5": 2048,
|
|
||||||
"text_len": 77,
|
|
||||||
"text_len_t5": 256,
|
|
||||||
})
|
|
||||||
|
|
||||||
hydit_conf = {
|
|
||||||
"G/2": { # Seems to be the main one
|
|
||||||
"unet_config": {
|
|
||||||
"depth" : 40,
|
|
||||||
"num_heads" : 16,
|
|
||||||
"patch_size" : 2,
|
|
||||||
"hidden_size" : 1408,
|
|
||||||
"mlp_ratio" : 4.3637,
|
|
||||||
"input_size": (1024//8, 1024//8),
|
|
||||||
"args": hydit_args,
|
|
||||||
},
|
|
||||||
"sampling_settings" : sampling_settings,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
dtypes = ["default", "auto (comfy)", "FP32", "FP16", "BF16"]
|
|
||||||
devices = ["auto", "cpu", "gpu"]
|
|
||||||
|
|
||||||
|
|
||||||
# these are the same as regular DiT, I think
|
|
||||||
from ..config import dit_conf
|
|
||||||
for name in ["XL/2", "L/2", "B/2"]:
|
|
||||||
hydit_conf[name] = {
|
|
||||||
"unet_config": dit_conf[name]["unet_config"].copy(),
|
|
||||||
"sampling_settings": sampling_settings,
|
|
||||||
}
|
|
||||||
hydit_conf[name]["unet_config"]["args"] = hydit_args
|
|
||||||
@@ -1,34 +0,0 @@
|
|||||||
{
|
|
||||||
"_name_or_path": "hfl/chinese-roberta-wwm-ext-large",
|
|
||||||
"architectures": [
|
|
||||||
"BertModel"
|
|
||||||
],
|
|
||||||
"attention_probs_dropout_prob": 0.1,
|
|
||||||
"bos_token_id": 0,
|
|
||||||
"classifier_dropout": null,
|
|
||||||
"directionality": "bidi",
|
|
||||||
"eos_token_id": 2,
|
|
||||||
"hidden_act": "gelu",
|
|
||||||
"hidden_dropout_prob": 0.1,
|
|
||||||
"hidden_size": 1024,
|
|
||||||
"initializer_range": 0.02,
|
|
||||||
"intermediate_size": 4096,
|
|
||||||
"layer_norm_eps": 1e-12,
|
|
||||||
"max_position_embeddings": 512,
|
|
||||||
"model_type": "bert",
|
|
||||||
"num_attention_heads": 16,
|
|
||||||
"num_hidden_layers": 24,
|
|
||||||
"output_past": true,
|
|
||||||
"pad_token_id": 0,
|
|
||||||
"pooler_fc_size": 768,
|
|
||||||
"pooler_num_attention_heads": 12,
|
|
||||||
"pooler_num_fc_layers": 3,
|
|
||||||
"pooler_size_per_head": 128,
|
|
||||||
"pooler_type": "first_token_transform",
|
|
||||||
"position_embedding_type": "absolute",
|
|
||||||
"torch_dtype": "float32",
|
|
||||||
"transformers_version": "4.22.1",
|
|
||||||
"type_vocab_size": 2,
|
|
||||||
"use_cache": true,
|
|
||||||
"vocab_size": 47020
|
|
||||||
}
|
|
||||||
@@ -1,33 +0,0 @@
|
|||||||
{
|
|
||||||
"_name_or_path": "mt5",
|
|
||||||
"architectures": [
|
|
||||||
"MT5EncoderModel"
|
|
||||||
],
|
|
||||||
"classifier_dropout": 0.0,
|
|
||||||
"d_ff": 5120,
|
|
||||||
"d_kv": 64,
|
|
||||||
"d_model": 2048,
|
|
||||||
"decoder_start_token_id": 0,
|
|
||||||
"dense_act_fn": "gelu_new",
|
|
||||||
"dropout_rate": 0.1,
|
|
||||||
"eos_token_id": 1,
|
|
||||||
"feed_forward_proj": "gated-gelu",
|
|
||||||
"initializer_factor": 1.0,
|
|
||||||
"is_encoder_decoder": true,
|
|
||||||
"is_gated_act": true,
|
|
||||||
"layer_norm_epsilon": 1e-06,
|
|
||||||
"model_type": "mt5",
|
|
||||||
"num_decoder_layers": 24,
|
|
||||||
"num_heads": 32,
|
|
||||||
"num_layers": 24,
|
|
||||||
"output_past": true,
|
|
||||||
"pad_token_id": 0,
|
|
||||||
"relative_attention_max_distance": 128,
|
|
||||||
"relative_attention_num_buckets": 32,
|
|
||||||
"tie_word_embeddings": false,
|
|
||||||
"tokenizer_class": "T5Tokenizer",
|
|
||||||
"torch_dtype": "float16",
|
|
||||||
"transformers_version": "4.40.2",
|
|
||||||
"use_cache": true,
|
|
||||||
"vocab_size": 250112
|
|
||||||
}
|
|
||||||
@@ -1,240 +0,0 @@
|
|||||||
|
|
||||||
|
|
||||||
import os
|
|
||||||
import torch
|
|
||||||
import comfy.supported_models_base
|
|
||||||
import comfy.latent_formats
|
|
||||||
import comfy.model_patcher
|
|
||||||
import comfy.model_base
|
|
||||||
import comfy.utils
|
|
||||||
import comfy.conds
|
|
||||||
|
|
||||||
from comfy import model_management
|
|
||||||
from tqdm import tqdm
|
|
||||||
from transformers import AutoTokenizer, modeling_utils
|
|
||||||
from transformers import T5Config, T5EncoderModel, BertConfig, BertModel
|
|
||||||
|
|
||||||
class EXM_HYDiT(comfy.supported_models_base.BASE):
|
|
||||||
unet_config = {}
|
|
||||||
unet_extra_config = {}
|
|
||||||
latent_format = comfy.latent_formats.SDXL
|
|
||||||
|
|
||||||
def __init__(self, model_conf):
|
|
||||||
self.unet_config = model_conf.get("unet_config", {})
|
|
||||||
self.sampling_settings = model_conf.get("sampling_settings", {})
|
|
||||||
self.latent_format = self.latent_format()
|
|
||||||
# UNET is handled by extension
|
|
||||||
self.unet_config["disable_unet_model_creation"] = True
|
|
||||||
|
|
||||||
def model_type(self, state_dict, prefix=""):
|
|
||||||
return comfy.model_base.ModelType.V_PREDICTION
|
|
||||||
|
|
||||||
|
|
||||||
class EXM_HYDiT_Model(comfy.model_base.BaseModel):
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
|
|
||||||
def extra_conds(self, **kwargs):
|
|
||||||
out = super().extra_conds(**kwargs)
|
|
||||||
|
|
||||||
for name in ["context_t5", "context_mask", "context_t5_mask"]:
|
|
||||||
out[name] = comfy.conds.CONDRegular(kwargs[name])
|
|
||||||
|
|
||||||
src_size_cond = kwargs.get("src_size_cond", None)
|
|
||||||
if src_size_cond is not None:
|
|
||||||
out["src_size_cond"] = comfy.conds.CONDRegular(torch.tensor(src_size_cond))
|
|
||||||
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def load_hydit(model_path, model_conf):
|
|
||||||
state_dict = comfy.utils.load_torch_file(model_path)
|
|
||||||
state_dict = state_dict.get("model", state_dict)
|
|
||||||
|
|
||||||
parameters = comfy.utils.calculate_parameters(state_dict)
|
|
||||||
unet_dtype = model_management.unet_dtype(model_params=parameters)
|
|
||||||
load_device = comfy.model_management.get_torch_device()
|
|
||||||
offload_device = comfy.model_management.unet_offload_device()
|
|
||||||
|
|
||||||
# ignore fp8/etc and use directly for now
|
|
||||||
manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device)
|
|
||||||
if manual_cast_dtype:
|
|
||||||
print(f"HunYuanDiT: falling back to {manual_cast_dtype}")
|
|
||||||
unet_dtype = manual_cast_dtype
|
|
||||||
|
|
||||||
model_conf = EXM_HYDiT(model_conf)
|
|
||||||
model = EXM_HYDiT_Model(
|
|
||||||
model_conf,
|
|
||||||
model_type=comfy.model_base.ModelType.V_PREDICTION,
|
|
||||||
device=model_management.get_torch_device()
|
|
||||||
)
|
|
||||||
|
|
||||||
from .models.models import HunYuanDiT
|
|
||||||
model.diffusion_model = HunYuanDiT(
|
|
||||||
**model_conf.unet_config,
|
|
||||||
log_fn=tqdm.write,
|
|
||||||
)
|
|
||||||
|
|
||||||
model.diffusion_model.load_state_dict(state_dict)
|
|
||||||
model.diffusion_model.dtype = unet_dtype
|
|
||||||
model.diffusion_model.eval()
|
|
||||||
model.diffusion_model.to(unet_dtype)
|
|
||||||
|
|
||||||
model_patcher = comfy.model_patcher.ModelPatcher(
|
|
||||||
model,
|
|
||||||
load_device=load_device,
|
|
||||||
offload_device=offload_device,
|
|
||||||
current_device="cpu",
|
|
||||||
)
|
|
||||||
return model_patcher
|
|
||||||
|
|
||||||
|
|
||||||
# CLIP Model
|
|
||||||
class hyCLIPModel(torch.nn.Module):
|
|
||||||
def __init__(self, textmodel_json_config=None, device="cpu", max_length=77, freeze=True, dtype=None):
|
|
||||||
super().__init__()
|
|
||||||
self.device = device
|
|
||||||
self.dtype = dtype
|
|
||||||
self.max_length = max_length
|
|
||||||
if textmodel_json_config is None:
|
|
||||||
textmodel_json_config = os.path.join(
|
|
||||||
os.path.dirname(os.path.realpath(__file__)),
|
|
||||||
f"config_clip.json"
|
|
||||||
)
|
|
||||||
config = BertConfig.from_json_file(textmodel_json_config)
|
|
||||||
with modeling_utils.no_init_weights():
|
|
||||||
self.transformer = BertModel(config)
|
|
||||||
self.to(dtype)
|
|
||||||
if freeze:
|
|
||||||
self.freeze()
|
|
||||||
|
|
||||||
def freeze(self):
|
|
||||||
self.transformer = self.transformer.eval()
|
|
||||||
for param in self.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
|
|
||||||
def load_sd(self, sd):
|
|
||||||
return self.transformer.load_state_dict(sd, strict=False)
|
|
||||||
|
|
||||||
def to(self, *args, **kwargs):
|
|
||||||
return self.transformer.to(*args, **kwargs)
|
|
||||||
|
|
||||||
class EXM_HyDiT_Tenc_Temp:
|
|
||||||
def __init__(self, no_init=False, device="cpu", dtype=None, model_class="mT5", *kwargs):
|
|
||||||
if no_init:
|
|
||||||
return
|
|
||||||
|
|
||||||
size = 8 if model_class == "mT5" else 2
|
|
||||||
if dtype == torch.float32:
|
|
||||||
size *= 2
|
|
||||||
size *= (1024**3)
|
|
||||||
|
|
||||||
if device == "auto":
|
|
||||||
self.load_device = model_management.text_encoder_device()
|
|
||||||
self.offload_device = model_management.text_encoder_offload_device()
|
|
||||||
self.init_device = "cpu"
|
|
||||||
elif device == "cpu":
|
|
||||||
size = 0 # doesn't matter
|
|
||||||
self.load_device = "cpu"
|
|
||||||
self.offload_device = "cpu"
|
|
||||||
self.init_device="cpu"
|
|
||||||
elif device.startswith("cuda"):
|
|
||||||
print("Direct CUDA device override!\nVRAM will not be freed by default.")
|
|
||||||
size = 0 # not used
|
|
||||||
self.load_device = device
|
|
||||||
self.offload_device = device
|
|
||||||
self.init_device = device
|
|
||||||
else:
|
|
||||||
self.load_device = model_management.get_torch_device()
|
|
||||||
self.offload_device = "cpu"
|
|
||||||
self.init_device="cpu"
|
|
||||||
|
|
||||||
self.dtype = dtype
|
|
||||||
self.device = self.load_device
|
|
||||||
if model_class == "mT5":
|
|
||||||
self.cond_stage_model = mT5Model(
|
|
||||||
device = self.load_device,
|
|
||||||
dtype = self.dtype,
|
|
||||||
)
|
|
||||||
tokenizer_args = {"subfolder": "t2i/mt5"} # web
|
|
||||||
tokenizer_path = os.path.join( # local
|
|
||||||
os.path.dirname(os.path.realpath(__file__)),
|
|
||||||
"mt5_tokenizer",
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
self.cond_stage_model = hyCLIPModel(
|
|
||||||
device = self.load_device,
|
|
||||||
dtype = self.dtype,
|
|
||||||
)
|
|
||||||
tokenizer_args = {"subfolder": "t2i/tokenizer",} # web
|
|
||||||
tokenizer_path = os.path.join( # local
|
|
||||||
os.path.dirname(os.path.realpath(__file__)),
|
|
||||||
"tokenizer",
|
|
||||||
)
|
|
||||||
# self.tokenizer = AutoTokenizer.from_pretrained(
|
|
||||||
# "Tencent-Hunyuan/HunyuanDiT",
|
|
||||||
# **tokenizer_args
|
|
||||||
# )
|
|
||||||
self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
|
||||||
self.patcher = comfy.model_patcher.ModelPatcher(
|
|
||||||
self.cond_stage_model,
|
|
||||||
load_device = self.load_device,
|
|
||||||
offload_device = self.offload_device,
|
|
||||||
current_device = self.load_device,
|
|
||||||
size = size,
|
|
||||||
)
|
|
||||||
|
|
||||||
def clone(self):
|
|
||||||
n = EXM_HyDiT_Tenc_Temp(no_init=True)
|
|
||||||
n.patcher = self.patcher.clone()
|
|
||||||
n.cond_stage_model = self.cond_stage_model
|
|
||||||
n.tokenizer = self.tokenizer
|
|
||||||
return n
|
|
||||||
|
|
||||||
def load_sd(self, sd):
|
|
||||||
return self.cond_stage_model.load_sd(sd)
|
|
||||||
|
|
||||||
def get_sd(self):
|
|
||||||
return self.cond_stage_model.state_dict()
|
|
||||||
|
|
||||||
def load_model(self):
|
|
||||||
if self.load_device != "cpu":
|
|
||||||
model_management.load_model_gpu(self.patcher)
|
|
||||||
return self.patcher
|
|
||||||
|
|
||||||
def add_patches(self, patches, strength_patch=1.0, strength_model=1.0):
|
|
||||||
return self.patcher.add_patches(patches, strength_patch, strength_model)
|
|
||||||
|
|
||||||
def get_key_patches(self):
|
|
||||||
return self.patcher.get_key_patches()
|
|
||||||
|
|
||||||
# MT5 model
|
|
||||||
class mT5Model(torch.nn.Module):
|
|
||||||
def __init__(self, textmodel_json_config=None, device="cpu", max_length=256, freeze=True, dtype=None):
|
|
||||||
super().__init__()
|
|
||||||
self.device = device
|
|
||||||
self.dtype = dtype
|
|
||||||
self.max_length = max_length
|
|
||||||
if textmodel_json_config is None:
|
|
||||||
textmodel_json_config = os.path.join(
|
|
||||||
os.path.dirname(os.path.realpath(__file__)),
|
|
||||||
f"config_mt5.json"
|
|
||||||
)
|
|
||||||
config = T5Config.from_json_file(textmodel_json_config)
|
|
||||||
with modeling_utils.no_init_weights():
|
|
||||||
self.transformer = T5EncoderModel(config)
|
|
||||||
self.to(dtype)
|
|
||||||
if freeze:
|
|
||||||
self.freeze()
|
|
||||||
|
|
||||||
def freeze(self):
|
|
||||||
self.transformer = self.transformer.eval()
|
|
||||||
for param in self.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
|
|
||||||
def load_sd(self, sd):
|
|
||||||
return self.transformer.load_state_dict(sd, strict=False)
|
|
||||||
|
|
||||||
def to(self, *args, **kwargs):
|
|
||||||
return self.transformer.to(*args, **kwargs)
|
|
||||||
|
|
||||||
@@ -1,377 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
from typing import Tuple, Union, Optional
|
|
||||||
|
|
||||||
try:
|
|
||||||
import flash_attn
|
|
||||||
if hasattr(flash_attn, '__version__') and int(flash_attn.__version__[0]) == 2:
|
|
||||||
from flash_attn.flash_attn_interface import flash_attn_kvpacked_func
|
|
||||||
from flash_attn.modules.mha import FlashSelfAttention, FlashCrossAttention
|
|
||||||
else:
|
|
||||||
from flash_attn.flash_attn_interface import flash_attn_unpadded_kvpacked_func
|
|
||||||
from flash_attn.modules.mha import FlashSelfAttention, FlashCrossAttention
|
|
||||||
except Exception as e:
|
|
||||||
print(f'flash_attn import failed: {e}')
|
|
||||||
|
|
||||||
|
|
||||||
def reshape_for_broadcast(freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]], x: torch.Tensor, head_first=False):
|
|
||||||
"""
|
|
||||||
Reshape frequency tensor for broadcasting it with another tensor.
|
|
||||||
|
|
||||||
This function reshapes the frequency tensor to have the same shape as the target tensor 'x'
|
|
||||||
for the purpose of broadcasting the frequency tensor during element-wise operations.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
freqs_cis (Union[torch.Tensor, Tuple[torch.Tensor]]): Frequency tensor to be reshaped.
|
|
||||||
x (torch.Tensor): Target tensor for broadcasting compatibility.
|
|
||||||
head_first (bool): head dimension first (except batch dim) or not.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
torch.Tensor: Reshaped frequency tensor.
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
AssertionError: If the frequency tensor doesn't match the expected shape.
|
|
||||||
AssertionError: If the target tensor 'x' doesn't have the expected number of dimensions.
|
|
||||||
"""
|
|
||||||
ndim = x.ndim
|
|
||||||
assert 0 <= 1 < ndim
|
|
||||||
|
|
||||||
if isinstance(freqs_cis, tuple):
|
|
||||||
# freqs_cis: (cos, sin) in real space
|
|
||||||
if head_first:
|
|
||||||
assert freqs_cis[0].shape == (x.shape[-2], x.shape[-1]), f'freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}'
|
|
||||||
shape = [d if i == ndim - 2 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
|
||||||
else:
|
|
||||||
assert freqs_cis[0].shape == (x.shape[1], x.shape[-1]), f'freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}'
|
|
||||||
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
|
||||||
return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)
|
|
||||||
else:
|
|
||||||
# freqs_cis: values in complex space
|
|
||||||
if head_first:
|
|
||||||
assert freqs_cis.shape == (x.shape[-2], x.shape[-1]), f'freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}'
|
|
||||||
shape = [d if i == ndim - 2 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
|
||||||
else:
|
|
||||||
assert freqs_cis.shape == (x.shape[1], x.shape[-1]), f'freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}'
|
|
||||||
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
|
||||||
return freqs_cis.view(*shape)
|
|
||||||
|
|
||||||
|
|
||||||
def rotate_half(x):
|
|
||||||
x_real, x_imag = x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
|
|
||||||
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
|
||||||
|
|
||||||
|
|
||||||
def apply_rotary_emb(
|
|
||||||
xq: torch.Tensor,
|
|
||||||
xk: Optional[torch.Tensor],
|
|
||||||
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
|
|
||||||
head_first: bool = False,
|
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
|
||||||
"""
|
|
||||||
Apply rotary embeddings to input tensors using the given frequency tensor.
|
|
||||||
|
|
||||||
This function applies rotary embeddings to the given query 'xq' and key 'xk' tensors using the provided
|
|
||||||
frequency tensor 'freqs_cis'. The input tensors are reshaped as complex numbers, and the frequency tensor
|
|
||||||
is reshaped for broadcasting compatibility. The resulting tensors contain rotary embeddings and are
|
|
||||||
returned as real tensors.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
xq (torch.Tensor): Query tensor to apply rotary embeddings. [B, S, H, D]
|
|
||||||
xk (torch.Tensor): Key tensor to apply rotary embeddings. [B, S, H, D]
|
|
||||||
freqs_cis (Union[torch.Tensor, Tuple[torch.Tensor]]): Precomputed frequency tensor for complex exponentials.
|
|
||||||
head_first (bool): head dimension first (except batch dim) or not.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
|
|
||||||
|
|
||||||
"""
|
|
||||||
xk_out = None
|
|
||||||
if isinstance(freqs_cis, tuple):
|
|
||||||
cos, sin = reshape_for_broadcast(freqs_cis, xq, head_first) # [S, D]
|
|
||||||
cos, sin = cos.to(xq.device), sin.to(xq.device)
|
|
||||||
xq_out = (xq.float() * cos + rotate_half(xq.float()) * sin).type_as(xq)
|
|
||||||
if xk is not None:
|
|
||||||
xk_out = (xk.float() * cos + rotate_half(xk.float()) * sin).type_as(xk)
|
|
||||||
else:
|
|
||||||
xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2)) # [B, S, H, D//2]
|
|
||||||
freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to(xq.device) # [S, D//2] --> [1, S, 1, D//2]
|
|
||||||
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3).type_as(xq)
|
|
||||||
if xk is not None:
|
|
||||||
xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2)) # [B, S, H, D//2]
|
|
||||||
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3).type_as(xk)
|
|
||||||
|
|
||||||
return xq_out, xk_out
|
|
||||||
|
|
||||||
|
|
||||||
class FlashSelfMHAModified(nn.Module):
|
|
||||||
"""
|
|
||||||
Use QK Normalization.
|
|
||||||
"""
|
|
||||||
def __init__(self,
|
|
||||||
dim,
|
|
||||||
num_heads,
|
|
||||||
qkv_bias=True,
|
|
||||||
qk_norm=False,
|
|
||||||
attn_drop=0.0,
|
|
||||||
proj_drop=0.0,
|
|
||||||
device=None,
|
|
||||||
dtype=None,
|
|
||||||
norm_layer=nn.LayerNorm,
|
|
||||||
):
|
|
||||||
factory_kwargs = {'device': device, 'dtype': dtype}
|
|
||||||
super().__init__()
|
|
||||||
self.dim = dim
|
|
||||||
self.num_heads = num_heads
|
|
||||||
assert self.dim % num_heads == 0, "self.kdim must be divisible by num_heads"
|
|
||||||
self.head_dim = self.dim // num_heads
|
|
||||||
assert self.head_dim % 8 == 0 and self.head_dim <= 128, "Only support head_dim <= 128 and divisible by 8"
|
|
||||||
|
|
||||||
self.Wqkv = nn.Linear(dim, 3 * dim, bias=qkv_bias, **factory_kwargs)
|
|
||||||
# TODO: eps should be 1 / 65530 if using fp16
|
|
||||||
self.q_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
self.k_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
self.inner_attn = FlashSelfAttention(attention_dropout=attn_drop)
|
|
||||||
self.out_proj = nn.Linear(dim, dim, bias=qkv_bias, **factory_kwargs)
|
|
||||||
self.proj_drop = nn.Dropout(proj_drop)
|
|
||||||
|
|
||||||
def forward(self, x, freqs_cis_img=None):
|
|
||||||
"""
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
x: torch.Tensor
|
|
||||||
(batch, seqlen, hidden_dim) (where hidden_dim = num heads * head dim)
|
|
||||||
freqs_cis_img: torch.Tensor
|
|
||||||
(batch, hidden_dim // 2), RoPE for image
|
|
||||||
"""
|
|
||||||
b, s, d = x.shape
|
|
||||||
|
|
||||||
qkv = self.Wqkv(x)
|
|
||||||
qkv = qkv.view(b, s, 3, self.num_heads, self.head_dim) # [b, s, 3, h, d]
|
|
||||||
q, k, v = qkv.unbind(dim=2) # [b, s, h, d]
|
|
||||||
q = self.q_norm(q).half() # [b, s, h, d]
|
|
||||||
k = self.k_norm(k).half()
|
|
||||||
|
|
||||||
# Apply RoPE if needed
|
|
||||||
if freqs_cis_img is not None:
|
|
||||||
qq, kk = apply_rotary_emb(q, k, freqs_cis_img)
|
|
||||||
assert qq.shape == q.shape and kk.shape == k.shape, f'qq: {qq.shape}, q: {q.shape}, kk: {kk.shape}, k: {k.shape}'
|
|
||||||
q, k = qq, kk
|
|
||||||
|
|
||||||
qkv = torch.stack([q, k, v], dim=2) # [b, s, 3, h, d]
|
|
||||||
context = self.inner_attn(qkv)
|
|
||||||
out = self.out_proj(context.view(b, s, d))
|
|
||||||
out = self.proj_drop(out)
|
|
||||||
|
|
||||||
out_tuple = (out,)
|
|
||||||
|
|
||||||
return out_tuple
|
|
||||||
|
|
||||||
|
|
||||||
class FlashCrossMHAModified(nn.Module):
|
|
||||||
"""
|
|
||||||
Use QK Normalization.
|
|
||||||
"""
|
|
||||||
def __init__(self,
|
|
||||||
qdim,
|
|
||||||
kdim,
|
|
||||||
num_heads,
|
|
||||||
qkv_bias=True,
|
|
||||||
qk_norm=False,
|
|
||||||
attn_drop=0.0,
|
|
||||||
proj_drop=0.0,
|
|
||||||
device=None,
|
|
||||||
dtype=None,
|
|
||||||
norm_layer=nn.LayerNorm,
|
|
||||||
):
|
|
||||||
factory_kwargs = {'device': device, 'dtype': dtype}
|
|
||||||
super().__init__()
|
|
||||||
self.qdim = qdim
|
|
||||||
self.kdim = kdim
|
|
||||||
self.num_heads = num_heads
|
|
||||||
assert self.qdim % num_heads == 0, "self.qdim must be divisible by num_heads"
|
|
||||||
self.head_dim = self.qdim // num_heads
|
|
||||||
assert self.head_dim % 8 == 0 and self.head_dim <= 128, "Only support head_dim <= 128 and divisible by 8"
|
|
||||||
|
|
||||||
self.scale = self.head_dim ** -0.5
|
|
||||||
|
|
||||||
self.q_proj = nn.Linear(qdim, qdim, bias=qkv_bias, **factory_kwargs)
|
|
||||||
self.kv_proj = nn.Linear(kdim, 2 * qdim, bias=qkv_bias, **factory_kwargs)
|
|
||||||
|
|
||||||
# TODO: eps should be 1 / 65530 if using fp16
|
|
||||||
self.q_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
self.k_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
|
|
||||||
self.inner_attn = FlashCrossAttention(attention_dropout=attn_drop)
|
|
||||||
self.out_proj = nn.Linear(qdim, qdim, bias=qkv_bias, **factory_kwargs)
|
|
||||||
self.proj_drop = nn.Dropout(proj_drop)
|
|
||||||
|
|
||||||
def forward(self, x, y, freqs_cis_img=None):
|
|
||||||
"""
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
x: torch.Tensor
|
|
||||||
(batch, seqlen1, hidden_dim) (where hidden_dim = num_heads * head_dim)
|
|
||||||
y: torch.Tensor
|
|
||||||
(batch, seqlen2, hidden_dim2)
|
|
||||||
freqs_cis_img: torch.Tensor
|
|
||||||
(batch, hidden_dim // num_heads), RoPE for image
|
|
||||||
"""
|
|
||||||
b, s1, _ = x.shape # [b, s1, D]
|
|
||||||
_, s2, _ = y.shape # [b, s2, 1024]
|
|
||||||
|
|
||||||
q = self.q_proj(x).view(b, s1, self.num_heads, self.head_dim) # [b, s1, h, d]
|
|
||||||
kv = self.kv_proj(y).view(b, s2, 2, self.num_heads, self.head_dim) # [b, s2, 2, h, d]
|
|
||||||
k, v = kv.unbind(dim=2) # [b, s2, h, d]
|
|
||||||
q = self.q_norm(q).half() # [b, s1, h, d]
|
|
||||||
k = self.k_norm(k).half() # [b, s2, h, d]
|
|
||||||
|
|
||||||
# Apply RoPE if needed
|
|
||||||
if freqs_cis_img is not None:
|
|
||||||
qq, _ = apply_rotary_emb(q, None, freqs_cis_img)
|
|
||||||
assert qq.shape == q.shape, f'qq: {qq.shape}, q: {q.shape}'
|
|
||||||
q = qq # [b, s1, h, d]
|
|
||||||
kv = torch.stack([k, v], dim=2) # [b, s1, 2, h, d]
|
|
||||||
context = self.inner_attn(q, kv) # [b, s1, h, d]
|
|
||||||
context = context.view(b, s1, -1) # [b, s1, D]
|
|
||||||
|
|
||||||
out = self.out_proj(context)
|
|
||||||
out = self.proj_drop(out)
|
|
||||||
|
|
||||||
out_tuple = (out,)
|
|
||||||
|
|
||||||
return out_tuple
|
|
||||||
|
|
||||||
|
|
||||||
class CrossAttention(nn.Module):
|
|
||||||
"""
|
|
||||||
Use QK Normalization.
|
|
||||||
"""
|
|
||||||
def __init__(self,
|
|
||||||
qdim,
|
|
||||||
kdim,
|
|
||||||
num_heads,
|
|
||||||
qkv_bias=True,
|
|
||||||
qk_norm=False,
|
|
||||||
attn_drop=0.0,
|
|
||||||
proj_drop=0.0,
|
|
||||||
device=None,
|
|
||||||
dtype=None,
|
|
||||||
norm_layer=nn.LayerNorm,
|
|
||||||
):
|
|
||||||
factory_kwargs = {'device': device, 'dtype': dtype}
|
|
||||||
super().__init__()
|
|
||||||
self.qdim = qdim
|
|
||||||
self.kdim = kdim
|
|
||||||
self.num_heads = num_heads
|
|
||||||
assert self.qdim % num_heads == 0, "self.qdim must be divisible by num_heads"
|
|
||||||
self.head_dim = self.qdim // num_heads
|
|
||||||
assert self.head_dim % 8 == 0 and self.head_dim <= 128, "Only support head_dim <= 128 and divisible by 8"
|
|
||||||
self.scale = self.head_dim ** -0.5
|
|
||||||
|
|
||||||
self.q_proj = nn.Linear(qdim, qdim, bias=qkv_bias, **factory_kwargs)
|
|
||||||
self.kv_proj = nn.Linear(kdim, 2 * qdim, bias=qkv_bias, **factory_kwargs)
|
|
||||||
|
|
||||||
# TODO: eps should be 1 / 65530 if using fp16
|
|
||||||
self.q_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
self.k_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
self.attn_drop = nn.Dropout(attn_drop)
|
|
||||||
self.out_proj = nn.Linear(qdim, qdim, bias=qkv_bias, **factory_kwargs)
|
|
||||||
self.proj_drop = nn.Dropout(proj_drop)
|
|
||||||
|
|
||||||
def forward(self, x, y, freqs_cis_img=None):
|
|
||||||
"""
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
x: torch.Tensor
|
|
||||||
(batch, seqlen1, hidden_dim) (where hidden_dim = num heads * head dim)
|
|
||||||
y: torch.Tensor
|
|
||||||
(batch, seqlen2, hidden_dim2)
|
|
||||||
freqs_cis_img: torch.Tensor
|
|
||||||
(batch, hidden_dim // 2), RoPE for image
|
|
||||||
"""
|
|
||||||
b, s1, c = x.shape # [b, s1, D]
|
|
||||||
_, s2, c = y.shape # [b, s2, 1024]
|
|
||||||
|
|
||||||
q = self.q_proj(x).view(b, s1, self.num_heads, self.head_dim) # [b, s1, h, d]
|
|
||||||
kv = self.kv_proj(y).view(b, s2, 2, self.num_heads, self.head_dim) # [b, s2, 2, h, d]
|
|
||||||
k, v = kv.unbind(dim=2) # [b, s, h, d]
|
|
||||||
q = self.q_norm(q)
|
|
||||||
k = self.k_norm(k)
|
|
||||||
|
|
||||||
# Apply RoPE if needed
|
|
||||||
if freqs_cis_img is not None:
|
|
||||||
qq, _ = apply_rotary_emb(q, None, freqs_cis_img)
|
|
||||||
assert qq.shape == q.shape, f'qq: {qq.shape}, q: {q.shape}'
|
|
||||||
q = qq
|
|
||||||
|
|
||||||
q = q * self.scale
|
|
||||||
q = q.transpose(-2, -3).contiguous() # q -> B, L1, H, C - B, H, L1, C
|
|
||||||
k = k.permute(0, 2, 3, 1).contiguous() # k -> B, L2, H, C - B, H, C, L2
|
|
||||||
attn = q @ k # attn -> B, H, L1, L2
|
|
||||||
attn = attn.softmax(dim=-1) # attn -> B, H, L1, L2
|
|
||||||
attn = self.attn_drop(attn)
|
|
||||||
x = attn @ v.transpose(-2, -3) # v -> B, L2, H, C - B, H, L2, C x-> B, H, L1, C
|
|
||||||
context = x.transpose(1, 2) # context -> B, H, L1, C - B, L1, H, C
|
|
||||||
|
|
||||||
context = context.contiguous().view(b, s1, -1)
|
|
||||||
|
|
||||||
out = self.out_proj(context) # context.reshape - B, L1, -1
|
|
||||||
out = self.proj_drop(out)
|
|
||||||
|
|
||||||
out_tuple = (out,)
|
|
||||||
|
|
||||||
return out_tuple
|
|
||||||
|
|
||||||
|
|
||||||
class Attention(nn.Module):
|
|
||||||
"""
|
|
||||||
We rename some layer names to align with flash attention
|
|
||||||
"""
|
|
||||||
def __init__(self, dim, num_heads, qkv_bias=True, qk_norm=False, attn_drop=0., proj_drop=0.,
|
|
||||||
norm_layer=nn.LayerNorm,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.dim = dim
|
|
||||||
self.num_heads = num_heads
|
|
||||||
assert self.dim % num_heads == 0, 'dim should be divisible by num_heads'
|
|
||||||
self.head_dim = self.dim // num_heads
|
|
||||||
# This assertion is aligned with flash attention
|
|
||||||
assert self.head_dim % 8 == 0 and self.head_dim <= 128, "Only support head_dim <= 128 and divisible by 8"
|
|
||||||
self.scale = self.head_dim ** -0.5
|
|
||||||
|
|
||||||
# qkv --> Wqkv
|
|
||||||
self.Wqkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
|
||||||
# TODO: eps should be 1 / 65530 if using fp16
|
|
||||||
self.q_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
self.k_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity()
|
|
||||||
self.attn_drop = nn.Dropout(attn_drop)
|
|
||||||
self.out_proj = nn.Linear(dim, dim)
|
|
||||||
self.proj_drop = nn.Dropout(proj_drop)
|
|
||||||
|
|
||||||
def forward(self, x, freqs_cis_img=None):
|
|
||||||
B, N, C = x.shape
|
|
||||||
qkv = self.Wqkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) # [3, b, h, s, d]
|
|
||||||
q, k, v = qkv.unbind(0) # [b, h, s, d]
|
|
||||||
q = self.q_norm(q) # [b, h, s, d]
|
|
||||||
k = self.k_norm(k) # [b, h, s, d]
|
|
||||||
|
|
||||||
# Apply RoPE if needed
|
|
||||||
if freqs_cis_img is not None:
|
|
||||||
qq, kk = apply_rotary_emb(q, k, freqs_cis_img, head_first=True)
|
|
||||||
assert qq.shape == q.shape and kk.shape == k.shape, \
|
|
||||||
f'qq: {qq.shape}, q: {q.shape}, kk: {kk.shape}, k: {k.shape}'
|
|
||||||
q, k = qq, kk
|
|
||||||
|
|
||||||
q = q * self.scale
|
|
||||||
attn = q @ k.transpose(-2, -1) # [b, h, s, d] @ [b, h, d, s]
|
|
||||||
attn = attn.softmax(dim=-1) # [b, h, s, s]
|
|
||||||
attn = self.attn_drop(attn)
|
|
||||||
x = attn @ v # [b, h, s, d]
|
|
||||||
|
|
||||||
x = x.transpose(1, 2).reshape(B, N, C) # [b, s, h, d]
|
|
||||||
x = self.out_proj(x)
|
|
||||||
x = self.proj_drop(x)
|
|
||||||
|
|
||||||
out_tuple = (x,)
|
|
||||||
|
|
||||||
return out_tuple
|
|
||||||
@@ -1,111 +0,0 @@
|
|||||||
import math
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
from einops import repeat
|
|
||||||
|
|
||||||
from timm.models.layers import to_2tuple
|
|
||||||
|
|
||||||
|
|
||||||
class PatchEmbed(nn.Module):
|
|
||||||
""" 2D Image to Patch Embedding
|
|
||||||
|
|
||||||
Image to Patch Embedding using Conv2d
|
|
||||||
|
|
||||||
A convolution based approach to patchifying a 2D image w/ embedding projection.
|
|
||||||
|
|
||||||
Based on the impl in https://github.com/google-research/vision_transformer
|
|
||||||
|
|
||||||
Hacked together by / Copyright 2020 Ross Wightman
|
|
||||||
|
|
||||||
Remove the _assert function in forward function to be compatible with multi-resolution images.
|
|
||||||
"""
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
img_size=224,
|
|
||||||
patch_size=16,
|
|
||||||
in_chans=3,
|
|
||||||
embed_dim=768,
|
|
||||||
norm_layer=None,
|
|
||||||
flatten=True,
|
|
||||||
bias=True,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
if isinstance(img_size, int):
|
|
||||||
img_size = to_2tuple(img_size)
|
|
||||||
elif isinstance(img_size, (tuple, list)) and len(img_size) == 2:
|
|
||||||
img_size = tuple(img_size)
|
|
||||||
else:
|
|
||||||
raise ValueError(f"img_size must be int or tuple/list of length 2. Got {img_size}")
|
|
||||||
patch_size = to_2tuple(patch_size)
|
|
||||||
self.img_size = img_size
|
|
||||||
self.patch_size = patch_size
|
|
||||||
self.grid_size = (img_size[0] // patch_size[0], img_size[1] // patch_size[1])
|
|
||||||
self.num_patches = self.grid_size[0] * self.grid_size[1]
|
|
||||||
self.flatten = flatten
|
|
||||||
|
|
||||||
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias)
|
|
||||||
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
|
||||||
|
|
||||||
def update_image_size(self, img_size):
|
|
||||||
self.img_size = img_size
|
|
||||||
self.grid_size = (img_size[0] // self.patch_size[0], img_size[1] // self.patch_size[1])
|
|
||||||
self.num_patches = self.grid_size[0] * self.grid_size[1]
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
# B, C, H, W = x.shape
|
|
||||||
# _assert(H == self.img_size[0], f"Input image height ({H}) doesn't match model ({self.img_size[0]}).")
|
|
||||||
# _assert(W == self.img_size[1], f"Input image width ({W}) doesn't match model ({self.img_size[1]}).")
|
|
||||||
x = self.proj(x)
|
|
||||||
if self.flatten:
|
|
||||||
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
|
|
||||||
x = self.norm(x)
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
def timestep_embedding(t, dim, max_period=10000, repeat_only=False):
|
|
||||||
"""
|
|
||||||
Create sinusoidal timestep embeddings.
|
|
||||||
:param t: a 1-D Tensor of N indices, one per batch element.
|
|
||||||
These may be fractional.
|
|
||||||
:param dim: the dimension of the output.
|
|
||||||
:param max_period: controls the minimum frequency of the embeddings.
|
|
||||||
:return: an (N, D) Tensor of positional embeddings.
|
|
||||||
"""
|
|
||||||
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
|
||||||
if not repeat_only:
|
|
||||||
half = dim // 2
|
|
||||||
freqs = torch.exp(
|
|
||||||
-math.log(max_period)
|
|
||||||
* torch.arange(start=0, end=half, dtype=torch.float32)
|
|
||||||
/ half
|
|
||||||
).to(device=t.device) # size: [dim/2], 一个指数衰减的曲线
|
|
||||||
args = t[:, None].float() * freqs[None]
|
|
||||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
|
||||||
if dim % 2:
|
|
||||||
embedding = torch.cat(
|
|
||||||
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
embedding = repeat(t, "b -> b d", d=dim)
|
|
||||||
return embedding
|
|
||||||
|
|
||||||
|
|
||||||
class TimestepEmbedder(nn.Module):
|
|
||||||
"""
|
|
||||||
Embeds scalar timesteps into vector representations.
|
|
||||||
"""
|
|
||||||
def __init__(self, hidden_size, frequency_embedding_size=256, out_size=None):
|
|
||||||
super().__init__()
|
|
||||||
if out_size is None:
|
|
||||||
out_size = hidden_size
|
|
||||||
self.mlp = nn.Sequential(
|
|
||||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
|
||||||
nn.SiLU(),
|
|
||||||
nn.Linear(hidden_size, out_size, bias=True),
|
|
||||||
)
|
|
||||||
self.frequency_embedding_size = frequency_embedding_size
|
|
||||||
|
|
||||||
def forward(self, t):
|
|
||||||
t_freq = timestep_embedding(t, self.frequency_embedding_size).type(self.mlp[0].weight.dtype)
|
|
||||||
t_emb = self.mlp(t_freq)
|
|
||||||
return t_emb
|
|
||||||
@@ -1,428 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from timm.models.vision_transformer import Mlp
|
|
||||||
|
|
||||||
from .attn_layers import Attention, FlashCrossMHAModified, FlashSelfMHAModified, CrossAttention
|
|
||||||
from .embedders import TimestepEmbedder, PatchEmbed, timestep_embedding
|
|
||||||
from .norm_layers import RMSNorm
|
|
||||||
from .poolers import AttentionPool
|
|
||||||
from .posemb_layers import get_2d_rotary_pos_embed, get_fill_resize_and_crop
|
|
||||||
|
|
||||||
def modulate(x, shift, scale):
|
|
||||||
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
|
||||||
|
|
||||||
|
|
||||||
class FP32_Layernorm(nn.LayerNorm):
|
|
||||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
|
||||||
origin_dtype = inputs.dtype
|
|
||||||
return F.layer_norm(inputs.float(), self.normalized_shape, self.weight.float(), self.bias.float(),
|
|
||||||
self.eps).to(origin_dtype)
|
|
||||||
|
|
||||||
|
|
||||||
class FP32_SiLU(nn.SiLU):
|
|
||||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
|
||||||
return torch.nn.functional.silu(inputs.float(), inplace=False).to(inputs.dtype)
|
|
||||||
|
|
||||||
|
|
||||||
class HunYuanDiTBlock(nn.Module):
|
|
||||||
"""
|
|
||||||
A HunYuanDiT block with `add` conditioning.
|
|
||||||
"""
|
|
||||||
def __init__(self,
|
|
||||||
hidden_size,
|
|
||||||
c_emb_size,
|
|
||||||
num_heads,
|
|
||||||
mlp_ratio=4.0,
|
|
||||||
text_states_dim=1024,
|
|
||||||
use_flash_attn=False,
|
|
||||||
qk_norm=False,
|
|
||||||
norm_type="layer",
|
|
||||||
skip=False,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.use_flash_attn = use_flash_attn
|
|
||||||
use_ele_affine = True
|
|
||||||
|
|
||||||
if norm_type == "layer":
|
|
||||||
norm_layer = FP32_Layernorm
|
|
||||||
elif norm_type == "rms":
|
|
||||||
norm_layer = RMSNorm
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unknown norm_type: {norm_type}")
|
|
||||||
|
|
||||||
# ========================= Self-Attention =========================
|
|
||||||
self.norm1 = norm_layer(hidden_size, elementwise_affine=use_ele_affine, eps=1e-6)
|
|
||||||
if use_flash_attn:
|
|
||||||
self.attn1 = FlashSelfMHAModified(hidden_size, num_heads=num_heads, qkv_bias=True, qk_norm=qk_norm)
|
|
||||||
else:
|
|
||||||
self.attn1 = Attention(hidden_size, num_heads=num_heads, qkv_bias=True, qk_norm=qk_norm)
|
|
||||||
|
|
||||||
# ========================= FFN =========================
|
|
||||||
self.norm2 = norm_layer(hidden_size, elementwise_affine=use_ele_affine, eps=1e-6)
|
|
||||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
|
||||||
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
|
||||||
self.mlp = Mlp(in_features=hidden_size, hidden_features=mlp_hidden_dim, act_layer=approx_gelu, drop=0)
|
|
||||||
|
|
||||||
# ========================= Add =========================
|
|
||||||
# Simply use add like SDXL.
|
|
||||||
self.default_modulation = nn.Sequential(
|
|
||||||
FP32_SiLU(),
|
|
||||||
nn.Linear(c_emb_size, hidden_size, bias=True)
|
|
||||||
)
|
|
||||||
|
|
||||||
# ========================= Cross-Attention =========================
|
|
||||||
if use_flash_attn:
|
|
||||||
self.attn2 = FlashCrossMHAModified(hidden_size, text_states_dim, num_heads=num_heads, qkv_bias=True,
|
|
||||||
qk_norm=qk_norm)
|
|
||||||
else:
|
|
||||||
self.attn2 = CrossAttention(hidden_size, text_states_dim, num_heads=num_heads, qkv_bias=True,
|
|
||||||
qk_norm=qk_norm)
|
|
||||||
self.norm3 = norm_layer(hidden_size, elementwise_affine=True, eps=1e-6)
|
|
||||||
|
|
||||||
# ========================= Skip Connection =========================
|
|
||||||
if skip:
|
|
||||||
self.skip_norm = norm_layer(2 * hidden_size, elementwise_affine=True, eps=1e-6)
|
|
||||||
self.skip_linear = nn.Linear(2 * hidden_size, hidden_size)
|
|
||||||
else:
|
|
||||||
self.skip_linear = None
|
|
||||||
|
|
||||||
def forward(self, x, c=None, text_states=None, freq_cis_img=None, skip=None):
|
|
||||||
# Long Skip Connection
|
|
||||||
if self.skip_linear is not None:
|
|
||||||
cat = torch.cat([x, skip], dim=-1)
|
|
||||||
cat = self.skip_norm(cat)
|
|
||||||
x = self.skip_linear(cat)
|
|
||||||
|
|
||||||
# Self-Attention
|
|
||||||
shift_msa = self.default_modulation(c).unsqueeze(dim=1)
|
|
||||||
attn_inputs = (
|
|
||||||
self.norm1(x) + shift_msa, freq_cis_img,
|
|
||||||
)
|
|
||||||
x = x + self.attn1(*attn_inputs)[0]
|
|
||||||
|
|
||||||
# Cross-Attention
|
|
||||||
cross_inputs = (
|
|
||||||
self.norm3(x), text_states, freq_cis_img
|
|
||||||
)
|
|
||||||
x = x + self.attn2(*cross_inputs)[0]
|
|
||||||
|
|
||||||
# FFN Layer
|
|
||||||
mlp_inputs = self.norm2(x)
|
|
||||||
x = x + self.mlp(mlp_inputs)
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class FinalLayer(nn.Module):
|
|
||||||
"""
|
|
||||||
The final layer of HunYuanDiT.
|
|
||||||
"""
|
|
||||||
def __init__(self, final_hidden_size, c_emb_size, patch_size, out_channels):
|
|
||||||
super().__init__()
|
|
||||||
self.norm_final = nn.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6)
|
|
||||||
self.linear = nn.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True)
|
|
||||||
self.adaLN_modulation = nn.Sequential(
|
|
||||||
FP32_SiLU(),
|
|
||||||
nn.Linear(c_emb_size, 2 * final_hidden_size, bias=True)
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(self, x, c):
|
|
||||||
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
|
||||||
x = modulate(self.norm_final(x), shift, scale)
|
|
||||||
x = self.linear(x)
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class HunYuanDiT(nn.Module):
|
|
||||||
"""
|
|
||||||
HunYuanDiT: Diffusion model with a Transformer backbone.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
args: argparse.Namespace
|
|
||||||
The arguments parsed by argparse.
|
|
||||||
input_size: tuple
|
|
||||||
The size of the input image.
|
|
||||||
patch_size: int
|
|
||||||
The size of the patch.
|
|
||||||
in_channels: int
|
|
||||||
The number of input channels.
|
|
||||||
hidden_size: int
|
|
||||||
The hidden size of the transformer backbone.
|
|
||||||
depth: int
|
|
||||||
The number of transformer blocks.
|
|
||||||
num_heads: int
|
|
||||||
The number of attention heads.
|
|
||||||
mlp_ratio: float
|
|
||||||
The ratio of the hidden size of the MLP in the transformer block.
|
|
||||||
log_fn: callable
|
|
||||||
The logging function.
|
|
||||||
"""
|
|
||||||
def __init__(
|
|
||||||
self, args,
|
|
||||||
input_size=(32, 32),
|
|
||||||
patch_size=2,
|
|
||||||
in_channels=4,
|
|
||||||
hidden_size=1152,
|
|
||||||
depth=28,
|
|
||||||
num_heads=16,
|
|
||||||
mlp_ratio=4.0,
|
|
||||||
log_fn=print,
|
|
||||||
**kwargs,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.args = args
|
|
||||||
self.log_fn = log_fn
|
|
||||||
self.depth = depth
|
|
||||||
self.learn_sigma = args.learn_sigma
|
|
||||||
self.in_channels = in_channels
|
|
||||||
self.out_channels = in_channels * 2 if args.learn_sigma else in_channels
|
|
||||||
self.patch_size = patch_size
|
|
||||||
self.num_heads = num_heads
|
|
||||||
self.hidden_size = hidden_size
|
|
||||||
self.head_size = hidden_size // num_heads
|
|
||||||
self.text_states_dim = args.text_states_dim
|
|
||||||
self.text_states_dim_t5 = args.text_states_dim_t5
|
|
||||||
self.text_len = args.text_len
|
|
||||||
self.text_len_t5 = args.text_len_t5
|
|
||||||
self.norm = args.norm
|
|
||||||
|
|
||||||
use_flash_attn = args.infer_mode == 'fa'
|
|
||||||
if use_flash_attn:
|
|
||||||
log_fn(f" Enable Flash Attention.")
|
|
||||||
qk_norm = True # See http://arxiv.org/abs/2302.05442 for details.
|
|
||||||
|
|
||||||
self.mlp_t5 = nn.Sequential(
|
|
||||||
nn.Linear(self.text_states_dim_t5, self.text_states_dim_t5 * 4, bias=True),
|
|
||||||
FP32_SiLU(),
|
|
||||||
nn.Linear(self.text_states_dim_t5 * 4, self.text_states_dim, bias=True),
|
|
||||||
)
|
|
||||||
# learnable replace
|
|
||||||
self.text_embedding_padding = nn.Parameter(
|
|
||||||
torch.randn(self.text_len + self.text_len_t5, self.text_states_dim, dtype=torch.float32))
|
|
||||||
|
|
||||||
# Attention pooling
|
|
||||||
self.pooler = AttentionPool(self.text_len_t5, self.text_states_dim_t5, num_heads=8, output_dim=1024)
|
|
||||||
|
|
||||||
# Here we use a default learned embedder layer for future extension.
|
|
||||||
self.style_embedder = nn.Embedding(1, hidden_size)
|
|
||||||
|
|
||||||
# Image size and crop size conditions
|
|
||||||
self.extra_in_dim = 256 * 6 + hidden_size
|
|
||||||
|
|
||||||
# Text embedding for `add`
|
|
||||||
self.last_size = input_size
|
|
||||||
self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size)
|
|
||||||
self.t_embedder = TimestepEmbedder(hidden_size)
|
|
||||||
self.extra_in_dim += 1024
|
|
||||||
self.extra_embedder = nn.Sequential(
|
|
||||||
nn.Linear(self.extra_in_dim, hidden_size * 4),
|
|
||||||
FP32_SiLU(),
|
|
||||||
nn.Linear(hidden_size * 4, hidden_size, bias=True),
|
|
||||||
)
|
|
||||||
|
|
||||||
# Image embedding
|
|
||||||
num_patches = self.x_embedder.num_patches
|
|
||||||
log_fn(f" Number of tokens: {num_patches}")
|
|
||||||
|
|
||||||
# HUnYuanDiT Blocks
|
|
||||||
self.blocks = nn.ModuleList([
|
|
||||||
HunYuanDiTBlock(hidden_size=hidden_size,
|
|
||||||
c_emb_size=hidden_size,
|
|
||||||
num_heads=num_heads,
|
|
||||||
mlp_ratio=mlp_ratio,
|
|
||||||
text_states_dim=self.text_states_dim,
|
|
||||||
use_flash_attn=use_flash_attn,
|
|
||||||
qk_norm=qk_norm,
|
|
||||||
norm_type=self.norm,
|
|
||||||
skip=layer > depth // 2,
|
|
||||||
)
|
|
||||||
for layer in range(depth)
|
|
||||||
])
|
|
||||||
|
|
||||||
self.final_layer = FinalLayer(hidden_size, hidden_size, patch_size, self.out_channels)
|
|
||||||
self.unpatchify_channels = self.out_channels
|
|
||||||
|
|
||||||
def forward_raw(self,
|
|
||||||
x,
|
|
||||||
t,
|
|
||||||
encoder_hidden_states=None,
|
|
||||||
text_embedding_mask=None,
|
|
||||||
encoder_hidden_states_t5=None,
|
|
||||||
text_embedding_mask_t5=None,
|
|
||||||
image_meta_size=None,
|
|
||||||
style=None,
|
|
||||||
cos_cis_img=None,
|
|
||||||
sin_cis_img=None,
|
|
||||||
return_dict=False,
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
Forward pass of the encoder.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
x: torch.Tensor
|
|
||||||
(B, D, H, W)
|
|
||||||
t: torch.Tensor
|
|
||||||
(B)
|
|
||||||
encoder_hidden_states: torch.Tensor
|
|
||||||
CLIP text embedding, (B, L_clip, D)
|
|
||||||
text_embedding_mask: torch.Tensor
|
|
||||||
CLIP text embedding mask, (B, L_clip)
|
|
||||||
encoder_hidden_states_t5: torch.Tensor
|
|
||||||
T5 text embedding, (B, L_t5, D)
|
|
||||||
text_embedding_mask_t5: torch.Tensor
|
|
||||||
T5 text embedding mask, (B, L_t5)
|
|
||||||
image_meta_size: torch.Tensor
|
|
||||||
(B, 6)
|
|
||||||
style: torch.Tensor
|
|
||||||
(B)
|
|
||||||
cos_cis_img: torch.Tensor
|
|
||||||
sin_cis_img: torch.Tensor
|
|
||||||
return_dict: bool
|
|
||||||
Whether to return a dictionary.
|
|
||||||
"""
|
|
||||||
|
|
||||||
text_states = encoder_hidden_states # 2,77,1024
|
|
||||||
text_states_t5 = encoder_hidden_states_t5 # 2,256,2048
|
|
||||||
text_states_mask = text_embedding_mask.bool() # 2,77
|
|
||||||
text_states_t5_mask = text_embedding_mask_t5.bool() # 2,256
|
|
||||||
b_t5, l_t5, c_t5 = text_states_t5.shape
|
|
||||||
text_states_t5 = self.mlp_t5(text_states_t5.view(-1, c_t5))
|
|
||||||
text_states = torch.cat([text_states, text_states_t5.view(b_t5, l_t5, -1)], dim=1) # 2,205,1024
|
|
||||||
clip_t5_mask = torch.cat([text_states_mask, text_states_t5_mask], dim=-1)
|
|
||||||
|
|
||||||
clip_t5_mask = clip_t5_mask
|
|
||||||
text_states = torch.where(clip_t5_mask.unsqueeze(2), text_states, self.text_embedding_padding.to(text_states))
|
|
||||||
|
|
||||||
_, _, oh, ow = x.shape
|
|
||||||
th, tw = oh // self.patch_size, ow // self.patch_size
|
|
||||||
|
|
||||||
# ========================= Build time and image embedding =========================
|
|
||||||
t = self.t_embedder(t)
|
|
||||||
x = self.x_embedder(x)
|
|
||||||
|
|
||||||
# Get image RoPE embedding according to `reso`lution.
|
|
||||||
freqs_cis_img = (cos_cis_img, sin_cis_img)
|
|
||||||
|
|
||||||
# ========================= Concatenate all extra vectors =========================
|
|
||||||
# Build text tokens with pooling
|
|
||||||
extra_vec = self.pooler(encoder_hidden_states_t5)
|
|
||||||
|
|
||||||
# Build image meta size tokens
|
|
||||||
image_meta_size = timestep_embedding(image_meta_size.view(-1), 256) # [B * 6, 256]
|
|
||||||
# if self.args.use_fp16:
|
|
||||||
# image_meta_size = image_meta_size.half()
|
|
||||||
image_meta_size = image_meta_size.view(-1, 6 * 256)
|
|
||||||
extra_vec = torch.cat([extra_vec, image_meta_size], dim=1) # [B, D + 6 * 256]
|
|
||||||
|
|
||||||
# Build style tokens
|
|
||||||
style_embedding = self.style_embedder(style)
|
|
||||||
extra_vec = torch.cat([extra_vec, style_embedding], dim=1)
|
|
||||||
|
|
||||||
# Concatenate all extra vectors
|
|
||||||
c = t + self.extra_embedder(extra_vec.to(self.dtype)) # [B, D]
|
|
||||||
|
|
||||||
# ========================= Forward pass through HunYuanDiT blocks =========================
|
|
||||||
skips = []
|
|
||||||
for layer, block in enumerate(self.blocks):
|
|
||||||
if layer > self.depth // 2:
|
|
||||||
skip = skips.pop()
|
|
||||||
x = block(x, c, text_states, freqs_cis_img, skip) # (N, L, D)
|
|
||||||
else:
|
|
||||||
x = block(x, c, text_states, freqs_cis_img) # (N, L, D)
|
|
||||||
|
|
||||||
if layer < (self.depth // 2 - 1):
|
|
||||||
skips.append(x)
|
|
||||||
|
|
||||||
# ========================= Final layer =========================
|
|
||||||
x = self.final_layer(x, c) # (N, L, patch_size ** 2 * out_channels)
|
|
||||||
x = self.unpatchify(x, th, tw) # (N, out_channels, H, W)
|
|
||||||
|
|
||||||
if return_dict:
|
|
||||||
return {'x': x}
|
|
||||||
return x
|
|
||||||
|
|
||||||
def calc_rope(self, height, width):
|
|
||||||
"""
|
|
||||||
Probably not the best in terms of perf to have this here
|
|
||||||
"""
|
|
||||||
th = height // 8 // self.patch_size
|
|
||||||
tw = width // 8 // self.patch_size
|
|
||||||
base_size = 512 // 8 // self.patch_size
|
|
||||||
start, stop = get_fill_resize_and_crop((th, tw), base_size)
|
|
||||||
sub_args = [start, stop, (th, tw)]
|
|
||||||
rope = get_2d_rotary_pos_embed(self.head_size, *sub_args)
|
|
||||||
return rope
|
|
||||||
|
|
||||||
def forward(self, x, timesteps, context, context_mask=None, context_t5=None, context_t5_mask=None, src_size_cond=(1024,1024), **kwargs):
|
|
||||||
"""
|
|
||||||
Forward pass that adapts comfy input to original forward function
|
|
||||||
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
|
||||||
timesteps: (N,) tensor of diffusion timesteps
|
|
||||||
context: (N, 1, 77, C) CLIP conditioning
|
|
||||||
context_t5: (N, 1, 256, C) MT5 conditioning
|
|
||||||
"""
|
|
||||||
# context_mask = torch.zeros(x.shape[0], 77, device=x.device)
|
|
||||||
# context_t5_mask = torch.zeros(x.shape[0], 256, device=x.device)
|
|
||||||
|
|
||||||
# style
|
|
||||||
style = torch.as_tensor([0] * (x.shape[0]), device=x.device)
|
|
||||||
|
|
||||||
# image size - todo separate for cond/uncond when batched
|
|
||||||
if torch.is_tensor(src_size_cond):
|
|
||||||
src_size_cond = (int(src_size_cond[0][0]), int(src_size_cond[0][1]))
|
|
||||||
|
|
||||||
image_size = (x.shape[2]//2*16, x.shape[3]//2*16)
|
|
||||||
size_cond = list(src_size_cond) + [image_size[1], image_size[0], 0, 0]
|
|
||||||
image_meta_size = torch.as_tensor([size_cond] * x.shape[0], device=x.device)
|
|
||||||
|
|
||||||
# RoPE
|
|
||||||
rope = self.calc_rope(*image_size)
|
|
||||||
|
|
||||||
# Update x_embedder if image size changed
|
|
||||||
if self.last_size != image_size:
|
|
||||||
from tqdm import tqdm
|
|
||||||
tqdm.write(f"HyDiT: New image size {image_size}")
|
|
||||||
self.x_embedder.update_image_size(
|
|
||||||
(image_size[0]//8, image_size[1]//8),
|
|
||||||
)
|
|
||||||
self.last_size = image_size
|
|
||||||
|
|
||||||
# Run original forward pass
|
|
||||||
out = self.forward_raw(
|
|
||||||
x = x.to(self.dtype),
|
|
||||||
t = timesteps.to(self.dtype),
|
|
||||||
encoder_hidden_states = context.to(self.dtype),
|
|
||||||
text_embedding_mask = context_mask.to(self.dtype),
|
|
||||||
encoder_hidden_states_t5 = context_t5.to(self.dtype),
|
|
||||||
text_embedding_mask_t5 = context_t5_mask.to(self.dtype),
|
|
||||||
image_meta_size = image_meta_size.to(self.dtype),
|
|
||||||
style = style,
|
|
||||||
cos_cis_img = rope[0],
|
|
||||||
sin_cis_img = rope[1],
|
|
||||||
)
|
|
||||||
|
|
||||||
# return
|
|
||||||
out = out.to(torch.float)
|
|
||||||
if self.learn_sigma:
|
|
||||||
eps, rest = out[:, :self.in_channels], out[:, self.in_channels:]
|
|
||||||
return eps
|
|
||||||
else:
|
|
||||||
return out
|
|
||||||
|
|
||||||
def unpatchify(self, x, h, w):
|
|
||||||
"""
|
|
||||||
x: (N, T, patch_size**2 * C)
|
|
||||||
imgs: (N, H, W, C)
|
|
||||||
"""
|
|
||||||
c = self.unpatchify_channels
|
|
||||||
p = self.x_embedder.patch_size[0]
|
|
||||||
# h = w = int(x.shape[1] ** 0.5)
|
|
||||||
assert h * w == x.shape[1]
|
|
||||||
|
|
||||||
x = x.reshape(shape=(x.shape[0], h, w, p, p, c))
|
|
||||||
x = torch.einsum('nhwpqc->nchpwq', x)
|
|
||||||
imgs = x.reshape(shape=(x.shape[0], c, h * p, w * p))
|
|
||||||
return imgs
|
|
||||||
@@ -1,68 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
|
|
||||||
|
|
||||||
class RMSNorm(nn.Module):
|
|
||||||
def __init__(self, dim: int, elementwise_affine=True, eps: float = 1e-6):
|
|
||||||
"""
|
|
||||||
Initialize the RMSNorm normalization layer.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
dim (int): The dimension of the input tensor.
|
|
||||||
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
|
|
||||||
|
|
||||||
Attributes:
|
|
||||||
eps (float): A small value added to the denominator for numerical stability.
|
|
||||||
weight (nn.Parameter): Learnable scaling parameter.
|
|
||||||
|
|
||||||
"""
|
|
||||||
super().__init__()
|
|
||||||
self.eps = eps
|
|
||||||
if elementwise_affine:
|
|
||||||
self.weight = nn.Parameter(torch.ones(dim))
|
|
||||||
|
|
||||||
def _norm(self, x):
|
|
||||||
"""
|
|
||||||
Apply the RMSNorm normalization to the input tensor.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
x (torch.Tensor): The input tensor.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
torch.Tensor: The normalized tensor.
|
|
||||||
|
|
||||||
"""
|
|
||||||
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
"""
|
|
||||||
Forward pass through the RMSNorm layer.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
x (torch.Tensor): The input tensor.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
torch.Tensor: The output tensor after applying RMSNorm.
|
|
||||||
|
|
||||||
"""
|
|
||||||
output = self._norm(x.float()).type_as(x)
|
|
||||||
if hasattr(self, "weight"):
|
|
||||||
output = output * self.weight
|
|
||||||
return output
|
|
||||||
|
|
||||||
|
|
||||||
class GroupNorm32(nn.GroupNorm):
|
|
||||||
def __init__(self, num_groups, num_channels, eps=1e-5, dtype=None):
|
|
||||||
super().__init__(num_groups=num_groups, num_channels=num_channels, eps=eps, dtype=dtype)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
y = super().forward(x).to(x.dtype)
|
|
||||||
return y
|
|
||||||
|
|
||||||
def normalization(channels, dtype=None):
|
|
||||||
"""
|
|
||||||
Make a standard normalization layer.
|
|
||||||
:param channels: number of input channels.
|
|
||||||
:return: an nn.Module for normalization.
|
|
||||||
"""
|
|
||||||
return GroupNorm32(num_channels=channels, num_groups=32, dtype=dtype)
|
|
||||||
@@ -1,39 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
|
|
||||||
class AttentionPool(nn.Module):
|
|
||||||
def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None):
|
|
||||||
super().__init__()
|
|
||||||
self.positional_embedding = nn.Parameter(torch.randn(spacial_dim + 1, embed_dim) / embed_dim ** 0.5)
|
|
||||||
self.k_proj = nn.Linear(embed_dim, embed_dim)
|
|
||||||
self.q_proj = nn.Linear(embed_dim, embed_dim)
|
|
||||||
self.v_proj = nn.Linear(embed_dim, embed_dim)
|
|
||||||
self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim)
|
|
||||||
self.num_heads = num_heads
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
x = x.permute(1, 0, 2) # NLC -> LNC
|
|
||||||
x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (L+1)NC
|
|
||||||
x = x + self.positional_embedding[:, None, :].to(x.dtype) # (L+1)NC
|
|
||||||
x, _ = F.multi_head_attention_forward(
|
|
||||||
query=x[:1], key=x, value=x,
|
|
||||||
embed_dim_to_check=x.shape[-1],
|
|
||||||
num_heads=self.num_heads,
|
|
||||||
q_proj_weight=self.q_proj.weight,
|
|
||||||
k_proj_weight=self.k_proj.weight,
|
|
||||||
v_proj_weight=self.v_proj.weight,
|
|
||||||
in_proj_weight=None,
|
|
||||||
in_proj_bias=torch.cat([self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]),
|
|
||||||
bias_k=None,
|
|
||||||
bias_v=None,
|
|
||||||
add_zero_attn=False,
|
|
||||||
dropout_p=0,
|
|
||||||
out_proj_weight=self.c_proj.weight,
|
|
||||||
out_proj_bias=self.c_proj.bias,
|
|
||||||
use_separate_proj_weight=True,
|
|
||||||
training=self.training,
|
|
||||||
need_weights=False
|
|
||||||
)
|
|
||||||
return x.squeeze(0)
|
|
||||||
@@ -1,225 +0,0 @@
|
|||||||
import torch
|
|
||||||
import numpy as np
|
|
||||||
from typing import Union
|
|
||||||
|
|
||||||
|
|
||||||
def _to_tuple(x):
|
|
||||||
if isinstance(x, int):
|
|
||||||
return x, x
|
|
||||||
else:
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
def get_fill_resize_and_crop(src, tgt): # src 来源的分辨率 tgt base 分辨率
|
|
||||||
th, tw = _to_tuple(tgt)
|
|
||||||
h, w = _to_tuple(src)
|
|
||||||
|
|
||||||
tr = th / tw # base 分辨率
|
|
||||||
r = h / w # 目标分辨率
|
|
||||||
|
|
||||||
# resize
|
|
||||||
if r > tr:
|
|
||||||
resize_height = th
|
|
||||||
resize_width = int(round(th / h * w))
|
|
||||||
else:
|
|
||||||
resize_width = tw
|
|
||||||
resize_height = int(round(tw / w * h)) # 根据base分辨率,将目标分辨率resize下来
|
|
||||||
|
|
||||||
crop_top = int(round((th - resize_height) / 2.0))
|
|
||||||
crop_left = int(round((tw - resize_width) / 2.0))
|
|
||||||
|
|
||||||
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
|
|
||||||
|
|
||||||
|
|
||||||
def get_meshgrid(start, *args):
|
|
||||||
if len(args) == 0:
|
|
||||||
# start is grid_size
|
|
||||||
num = _to_tuple(start)
|
|
||||||
start = (0, 0)
|
|
||||||
stop = num
|
|
||||||
elif len(args) == 1:
|
|
||||||
# start is start, args[0] is stop, step is 1
|
|
||||||
start = _to_tuple(start)
|
|
||||||
stop = _to_tuple(args[0])
|
|
||||||
num = (stop[0] - start[0], stop[1] - start[1])
|
|
||||||
elif len(args) == 2:
|
|
||||||
# start is start, args[0] is stop, args[1] is num
|
|
||||||
start = _to_tuple(start) # 左上角 eg: 12,0
|
|
||||||
stop = _to_tuple(args[0]) # 右下角 eg: 20,32
|
|
||||||
num = _to_tuple(args[1]) # 目标大小 eg: 32,124
|
|
||||||
else:
|
|
||||||
raise ValueError(f"len(args) should be 0, 1 or 2, but got {len(args)}")
|
|
||||||
|
|
||||||
grid_h = np.linspace(start[0], stop[0], num[0], endpoint=False, dtype=np.float32) # 12-20 中间差值32份 0-32 中间差值124份
|
|
||||||
grid_w = np.linspace(start[1], stop[1], num[1], endpoint=False, dtype=np.float32)
|
|
||||||
grid = np.meshgrid(grid_w, grid_h) # here w goes first
|
|
||||||
grid = np.stack(grid, axis=0) # [2, W, H]
|
|
||||||
return grid
|
|
||||||
|
|
||||||
#################################################################################
|
|
||||||
# Sine/Cosine Positional Embedding Functions #
|
|
||||||
#################################################################################
|
|
||||||
# https://github.com/facebookresearch/mae/blob/main/util/pos_embed.py
|
|
||||||
|
|
||||||
def get_2d_sincos_pos_embed(embed_dim, start, *args, cls_token=False, extra_tokens=0):
|
|
||||||
"""
|
|
||||||
grid_size: int of the grid height and width
|
|
||||||
return:
|
|
||||||
pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
|
|
||||||
"""
|
|
||||||
grid = get_meshgrid(start, *args) # [2, H, w]
|
|
||||||
# grid_h = np.arange(grid_size, dtype=np.float32)
|
|
||||||
# grid_w = np.arange(grid_size, dtype=np.float32)
|
|
||||||
# grid = np.meshgrid(grid_w, grid_h) # here w goes first
|
|
||||||
# grid = np.stack(grid, axis=0) # [2, W, H]
|
|
||||||
|
|
||||||
grid = grid.reshape([2, 1, *grid.shape[1:]])
|
|
||||||
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
|
|
||||||
if cls_token and extra_tokens > 0:
|
|
||||||
pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)
|
|
||||||
return pos_embed
|
|
||||||
|
|
||||||
|
|
||||||
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
|
|
||||||
assert embed_dim % 2 == 0
|
|
||||||
|
|
||||||
# use half of dimensions to encode grid_h
|
|
||||||
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
|
|
||||||
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
|
|
||||||
|
|
||||||
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
|
|
||||||
return emb
|
|
||||||
|
|
||||||
|
|
||||||
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
|
|
||||||
"""
|
|
||||||
embed_dim: output dimension for each position
|
|
||||||
pos: a list of positions to be encoded: size (W,H)
|
|
||||||
out: (M, D)
|
|
||||||
"""
|
|
||||||
assert embed_dim % 2 == 0
|
|
||||||
omega = np.arange(embed_dim // 2, dtype=np.float64)
|
|
||||||
omega /= embed_dim / 2.
|
|
||||||
omega = 1. / 10000**omega # (D/2,)
|
|
||||||
|
|
||||||
pos = pos.reshape(-1) # (M,)
|
|
||||||
out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
|
|
||||||
|
|
||||||
emb_sin = np.sin(out) # (M, D/2)
|
|
||||||
emb_cos = np.cos(out) # (M, D/2)
|
|
||||||
|
|
||||||
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
|
|
||||||
return emb
|
|
||||||
|
|
||||||
|
|
||||||
#################################################################################
|
|
||||||
# Rotary Positional Embedding Functions #
|
|
||||||
#################################################################################
|
|
||||||
# https://github.com/facebookresearch/llama/blob/main/llama/model.py#L443
|
|
||||||
|
|
||||||
def get_2d_rotary_pos_embed(embed_dim, start, *args, use_real=True):
|
|
||||||
"""
|
|
||||||
This is a 2d version of precompute_freqs_cis, which is a RoPE for image tokens with 2d structure.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
embed_dim: int
|
|
||||||
embedding dimension size
|
|
||||||
start: int or tuple of int
|
|
||||||
If len(args) == 0, start is num; If len(args) == 1, start is start, args[0] is stop, step is 1;
|
|
||||||
If len(args) == 2, start is start, args[0] is stop, args[1] is num.
|
|
||||||
use_real: bool
|
|
||||||
If True, return real part and imaginary part separately. Otherwise, return complex numbers.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
pos_embed: torch.Tensor
|
|
||||||
[HW, D/2]
|
|
||||||
"""
|
|
||||||
grid = get_meshgrid(start, *args) # [2, H, w]
|
|
||||||
grid = grid.reshape([2, 1, *grid.shape[1:]]) # 返回一个采样矩阵 分辨率与目标分辨率一致
|
|
||||||
pos_embed = get_2d_rotary_pos_embed_from_grid(embed_dim, grid, use_real=use_real)
|
|
||||||
return pos_embed
|
|
||||||
|
|
||||||
|
|
||||||
def get_2d_rotary_pos_embed_from_grid(embed_dim, grid, use_real=False):
|
|
||||||
assert embed_dim % 4 == 0
|
|
||||||
|
|
||||||
# use half of dimensions to encode grid_h
|
|
||||||
emb_h = get_1d_rotary_pos_embed(embed_dim // 2, grid[0].reshape(-1), use_real=use_real) # (H*W, D/4)
|
|
||||||
emb_w = get_1d_rotary_pos_embed(embed_dim // 2, grid[1].reshape(-1), use_real=use_real) # (H*W, D/4)
|
|
||||||
|
|
||||||
if use_real:
|
|
||||||
cos = torch.cat([emb_h[0], emb_w[0]], dim=1) # (H*W, D/2)
|
|
||||||
sin = torch.cat([emb_h[1], emb_w[1]], dim=1) # (H*W, D/2)
|
|
||||||
return cos, sin
|
|
||||||
else:
|
|
||||||
emb = torch.cat([emb_h, emb_w], dim=1) # (H*W, D/2)
|
|
||||||
return emb
|
|
||||||
|
|
||||||
|
|
||||||
def get_1d_rotary_pos_embed(dim: int, pos: Union[np.ndarray, int], theta: float = 10000.0, use_real=False):
|
|
||||||
"""
|
|
||||||
Precompute the frequency tensor for complex exponentials (cis) with given dimensions.
|
|
||||||
|
|
||||||
This function calculates a frequency tensor with complex exponentials using the given dimension 'dim'
|
|
||||||
and the end index 'end'. The 'theta' parameter scales the frequencies.
|
|
||||||
The returned tensor contains complex values in complex64 data type.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
dim (int): Dimension of the frequency tensor.
|
|
||||||
pos (np.ndarray, int): Position indices for the frequency tensor. [S] or scalar
|
|
||||||
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
|
|
||||||
use_real (bool, optional): If True, return real part and imaginary part separately.
|
|
||||||
Otherwise, return complex numbers.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
torch.Tensor: Precomputed frequency tensor with complex exponentials. [S, D/2]
|
|
||||||
|
|
||||||
"""
|
|
||||||
if isinstance(pos, int):
|
|
||||||
pos = np.arange(pos)
|
|
||||||
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)) # [D/2]
|
|
||||||
t = torch.from_numpy(pos).to(freqs.device) # type: ignore # [S]
|
|
||||||
freqs = torch.outer(t, freqs).float() # type: ignore # [S, D/2]
|
|
||||||
if use_real:
|
|
||||||
freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D]
|
|
||||||
freqs_sin = freqs.sin().repeat_interleave(2, dim=1) # [S, D]
|
|
||||||
return freqs_cos, freqs_sin
|
|
||||||
else:
|
|
||||||
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
|
|
||||||
return freqs_cis
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
def calc_sizes(rope_img, patch_size, th, tw):
|
|
||||||
""" 计算 RoPE 的尺寸. """
|
|
||||||
if rope_img == 'extend':
|
|
||||||
# 拓展模式
|
|
||||||
sub_args = [(th, tw)]
|
|
||||||
elif rope_img.startswith('base'):
|
|
||||||
# 基于一个尺寸, 其他尺寸插值获得.
|
|
||||||
base_size = int(rope_img[4:]) // 8 // patch_size # 基于512作为base,其他根据512差值得到
|
|
||||||
start, stop = get_fill_resize_and_crop((th, tw), base_size) # 需要在32x32里面 crop的左上角和右下角
|
|
||||||
sub_args = [start, stop, (th, tw)]
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unknown rope_img: {rope_img}")
|
|
||||||
return sub_args
|
|
||||||
|
|
||||||
|
|
||||||
def init_image_posemb(rope_img,
|
|
||||||
resolutions,
|
|
||||||
patch_size,
|
|
||||||
hidden_size,
|
|
||||||
num_heads,
|
|
||||||
log_fn,
|
|
||||||
rope_real=True,
|
|
||||||
):
|
|
||||||
freqs_cis_img = {}
|
|
||||||
for reso in resolutions:
|
|
||||||
th, tw = reso.height // 8 // patch_size, reso.width // 8 // patch_size
|
|
||||||
sub_args = calc_sizes(rope_img, patch_size, th, tw) # [左上角, 右下角, 目标高宽] 需要在32x32里面 crop的左上角和右下角
|
|
||||||
freqs_cis_img[str(reso)] = get_2d_rotary_pos_embed(hidden_size // num_heads, *sub_args, use_real=rope_real)
|
|
||||||
log_fn(f" Using image RoPE ({rope_img}) ({'real' if rope_real else 'complex'}): {sub_args} | ({reso}) "
|
|
||||||
f"{freqs_cis_img[str(reso)][0].shape if rope_real else freqs_cis_img[str(reso)].shape}")
|
|
||||||
return freqs_cis_img
|
|
||||||
@@ -1,33 +0,0 @@
|
|||||||
{
|
|
||||||
"_name_or_path": "mt5",
|
|
||||||
"architectures": [
|
|
||||||
"MT5ForConditionalGeneration"
|
|
||||||
],
|
|
||||||
"classifier_dropout": 0.0,
|
|
||||||
"d_ff": 5120,
|
|
||||||
"d_kv": 64,
|
|
||||||
"d_model": 2048,
|
|
||||||
"decoder_start_token_id": 0,
|
|
||||||
"dense_act_fn": "gelu_new",
|
|
||||||
"dropout_rate": 0.1,
|
|
||||||
"eos_token_id": 1,
|
|
||||||
"feed_forward_proj": "gated-gelu",
|
|
||||||
"initializer_factor": 1.0,
|
|
||||||
"is_encoder_decoder": true,
|
|
||||||
"is_gated_act": true,
|
|
||||||
"layer_norm_epsilon": 1e-06,
|
|
||||||
"model_type": "mt5",
|
|
||||||
"num_decoder_layers": 24,
|
|
||||||
"num_heads": 32,
|
|
||||||
"num_layers": 24,
|
|
||||||
"output_past": true,
|
|
||||||
"pad_token_id": 0,
|
|
||||||
"relative_attention_max_distance": 128,
|
|
||||||
"relative_attention_num_buckets": 32,
|
|
||||||
"tie_word_embeddings": false,
|
|
||||||
"tokenizer_class": "T5Tokenizer",
|
|
||||||
"torch_dtype": "float16",
|
|
||||||
"transformers_version": "4.40.2",
|
|
||||||
"use_cache": true,
|
|
||||||
"vocab_size": 250112
|
|
||||||
}
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
{"eos_token": "</s>", "unk_token": "<unk>", "pad_token": "<pad>"}
|
|
||||||
Binary file not shown.
@@ -1 +0,0 @@
|
|||||||
{"eos_token": "</s>", "unk_token": "<unk>", "pad_token": "<pad>", "extra_ids": 0, "additional_special_tokens": null, "special_tokens_map_file": "/home/patrick/.cache/torch/transformers/685ac0ca8568ec593a48b61b0a3c272beee9bc194a3c7241d15dcadb5f875e53.f76030f3ec1b96a8199b2593390c610e76ca8028ef3d24680000619ffb646276", "tokenizer_file": null, "name_or_path": "google/mt5-small"}
|
|
||||||
@@ -1,34 +0,0 @@
|
|||||||
{
|
|
||||||
"_name_or_path": "hfl/chinese-roberta-wwm-ext-large",
|
|
||||||
"architectures": [
|
|
||||||
"BertModel"
|
|
||||||
],
|
|
||||||
"attention_probs_dropout_prob": 0.1,
|
|
||||||
"bos_token_id": 0,
|
|
||||||
"classifier_dropout": null,
|
|
||||||
"directionality": "bidi",
|
|
||||||
"eos_token_id": 2,
|
|
||||||
"hidden_act": "gelu",
|
|
||||||
"hidden_dropout_prob": 0.1,
|
|
||||||
"hidden_size": 1024,
|
|
||||||
"initializer_range": 0.02,
|
|
||||||
"intermediate_size": 4096,
|
|
||||||
"layer_norm_eps": 1e-12,
|
|
||||||
"max_position_embeddings": 512,
|
|
||||||
"model_type": "bert",
|
|
||||||
"num_attention_heads": 16,
|
|
||||||
"num_hidden_layers": 24,
|
|
||||||
"output_past": true,
|
|
||||||
"pad_token_id": 0,
|
|
||||||
"pooler_fc_size": 768,
|
|
||||||
"pooler_num_attention_heads": 12,
|
|
||||||
"pooler_num_fc_layers": 3,
|
|
||||||
"pooler_size_per_head": 128,
|
|
||||||
"pooler_type": "first_token_transform",
|
|
||||||
"position_embedding_type": "absolute",
|
|
||||||
"torch_dtype": "float32",
|
|
||||||
"transformers_version": "4.22.1",
|
|
||||||
"type_vocab_size": 2,
|
|
||||||
"use_cache": true,
|
|
||||||
"vocab_size": 47020
|
|
||||||
}
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
{
|
|
||||||
"cls_token": "[CLS]",
|
|
||||||
"mask_token": "[MASK]",
|
|
||||||
"pad_token": "[PAD]",
|
|
||||||
"sep_token": "[SEP]",
|
|
||||||
"unk_token": "[UNK]"
|
|
||||||
}
|
|
||||||
@@ -1,16 +0,0 @@
|
|||||||
{
|
|
||||||
"cls_token": "[CLS]",
|
|
||||||
"do_basic_tokenize": true,
|
|
||||||
"do_lower_case": true,
|
|
||||||
"mask_token": "[MASK]",
|
|
||||||
"name_or_path": "hfl/chinese-roberta-wwm-ext",
|
|
||||||
"never_split": null,
|
|
||||||
"pad_token": "[PAD]",
|
|
||||||
"sep_token": "[SEP]",
|
|
||||||
"special_tokens_map_file": "/home/chenweifeng/.cache/huggingface/hub/models--hfl--chinese-roberta-wwm-ext/snapshots/5c58d0b8ec1d9014354d691c538661bf00bfdb44/special_tokens_map.json",
|
|
||||||
"strip_accents": null,
|
|
||||||
"tokenize_chinese_chars": true,
|
|
||||||
"tokenizer_class": "BertTokenizer",
|
|
||||||
"unk_token": "[UNK]",
|
|
||||||
"model_max_length": 77
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,334 +0,0 @@
|
|||||||
#credit to ExponentialML for this module
|
|
||||||
#from https://github.com/ExponentialML/ComfyUI_Native_DynamiCrafter
|
|
||||||
import os
|
|
||||||
import torch
|
|
||||||
import comfy
|
|
||||||
|
|
||||||
from einops import rearrange
|
|
||||||
from comfy import model_base, model_management
|
|
||||||
from .lvdm.modules.networks.openaimodel3d import UNetModel as DynamiCrafterUNetModel
|
|
||||||
|
|
||||||
from .utils.model_utils import DynamiCrafterBase, DYNAMICRAFTER_CONFIG, load_image_proj_dict, load_dynamicrafter_dict, get_image_proj_model
|
|
||||||
|
|
||||||
class DynamiCrafter:
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
self.model_patcher = None
|
|
||||||
|
|
||||||
# There is probably a better way to do this, but with the apply_model callback, this seems necessary.
|
|
||||||
# The model gets wrapped around a CFG Denoiser class, and handles the conditioning parts there.
|
|
||||||
# We cannot access it, so we must find the conditioning according to how ComfyUI handles it.
|
|
||||||
def get_conditioning_pair(self, c_crossattn, use_cfg: bool):
|
|
||||||
if not use_cfg:
|
|
||||||
return c_crossattn
|
|
||||||
|
|
||||||
conditioning_group = []
|
|
||||||
|
|
||||||
for i in range(c_crossattn.shape[0]):
|
|
||||||
# Get the positive and negative conditioning.
|
|
||||||
positive_idx = i + 1
|
|
||||||
negative_idx = i
|
|
||||||
|
|
||||||
if positive_idx >= c_crossattn.shape[0]:
|
|
||||||
break
|
|
||||||
|
|
||||||
if not torch.equal(c_crossattn[[positive_idx]], c_crossattn[[negative_idx]]):
|
|
||||||
conditioning_group = [
|
|
||||||
c_crossattn[[positive_idx]],
|
|
||||||
c_crossattn[[negative_idx]]
|
|
||||||
]
|
|
||||||
break
|
|
||||||
|
|
||||||
if len(conditioning_group) == 0:
|
|
||||||
raise ValueError("Could not get the appropriate conditioning group.")
|
|
||||||
|
|
||||||
return torch.cat(conditioning_group)
|
|
||||||
|
|
||||||
# apply_model, {"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond}
|
|
||||||
def _forward(self, *args):
|
|
||||||
transformer_options = self.model_patcher.model_options['transformer_options']
|
|
||||||
conditioning = transformer_options['conditioning']
|
|
||||||
|
|
||||||
apply_model = args[0]
|
|
||||||
|
|
||||||
# forward_dict
|
|
||||||
fd = args[1]
|
|
||||||
|
|
||||||
x, t, model_in_kwargs, _ = fd['input'], fd['timestep'], fd['c'], fd['cond_or_uncond']
|
|
||||||
|
|
||||||
c_crossattn = model_in_kwargs.pop("c_crossattn")
|
|
||||||
c_concat = conditioning['c_concat']
|
|
||||||
num_video_frames = conditioning['num_video_frames']
|
|
||||||
fs = conditioning['fs']
|
|
||||||
|
|
||||||
original_num_frames = num_video_frames
|
|
||||||
|
|
||||||
# Better way to determine if we're using CFG
|
|
||||||
# The cond batch will always be num_frames >= 2 since we're doing video,
|
|
||||||
# so we need get this condition differently here.
|
|
||||||
if x.shape[0] > num_video_frames:
|
|
||||||
num_video_frames *= 2
|
|
||||||
batch_size = 2
|
|
||||||
use_cfg = True
|
|
||||||
else:
|
|
||||||
use_cfg = False
|
|
||||||
batch_size = 1
|
|
||||||
|
|
||||||
if use_cfg:
|
|
||||||
c_concat = torch.cat([c_concat] * 2)
|
|
||||||
|
|
||||||
self.validate_forwardable_latent(x, c_concat, num_video_frames, use_cfg)
|
|
||||||
|
|
||||||
x_in, c_concat = map(lambda xc: rearrange(xc, '(b t) c h w -> b c t h w', b=batch_size), (x, c_concat))
|
|
||||||
|
|
||||||
# We always assume video, so there will always be batched conditionings.
|
|
||||||
c_crossattn = self.get_conditioning_pair(c_crossattn, use_cfg)
|
|
||||||
c_crossattn = c_crossattn[:2] if use_cfg else c_crossattn[:1]
|
|
||||||
context_in = c_crossattn
|
|
||||||
|
|
||||||
img_embs = conditioning['image_emb']
|
|
||||||
|
|
||||||
if use_cfg:
|
|
||||||
img_emb_uncond = conditioning['image_emb_uncond']
|
|
||||||
img_embs = torch.cat([img_embs, img_emb_uncond])
|
|
||||||
|
|
||||||
fs = torch.cat([fs] * x_in.shape[0])
|
|
||||||
|
|
||||||
outs = []
|
|
||||||
for i in range(batch_size):
|
|
||||||
model_in_kwargs['transformer_options']['cond_idx'] = i
|
|
||||||
x_out = apply_model(
|
|
||||||
x_in[[i]],
|
|
||||||
t=torch.cat([t[:1]]),
|
|
||||||
context_in=context_in[[i]],
|
|
||||||
c_crossattn=c_crossattn,
|
|
||||||
cc_concat=c_concat[[i]], # "cc" is to handle naming conflict with apply_model wrapper.
|
|
||||||
# We want to handle this in the UNet forward.
|
|
||||||
num_video_frames=num_video_frames // 2 if batch_size > 1 else num_video_frames,
|
|
||||||
img_emb=img_embs[[i]],
|
|
||||||
fs=fs[[i]],
|
|
||||||
**model_in_kwargs
|
|
||||||
)
|
|
||||||
outs.append(x_out)
|
|
||||||
|
|
||||||
x_out = torch.cat(list(reversed(outs)))
|
|
||||||
x_out = rearrange(x_out, 'b c t h w -> (b t) c h w')
|
|
||||||
|
|
||||||
return x_out
|
|
||||||
|
|
||||||
def assign_forward_args(
|
|
||||||
self,
|
|
||||||
model,
|
|
||||||
c_concat,
|
|
||||||
image_emb,
|
|
||||||
image_emb_uncond,
|
|
||||||
fs,
|
|
||||||
frames,
|
|
||||||
):
|
|
||||||
model.model_options['transformer_options']['conditioning'] = {
|
|
||||||
"c_concat": c_concat,
|
|
||||||
"image_emb": image_emb,
|
|
||||||
'image_emb_uncond': image_emb_uncond,
|
|
||||||
"fs": fs,
|
|
||||||
"num_video_frames": frames,
|
|
||||||
}
|
|
||||||
|
|
||||||
def validate_forwardable_latent(self, latent, c_concat, num_video_frames, use_cfg):
|
|
||||||
check_no_cfg = latent.shape[0] != num_video_frames
|
|
||||||
check_with_cfg = latent.shape[0] != (num_video_frames * 2)
|
|
||||||
|
|
||||||
latent_batch_size = latent.shape[0] if not use_cfg else latent.shape[0] // 2
|
|
||||||
num_frames = num_video_frames if not use_cfg else num_video_frames // 2
|
|
||||||
|
|
||||||
if all([check_no_cfg, check_with_cfg]):
|
|
||||||
raise ValueError(
|
|
||||||
"Please make sure your latent inputs match the number of frames in the DynamiCrafter Processor."
|
|
||||||
f"Got a latent batch size of ({latent_batch_size}) with number of frames being ({num_frames})."
|
|
||||||
)
|
|
||||||
|
|
||||||
latent_h, latent_w = latent.shape[-2:]
|
|
||||||
c_concat_h, c_concat_w = c_concat.shape[-2:]
|
|
||||||
|
|
||||||
if not all([latent_h == c_concat_h, latent_w == c_concat_w]):
|
|
||||||
raise ValueError(
|
|
||||||
"Please make sure that your input latent and image frames are the same height and width.",
|
|
||||||
f"Image Size: {c_concat_w * 8}, {c_concat_h * 8}, Latent Size: {latent_h * 8}, {latent_w * 8}"
|
|
||||||
)
|
|
||||||
|
|
||||||
def process_image_conditioning(
|
|
||||||
self,
|
|
||||||
model,
|
|
||||||
clip_vision,
|
|
||||||
vae,
|
|
||||||
image_proj_model,
|
|
||||||
images,
|
|
||||||
use_interpolate,
|
|
||||||
fps: int,
|
|
||||||
frames: int,
|
|
||||||
scale_latents: bool
|
|
||||||
):
|
|
||||||
self.model_patcher = model
|
|
||||||
encoded_latent = vae.encode(images[:, :, :, :3])
|
|
||||||
|
|
||||||
encoded_image = clip_vision.encode_image(images[:1])['last_hidden_state']
|
|
||||||
image_emb = image_proj_model(encoded_image)
|
|
||||||
|
|
||||||
encoded_image_uncond = clip_vision.encode_image(torch.zeros_like(images)[:1])['last_hidden_state']
|
|
||||||
image_emb_uncond = image_proj_model(encoded_image_uncond)
|
|
||||||
|
|
||||||
c_concat = encoded_latent
|
|
||||||
|
|
||||||
if scale_latents:
|
|
||||||
vae_process_input = vae.process_input
|
|
||||||
vae.process_input = lambda image: (image - .5) * 2
|
|
||||||
c_concat = vae.encode(images[:, :, :, :3])
|
|
||||||
vae.process_input = vae_process_input
|
|
||||||
c_concat = model.model.process_latent_in(c_concat) * 1.3
|
|
||||||
else:
|
|
||||||
c_concat = model.model.process_latent_in(c_concat)
|
|
||||||
|
|
||||||
fs = torch.tensor([fps], dtype=torch.long, device=model_management.intermediate_device())
|
|
||||||
|
|
||||||
model.set_model_unet_function_wrapper(self._forward)
|
|
||||||
|
|
||||||
used_interpolate_processing = False
|
|
||||||
|
|
||||||
if use_interpolate and frames > 16:
|
|
||||||
raise ValueError(
|
|
||||||
"When using interpolation mode, the maximum amount of frames are 16."
|
|
||||||
"If you're doing long video generation, consider using the last frame\
|
|
||||||
from the first generation for the next one (autoregressive)."
|
|
||||||
)
|
|
||||||
if encoded_latent.shape[0] == 1:
|
|
||||||
c_concat = torch.cat([c_concat] * frames, dim=0)[:frames]
|
|
||||||
|
|
||||||
if use_interpolate:
|
|
||||||
mask = torch.zeros_like(c_concat)
|
|
||||||
mask[:1] = c_concat[:1]
|
|
||||||
c_concat = mask
|
|
||||||
|
|
||||||
used_interpolate_processing = True
|
|
||||||
else:
|
|
||||||
if use_interpolate and c_concat.shape[0] in [2, 3]:
|
|
||||||
input_frame_count = c_concat.shape[0]
|
|
||||||
|
|
||||||
# We're just padding to the same type an size of the concat
|
|
||||||
masked_frames = torch.zeros_like(torch.cat([c_concat[:1]] * frames))[:frames]
|
|
||||||
|
|
||||||
# Start frame
|
|
||||||
masked_frames[:1] = c_concat[:1]
|
|
||||||
|
|
||||||
end_frame_idx = -1
|
|
||||||
|
|
||||||
# TODO
|
|
||||||
speed = 1.0
|
|
||||||
if speed < 1.0:
|
|
||||||
possible_speeds = list(torch.linspace(0, 1.0, c_concat.shape[0]))
|
|
||||||
speed_from_frames = enumerate(possible_speeds)
|
|
||||||
speed_idx = min(speed_from_frames, key=lambda n: n[1] - speed)[0]
|
|
||||||
end_frame_idx = speed_idx
|
|
||||||
|
|
||||||
# End frame
|
|
||||||
masked_frames[-1:] = c_concat[[end_frame_idx]]
|
|
||||||
|
|
||||||
# Possible middle frame, but not working at the moment.
|
|
||||||
if input_frame_count == 3:
|
|
||||||
middle_idx = masked_frames.shape[0] // 2
|
|
||||||
middle_idx_frame = c_concat.shape[0] // 2
|
|
||||||
masked_frames[[middle_idx]] = c_concat[[middle_idx_frame]]
|
|
||||||
|
|
||||||
c_concat = masked_frames
|
|
||||||
used_interpolate_processing = True
|
|
||||||
|
|
||||||
print(f"Using interpolation mode with {input_frame_count} frames.")
|
|
||||||
|
|
||||||
if c_concat.shape[0] < frames and not used_interpolate_processing:
|
|
||||||
print(
|
|
||||||
"Multiple images found, but interpolation mode is unset. Using the first frame as condition.",
|
|
||||||
)
|
|
||||||
c_concat = torch.cat([c_concat[:1]] * frames)
|
|
||||||
|
|
||||||
c_concat = c_concat[:frames]
|
|
||||||
|
|
||||||
if encoded_latent.shape[0] == 1:
|
|
||||||
encoded_latent = torch.cat([encoded_latent] * frames)[:frames]
|
|
||||||
|
|
||||||
if encoded_latent.shape[0] < frames and encoded_latent.shape[0] != 1:
|
|
||||||
encoded_latent = torch.cat(
|
|
||||||
[encoded_latent] + [encoded_latent[-1:]] * abs(encoded_latent.shape[0] - frames)
|
|
||||||
)[:frames]
|
|
||||||
|
|
||||||
# We could store this as a state in this Node Class Instance, but to prevent any weird edge cases,
|
|
||||||
# this should always be passed through the 'stateless' way, and let ComfyUI handle the transformer_options state.
|
|
||||||
self.assign_forward_args(model, c_concat, image_emb, image_emb_uncond, fs, frames)
|
|
||||||
|
|
||||||
return (model, {"samples": torch.zeros_like(c_concat)}, {"samples": encoded_latent},)
|
|
||||||
|
|
||||||
|
|
||||||
# Loader for the DynamiCrafter model.
|
|
||||||
def load_model_sicts(self, model_path: str):
|
|
||||||
model_state_dict = comfy.utils.load_torch_file(model_path)
|
|
||||||
dynamicrafter_dict = load_dynamicrafter_dict(model_state_dict)
|
|
||||||
image_proj_dict = load_image_proj_dict(model_state_dict)
|
|
||||||
|
|
||||||
return dynamicrafter_dict, image_proj_dict
|
|
||||||
|
|
||||||
def get_prediction_type(self, is_eps: bool, model_config):
|
|
||||||
if not is_eps and "image_cross_attention_scale_learnable" in model_config.unet_config.keys():
|
|
||||||
model_config.unet_config["image_cross_attention_scale_learnable"] = False
|
|
||||||
|
|
||||||
return model_base.ModelType.EPS if is_eps else model_base.ModelType.V_PREDICTION
|
|
||||||
|
|
||||||
def handle_model_management(self, dynamicrafter_dict: dict, model_config):
|
|
||||||
parameters = comfy.utils.calculate_parameters(dynamicrafter_dict, "model.diffusion_model.")
|
|
||||||
load_device = model_management.get_torch_device()
|
|
||||||
unet_dtype = model_management.unet_dtype(
|
|
||||||
model_params=parameters,
|
|
||||||
supported_dtypes=model_config.supported_inference_dtypes
|
|
||||||
)
|
|
||||||
manual_cast_dtype = model_management.unet_manual_cast(
|
|
||||||
unet_dtype,
|
|
||||||
load_device,
|
|
||||||
model_config.supported_inference_dtypes
|
|
||||||
)
|
|
||||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype)
|
|
||||||
inital_load_device = model_management.unet_inital_load_device(parameters, unet_dtype)
|
|
||||||
offload_device = model_management.unet_offload_device()
|
|
||||||
|
|
||||||
return load_device, inital_load_device
|
|
||||||
|
|
||||||
def check_leftover_keys(self, state_dict: dict):
|
|
||||||
left_over = state_dict.keys()
|
|
||||||
if len(left_over) > 0:
|
|
||||||
print("left over keys:", left_over)
|
|
||||||
|
|
||||||
def load_dynamicrafter(self, model_path):
|
|
||||||
|
|
||||||
if os.path.exists(model_path):
|
|
||||||
dynamicrafter_dict, image_proj_dict = self.load_model_sicts(model_path)
|
|
||||||
model_config = DynamiCrafterBase(DYNAMICRAFTER_CONFIG)
|
|
||||||
|
|
||||||
dynamicrafter_dict, is_eps = model_config.process_dict_version(state_dict=dynamicrafter_dict)
|
|
||||||
|
|
||||||
MODEL_TYPE = self.get_prediction_type(is_eps, model_config)
|
|
||||||
load_device, inital_load_device = self.handle_model_management(dynamicrafter_dict, model_config)
|
|
||||||
|
|
||||||
model = model_base.BaseModel(
|
|
||||||
model_config,
|
|
||||||
model_type=MODEL_TYPE,
|
|
||||||
device=inital_load_device,
|
|
||||||
unet_model=DynamiCrafterUNetModel
|
|
||||||
)
|
|
||||||
|
|
||||||
image_proj_model = get_image_proj_model(image_proj_dict)
|
|
||||||
model.load_model_weights(dynamicrafter_dict, "model.diffusion_model.")
|
|
||||||
self.check_leftover_keys(dynamicrafter_dict)
|
|
||||||
|
|
||||||
model_patcher = comfy.model_patcher.ModelPatcher(
|
|
||||||
model,
|
|
||||||
load_device=load_device,
|
|
||||||
offload_device=model_management.unet_offload_device(),
|
|
||||||
current_device=inital_load_device
|
|
||||||
)
|
|
||||||
|
|
||||||
return (model_patcher, image_proj_model,)
|
|
||||||
@@ -1,102 +0,0 @@
|
|||||||
# adopted from
|
|
||||||
# https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
|
|
||||||
# and
|
|
||||||
# https://github.com/lucidrains/denoising-diffusion-pytorch/blob/7706bdfc6f527f58d33f84b7b522e61e6e3164b3/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py
|
|
||||||
# and
|
|
||||||
# https://github.com/openai/guided-diffusion/blob/0ba878e517b276c45d1195eb29f6f5f72659a05b/guided_diffusion/nn.py
|
|
||||||
#
|
|
||||||
# thanks!
|
|
||||||
|
|
||||||
import torch.nn as nn
|
|
||||||
import comfy.ops
|
|
||||||
ops = comfy.ops.disable_weight_init
|
|
||||||
|
|
||||||
from ..utils.utils import instantiate_from_config
|
|
||||||
|
|
||||||
def disabled_train(self, mode=True):
|
|
||||||
"""Overwrite model.train with this function to make sure train/eval mode
|
|
||||||
does not change anymore."""
|
|
||||||
return self
|
|
||||||
|
|
||||||
def zero_module(module):
|
|
||||||
"""
|
|
||||||
Zero out the parameters of a module and return it.
|
|
||||||
"""
|
|
||||||
for p in module.parameters():
|
|
||||||
p.detach().zero_()
|
|
||||||
return module
|
|
||||||
|
|
||||||
def scale_module(module, scale):
|
|
||||||
"""
|
|
||||||
Scale the parameters of a module and return it.
|
|
||||||
"""
|
|
||||||
for p in module.parameters():
|
|
||||||
p.detach().mul_(scale)
|
|
||||||
return module
|
|
||||||
|
|
||||||
|
|
||||||
def conv_nd(dims, *args, **kwargs):
|
|
||||||
"""
|
|
||||||
Create a 1D, 2D, or 3D convolution module.
|
|
||||||
"""
|
|
||||||
if dims == 1:
|
|
||||||
return nn.Conv1d(*args, **kwargs)
|
|
||||||
elif dims == 2:
|
|
||||||
return ops.Conv2d(*args, **kwargs)
|
|
||||||
elif dims == 3:
|
|
||||||
return ops.Conv3d(*args, **kwargs)
|
|
||||||
raise ValueError(f"unsupported dimensions: {dims}")
|
|
||||||
|
|
||||||
|
|
||||||
def linear(*args, **kwargs):
|
|
||||||
"""
|
|
||||||
Create a linear module.
|
|
||||||
"""
|
|
||||||
return ops.Linear(*args, **kwargs)
|
|
||||||
|
|
||||||
|
|
||||||
def avg_pool_nd(dims, *args, **kwargs):
|
|
||||||
"""
|
|
||||||
Create a 1D, 2D, or 3D average pooling module.
|
|
||||||
"""
|
|
||||||
if dims == 1:
|
|
||||||
return nn.AvgPool1d(*args, **kwargs)
|
|
||||||
elif dims == 2:
|
|
||||||
return nn.AvgPool2d(*args, **kwargs)
|
|
||||||
elif dims == 3:
|
|
||||||
return nn.AvgPool3d(*args, **kwargs)
|
|
||||||
raise ValueError(f"unsupported dimensions: {dims}")
|
|
||||||
|
|
||||||
|
|
||||||
def nonlinearity(type='silu'):
|
|
||||||
if type == 'silu':
|
|
||||||
return nn.SiLU()
|
|
||||||
elif type == 'leaky_relu':
|
|
||||||
return nn.LeakyReLU()
|
|
||||||
|
|
||||||
|
|
||||||
class GroupNormSpecific(ops.GroupNorm):
|
|
||||||
def forward(self, x):
|
|
||||||
return super().forward(x.float()).type(x.dtype)
|
|
||||||
|
|
||||||
|
|
||||||
def normalization(channels, num_groups=32, dtype=None, device=None):
|
|
||||||
"""
|
|
||||||
Make a standard normalization layer.
|
|
||||||
:param channels: number of input channels.
|
|
||||||
:return: an nn.Module for normalization.
|
|
||||||
"""
|
|
||||||
return GroupNormSpecific(num_groups, channels, dtype=dtype, device=device)
|
|
||||||
|
|
||||||
|
|
||||||
class HybridConditioner(nn.Module):
|
|
||||||
|
|
||||||
def __init__(self, c_concat_config, c_crossattn_config):
|
|
||||||
super().__init__()
|
|
||||||
self.concat_conditioner = instantiate_from_config(c_concat_config)
|
|
||||||
self.crossattn_conditioner = instantiate_from_config(c_crossattn_config)
|
|
||||||
|
|
||||||
def forward(self, c_concat, c_crossattn):
|
|
||||||
c_concat = self.concat_conditioner(c_concat)
|
|
||||||
c_crossattn = self.crossattn_conditioner(c_crossattn)
|
|
||||||
return {'c_concat': [c_concat], 'c_crossattn': [c_crossattn]}
|
|
||||||
@@ -1,94 +0,0 @@
|
|||||||
import math
|
|
||||||
from inspect import isfunction
|
|
||||||
import torch
|
|
||||||
from torch import nn
|
|
||||||
import torch.distributed as dist
|
|
||||||
|
|
||||||
|
|
||||||
def gather_data(data, return_np=True):
|
|
||||||
''' gather data from multiple processes to one list '''
|
|
||||||
data_list = [torch.zeros_like(data) for _ in range(dist.get_world_size())]
|
|
||||||
dist.all_gather(data_list, data) # gather not supported with NCCL
|
|
||||||
if return_np:
|
|
||||||
data_list = [data.cpu().numpy() for data in data_list]
|
|
||||||
return data_list
|
|
||||||
|
|
||||||
def autocast(f):
|
|
||||||
def do_autocast(*args, **kwargs):
|
|
||||||
with torch.cuda.amp.autocast(enabled=True,
|
|
||||||
dtype=torch.get_autocast_gpu_dtype(),
|
|
||||||
cache_enabled=torch.is_autocast_cache_enabled()):
|
|
||||||
return f(*args, **kwargs)
|
|
||||||
return do_autocast
|
|
||||||
|
|
||||||
|
|
||||||
def extract_into_tensor(a, t, x_shape):
|
|
||||||
b, *_ = t.shape
|
|
||||||
out = a.gather(-1, t)
|
|
||||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
|
||||||
|
|
||||||
|
|
||||||
def noise_like(shape, device, repeat=False):
|
|
||||||
repeat_noise = lambda: torch.randn((1, *shape[1:]), device=device).repeat(shape[0], *((1,) * (len(shape) - 1)))
|
|
||||||
noise = lambda: torch.randn(shape, device=device)
|
|
||||||
return repeat_noise() if repeat else noise()
|
|
||||||
|
|
||||||
|
|
||||||
def default(val, d):
|
|
||||||
if exists(val):
|
|
||||||
return val
|
|
||||||
return d() if isfunction(d) else d
|
|
||||||
|
|
||||||
def exists(val):
|
|
||||||
return val is not None
|
|
||||||
|
|
||||||
def identity(*args, **kwargs):
|
|
||||||
return nn.Identity()
|
|
||||||
|
|
||||||
def uniq(arr):
|
|
||||||
return{el: True for el in arr}.keys()
|
|
||||||
|
|
||||||
def mean_flat(tensor):
|
|
||||||
"""
|
|
||||||
Take the mean over all non-batch dimensions.
|
|
||||||
"""
|
|
||||||
return tensor.mean(dim=list(range(1, len(tensor.shape))))
|
|
||||||
|
|
||||||
def ismap(x):
|
|
||||||
if not isinstance(x, torch.Tensor):
|
|
||||||
return False
|
|
||||||
return (len(x.shape) == 4) and (x.shape[1] > 3)
|
|
||||||
|
|
||||||
def isimage(x):
|
|
||||||
if not isinstance(x,torch.Tensor):
|
|
||||||
return False
|
|
||||||
return (len(x.shape) == 4) and (x.shape[1] == 3 or x.shape[1] == 1)
|
|
||||||
|
|
||||||
def max_neg_value(t):
|
|
||||||
return -torch.finfo(t.dtype).max
|
|
||||||
|
|
||||||
def shape_to_str(x):
|
|
||||||
shape_str = "x".join([str(x) for x in x.shape])
|
|
||||||
return shape_str
|
|
||||||
|
|
||||||
def init_(tensor):
|
|
||||||
dim = tensor.shape[-1]
|
|
||||||
std = 1 / math.sqrt(dim)
|
|
||||||
tensor.uniform_(-std, std)
|
|
||||||
return tensor
|
|
||||||
|
|
||||||
ckpt = torch.utils.checkpoint.checkpoint
|
|
||||||
def checkpoint(func, inputs, params, flag):
|
|
||||||
"""
|
|
||||||
Evaluate a function without caching intermediate activations, allowing for
|
|
||||||
reduced memory at the expense of extra compute in the backward pass.
|
|
||||||
:param func: the function to evaluate.
|
|
||||||
:param inputs: the argument sequence to pass to `func`.
|
|
||||||
:param params: a sequence of parameters `func` depends on but does not
|
|
||||||
explicitly take as arguments.
|
|
||||||
:param flag: if False, disable gradient checkpointing.
|
|
||||||
"""
|
|
||||||
if flag:
|
|
||||||
return ckpt(func, *inputs, use_reentrant=False)
|
|
||||||
else:
|
|
||||||
return func(*inputs)
|
|
||||||
@@ -1,95 +0,0 @@
|
|||||||
import torch
|
|
||||||
import numpy as np
|
|
||||||
|
|
||||||
|
|
||||||
class AbstractDistribution:
|
|
||||||
def sample(self):
|
|
||||||
raise NotImplementedError()
|
|
||||||
|
|
||||||
def mode(self):
|
|
||||||
raise NotImplementedError()
|
|
||||||
|
|
||||||
|
|
||||||
class DiracDistribution(AbstractDistribution):
|
|
||||||
def __init__(self, value):
|
|
||||||
self.value = value
|
|
||||||
|
|
||||||
def sample(self):
|
|
||||||
return self.value
|
|
||||||
|
|
||||||
def mode(self):
|
|
||||||
return self.value
|
|
||||||
|
|
||||||
|
|
||||||
class DiagonalGaussianDistribution(object):
|
|
||||||
def __init__(self, parameters, deterministic=False):
|
|
||||||
self.parameters = parameters
|
|
||||||
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
|
|
||||||
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
|
||||||
self.deterministic = deterministic
|
|
||||||
self.std = torch.exp(0.5 * self.logvar)
|
|
||||||
self.var = torch.exp(self.logvar)
|
|
||||||
if self.deterministic:
|
|
||||||
self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
|
|
||||||
|
|
||||||
def sample(self, noise=None):
|
|
||||||
if noise is None:
|
|
||||||
noise = torch.randn(self.mean.shape)
|
|
||||||
|
|
||||||
x = self.mean + self.std * noise.to(device=self.parameters.device)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def kl(self, other=None):
|
|
||||||
if self.deterministic:
|
|
||||||
return torch.Tensor([0.])
|
|
||||||
else:
|
|
||||||
if other is None:
|
|
||||||
return 0.5 * torch.sum(torch.pow(self.mean, 2)
|
|
||||||
+ self.var - 1.0 - self.logvar,
|
|
||||||
dim=[1, 2, 3])
|
|
||||||
else:
|
|
||||||
return 0.5 * torch.sum(
|
|
||||||
torch.pow(self.mean - other.mean, 2) / other.var
|
|
||||||
+ self.var / other.var - 1.0 - self.logvar + other.logvar,
|
|
||||||
dim=[1, 2, 3])
|
|
||||||
|
|
||||||
def nll(self, sample, dims=[1,2,3]):
|
|
||||||
if self.deterministic:
|
|
||||||
return torch.Tensor([0.])
|
|
||||||
logtwopi = np.log(2.0 * np.pi)
|
|
||||||
return 0.5 * torch.sum(
|
|
||||||
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
|
||||||
dim=dims)
|
|
||||||
|
|
||||||
def mode(self):
|
|
||||||
return self.mean
|
|
||||||
|
|
||||||
|
|
||||||
def normal_kl(mean1, logvar1, mean2, logvar2):
|
|
||||||
"""
|
|
||||||
source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12
|
|
||||||
Compute the KL divergence between two gaussians.
|
|
||||||
Shapes are automatically broadcasted, so batches can be compared to
|
|
||||||
scalars, among other use cases.
|
|
||||||
"""
|
|
||||||
tensor = None
|
|
||||||
for obj in (mean1, logvar1, mean2, logvar2):
|
|
||||||
if isinstance(obj, torch.Tensor):
|
|
||||||
tensor = obj
|
|
||||||
break
|
|
||||||
assert tensor is not None, "at least one argument must be a Tensor"
|
|
||||||
|
|
||||||
# Force variances to be Tensors. Broadcasting helps convert scalars to
|
|
||||||
# Tensors, but it does not work for torch.exp().
|
|
||||||
logvar1, logvar2 = [
|
|
||||||
x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor)
|
|
||||||
for x in (logvar1, logvar2)
|
|
||||||
]
|
|
||||||
|
|
||||||
return 0.5 * (
|
|
||||||
-1.0
|
|
||||||
+ logvar2
|
|
||||||
- logvar1
|
|
||||||
+ torch.exp(logvar1 - logvar2)
|
|
||||||
+ ((mean1 - mean2) ** 2) * torch.exp(-logvar2)
|
|
||||||
)
|
|
||||||
@@ -1,76 +0,0 @@
|
|||||||
import torch
|
|
||||||
from torch import nn
|
|
||||||
|
|
||||||
|
|
||||||
class LitEma(nn.Module):
|
|
||||||
def __init__(self, model, decay=0.9999, use_num_upates=True):
|
|
||||||
super().__init__()
|
|
||||||
if decay < 0.0 or decay > 1.0:
|
|
||||||
raise ValueError('Decay must be between 0 and 1')
|
|
||||||
|
|
||||||
self.m_name2s_name = {}
|
|
||||||
self.register_buffer('decay', torch.tensor(decay, dtype=torch.float32))
|
|
||||||
self.register_buffer('num_updates', torch.tensor(0,dtype=torch.int) if use_num_upates
|
|
||||||
else torch.tensor(-1,dtype=torch.int))
|
|
||||||
|
|
||||||
for name, p in model.named_parameters():
|
|
||||||
if p.requires_grad:
|
|
||||||
#remove as '.'-character is not allowed in buffers
|
|
||||||
s_name = name.replace('.','')
|
|
||||||
self.m_name2s_name.update({name:s_name})
|
|
||||||
self.register_buffer(s_name,p.clone().detach().data)
|
|
||||||
|
|
||||||
self.collected_params = []
|
|
||||||
|
|
||||||
def forward(self,model):
|
|
||||||
decay = self.decay
|
|
||||||
|
|
||||||
if self.num_updates >= 0:
|
|
||||||
self.num_updates += 1
|
|
||||||
decay = min(self.decay,(1 + self.num_updates) / (10 + self.num_updates))
|
|
||||||
|
|
||||||
one_minus_decay = 1.0 - decay
|
|
||||||
|
|
||||||
with torch.no_grad():
|
|
||||||
m_param = dict(model.named_parameters())
|
|
||||||
shadow_params = dict(self.named_buffers())
|
|
||||||
|
|
||||||
for key in m_param:
|
|
||||||
if m_param[key].requires_grad:
|
|
||||||
sname = self.m_name2s_name[key]
|
|
||||||
shadow_params[sname] = shadow_params[sname].type_as(m_param[key])
|
|
||||||
shadow_params[sname].sub_(one_minus_decay * (shadow_params[sname] - m_param[key]))
|
|
||||||
else:
|
|
||||||
assert not key in self.m_name2s_name
|
|
||||||
|
|
||||||
def copy_to(self, model):
|
|
||||||
m_param = dict(model.named_parameters())
|
|
||||||
shadow_params = dict(self.named_buffers())
|
|
||||||
for key in m_param:
|
|
||||||
if m_param[key].requires_grad:
|
|
||||||
m_param[key].data.copy_(shadow_params[self.m_name2s_name[key]].data)
|
|
||||||
else:
|
|
||||||
assert not key in self.m_name2s_name
|
|
||||||
|
|
||||||
def store(self, parameters):
|
|
||||||
"""
|
|
||||||
Save the current parameters for restoring later.
|
|
||||||
Args:
|
|
||||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
|
||||||
temporarily stored.
|
|
||||||
"""
|
|
||||||
self.collected_params = [param.clone() for param in parameters]
|
|
||||||
|
|
||||||
def restore(self, parameters):
|
|
||||||
"""
|
|
||||||
Restore the parameters stored with the `store` method.
|
|
||||||
Useful to validate the model with EMA parameters without affecting the
|
|
||||||
original optimization process. Store the parameters before the
|
|
||||||
`copy_to` method. After validation (or model saving), use this to
|
|
||||||
restore the former parameters.
|
|
||||||
Args:
|
|
||||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
|
||||||
updated with the stored parameters.
|
|
||||||
"""
|
|
||||||
for c_param, param in zip(self.collected_params, parameters):
|
|
||||||
param.data.copy_(c_param.data)
|
|
||||||
@@ -1,219 +0,0 @@
|
|||||||
import os
|
|
||||||
from contextlib import contextmanager
|
|
||||||
import torch
|
|
||||||
import numpy as np
|
|
||||||
from einops import rearrange
|
|
||||||
import torch.nn.functional as F
|
|
||||||
import pytorch_lightning as pl
|
|
||||||
from ...modules.networks.ae_modules import Encoder, Decoder
|
|
||||||
from ...distributions import DiagonalGaussianDistribution
|
|
||||||
from utils.utils import instantiate_from_config
|
|
||||||
|
|
||||||
|
|
||||||
class AutoencoderKL(pl.LightningModule):
|
|
||||||
def __init__(self,
|
|
||||||
ddconfig,
|
|
||||||
lossconfig,
|
|
||||||
embed_dim,
|
|
||||||
ckpt_path=None,
|
|
||||||
ignore_keys=[],
|
|
||||||
image_key="image",
|
|
||||||
colorize_nlabels=None,
|
|
||||||
monitor=None,
|
|
||||||
test=False,
|
|
||||||
logdir=None,
|
|
||||||
input_dim=4,
|
|
||||||
test_args=None,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.image_key = image_key
|
|
||||||
self.encoder = Encoder(**ddconfig)
|
|
||||||
self.decoder = Decoder(**ddconfig)
|
|
||||||
self.loss = instantiate_from_config(lossconfig)
|
|
||||||
assert ddconfig["double_z"]
|
|
||||||
self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
|
|
||||||
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
|
|
||||||
self.embed_dim = embed_dim
|
|
||||||
self.input_dim = input_dim
|
|
||||||
self.test = test
|
|
||||||
self.test_args = test_args
|
|
||||||
self.logdir = logdir
|
|
||||||
if colorize_nlabels is not None:
|
|
||||||
assert type(colorize_nlabels)==int
|
|
||||||
self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))
|
|
||||||
if monitor is not None:
|
|
||||||
self.monitor = monitor
|
|
||||||
if ckpt_path is not None:
|
|
||||||
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
|
|
||||||
if self.test:
|
|
||||||
self.init_test()
|
|
||||||
|
|
||||||
def init_test(self,):
|
|
||||||
self.test = True
|
|
||||||
save_dir = os.path.join(self.logdir, "test")
|
|
||||||
if 'ckpt' in self.test_args:
|
|
||||||
ckpt_name = os.path.basename(self.test_args.ckpt).split('.ckpt')[0] + f'_epoch{self._cur_epoch}'
|
|
||||||
self.root = os.path.join(save_dir, ckpt_name)
|
|
||||||
else:
|
|
||||||
self.root = save_dir
|
|
||||||
if 'test_subdir' in self.test_args:
|
|
||||||
self.root = os.path.join(save_dir, self.test_args.test_subdir)
|
|
||||||
|
|
||||||
self.root_zs = os.path.join(self.root, "zs")
|
|
||||||
self.root_dec = os.path.join(self.root, "reconstructions")
|
|
||||||
self.root_inputs = os.path.join(self.root, "inputs")
|
|
||||||
os.makedirs(self.root, exist_ok=True)
|
|
||||||
|
|
||||||
if self.test_args.save_z:
|
|
||||||
os.makedirs(self.root_zs, exist_ok=True)
|
|
||||||
if self.test_args.save_reconstruction:
|
|
||||||
os.makedirs(self.root_dec, exist_ok=True)
|
|
||||||
if self.test_args.save_input:
|
|
||||||
os.makedirs(self.root_inputs, exist_ok=True)
|
|
||||||
assert(self.test_args is not None)
|
|
||||||
self.test_maximum = getattr(self.test_args, 'test_maximum', None)
|
|
||||||
self.count = 0
|
|
||||||
self.eval_metrics = {}
|
|
||||||
self.decodes = []
|
|
||||||
self.save_decode_samples = 2048
|
|
||||||
|
|
||||||
def init_from_ckpt(self, path, ignore_keys=list()):
|
|
||||||
sd = torch.load(path, map_location="cpu")
|
|
||||||
try:
|
|
||||||
self._cur_epoch = sd['epoch']
|
|
||||||
sd = sd["state_dict"]
|
|
||||||
except:
|
|
||||||
self._cur_epoch = 'null'
|
|
||||||
keys = list(sd.keys())
|
|
||||||
for k in keys:
|
|
||||||
for ik in ignore_keys:
|
|
||||||
if k.startswith(ik):
|
|
||||||
print("Deleting key {} from state_dict.".format(k))
|
|
||||||
del sd[k]
|
|
||||||
self.load_state_dict(sd, strict=False)
|
|
||||||
# self.load_state_dict(sd, strict=True)
|
|
||||||
print(f"Restored from {path}")
|
|
||||||
|
|
||||||
def encode(self, x, **kwargs):
|
|
||||||
|
|
||||||
h = self.encoder(x)
|
|
||||||
moments = self.quant_conv(h)
|
|
||||||
posterior = DiagonalGaussianDistribution(moments)
|
|
||||||
return posterior
|
|
||||||
|
|
||||||
def decode(self, z, **kwargs):
|
|
||||||
z = self.post_quant_conv(z)
|
|
||||||
dec = self.decoder(z)
|
|
||||||
return dec
|
|
||||||
|
|
||||||
def forward(self, input, sample_posterior=True):
|
|
||||||
posterior = self.encode(input)
|
|
||||||
if sample_posterior:
|
|
||||||
z = posterior.sample()
|
|
||||||
else:
|
|
||||||
z = posterior.mode()
|
|
||||||
dec = self.decode(z)
|
|
||||||
return dec, posterior
|
|
||||||
|
|
||||||
def get_input(self, batch, k):
|
|
||||||
x = batch[k]
|
|
||||||
if x.dim() == 5 and self.input_dim == 4:
|
|
||||||
b,c,t,h,w = x.shape
|
|
||||||
self.b = b
|
|
||||||
self.t = t
|
|
||||||
x = rearrange(x, 'b c t h w -> (b t) c h w')
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
def training_step(self, batch, batch_idx, optimizer_idx):
|
|
||||||
inputs = self.get_input(batch, self.image_key)
|
|
||||||
reconstructions, posterior = self(inputs)
|
|
||||||
|
|
||||||
if optimizer_idx == 0:
|
|
||||||
# train encoder+decoder+logvar
|
|
||||||
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
|
|
||||||
last_layer=self.get_last_layer(), split="train")
|
|
||||||
self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
|
|
||||||
self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False)
|
|
||||||
return aeloss
|
|
||||||
|
|
||||||
if optimizer_idx == 1:
|
|
||||||
# train the discriminator
|
|
||||||
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
|
|
||||||
last_layer=self.get_last_layer(), split="train")
|
|
||||||
|
|
||||||
self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
|
|
||||||
self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False)
|
|
||||||
return discloss
|
|
||||||
|
|
||||||
def validation_step(self, batch, batch_idx):
|
|
||||||
inputs = self.get_input(batch, self.image_key)
|
|
||||||
reconstructions, posterior = self(inputs)
|
|
||||||
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step,
|
|
||||||
last_layer=self.get_last_layer(), split="val")
|
|
||||||
|
|
||||||
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step,
|
|
||||||
last_layer=self.get_last_layer(), split="val")
|
|
||||||
|
|
||||||
self.log("val/rec_loss", log_dict_ae["val/rec_loss"])
|
|
||||||
self.log_dict(log_dict_ae)
|
|
||||||
self.log_dict(log_dict_disc)
|
|
||||||
return self.log_dict
|
|
||||||
|
|
||||||
def configure_optimizers(self):
|
|
||||||
lr = self.learning_rate
|
|
||||||
opt_ae = torch.optim.Adam(list(self.encoder.parameters())+
|
|
||||||
list(self.decoder.parameters())+
|
|
||||||
list(self.quant_conv.parameters())+
|
|
||||||
list(self.post_quant_conv.parameters()),
|
|
||||||
lr=lr, betas=(0.5, 0.9))
|
|
||||||
opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(),
|
|
||||||
lr=lr, betas=(0.5, 0.9))
|
|
||||||
return [opt_ae, opt_disc], []
|
|
||||||
|
|
||||||
def get_last_layer(self):
|
|
||||||
return self.decoder.conv_out.weight
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def log_images(self, batch, only_inputs=False, **kwargs):
|
|
||||||
log = dict()
|
|
||||||
x = self.get_input(batch, self.image_key)
|
|
||||||
x = x.to(self.device)
|
|
||||||
if not only_inputs:
|
|
||||||
xrec, posterior = self(x)
|
|
||||||
if x.shape[1] > 3:
|
|
||||||
# colorize with random projection
|
|
||||||
assert xrec.shape[1] > 3
|
|
||||||
x = self.to_rgb(x)
|
|
||||||
xrec = self.to_rgb(xrec)
|
|
||||||
log["samples"] = self.decode(torch.randn_like(posterior.sample()))
|
|
||||||
log["reconstructions"] = xrec
|
|
||||||
log["inputs"] = x
|
|
||||||
return log
|
|
||||||
|
|
||||||
def to_rgb(self, x):
|
|
||||||
assert self.image_key == "segmentation"
|
|
||||||
if not hasattr(self, "colorize"):
|
|
||||||
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
|
|
||||||
x = F.conv2d(x, weight=self.colorize)
|
|
||||||
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
|
|
||||||
return x
|
|
||||||
|
|
||||||
class IdentityFirstStage(torch.nn.Module):
|
|
||||||
def __init__(self, *args, vq_interface=False, **kwargs):
|
|
||||||
self.vq_interface = vq_interface # TODO: Should be true by default but check to not break older stuff
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
def encode(self, x, *args, **kwargs):
|
|
||||||
return x
|
|
||||||
|
|
||||||
def decode(self, x, *args, **kwargs):
|
|
||||||
return x
|
|
||||||
|
|
||||||
def quantize(self, x, *args, **kwargs):
|
|
||||||
if self.vq_interface:
|
|
||||||
return x, None, [None, None, None]
|
|
||||||
return x
|
|
||||||
|
|
||||||
def forward(self, x, *args, **kwargs):
|
|
||||||
return x
|
|
||||||
@@ -1,762 +0,0 @@
|
|||||||
"""
|
|
||||||
wild mixture of
|
|
||||||
https://github.com/openai/improved-diffusion/blob/e94489283bb876ac1477d5dd7709bbbd2d9902ce/improved_diffusion/gaussian_diffusion.py
|
|
||||||
https://github.com/lucidrains/denoising-diffusion-pytorch/blob/7706bdfc6f527f58d33f84b7b522e61e6e3164b3/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py
|
|
||||||
https://github.com/CompVis/taming-transformers
|
|
||||||
-- merci
|
|
||||||
"""
|
|
||||||
|
|
||||||
from functools import partial
|
|
||||||
from contextlib import contextmanager
|
|
||||||
import numpy as np
|
|
||||||
from tqdm import tqdm
|
|
||||||
from einops import rearrange, repeat
|
|
||||||
import logging
|
|
||||||
mainlogger = logging.getLogger('mainlogger')
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
from torchvision.utils import make_grid
|
|
||||||
|
|
||||||
from ...utils.utils import instantiate_from_config
|
|
||||||
from ..ema import LitEma
|
|
||||||
from ..distributions import DiagonalGaussianDistribution
|
|
||||||
from ..models.utils_diffusion import make_beta_schedule, rescale_zero_terminal_snr
|
|
||||||
from ..basics import disabled_train
|
|
||||||
from ..common import (
|
|
||||||
extract_into_tensor,
|
|
||||||
noise_like,
|
|
||||||
exists,
|
|
||||||
default
|
|
||||||
)
|
|
||||||
|
|
||||||
__conditioning_keys__ = {'concat': 'c_concat',
|
|
||||||
'crossattn': 'c_crossattn',
|
|
||||||
'adm': 'y'}
|
|
||||||
|
|
||||||
class DDPM(nn.Module):
|
|
||||||
# classic DDPM with Gaussian diffusion, in image space
|
|
||||||
def __init__(self,
|
|
||||||
unet_config,
|
|
||||||
timesteps=1000,
|
|
||||||
beta_schedule="linear",
|
|
||||||
loss_type="l2",
|
|
||||||
ckpt_path=None,
|
|
||||||
ignore_keys=[],
|
|
||||||
load_only_unet=False,
|
|
||||||
monitor=None,
|
|
||||||
use_ema=True,
|
|
||||||
first_stage_key="image",
|
|
||||||
image_size=256,
|
|
||||||
channels=3,
|
|
||||||
log_every_t=100,
|
|
||||||
clip_denoised=True,
|
|
||||||
linear_start=1e-4,
|
|
||||||
linear_end=2e-2,
|
|
||||||
cosine_s=8e-3,
|
|
||||||
given_betas=None,
|
|
||||||
original_elbo_weight=0.,
|
|
||||||
v_posterior=0., # weight for choosing posterior variance as sigma = (1-v) * beta_tilde + v * beta
|
|
||||||
l_simple_weight=1.,
|
|
||||||
conditioning_key=None,
|
|
||||||
parameterization="eps", # all assuming fixed variance schedules
|
|
||||||
scheduler_config=None,
|
|
||||||
use_positional_encodings=False,
|
|
||||||
learn_logvar=False,
|
|
||||||
logvar_init=0.,
|
|
||||||
rescale_betas_zero_snr=False,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
assert parameterization in ["eps", "x0", "v"], 'currently only supporting "eps" and "x0" and "v"'
|
|
||||||
self.parameterization = parameterization
|
|
||||||
mainlogger.info(f"{self.__class__.__name__}: Running in {self.parameterization}-prediction mode")
|
|
||||||
self.cond_stage_model = None
|
|
||||||
self.clip_denoised = clip_denoised
|
|
||||||
self.log_every_t = log_every_t
|
|
||||||
self.first_stage_key = first_stage_key
|
|
||||||
self.channels = channels
|
|
||||||
self.temporal_length = unet_config.params.temporal_length
|
|
||||||
self.image_size = image_size # try conv?
|
|
||||||
if isinstance(self.image_size, int):
|
|
||||||
self.image_size = [self.image_size, self.image_size]
|
|
||||||
self.use_positional_encodings = use_positional_encodings
|
|
||||||
self.model = DiffusionWrapper(unet_config, conditioning_key)
|
|
||||||
#count_params(self.model, verbose=True)
|
|
||||||
self.use_ema = use_ema
|
|
||||||
self.rescale_betas_zero_snr = rescale_betas_zero_snr
|
|
||||||
if self.use_ema:
|
|
||||||
self.model_ema = LitEma(self.model)
|
|
||||||
mainlogger.info(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.")
|
|
||||||
|
|
||||||
self.use_scheduler = scheduler_config is not None
|
|
||||||
if self.use_scheduler:
|
|
||||||
self.scheduler_config = scheduler_config
|
|
||||||
|
|
||||||
self.v_posterior = v_posterior
|
|
||||||
self.original_elbo_weight = original_elbo_weight
|
|
||||||
self.l_simple_weight = l_simple_weight
|
|
||||||
|
|
||||||
if monitor is not None:
|
|
||||||
self.monitor = monitor
|
|
||||||
if ckpt_path is not None:
|
|
||||||
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys, only_model=load_only_unet)
|
|
||||||
|
|
||||||
self.register_schedule(given_betas=given_betas, beta_schedule=beta_schedule, timesteps=timesteps,
|
|
||||||
linear_start=linear_start, linear_end=linear_end, cosine_s=cosine_s)
|
|
||||||
|
|
||||||
self.loss_type = loss_type
|
|
||||||
|
|
||||||
self.learn_logvar = learn_logvar
|
|
||||||
self.logvar = torch.full(fill_value=logvar_init, size=(self.num_timesteps,))
|
|
||||||
if self.learn_logvar:
|
|
||||||
self.logvar = nn.Parameter(self.logvar, requires_grad=True)
|
|
||||||
|
|
||||||
def register_schedule(self, given_betas=None, beta_schedule="linear", timesteps=1000,
|
|
||||||
linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
|
|
||||||
if exists(given_betas):
|
|
||||||
betas = given_betas
|
|
||||||
else:
|
|
||||||
betas = make_beta_schedule(beta_schedule, timesteps, linear_start=linear_start, linear_end=linear_end,
|
|
||||||
cosine_s=cosine_s)
|
|
||||||
if self.rescale_betas_zero_snr:
|
|
||||||
betas = rescale_zero_terminal_snr(betas)
|
|
||||||
|
|
||||||
alphas = 1. - betas
|
|
||||||
alphas_cumprod = np.cumprod(alphas, axis=0)
|
|
||||||
alphas_cumprod_prev = np.append(1., alphas_cumprod[:-1])
|
|
||||||
|
|
||||||
timesteps, = betas.shape
|
|
||||||
self.num_timesteps = int(timesteps)
|
|
||||||
self.linear_start = linear_start
|
|
||||||
self.linear_end = linear_end
|
|
||||||
assert alphas_cumprod.shape[0] == self.num_timesteps, 'alphas have to be defined for each timestep'
|
|
||||||
|
|
||||||
to_torch = partial(torch.tensor, dtype=torch.float32)
|
|
||||||
|
|
||||||
self.register_buffer('betas', to_torch(betas))
|
|
||||||
self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod))
|
|
||||||
self.register_buffer('alphas_cumprod_prev', to_torch(alphas_cumprod_prev))
|
|
||||||
|
|
||||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
|
||||||
self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod)))
|
|
||||||
self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod)))
|
|
||||||
self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod)))
|
|
||||||
|
|
||||||
if self.parameterization != 'v':
|
|
||||||
self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod)))
|
|
||||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod - 1)))
|
|
||||||
else:
|
|
||||||
self.register_buffer('sqrt_recip_alphas_cumprod', torch.zeros_like(to_torch(alphas_cumprod)))
|
|
||||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', torch.zeros_like(to_torch(alphas_cumprod)))
|
|
||||||
|
|
||||||
# calculations for posterior q(x_{t-1} | x_t, x_0)
|
|
||||||
posterior_variance = (1 - self.v_posterior) * betas * (1. - alphas_cumprod_prev) / (
|
|
||||||
1. - alphas_cumprod) + self.v_posterior * betas
|
|
||||||
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
|
|
||||||
self.register_buffer('posterior_variance', to_torch(posterior_variance))
|
|
||||||
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
|
||||||
self.register_buffer('posterior_log_variance_clipped', to_torch(np.log(np.maximum(posterior_variance, 1e-20))))
|
|
||||||
self.register_buffer('posterior_mean_coef1', to_torch(
|
|
||||||
betas * np.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod)))
|
|
||||||
self.register_buffer('posterior_mean_coef2', to_torch(
|
|
||||||
(1. - alphas_cumprod_prev) * np.sqrt(alphas) / (1. - alphas_cumprod)))
|
|
||||||
|
|
||||||
if self.parameterization == "eps":
|
|
||||||
lvlb_weights = self.betas ** 2 / (
|
|
||||||
2 * self.posterior_variance * to_torch(alphas) * (1 - self.alphas_cumprod))
|
|
||||||
elif self.parameterization == "x0":
|
|
||||||
lvlb_weights = 0.5 * np.sqrt(torch.Tensor(alphas_cumprod)) / (2. * 1 - torch.Tensor(alphas_cumprod))
|
|
||||||
elif self.parameterization == "v":
|
|
||||||
lvlb_weights = torch.ones_like(self.betas ** 2 / (
|
|
||||||
2 * self.posterior_variance * to_torch(alphas) * (1 - self.alphas_cumprod)))
|
|
||||||
else:
|
|
||||||
raise NotImplementedError("mu not supported")
|
|
||||||
# TODO how to choose this term
|
|
||||||
lvlb_weights[0] = lvlb_weights[1]
|
|
||||||
self.register_buffer('lvlb_weights', lvlb_weights, persistent=False)
|
|
||||||
assert not torch.isnan(self.lvlb_weights).all()
|
|
||||||
|
|
||||||
@contextmanager
|
|
||||||
def ema_scope(self, context=None):
|
|
||||||
if self.use_ema:
|
|
||||||
self.model_ema.store(self.model.parameters())
|
|
||||||
self.model_ema.copy_to(self.model)
|
|
||||||
if context is not None:
|
|
||||||
mainlogger.info(f"{context}: Switched to EMA weights")
|
|
||||||
try:
|
|
||||||
yield None
|
|
||||||
finally:
|
|
||||||
if self.use_ema:
|
|
||||||
self.model_ema.restore(self.model.parameters())
|
|
||||||
if context is not None:
|
|
||||||
mainlogger.info(f"{context}: Restored training weights")
|
|
||||||
|
|
||||||
def init_from_ckpt(self, path, ignore_keys=list(), only_model=False):
|
|
||||||
sd = torch.load(path, map_location="cpu")
|
|
||||||
if "state_dict" in list(sd.keys()):
|
|
||||||
sd = sd["state_dict"]
|
|
||||||
keys = list(sd.keys())
|
|
||||||
for k in keys:
|
|
||||||
for ik in ignore_keys:
|
|
||||||
if k.startswith(ik):
|
|
||||||
mainlogger.info("Deleting key {} from state_dict.".format(k))
|
|
||||||
del sd[k]
|
|
||||||
missing, unexpected = self.load_state_dict(sd, strict=False) if not only_model else self.model.load_state_dict(
|
|
||||||
sd, strict=False)
|
|
||||||
mainlogger.info(f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys")
|
|
||||||
if len(missing) > 0:
|
|
||||||
mainlogger.info(f"Missing Keys: {missing}")
|
|
||||||
if len(unexpected) > 0:
|
|
||||||
mainlogger.info(f"Unexpected Keys: {unexpected}")
|
|
||||||
|
|
||||||
def q_mean_variance(self, x_start, t):
|
|
||||||
"""
|
|
||||||
Get the distribution q(x_t | x_0).
|
|
||||||
:param x_start: the [N x C x ...] tensor of noiseless inputs.
|
|
||||||
:param t: the number of diffusion steps (minus 1). Here, 0 means one step.
|
|
||||||
:return: A tuple (mean, variance, log_variance), all of x_start's shape.
|
|
||||||
"""
|
|
||||||
mean = (extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start)
|
|
||||||
variance = extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape)
|
|
||||||
log_variance = extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape)
|
|
||||||
return mean, variance, log_variance
|
|
||||||
|
|
||||||
def predict_start_from_noise(self, x_t, t, noise):
|
|
||||||
return (
|
|
||||||
extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t -
|
|
||||||
extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise
|
|
||||||
)
|
|
||||||
|
|
||||||
def predict_start_from_z_and_v(self, x_t, t, v):
|
|
||||||
# self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod)))
|
|
||||||
# self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod)))
|
|
||||||
return (
|
|
||||||
extract_into_tensor(self.sqrt_alphas_cumprod, t, x_t.shape) * x_t -
|
|
||||||
extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * v
|
|
||||||
)
|
|
||||||
|
|
||||||
def predict_eps_from_z_and_v(self, x_t, t, v):
|
|
||||||
return (
|
|
||||||
extract_into_tensor(self.sqrt_alphas_cumprod, t, x_t.shape) * v +
|
|
||||||
extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * x_t
|
|
||||||
)
|
|
||||||
|
|
||||||
def q_posterior(self, x_start, x_t, t):
|
|
||||||
posterior_mean = (
|
|
||||||
extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start +
|
|
||||||
extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t
|
|
||||||
)
|
|
||||||
posterior_variance = extract_into_tensor(self.posterior_variance, t, x_t.shape)
|
|
||||||
posterior_log_variance_clipped = extract_into_tensor(self.posterior_log_variance_clipped, t, x_t.shape)
|
|
||||||
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
|
||||||
|
|
||||||
def p_mean_variance(self, x, t, clip_denoised: bool):
|
|
||||||
model_out = self.model(x, t)
|
|
||||||
if self.parameterization == "eps":
|
|
||||||
x_recon = self.predict_start_from_noise(x, t=t, noise=model_out)
|
|
||||||
elif self.parameterization == "x0":
|
|
||||||
x_recon = model_out
|
|
||||||
if clip_denoised:
|
|
||||||
x_recon.clamp_(-1., 1.)
|
|
||||||
|
|
||||||
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start=x_recon, x_t=x, t=t)
|
|
||||||
return model_mean, posterior_variance, posterior_log_variance
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def p_sample(self, x, t, clip_denoised=True, repeat_noise=False):
|
|
||||||
b, *_, device = *x.shape, x.device
|
|
||||||
model_mean, _, model_log_variance = self.p_mean_variance(x=x, t=t, clip_denoised=clip_denoised)
|
|
||||||
noise = noise_like(x.shape, device, repeat_noise)
|
|
||||||
# no noise when t == 0
|
|
||||||
nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1)))
|
|
||||||
return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def p_sample_loop(self, shape, return_intermediates=False):
|
|
||||||
device = self.betas.device
|
|
||||||
b = shape[0]
|
|
||||||
img = torch.randn(shape, device=device)
|
|
||||||
intermediates = [img]
|
|
||||||
for i in tqdm(reversed(range(0, self.num_timesteps)), desc='Sampling t', total=self.num_timesteps):
|
|
||||||
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long),
|
|
||||||
clip_denoised=self.clip_denoised)
|
|
||||||
if i % self.log_every_t == 0 or i == self.num_timesteps - 1:
|
|
||||||
intermediates.append(img)
|
|
||||||
if return_intermediates:
|
|
||||||
return img, intermediates
|
|
||||||
return img
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def sample(self, batch_size=16, return_intermediates=False):
|
|
||||||
image_size = self.image_size
|
|
||||||
channels = self.channels
|
|
||||||
return self.p_sample_loop((batch_size, channels, image_size, image_size),
|
|
||||||
return_intermediates=return_intermediates)
|
|
||||||
|
|
||||||
def q_sample(self, x_start, t, noise=None):
|
|
||||||
noise = default(noise, lambda: torch.randn_like(x_start))
|
|
||||||
return (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)
|
|
||||||
|
|
||||||
def get_v(self, x, noise, t):
|
|
||||||
return (
|
|
||||||
extract_into_tensor(self.sqrt_alphas_cumprod, t, x.shape) * noise -
|
|
||||||
extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x.shape) * x
|
|
||||||
)
|
|
||||||
|
|
||||||
def get_input(self, batch, k):
|
|
||||||
x = batch[k]
|
|
||||||
x = x.to(memory_format=torch.contiguous_format).float()
|
|
||||||
return x
|
|
||||||
|
|
||||||
def _get_rows_from_list(self, samples):
|
|
||||||
n_imgs_per_row = len(samples)
|
|
||||||
denoise_grid = rearrange(samples, 'n b c h w -> b n c h w')
|
|
||||||
denoise_grid = rearrange(denoise_grid, 'b n c h w -> (b n) c h w')
|
|
||||||
denoise_grid = make_grid(denoise_grid, nrow=n_imgs_per_row)
|
|
||||||
return denoise_grid
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def log_images(self, batch, N=8, n_row=2, sample=True, return_keys=None, **kwargs):
|
|
||||||
log = dict()
|
|
||||||
x = self.get_input(batch, self.first_stage_key)
|
|
||||||
N = min(x.shape[0], N)
|
|
||||||
n_row = min(x.shape[0], n_row)
|
|
||||||
x = x.to(self.device)[:N]
|
|
||||||
log["inputs"] = x
|
|
||||||
|
|
||||||
# get diffusion row
|
|
||||||
diffusion_row = list()
|
|
||||||
x_start = x[:n_row]
|
|
||||||
|
|
||||||
for t in range(self.num_timesteps):
|
|
||||||
if t % self.log_every_t == 0 or t == self.num_timesteps - 1:
|
|
||||||
t = repeat(torch.tensor([t]), '1 -> b', b=n_row)
|
|
||||||
t = t.to(self.device).long()
|
|
||||||
noise = torch.randn_like(x_start)
|
|
||||||
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
|
|
||||||
diffusion_row.append(x_noisy)
|
|
||||||
|
|
||||||
log["diffusion_row"] = self._get_rows_from_list(diffusion_row)
|
|
||||||
|
|
||||||
if sample:
|
|
||||||
# get denoise row
|
|
||||||
with self.ema_scope("Plotting"):
|
|
||||||
samples, denoise_row = self.sample(batch_size=N, return_intermediates=True)
|
|
||||||
|
|
||||||
log["samples"] = samples
|
|
||||||
log["denoise_row"] = self._get_rows_from_list(denoise_row)
|
|
||||||
|
|
||||||
if return_keys:
|
|
||||||
if np.intersect1d(list(log.keys()), return_keys).shape[0] == 0:
|
|
||||||
return log
|
|
||||||
else:
|
|
||||||
return {key: log[key] for key in return_keys}
|
|
||||||
return log
|
|
||||||
|
|
||||||
|
|
||||||
class LatentDiffusion(DDPM):
|
|
||||||
"""main class"""
|
|
||||||
def __init__(self,
|
|
||||||
first_stage_config,
|
|
||||||
cond_stage_config,
|
|
||||||
num_timesteps_cond=None,
|
|
||||||
cond_stage_key="caption",
|
|
||||||
cond_stage_trainable=False,
|
|
||||||
cond_stage_forward=None,
|
|
||||||
conditioning_key=None,
|
|
||||||
uncond_prob=0.2,
|
|
||||||
uncond_type="empty_seq",
|
|
||||||
scale_factor=1.0,
|
|
||||||
scale_by_std=False,
|
|
||||||
encoder_type="2d",
|
|
||||||
only_model=False,
|
|
||||||
noise_strength=0,
|
|
||||||
use_dynamic_rescale=False,
|
|
||||||
base_scale=0.7,
|
|
||||||
turning_step=400,
|
|
||||||
loop_video=False,
|
|
||||||
fps_condition_type='fs',
|
|
||||||
perframe_ae=False,
|
|
||||||
*args, **kwargs):
|
|
||||||
self.num_timesteps_cond = default(num_timesteps_cond, 1)
|
|
||||||
self.scale_by_std = scale_by_std
|
|
||||||
assert self.num_timesteps_cond <= kwargs['timesteps']
|
|
||||||
# for backwards compatibility after implementation of DiffusionWrapper
|
|
||||||
ckpt_path = kwargs.pop("ckpt_path", None)
|
|
||||||
ignore_keys = kwargs.pop("ignore_keys", [])
|
|
||||||
conditioning_key = default(conditioning_key, 'crossattn')
|
|
||||||
super().__init__(conditioning_key=conditioning_key, *args, **kwargs)
|
|
||||||
|
|
||||||
self.cond_stage_trainable = cond_stage_trainable
|
|
||||||
self.cond_stage_key = cond_stage_key
|
|
||||||
self.noise_strength = noise_strength
|
|
||||||
self.use_dynamic_rescale = use_dynamic_rescale
|
|
||||||
self.loop_video = loop_video
|
|
||||||
self.fps_condition_type = fps_condition_type
|
|
||||||
self.perframe_ae = perframe_ae
|
|
||||||
try:
|
|
||||||
self.num_downs = len(first_stage_config.params.ddconfig.ch_mult) - 1
|
|
||||||
except:
|
|
||||||
self.num_downs = 0
|
|
||||||
if not scale_by_std:
|
|
||||||
self.scale_factor = scale_factor
|
|
||||||
else:
|
|
||||||
self.register_buffer('scale_factor', torch.tensor(scale_factor))
|
|
||||||
|
|
||||||
if use_dynamic_rescale:
|
|
||||||
scale_arr1 = np.linspace(1.0, base_scale, turning_step)
|
|
||||||
scale_arr2 = np.full(self.num_timesteps, base_scale)
|
|
||||||
scale_arr = np.concatenate((scale_arr1, scale_arr2))
|
|
||||||
to_torch = partial(torch.tensor, dtype=torch.float32)
|
|
||||||
self.register_buffer('scale_arr', to_torch(scale_arr))
|
|
||||||
|
|
||||||
self.instantiate_first_stage(first_stage_config)
|
|
||||||
self.instantiate_cond_stage(cond_stage_config)
|
|
||||||
self.first_stage_config = first_stage_config
|
|
||||||
self.cond_stage_config = cond_stage_config
|
|
||||||
self.clip_denoised = False
|
|
||||||
|
|
||||||
self.cond_stage_forward = cond_stage_forward
|
|
||||||
self.encoder_type = encoder_type
|
|
||||||
assert(encoder_type in ["2d", "3d"])
|
|
||||||
self.uncond_prob = uncond_prob
|
|
||||||
self.classifier_free_guidance = True if uncond_prob > 0 else False
|
|
||||||
assert(uncond_type in ["zero_embed", "empty_seq"])
|
|
||||||
self.uncond_type = uncond_type
|
|
||||||
|
|
||||||
self.restarted_from_ckpt = False
|
|
||||||
if ckpt_path is not None:
|
|
||||||
self.init_from_ckpt(ckpt_path, ignore_keys, only_model=only_model)
|
|
||||||
self.restarted_from_ckpt = True
|
|
||||||
|
|
||||||
|
|
||||||
def make_cond_schedule(self, ):
|
|
||||||
self.cond_ids = torch.full(size=(self.num_timesteps,), fill_value=self.num_timesteps - 1, dtype=torch.long)
|
|
||||||
ids = torch.round(torch.linspace(0, self.num_timesteps - 1, self.num_timesteps_cond)).long()
|
|
||||||
self.cond_ids[:self.num_timesteps_cond] = ids
|
|
||||||
|
|
||||||
def instantiate_first_stage(self, config):
|
|
||||||
model = instantiate_from_config(config)
|
|
||||||
self.first_stage_model = model.eval()
|
|
||||||
self.first_stage_model.train = disabled_train
|
|
||||||
for param in self.first_stage_model.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
|
|
||||||
def instantiate_cond_stage(self, config):
|
|
||||||
if not self.cond_stage_trainable:
|
|
||||||
model = instantiate_from_config(config)
|
|
||||||
self.cond_stage_model = model.eval()
|
|
||||||
self.cond_stage_model.train = disabled_train
|
|
||||||
for param in self.cond_stage_model.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
else:
|
|
||||||
model = instantiate_from_config(config)
|
|
||||||
self.cond_stage_model = model
|
|
||||||
|
|
||||||
def get_learned_conditioning(self, c):
|
|
||||||
if self.cond_stage_forward is None:
|
|
||||||
if hasattr(self.cond_stage_model, 'encode') and callable(self.cond_stage_model.encode):
|
|
||||||
c = self.cond_stage_model.encode(c)
|
|
||||||
if isinstance(c, DiagonalGaussianDistribution):
|
|
||||||
c = c.mode()
|
|
||||||
else:
|
|
||||||
c = self.cond_stage_model(c)
|
|
||||||
else:
|
|
||||||
assert hasattr(self.cond_stage_model, self.cond_stage_forward)
|
|
||||||
c = getattr(self.cond_stage_model, self.cond_stage_forward)(c)
|
|
||||||
return c
|
|
||||||
|
|
||||||
def get_first_stage_encoding(self, encoder_posterior, noise=None):
|
|
||||||
if isinstance(encoder_posterior, DiagonalGaussianDistribution):
|
|
||||||
z = encoder_posterior.sample(noise=noise)
|
|
||||||
elif isinstance(encoder_posterior, torch.Tensor):
|
|
||||||
z = encoder_posterior
|
|
||||||
else:
|
|
||||||
raise NotImplementedError(f"encoder_posterior of type '{type(encoder_posterior)}' not yet implemented")
|
|
||||||
return self.scale_factor * z
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def encode_first_stage(self, x):
|
|
||||||
if self.encoder_type == "2d" and x.dim() == 5:
|
|
||||||
b, _, t, _, _ = x.shape
|
|
||||||
x = rearrange(x, 'b c t h w -> (b t) c h w')
|
|
||||||
reshape_back = True
|
|
||||||
else:
|
|
||||||
reshape_back = False
|
|
||||||
|
|
||||||
## consume more GPU memory but faster
|
|
||||||
if not self.perframe_ae:
|
|
||||||
encoder_posterior = self.first_stage_model.encode(x)
|
|
||||||
results = self.get_first_stage_encoding(encoder_posterior).detach()
|
|
||||||
else: ## consume less GPU memory but slower
|
|
||||||
results = []
|
|
||||||
for index in range(x.shape[0]):
|
|
||||||
frame_batch = self.first_stage_model.encode(x[index:index+1,:,:,:])
|
|
||||||
frame_result = self.get_first_stage_encoding(frame_batch).detach()
|
|
||||||
results.append(frame_result)
|
|
||||||
results = torch.cat(results, dim=0)
|
|
||||||
|
|
||||||
if reshape_back:
|
|
||||||
results = rearrange(results, '(b t) c h w -> b c t h w', b=b,t=t)
|
|
||||||
|
|
||||||
return results
|
|
||||||
|
|
||||||
def decode_core(self, z, **kwargs):
|
|
||||||
if self.encoder_type == "2d" and z.dim() == 5:
|
|
||||||
b, _, t, _, _ = z.shape
|
|
||||||
z = rearrange(z, 'b c t h w -> (b t) c h w')
|
|
||||||
reshape_back = True
|
|
||||||
else:
|
|
||||||
reshape_back = False
|
|
||||||
|
|
||||||
if not self.perframe_ae:
|
|
||||||
z = 1. / self.scale_factor * z
|
|
||||||
results = self.first_stage_model.decode(z, **kwargs)
|
|
||||||
else:
|
|
||||||
results = []
|
|
||||||
for index in range(z.shape[0]):
|
|
||||||
frame_z = 1. / self.scale_factor * z[index:index+1,:,:,:]
|
|
||||||
frame_result = self.first_stage_model.decode(frame_z, **kwargs)
|
|
||||||
results.append(frame_result)
|
|
||||||
results = torch.cat(results, dim=0)
|
|
||||||
|
|
||||||
if reshape_back:
|
|
||||||
results = rearrange(results, '(b t) c h w -> b c t h w', b=b,t=t)
|
|
||||||
return results
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def decode_first_stage(self, z, **kwargs):
|
|
||||||
return self.decode_core(z, **kwargs)
|
|
||||||
|
|
||||||
# same as above but without decorator
|
|
||||||
def differentiable_decode_first_stage(self, z, **kwargs):
|
|
||||||
return self.decode_core(z, **kwargs)
|
|
||||||
|
|
||||||
def forward(self, x, c, **kwargs):
|
|
||||||
t = torch.randint(0, self.num_timesteps, (x.shape[0],), device=self.device).long()
|
|
||||||
if self.use_dynamic_rescale:
|
|
||||||
x = x * extract_into_tensor(self.scale_arr, t, x.shape)
|
|
||||||
return self.p_losses(x, c, t, **kwargs)
|
|
||||||
|
|
||||||
def apply_model(self, x_noisy, t, cond, **kwargs):
|
|
||||||
if isinstance(cond, dict):
|
|
||||||
# hybrid case, cond is exptected to be a dict
|
|
||||||
pass
|
|
||||||
else:
|
|
||||||
if not isinstance(cond, list):
|
|
||||||
cond = [cond]
|
|
||||||
key = 'c_concat' if self.model.conditioning_key == 'concat' else 'c_crossattn'
|
|
||||||
cond = {key: cond}
|
|
||||||
|
|
||||||
x_recon = self.model(x_noisy, t, **cond, **kwargs)
|
|
||||||
|
|
||||||
if isinstance(x_recon, tuple):
|
|
||||||
return x_recon[0]
|
|
||||||
else:
|
|
||||||
return x_recon
|
|
||||||
|
|
||||||
def _get_denoise_row_from_list(self, samples, desc=''):
|
|
||||||
denoise_row = []
|
|
||||||
for zd in tqdm(samples, desc=desc):
|
|
||||||
denoise_row.append(self.decode_first_stage(zd.to(self.device)))
|
|
||||||
n_log_timesteps = len(denoise_row)
|
|
||||||
|
|
||||||
denoise_row = torch.stack(denoise_row) # n_log_timesteps, b, C, H, W
|
|
||||||
|
|
||||||
if denoise_row.dim() == 5:
|
|
||||||
denoise_grid = rearrange(denoise_row, 'n b c h w -> b n c h w')
|
|
||||||
denoise_grid = rearrange(denoise_grid, 'b n c h w -> (b n) c h w')
|
|
||||||
denoise_grid = make_grid(denoise_grid, nrow=n_log_timesteps)
|
|
||||||
elif denoise_row.dim() == 6:
|
|
||||||
# video, grid_size=[n_log_timesteps*bs, t]
|
|
||||||
video_length = denoise_row.shape[3]
|
|
||||||
denoise_grid = rearrange(denoise_row, 'n b c t h w -> b n c t h w')
|
|
||||||
denoise_grid = rearrange(denoise_grid, 'b n c t h w -> (b n) c t h w')
|
|
||||||
denoise_grid = rearrange(denoise_grid, 'n c t h w -> (n t) c h w')
|
|
||||||
denoise_grid = make_grid(denoise_grid, nrow=video_length)
|
|
||||||
else:
|
|
||||||
raise ValueError
|
|
||||||
|
|
||||||
return denoise_grid
|
|
||||||
|
|
||||||
|
|
||||||
def p_mean_variance(self, x, c, t, clip_denoised: bool, return_x0=False, score_corrector=None, corrector_kwargs=None, **kwargs):
|
|
||||||
t_in = t
|
|
||||||
model_out = self.apply_model(x, t_in, c, **kwargs)
|
|
||||||
|
|
||||||
if score_corrector is not None:
|
|
||||||
assert self.parameterization == "eps"
|
|
||||||
model_out = score_corrector.modify_score(self, model_out, x, t, c, **corrector_kwargs)
|
|
||||||
|
|
||||||
if self.parameterization == "eps":
|
|
||||||
x_recon = self.predict_start_from_noise(x, t=t, noise=model_out)
|
|
||||||
elif self.parameterization == "x0":
|
|
||||||
x_recon = model_out
|
|
||||||
else:
|
|
||||||
raise NotImplementedError()
|
|
||||||
|
|
||||||
if clip_denoised:
|
|
||||||
x_recon.clamp_(-1., 1.)
|
|
||||||
|
|
||||||
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start=x_recon, x_t=x, t=t)
|
|
||||||
|
|
||||||
if return_x0:
|
|
||||||
return model_mean, posterior_variance, posterior_log_variance, x_recon
|
|
||||||
else:
|
|
||||||
return model_mean, posterior_variance, posterior_log_variance
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def p_sample(self, x, c, t, clip_denoised=False, repeat_noise=False, return_x0=False, \
|
|
||||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, **kwargs):
|
|
||||||
b, *_, device = *x.shape, x.device
|
|
||||||
outputs = self.p_mean_variance(x=x, c=c, t=t, clip_denoised=clip_denoised, return_x0=return_x0, \
|
|
||||||
score_corrector=score_corrector, corrector_kwargs=corrector_kwargs, **kwargs)
|
|
||||||
if return_x0:
|
|
||||||
model_mean, _, model_log_variance, x0 = outputs
|
|
||||||
else:
|
|
||||||
model_mean, _, model_log_variance = outputs
|
|
||||||
|
|
||||||
noise = noise_like(x.shape, device, repeat_noise) * temperature
|
|
||||||
if noise_dropout > 0.:
|
|
||||||
noise = torch.nn.functional.dropout(noise, p=noise_dropout)
|
|
||||||
# no noise when t == 0
|
|
||||||
nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1)))
|
|
||||||
|
|
||||||
if return_x0:
|
|
||||||
return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise, x0
|
|
||||||
else:
|
|
||||||
return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def p_sample_loop(self, cond, shape, return_intermediates=False, x_T=None, verbose=True, callback=None, \
|
|
||||||
timesteps=None, mask=None, x0=None, img_callback=None, start_T=None, log_every_t=None, **kwargs):
|
|
||||||
|
|
||||||
if not log_every_t:
|
|
||||||
log_every_t = self.log_every_t
|
|
||||||
device = self.betas.device
|
|
||||||
b = shape[0]
|
|
||||||
# sample an initial noise
|
|
||||||
if x_T is None:
|
|
||||||
img = torch.randn(shape, device=device)
|
|
||||||
else:
|
|
||||||
img = x_T
|
|
||||||
|
|
||||||
intermediates = [img]
|
|
||||||
if timesteps is None:
|
|
||||||
timesteps = self.num_timesteps
|
|
||||||
if start_T is not None:
|
|
||||||
timesteps = min(timesteps, start_T)
|
|
||||||
|
|
||||||
iterator = tqdm(reversed(range(0, timesteps)), desc='Sampling t', total=timesteps) if verbose else reversed(range(0, timesteps))
|
|
||||||
|
|
||||||
if mask is not None:
|
|
||||||
assert x0 is not None
|
|
||||||
assert x0.shape[2:3] == mask.shape[2:3] # spatial size has to match
|
|
||||||
|
|
||||||
for i in iterator:
|
|
||||||
ts = torch.full((b,), i, device=device, dtype=torch.long)
|
|
||||||
if self.shorten_cond_schedule:
|
|
||||||
assert self.model.conditioning_key != 'hybrid'
|
|
||||||
tc = self.cond_ids[ts].to(cond.device)
|
|
||||||
cond = self.q_sample(x_start=cond, t=tc, noise=torch.randn_like(cond))
|
|
||||||
|
|
||||||
img = self.p_sample(img, cond, ts, clip_denoised=self.clip_denoised, **kwargs)
|
|
||||||
if mask is not None:
|
|
||||||
img_orig = self.q_sample(x0, ts)
|
|
||||||
img = img_orig * mask + (1. - mask) * img
|
|
||||||
|
|
||||||
if i % log_every_t == 0 or i == timesteps - 1:
|
|
||||||
intermediates.append(img)
|
|
||||||
if callback: callback(i)
|
|
||||||
if img_callback: img_callback(img, i)
|
|
||||||
|
|
||||||
if return_intermediates:
|
|
||||||
return img, intermediates
|
|
||||||
return img
|
|
||||||
|
|
||||||
|
|
||||||
class LatentVisualDiffusion(LatentDiffusion):
|
|
||||||
def __init__(self, img_cond_stage_config, image_proj_stage_config, freeze_embedder=True, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self._init_embedder(img_cond_stage_config, freeze_embedder)
|
|
||||||
self.image_proj_model = instantiate_from_config(image_proj_stage_config)
|
|
||||||
|
|
||||||
def _init_embedder(self, config, freeze=True):
|
|
||||||
embedder = instantiate_from_config(config)
|
|
||||||
if freeze:
|
|
||||||
self.embedder = embedder.eval()
|
|
||||||
self.embedder.train = disabled_train
|
|
||||||
for param in self.embedder.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
|
|
||||||
|
|
||||||
class DiffusionWrapper(nn.Module):
|
|
||||||
def __init__(self, diff_model_config, conditioning_key):
|
|
||||||
super().__init__()
|
|
||||||
self.diffusion_model = instantiate_from_config(diff_model_config)
|
|
||||||
self.conditioning_key = conditioning_key
|
|
||||||
|
|
||||||
def forward(self, x, t, c_concat: list = None, c_crossattn: list = None,
|
|
||||||
c_adm=None, s=None, mask=None, **kwargs):
|
|
||||||
# temporal_context = fps is foNone
|
|
||||||
if self.conditioning_key is None:
|
|
||||||
out = self.diffusion_model(x, t)
|
|
||||||
elif self.conditioning_key == 'concat':
|
|
||||||
xc = torch.cat([x] + c_concat, dim=1)
|
|
||||||
out = self.diffusion_model(xc, t, **kwargs)
|
|
||||||
elif self.conditioning_key == 'crossattn':
|
|
||||||
cc = torch.cat(c_crossattn, 1)
|
|
||||||
out = self.diffusion_model(x, t, context=cc, **kwargs)
|
|
||||||
elif self.conditioning_key == 'hybrid':
|
|
||||||
## it is just right [b,c,t,h,w]: concatenate in channel dim
|
|
||||||
xc = torch.cat([x] + c_concat, dim=1)
|
|
||||||
cc = torch.cat(c_crossattn, 1)
|
|
||||||
out = self.diffusion_model(xc, t, context=cc, **kwargs)
|
|
||||||
elif self.conditioning_key == 'resblockcond':
|
|
||||||
cc = c_crossattn[0]
|
|
||||||
out = self.diffusion_model(x, t, context=cc)
|
|
||||||
elif self.conditioning_key == 'adm':
|
|
||||||
cc = c_crossattn[0]
|
|
||||||
out = self.diffusion_model(x, t, y=cc)
|
|
||||||
elif self.conditioning_key == 'hybrid-adm':
|
|
||||||
assert c_adm is not None
|
|
||||||
xc = torch.cat([x] + c_concat, dim=1)
|
|
||||||
cc = torch.cat(c_crossattn, 1)
|
|
||||||
out = self.diffusion_model(xc, t, context=cc, y=c_adm, **kwargs)
|
|
||||||
elif self.conditioning_key == 'hybrid-time':
|
|
||||||
assert s is not None
|
|
||||||
xc = torch.cat([x] + c_concat, dim=1)
|
|
||||||
cc = torch.cat(c_crossattn, 1)
|
|
||||||
out = self.diffusion_model(xc, t, context=cc, s=s)
|
|
||||||
elif self.conditioning_key == 'concat-time-mask':
|
|
||||||
# assert s is not None
|
|
||||||
xc = torch.cat([x] + c_concat, dim=1)
|
|
||||||
out = self.diffusion_model(xc, t, context=None, s=s, mask=mask)
|
|
||||||
elif self.conditioning_key == 'concat-adm-mask':
|
|
||||||
# assert s is not None
|
|
||||||
if c_concat is not None:
|
|
||||||
xc = torch.cat([x] + c_concat, dim=1)
|
|
||||||
else:
|
|
||||||
xc = x
|
|
||||||
out = self.diffusion_model(xc, t, context=None, y=s, mask=mask)
|
|
||||||
elif self.conditioning_key == 'hybrid-adm-mask':
|
|
||||||
cc = torch.cat(c_crossattn, 1)
|
|
||||||
if c_concat is not None:
|
|
||||||
xc = torch.cat([x] + c_concat, dim=1)
|
|
||||||
else:
|
|
||||||
xc = x
|
|
||||||
out = self.diffusion_model(xc, t, context=cc, y=s, mask=mask)
|
|
||||||
elif self.conditioning_key == 'hybrid-time-adm': # adm means y, e.g., class index
|
|
||||||
# assert s is not None
|
|
||||||
assert c_adm is not None
|
|
||||||
xc = torch.cat([x] + c_concat, dim=1)
|
|
||||||
cc = torch.cat(c_crossattn, 1)
|
|
||||||
out = self.diffusion_model(xc, t, context=cc, s=s, y=c_adm)
|
|
||||||
elif self.conditioning_key == 'crossattn-adm':
|
|
||||||
assert c_adm is not None
|
|
||||||
cc = torch.cat(c_crossattn, 1)
|
|
||||||
out = self.diffusion_model(x, t, context=cc, y=c_adm)
|
|
||||||
else:
|
|
||||||
raise NotImplementedError()
|
|
||||||
|
|
||||||
return out
|
|
||||||
@@ -1,317 +0,0 @@
|
|||||||
import numpy as np
|
|
||||||
from tqdm import tqdm
|
|
||||||
import torch
|
|
||||||
from ..models.utils_diffusion import make_ddim_sampling_parameters, make_ddim_timesteps, rescale_noise_cfg
|
|
||||||
from ..common import noise_like
|
|
||||||
from ..common import extract_into_tensor
|
|
||||||
import copy
|
|
||||||
|
|
||||||
|
|
||||||
class DDIMSampler(object):
|
|
||||||
def __init__(self, model, schedule="linear", **kwargs):
|
|
||||||
super().__init__()
|
|
||||||
self.model = model
|
|
||||||
self.ddpm_num_timesteps = model.num_timesteps
|
|
||||||
self.schedule = schedule
|
|
||||||
self.counter = 0
|
|
||||||
|
|
||||||
def register_buffer(self, name, attr):
|
|
||||||
if type(attr) == torch.Tensor:
|
|
||||||
if attr.device != torch.device("cuda"):
|
|
||||||
attr = attr.to(torch.device("cuda"))
|
|
||||||
setattr(self, name, attr)
|
|
||||||
|
|
||||||
def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True):
|
|
||||||
self.ddim_timesteps = make_ddim_timesteps(ddim_discr_method=ddim_discretize, num_ddim_timesteps=ddim_num_steps,
|
|
||||||
num_ddpm_timesteps=self.ddpm_num_timesteps,verbose=verbose)
|
|
||||||
alphas_cumprod = self.model.alphas_cumprod
|
|
||||||
assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, 'alphas have to be defined for each timestep'
|
|
||||||
to_torch = lambda x: x.clone().detach().to(torch.float32).to(self.model.device)
|
|
||||||
|
|
||||||
if self.model.use_dynamic_rescale:
|
|
||||||
self.ddim_scale_arr = self.model.scale_arr[self.ddim_timesteps]
|
|
||||||
self.ddim_scale_arr_prev = torch.cat([self.ddim_scale_arr[0:1], self.ddim_scale_arr[:-1]])
|
|
||||||
|
|
||||||
self.register_buffer('betas', to_torch(self.model.betas))
|
|
||||||
self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod))
|
|
||||||
self.register_buffer('alphas_cumprod_prev', to_torch(self.model.alphas_cumprod_prev))
|
|
||||||
|
|
||||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
|
||||||
self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod.cpu())))
|
|
||||||
self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod.cpu())))
|
|
||||||
self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod.cpu())))
|
|
||||||
self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu())))
|
|
||||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu() - 1)))
|
|
||||||
|
|
||||||
# ddim sampling parameters
|
|
||||||
ddim_sigmas, ddim_alphas, ddim_alphas_prev = make_ddim_sampling_parameters(alphacums=alphas_cumprod.cpu(),
|
|
||||||
ddim_timesteps=self.ddim_timesteps,
|
|
||||||
eta=ddim_eta,verbose=verbose)
|
|
||||||
self.register_buffer('ddim_sigmas', ddim_sigmas)
|
|
||||||
self.register_buffer('ddim_alphas', ddim_alphas)
|
|
||||||
self.register_buffer('ddim_alphas_prev', ddim_alphas_prev)
|
|
||||||
self.register_buffer('ddim_sqrt_one_minus_alphas', np.sqrt(1. - ddim_alphas))
|
|
||||||
sigmas_for_original_sampling_steps = ddim_eta * torch.sqrt(
|
|
||||||
(1 - self.alphas_cumprod_prev) / (1 - self.alphas_cumprod) * (
|
|
||||||
1 - self.alphas_cumprod / self.alphas_cumprod_prev))
|
|
||||||
self.register_buffer('ddim_sigmas_for_original_num_steps', sigmas_for_original_sampling_steps)
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def sample(self,
|
|
||||||
S,
|
|
||||||
batch_size,
|
|
||||||
shape,
|
|
||||||
conditioning=None,
|
|
||||||
callback=None,
|
|
||||||
normals_sequence=None,
|
|
||||||
img_callback=None,
|
|
||||||
quantize_x0=False,
|
|
||||||
eta=0.,
|
|
||||||
mask=None,
|
|
||||||
x0=None,
|
|
||||||
temperature=1.,
|
|
||||||
noise_dropout=0.,
|
|
||||||
score_corrector=None,
|
|
||||||
corrector_kwargs=None,
|
|
||||||
verbose=True,
|
|
||||||
schedule_verbose=False,
|
|
||||||
x_T=None,
|
|
||||||
log_every_t=100,
|
|
||||||
unconditional_guidance_scale=1.,
|
|
||||||
unconditional_conditioning=None,
|
|
||||||
precision=None,
|
|
||||||
fs=None,
|
|
||||||
timestep_spacing='uniform', #uniform_trailing for starting from last timestep
|
|
||||||
guidance_rescale=0.0,
|
|
||||||
**kwargs
|
|
||||||
):
|
|
||||||
|
|
||||||
# check condition bs
|
|
||||||
if conditioning is not None:
|
|
||||||
if isinstance(conditioning, dict):
|
|
||||||
try:
|
|
||||||
cbs = conditioning[list(conditioning.keys())[0]].shape[0]
|
|
||||||
except:
|
|
||||||
cbs = conditioning[list(conditioning.keys())[0]][0].shape[0]
|
|
||||||
|
|
||||||
if cbs != batch_size:
|
|
||||||
print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}")
|
|
||||||
else:
|
|
||||||
if conditioning.shape[0] != batch_size:
|
|
||||||
print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}")
|
|
||||||
|
|
||||||
self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, verbose=schedule_verbose)
|
|
||||||
|
|
||||||
# make shape
|
|
||||||
if len(shape) == 3:
|
|
||||||
C, H, W = shape
|
|
||||||
size = (batch_size, C, H, W)
|
|
||||||
elif len(shape) == 4:
|
|
||||||
C, T, H, W = shape
|
|
||||||
size = (batch_size, C, T, H, W)
|
|
||||||
|
|
||||||
samples, intermediates = self.ddim_sampling(conditioning, size,
|
|
||||||
callback=callback,
|
|
||||||
img_callback=img_callback,
|
|
||||||
quantize_denoised=quantize_x0,
|
|
||||||
mask=mask, x0=x0,
|
|
||||||
ddim_use_original_steps=False,
|
|
||||||
noise_dropout=noise_dropout,
|
|
||||||
temperature=temperature,
|
|
||||||
score_corrector=score_corrector,
|
|
||||||
corrector_kwargs=corrector_kwargs,
|
|
||||||
x_T=x_T,
|
|
||||||
log_every_t=log_every_t,
|
|
||||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
|
||||||
unconditional_conditioning=unconditional_conditioning,
|
|
||||||
verbose=verbose,
|
|
||||||
precision=precision,
|
|
||||||
fs=fs,
|
|
||||||
guidance_rescale=guidance_rescale,
|
|
||||||
**kwargs)
|
|
||||||
return samples, intermediates
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def ddim_sampling(self, cond, shape,
|
|
||||||
x_T=None, ddim_use_original_steps=False,
|
|
||||||
callback=None, timesteps=None, quantize_denoised=False,
|
|
||||||
mask=None, x0=None, img_callback=None, log_every_t=100,
|
|
||||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
|
||||||
unconditional_guidance_scale=1., unconditional_conditioning=None, verbose=True,precision=None,fs=None,guidance_rescale=0.0,
|
|
||||||
**kwargs):
|
|
||||||
device = self.model.betas.device
|
|
||||||
b = shape[0]
|
|
||||||
if x_T is None:
|
|
||||||
img = torch.randn(shape, device=device)
|
|
||||||
else:
|
|
||||||
img = x_T
|
|
||||||
if precision is not None:
|
|
||||||
if precision == 16:
|
|
||||||
img = img.to(dtype=torch.float16)
|
|
||||||
|
|
||||||
if timesteps is None:
|
|
||||||
timesteps = self.ddpm_num_timesteps if ddim_use_original_steps else self.ddim_timesteps
|
|
||||||
elif timesteps is not None and not ddim_use_original_steps:
|
|
||||||
subset_end = int(min(timesteps / self.ddim_timesteps.shape[0], 1) * self.ddim_timesteps.shape[0]) - 1
|
|
||||||
timesteps = self.ddim_timesteps[:subset_end]
|
|
||||||
|
|
||||||
intermediates = {'x_inter': [img], 'pred_x0': [img]}
|
|
||||||
time_range = reversed(range(0,timesteps)) if ddim_use_original_steps else np.flip(timesteps)
|
|
||||||
total_steps = timesteps if ddim_use_original_steps else timesteps.shape[0]
|
|
||||||
if verbose:
|
|
||||||
iterator = tqdm(time_range, desc='DDIM Sampler', total=total_steps)
|
|
||||||
else:
|
|
||||||
iterator = time_range
|
|
||||||
|
|
||||||
clean_cond = kwargs.pop("clean_cond", False)
|
|
||||||
|
|
||||||
# cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning)
|
|
||||||
for i, step in enumerate(iterator):
|
|
||||||
index = total_steps - i - 1
|
|
||||||
ts = torch.full((b,), step, device=device, dtype=torch.long)
|
|
||||||
|
|
||||||
## use mask to blend noised original latent (img_orig) & new sampled latent (img)
|
|
||||||
if mask is not None:
|
|
||||||
assert x0 is not None
|
|
||||||
if clean_cond:
|
|
||||||
img_orig = x0
|
|
||||||
else:
|
|
||||||
img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? <ddim inversion>
|
|
||||||
img = img_orig * mask + (1. - mask) * img # keep original & modify use img
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps,
|
|
||||||
quantize_denoised=quantize_denoised, temperature=temperature,
|
|
||||||
noise_dropout=noise_dropout, score_corrector=score_corrector,
|
|
||||||
corrector_kwargs=corrector_kwargs,
|
|
||||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
|
||||||
unconditional_conditioning=unconditional_conditioning,
|
|
||||||
mask=mask,x0=x0,fs=fs,guidance_rescale=guidance_rescale,
|
|
||||||
**kwargs)
|
|
||||||
|
|
||||||
|
|
||||||
img, pred_x0 = outs
|
|
||||||
if callback: callback(i)
|
|
||||||
if img_callback: img_callback(pred_x0, i)
|
|
||||||
|
|
||||||
if index % log_every_t == 0 or index == total_steps - 1:
|
|
||||||
intermediates['x_inter'].append(img)
|
|
||||||
intermediates['pred_x0'].append(pred_x0)
|
|
||||||
|
|
||||||
return img, intermediates
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def p_sample_ddim(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
|
|
||||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
|
||||||
unconditional_guidance_scale=1., unconditional_conditioning=None,
|
|
||||||
uc_type=None, conditional_guidance_scale_temporal=None,mask=None,x0=None,guidance_rescale=0.0,**kwargs):
|
|
||||||
b, *_, device = *x.shape, x.device
|
|
||||||
if x.dim() == 5:
|
|
||||||
is_video = True
|
|
||||||
else:
|
|
||||||
is_video = False
|
|
||||||
|
|
||||||
if unconditional_conditioning is None or unconditional_guidance_scale == 1.:
|
|
||||||
model_output = self.model.apply_model(x, t, c, **kwargs) # unet denoiser
|
|
||||||
else:
|
|
||||||
### do_classifier_free_guidance
|
|
||||||
if isinstance(c, torch.Tensor) or isinstance(c, dict):
|
|
||||||
e_t_cond = self.model.apply_model(x, t, c, **kwargs)
|
|
||||||
e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs)
|
|
||||||
else:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
model_output = e_t_uncond + unconditional_guidance_scale * (e_t_cond - e_t_uncond)
|
|
||||||
|
|
||||||
if guidance_rescale > 0.0:
|
|
||||||
model_output = rescale_noise_cfg(model_output, e_t_cond, guidance_rescale=guidance_rescale)
|
|
||||||
|
|
||||||
if self.model.parameterization == "v":
|
|
||||||
e_t = self.model.predict_eps_from_z_and_v(x, t, model_output)
|
|
||||||
else:
|
|
||||||
e_t = model_output
|
|
||||||
|
|
||||||
if score_corrector is not None:
|
|
||||||
assert self.model.parameterization == "eps", 'not implemented'
|
|
||||||
e_t = score_corrector.modify_score(self.model, e_t, x, t, c, **corrector_kwargs)
|
|
||||||
|
|
||||||
alphas = self.model.alphas_cumprod if use_original_steps else self.ddim_alphas
|
|
||||||
alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev
|
|
||||||
sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas
|
|
||||||
# sigmas = self.model.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
|
||||||
sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
|
||||||
# select parameters corresponding to the currently considered timestep
|
|
||||||
|
|
||||||
if is_video:
|
|
||||||
size = (b, 1, 1, 1, 1)
|
|
||||||
else:
|
|
||||||
size = (b, 1, 1, 1)
|
|
||||||
a_t = torch.full(size, alphas[index], device=device)
|
|
||||||
a_prev = torch.full(size, alphas_prev[index], device=device)
|
|
||||||
sigma_t = torch.full(size, sigmas[index], device=device)
|
|
||||||
sqrt_one_minus_at = torch.full(size, sqrt_one_minus_alphas[index],device=device)
|
|
||||||
|
|
||||||
# current prediction for x_0
|
|
||||||
if self.model.parameterization != "v":
|
|
||||||
pred_x0 = (x - sqrt_one_minus_at * e_t) / a_t.sqrt()
|
|
||||||
else:
|
|
||||||
pred_x0 = self.model.predict_start_from_z_and_v(x, t, model_output)
|
|
||||||
|
|
||||||
if self.model.use_dynamic_rescale:
|
|
||||||
scale_t = torch.full(size, self.ddim_scale_arr[index], device=device)
|
|
||||||
prev_scale_t = torch.full(size, self.ddim_scale_arr_prev[index], device=device)
|
|
||||||
rescale = (prev_scale_t / scale_t)
|
|
||||||
pred_x0 *= rescale
|
|
||||||
|
|
||||||
if quantize_denoised:
|
|
||||||
pred_x0, _, *_ = self.model.first_stage_model.quantize(pred_x0)
|
|
||||||
# direction pointing to x_t
|
|
||||||
dir_xt = (1. - a_prev - sigma_t**2).sqrt() * e_t
|
|
||||||
|
|
||||||
noise = sigma_t * noise_like(x.shape, device, repeat_noise) * temperature
|
|
||||||
if noise_dropout > 0.:
|
|
||||||
noise = torch.nn.functional.dropout(noise, p=noise_dropout)
|
|
||||||
|
|
||||||
x_prev = a_prev.sqrt() * pred_x0 + dir_xt + noise
|
|
||||||
|
|
||||||
return x_prev, pred_x0
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def decode(self, x_latent, cond, t_start, unconditional_guidance_scale=1.0, unconditional_conditioning=None,
|
|
||||||
use_original_steps=False, callback=None):
|
|
||||||
|
|
||||||
timesteps = np.arange(self.ddpm_num_timesteps) if use_original_steps else self.ddim_timesteps
|
|
||||||
timesteps = timesteps[:t_start]
|
|
||||||
|
|
||||||
time_range = np.flip(timesteps)
|
|
||||||
total_steps = timesteps.shape[0]
|
|
||||||
print(f"Running DDIM Sampling with {total_steps} timesteps")
|
|
||||||
|
|
||||||
iterator = tqdm(time_range, desc='Decoding image', total=total_steps)
|
|
||||||
x_dec = x_latent
|
|
||||||
for i, step in enumerate(iterator):
|
|
||||||
index = total_steps - i - 1
|
|
||||||
ts = torch.full((x_latent.shape[0],), step, device=x_latent.device, dtype=torch.long)
|
|
||||||
x_dec, _ = self.p_sample_ddim(x_dec, cond, ts, index=index, use_original_steps=use_original_steps,
|
|
||||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
|
||||||
unconditional_conditioning=unconditional_conditioning)
|
|
||||||
if callback: callback(i)
|
|
||||||
return x_dec
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def stochastic_encode(self, x0, t, use_original_steps=False, noise=None):
|
|
||||||
# fast, but does not allow for exact reconstruction
|
|
||||||
# t serves as an index to gather the correct alphas
|
|
||||||
if use_original_steps:
|
|
||||||
sqrt_alphas_cumprod = self.sqrt_alphas_cumprod
|
|
||||||
sqrt_one_minus_alphas_cumprod = self.sqrt_one_minus_alphas_cumprod
|
|
||||||
else:
|
|
||||||
sqrt_alphas_cumprod = torch.sqrt(self.ddim_alphas)
|
|
||||||
sqrt_one_minus_alphas_cumprod = self.ddim_sqrt_one_minus_alphas
|
|
||||||
|
|
||||||
if noise is None:
|
|
||||||
noise = torch.randn_like(x0)
|
|
||||||
return (extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
|
|
||||||
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) * noise)
|
|
||||||
@@ -1,323 +0,0 @@
|
|||||||
import numpy as np
|
|
||||||
from tqdm import tqdm
|
|
||||||
import torch
|
|
||||||
from ...models.utils_diffusion import make_ddim_sampling_parameters, make_ddim_timesteps, rescale_noise_cfg
|
|
||||||
from ..common import noise_like
|
|
||||||
from ..common import extract_into_tensor
|
|
||||||
import copy
|
|
||||||
|
|
||||||
|
|
||||||
class DDIMSampler(object):
|
|
||||||
def __init__(self, model, schedule="linear", **kwargs):
|
|
||||||
super().__init__()
|
|
||||||
self.model = model
|
|
||||||
self.ddpm_num_timesteps = model.num_timesteps
|
|
||||||
self.schedule = schedule
|
|
||||||
self.counter = 0
|
|
||||||
|
|
||||||
def register_buffer(self, name, attr):
|
|
||||||
if type(attr) == torch.Tensor:
|
|
||||||
if attr.device != torch.device("cuda"):
|
|
||||||
attr = attr.to(torch.device("cuda"))
|
|
||||||
setattr(self, name, attr)
|
|
||||||
|
|
||||||
def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True):
|
|
||||||
self.ddim_timesteps = make_ddim_timesteps(ddim_discr_method=ddim_discretize, num_ddim_timesteps=ddim_num_steps,
|
|
||||||
num_ddpm_timesteps=self.ddpm_num_timesteps,verbose=verbose)
|
|
||||||
alphas_cumprod = self.model.alphas_cumprod
|
|
||||||
assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, 'alphas have to be defined for each timestep'
|
|
||||||
to_torch = lambda x: x.clone().detach().to(torch.float32).to(self.model.device)
|
|
||||||
|
|
||||||
if self.model.use_dynamic_rescale:
|
|
||||||
self.ddim_scale_arr = self.model.scale_arr[self.ddim_timesteps]
|
|
||||||
self.ddim_scale_arr_prev = torch.cat([self.ddim_scale_arr[0:1], self.ddim_scale_arr[:-1]])
|
|
||||||
|
|
||||||
self.register_buffer('betas', to_torch(self.model.betas))
|
|
||||||
self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod))
|
|
||||||
self.register_buffer('alphas_cumprod_prev', to_torch(self.model.alphas_cumprod_prev))
|
|
||||||
|
|
||||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
|
||||||
self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod.cpu())))
|
|
||||||
self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod.cpu())))
|
|
||||||
self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod.cpu())))
|
|
||||||
self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu())))
|
|
||||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu() - 1)))
|
|
||||||
|
|
||||||
# ddim sampling parameters
|
|
||||||
ddim_sigmas, ddim_alphas, ddim_alphas_prev = make_ddim_sampling_parameters(alphacums=alphas_cumprod.cpu(),
|
|
||||||
ddim_timesteps=self.ddim_timesteps,
|
|
||||||
eta=ddim_eta,verbose=verbose)
|
|
||||||
self.register_buffer('ddim_sigmas', ddim_sigmas)
|
|
||||||
self.register_buffer('ddim_alphas', ddim_alphas)
|
|
||||||
self.register_buffer('ddim_alphas_prev', ddim_alphas_prev)
|
|
||||||
self.register_buffer('ddim_sqrt_one_minus_alphas', np.sqrt(1. - ddim_alphas))
|
|
||||||
sigmas_for_original_sampling_steps = ddim_eta * torch.sqrt(
|
|
||||||
(1 - self.alphas_cumprod_prev) / (1 - self.alphas_cumprod) * (
|
|
||||||
1 - self.alphas_cumprod / self.alphas_cumprod_prev))
|
|
||||||
self.register_buffer('ddim_sigmas_for_original_num_steps', sigmas_for_original_sampling_steps)
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def sample(self,
|
|
||||||
S,
|
|
||||||
batch_size,
|
|
||||||
shape,
|
|
||||||
conditioning=None,
|
|
||||||
callback=None,
|
|
||||||
normals_sequence=None,
|
|
||||||
img_callback=None,
|
|
||||||
quantize_x0=False,
|
|
||||||
eta=0.,
|
|
||||||
mask=None,
|
|
||||||
x0=None,
|
|
||||||
temperature=1.,
|
|
||||||
noise_dropout=0.,
|
|
||||||
score_corrector=None,
|
|
||||||
corrector_kwargs=None,
|
|
||||||
verbose=True,
|
|
||||||
schedule_verbose=False,
|
|
||||||
x_T=None,
|
|
||||||
log_every_t=100,
|
|
||||||
unconditional_guidance_scale=1.,
|
|
||||||
unconditional_conditioning=None,
|
|
||||||
precision=None,
|
|
||||||
fs=None,
|
|
||||||
timestep_spacing='uniform', #uniform_trailing for starting from last timestep
|
|
||||||
guidance_rescale=0.0,
|
|
||||||
# this has to come in the same format as the conditioning, # e.g. as encoded tokens, ...
|
|
||||||
**kwargs
|
|
||||||
):
|
|
||||||
|
|
||||||
# check condition bs
|
|
||||||
if conditioning is not None:
|
|
||||||
if isinstance(conditioning, dict):
|
|
||||||
try:
|
|
||||||
cbs = conditioning[list(conditioning.keys())[0]].shape[0]
|
|
||||||
except:
|
|
||||||
cbs = conditioning[list(conditioning.keys())[0]][0].shape[0]
|
|
||||||
|
|
||||||
if cbs != batch_size:
|
|
||||||
print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}")
|
|
||||||
else:
|
|
||||||
if conditioning.shape[0] != batch_size:
|
|
||||||
print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}")
|
|
||||||
|
|
||||||
# print('==> timestep_spacing: ', timestep_spacing, guidance_rescale)
|
|
||||||
self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, verbose=schedule_verbose)
|
|
||||||
|
|
||||||
# make shape
|
|
||||||
if len(shape) == 3:
|
|
||||||
C, H, W = shape
|
|
||||||
size = (batch_size, C, H, W)
|
|
||||||
elif len(shape) == 4:
|
|
||||||
C, T, H, W = shape
|
|
||||||
size = (batch_size, C, T, H, W)
|
|
||||||
# print(f'Data shape for DDIM sampling is {size}, eta {eta}')
|
|
||||||
|
|
||||||
samples, intermediates = self.ddim_sampling(conditioning, size,
|
|
||||||
callback=callback,
|
|
||||||
img_callback=img_callback,
|
|
||||||
quantize_denoised=quantize_x0,
|
|
||||||
mask=mask, x0=x0,
|
|
||||||
ddim_use_original_steps=False,
|
|
||||||
noise_dropout=noise_dropout,
|
|
||||||
temperature=temperature,
|
|
||||||
score_corrector=score_corrector,
|
|
||||||
corrector_kwargs=corrector_kwargs,
|
|
||||||
x_T=x_T,
|
|
||||||
log_every_t=log_every_t,
|
|
||||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
|
||||||
unconditional_conditioning=unconditional_conditioning,
|
|
||||||
verbose=verbose,
|
|
||||||
precision=precision,
|
|
||||||
fs=fs,
|
|
||||||
guidance_rescale=guidance_rescale,
|
|
||||||
**kwargs)
|
|
||||||
return samples, intermediates
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def ddim_sampling(self, cond, shape,
|
|
||||||
x_T=None, ddim_use_original_steps=False,
|
|
||||||
callback=None, timesteps=None, quantize_denoised=False,
|
|
||||||
mask=None, x0=None, img_callback=None, log_every_t=100,
|
|
||||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
|
||||||
unconditional_guidance_scale=1., unconditional_conditioning=None, verbose=True,precision=None,fs=None,guidance_rescale=0.0,
|
|
||||||
**kwargs):
|
|
||||||
device = self.model.betas.device
|
|
||||||
b = shape[0]
|
|
||||||
if x_T is None:
|
|
||||||
img = torch.randn(shape, device=device)
|
|
||||||
else:
|
|
||||||
img = x_T
|
|
||||||
if precision is not None:
|
|
||||||
if precision == 16:
|
|
||||||
img = img.to(dtype=torch.float16)
|
|
||||||
|
|
||||||
|
|
||||||
if timesteps is None:
|
|
||||||
timesteps = self.ddpm_num_timesteps if ddim_use_original_steps else self.ddim_timesteps
|
|
||||||
elif timesteps is not None and not ddim_use_original_steps:
|
|
||||||
subset_end = int(min(timesteps / self.ddim_timesteps.shape[0], 1) * self.ddim_timesteps.shape[0]) - 1
|
|
||||||
timesteps = self.ddim_timesteps[:subset_end]
|
|
||||||
|
|
||||||
intermediates = {'x_inter': [img], 'pred_x0': [img]}
|
|
||||||
time_range = reversed(range(0,timesteps)) if ddim_use_original_steps else np.flip(timesteps)
|
|
||||||
total_steps = timesteps if ddim_use_original_steps else timesteps.shape[0]
|
|
||||||
if verbose:
|
|
||||||
iterator = tqdm(time_range, desc='DDIM Sampler', total=total_steps)
|
|
||||||
else:
|
|
||||||
iterator = time_range
|
|
||||||
|
|
||||||
clean_cond = kwargs.pop("clean_cond", False)
|
|
||||||
|
|
||||||
# cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning)
|
|
||||||
for i, step in enumerate(iterator):
|
|
||||||
index = total_steps - i - 1
|
|
||||||
ts = torch.full((b,), step, device=device, dtype=torch.long)
|
|
||||||
|
|
||||||
## use mask to blend noised original latent (img_orig) & new sampled latent (img)
|
|
||||||
if mask is not None:
|
|
||||||
assert x0 is not None
|
|
||||||
if clean_cond:
|
|
||||||
img_orig = x0
|
|
||||||
else:
|
|
||||||
img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? <ddim inversion>
|
|
||||||
img = img_orig * mask + (1. - mask) * img # keep original & modify use img
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps,
|
|
||||||
quantize_denoised=quantize_denoised, temperature=temperature,
|
|
||||||
noise_dropout=noise_dropout, score_corrector=score_corrector,
|
|
||||||
corrector_kwargs=corrector_kwargs,
|
|
||||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
|
||||||
unconditional_conditioning=unconditional_conditioning,
|
|
||||||
mask=mask,x0=x0,fs=fs,guidance_rescale=guidance_rescale,
|
|
||||||
**kwargs)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
img, pred_x0 = outs
|
|
||||||
if callback: callback(i)
|
|
||||||
if img_callback: img_callback(pred_x0, i)
|
|
||||||
|
|
||||||
if index % log_every_t == 0 or index == total_steps - 1:
|
|
||||||
intermediates['x_inter'].append(img)
|
|
||||||
intermediates['pred_x0'].append(pred_x0)
|
|
||||||
|
|
||||||
return img, intermediates
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def p_sample_ddim(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
|
|
||||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
|
||||||
unconditional_guidance_scale=1., unconditional_conditioning=None,
|
|
||||||
uc_type=None, cfg_img=None,mask=None,x0=None,guidance_rescale=0.0, **kwargs):
|
|
||||||
b, *_, device = *x.shape, x.device
|
|
||||||
if x.dim() == 5:
|
|
||||||
is_video = True
|
|
||||||
else:
|
|
||||||
is_video = False
|
|
||||||
if cfg_img is None:
|
|
||||||
cfg_img = unconditional_guidance_scale
|
|
||||||
|
|
||||||
unconditional_conditioning_img_nonetext = kwargs['unconditional_conditioning_img_nonetext']
|
|
||||||
|
|
||||||
|
|
||||||
if unconditional_conditioning is None or unconditional_guidance_scale == 1.:
|
|
||||||
model_output = self.model.apply_model(x, t, c, **kwargs) # unet denoiser
|
|
||||||
else:
|
|
||||||
### with unconditional condition
|
|
||||||
e_t_cond = self.model.apply_model(x, t, c, **kwargs)
|
|
||||||
e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs)
|
|
||||||
e_t_uncond_img = self.model.apply_model(x, t, unconditional_conditioning_img_nonetext, **kwargs)
|
|
||||||
# text cfg
|
|
||||||
model_output = e_t_uncond + cfg_img * (e_t_uncond_img - e_t_uncond) + unconditional_guidance_scale * (e_t_cond - e_t_uncond_img)
|
|
||||||
if guidance_rescale > 0.0:
|
|
||||||
model_output = rescale_noise_cfg(model_output, e_t_cond, guidance_rescale=guidance_rescale)
|
|
||||||
|
|
||||||
if self.model.parameterization == "v":
|
|
||||||
e_t = self.model.predict_eps_from_z_and_v(x, t, model_output)
|
|
||||||
else:
|
|
||||||
e_t = model_output
|
|
||||||
|
|
||||||
if score_corrector is not None:
|
|
||||||
assert self.model.parameterization == "eps", 'not implemented'
|
|
||||||
e_t = score_corrector.modify_score(self.model, e_t, x, t, c, **corrector_kwargs)
|
|
||||||
|
|
||||||
alphas = self.model.alphas_cumprod if use_original_steps else self.ddim_alphas
|
|
||||||
alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev
|
|
||||||
sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas
|
|
||||||
sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
|
||||||
# select parameters corresponding to the currently considered timestep
|
|
||||||
|
|
||||||
if is_video:
|
|
||||||
size = (b, 1, 1, 1, 1)
|
|
||||||
else:
|
|
||||||
size = (b, 1, 1, 1)
|
|
||||||
a_t = torch.full(size, alphas[index], device=device)
|
|
||||||
a_prev = torch.full(size, alphas_prev[index], device=device)
|
|
||||||
sigma_t = torch.full(size, sigmas[index], device=device)
|
|
||||||
sqrt_one_minus_at = torch.full(size, sqrt_one_minus_alphas[index],device=device)
|
|
||||||
|
|
||||||
# current prediction for x_0
|
|
||||||
if self.model.parameterization != "v":
|
|
||||||
pred_x0 = (x - sqrt_one_minus_at * e_t) / a_t.sqrt()
|
|
||||||
else:
|
|
||||||
pred_x0 = self.model.predict_start_from_z_and_v(x, t, model_output)
|
|
||||||
|
|
||||||
if self.model.use_dynamic_rescale:
|
|
||||||
scale_t = torch.full(size, self.ddim_scale_arr[index], device=device)
|
|
||||||
prev_scale_t = torch.full(size, self.ddim_scale_arr_prev[index], device=device)
|
|
||||||
rescale = (prev_scale_t / scale_t)
|
|
||||||
pred_x0 *= rescale
|
|
||||||
|
|
||||||
if quantize_denoised:
|
|
||||||
pred_x0, _, *_ = self.model.first_stage_model.quantize(pred_x0)
|
|
||||||
# direction pointing to x_t
|
|
||||||
dir_xt = (1. - a_prev - sigma_t**2).sqrt() * e_t
|
|
||||||
|
|
||||||
noise = sigma_t * noise_like(x.shape, device, repeat_noise) * temperature
|
|
||||||
if noise_dropout > 0.:
|
|
||||||
noise = torch.nn.functional.dropout(noise, p=noise_dropout)
|
|
||||||
|
|
||||||
x_prev = a_prev.sqrt() * pred_x0 + dir_xt + noise
|
|
||||||
|
|
||||||
return x_prev, pred_x0
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def decode(self, x_latent, cond, t_start, unconditional_guidance_scale=1.0, unconditional_conditioning=None,
|
|
||||||
use_original_steps=False, callback=None):
|
|
||||||
|
|
||||||
timesteps = np.arange(self.ddpm_num_timesteps) if use_original_steps else self.ddim_timesteps
|
|
||||||
timesteps = timesteps[:t_start]
|
|
||||||
|
|
||||||
time_range = np.flip(timesteps)
|
|
||||||
total_steps = timesteps.shape[0]
|
|
||||||
print(f"Running DDIM Sampling with {total_steps} timesteps")
|
|
||||||
|
|
||||||
iterator = tqdm(time_range, desc='Decoding image', total=total_steps)
|
|
||||||
x_dec = x_latent
|
|
||||||
for i, step in enumerate(iterator):
|
|
||||||
index = total_steps - i - 1
|
|
||||||
ts = torch.full((x_latent.shape[0],), step, device=x_latent.device, dtype=torch.long)
|
|
||||||
x_dec, _ = self.p_sample_ddim(x_dec, cond, ts, index=index, use_original_steps=use_original_steps,
|
|
||||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
|
||||||
unconditional_conditioning=unconditional_conditioning)
|
|
||||||
if callback: callback(i)
|
|
||||||
return x_dec
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def stochastic_encode(self, x0, t, use_original_steps=False, noise=None):
|
|
||||||
# fast, but does not allow for exact reconstruction
|
|
||||||
# t serves as an index to gather the correct alphas
|
|
||||||
if use_original_steps:
|
|
||||||
sqrt_alphas_cumprod = self.sqrt_alphas_cumprod
|
|
||||||
sqrt_one_minus_alphas_cumprod = self.sqrt_one_minus_alphas_cumprod
|
|
||||||
else:
|
|
||||||
sqrt_alphas_cumprod = torch.sqrt(self.ddim_alphas)
|
|
||||||
sqrt_one_minus_alphas_cumprod = self.ddim_sqrt_one_minus_alphas
|
|
||||||
|
|
||||||
if noise is None:
|
|
||||||
noise = torch.randn_like(x0)
|
|
||||||
return (extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
|
|
||||||
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) * noise)
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
from .sampler import UniPCSampler
|
|
||||||
@@ -1,79 +0,0 @@
|
|||||||
"""SAMPLING ONLY."""
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from .uni_pc import NoiseScheduleVP, model_wrapper, UniPC
|
|
||||||
|
|
||||||
class UniPCSampler(object):
|
|
||||||
def __init__(self, model, **kwargs):
|
|
||||||
super().__init__()
|
|
||||||
self.model = model
|
|
||||||
to_torch = lambda x: x.clone().detach().to(torch.float32).to(model.device)
|
|
||||||
self.register_buffer('alphas_cumprod', to_torch(model.alphas_cumprod))
|
|
||||||
|
|
||||||
def register_buffer(self, name, attr):
|
|
||||||
if type(attr) == torch.Tensor:
|
|
||||||
if attr.device != torch.device("cuda"):
|
|
||||||
attr = attr.to(torch.device("cuda"))
|
|
||||||
setattr(self, name, attr)
|
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def sample(self,
|
|
||||||
S,
|
|
||||||
batch_size,
|
|
||||||
shape,
|
|
||||||
conditioning=None,
|
|
||||||
callback=None,
|
|
||||||
normals_sequence=None,
|
|
||||||
img_callback=None,
|
|
||||||
quantize_x0=False,
|
|
||||||
eta=0.,
|
|
||||||
mask=None,
|
|
||||||
x0=None,
|
|
||||||
temperature=1.,
|
|
||||||
noise_dropout=0.,
|
|
||||||
score_corrector=None,
|
|
||||||
corrector_kwargs=None,
|
|
||||||
verbose=True,
|
|
||||||
x_T=None,
|
|
||||||
log_every_t=100,
|
|
||||||
unconditional_guidance_scale=1.,
|
|
||||||
unconditional_conditioning=None,
|
|
||||||
# this has to come in the same format as the conditioning, # e.g. as encoded tokens, ...
|
|
||||||
**kwargs
|
|
||||||
):
|
|
||||||
if conditioning is not None:
|
|
||||||
if isinstance(conditioning, dict):
|
|
||||||
cbs = conditioning[list(conditioning.keys())[0]].shape[0]
|
|
||||||
if cbs != batch_size:
|
|
||||||
print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}")
|
|
||||||
else:
|
|
||||||
if conditioning.shape[0] != batch_size:
|
|
||||||
print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}")
|
|
||||||
|
|
||||||
# sampling
|
|
||||||
C, F, H, W = shape
|
|
||||||
size = (batch_size, C, H, W)
|
|
||||||
|
|
||||||
device = self.model.betas.device
|
|
||||||
if x_T is None:
|
|
||||||
img = torch.randn(size, device=device)
|
|
||||||
else:
|
|
||||||
img = x_T
|
|
||||||
|
|
||||||
ns = NoiseScheduleVP('discrete', alphas_cumprod=self.alphas_cumprod)
|
|
||||||
|
|
||||||
model_fn = model_wrapper(
|
|
||||||
lambda x, t, c: self.model.apply_model(x, t, c),
|
|
||||||
ns,
|
|
||||||
model_type="noise",
|
|
||||||
guidance_type="classifier-free",
|
|
||||||
condition=conditioning,
|
|
||||||
unconditional_condition=unconditional_conditioning,
|
|
||||||
guidance_scale=unconditional_guidance_scale,
|
|
||||||
)
|
|
||||||
|
|
||||||
uni_pc = UniPC(model_fn, ns, predict_x0=True, thresholding=False)
|
|
||||||
x = uni_pc.sample(img, steps=S, skip_type="time_uniform", method="multistep", order=3, lower_order_final=True)
|
|
||||||
|
|
||||||
return x.to(device), None
|
|
||||||
@@ -1,808 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
import math
|
|
||||||
|
|
||||||
|
|
||||||
class NoiseScheduleVP:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
schedule='discrete',
|
|
||||||
betas=None,
|
|
||||||
alphas_cumprod=None,
|
|
||||||
continuous_beta_0=0.1,
|
|
||||||
continuous_beta_1=20.,
|
|
||||||
):
|
|
||||||
"""Create a wrapper class for the forward SDE (VP type).
|
|
||||||
|
|
||||||
***
|
|
||||||
Update: We support discrete-time diffusion models by implementing a picewise linear interpolation for log_alpha_t.
|
|
||||||
We recommend to use schedule='discrete' for the discrete-time diffusion models, especially for high-resolution images.
|
|
||||||
***
|
|
||||||
|
|
||||||
The forward SDE ensures that the condition distribution q_{t|0}(x_t | x_0) = N ( alpha_t * x_0, sigma_t^2 * I ).
|
|
||||||
We further define lambda_t = log(alpha_t) - log(sigma_t), which is the half-logSNR (described in the DPM-Solver paper).
|
|
||||||
Therefore, we implement the functions for computing alpha_t, sigma_t and lambda_t. For t in [0, T], we have:
|
|
||||||
|
|
||||||
log_alpha_t = self.marginal_log_mean_coeff(t)
|
|
||||||
sigma_t = self.marginal_std(t)
|
|
||||||
lambda_t = self.marginal_lambda(t)
|
|
||||||
|
|
||||||
Moreover, as lambda(t) is an invertible function, we also support its inverse function:
|
|
||||||
|
|
||||||
t = self.inverse_lambda(lambda_t)
|
|
||||||
|
|
||||||
===============================================================
|
|
||||||
|
|
||||||
We support both discrete-time DPMs (trained on n = 0, 1, ..., N-1) and continuous-time DPMs (trained on t in [t_0, T]).
|
|
||||||
|
|
||||||
1. For discrete-time DPMs:
|
|
||||||
|
|
||||||
For discrete-time DPMs trained on n = 0, 1, ..., N-1, we convert the discrete steps to continuous time steps by:
|
|
||||||
t_i = (i + 1) / N
|
|
||||||
e.g. for N = 1000, we have t_0 = 1e-3 and T = t_{N-1} = 1.
|
|
||||||
We solve the corresponding diffusion ODE from time T = 1 to time t_0 = 1e-3.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
betas: A `torch.Tensor`. The beta array for the discrete-time DPM. (See the original DDPM paper for details)
|
|
||||||
alphas_cumprod: A `torch.Tensor`. The cumprod alphas for the discrete-time DPM. (See the original DDPM paper for details)
|
|
||||||
|
|
||||||
Note that we always have alphas_cumprod = cumprod(betas). Therefore, we only need to set one of `betas` and `alphas_cumprod`.
|
|
||||||
|
|
||||||
**Important**: Please pay special attention for the args for `alphas_cumprod`:
|
|
||||||
The `alphas_cumprod` is the \hat{alpha_n} arrays in the notations of DDPM. Specifically, DDPMs assume that
|
|
||||||
q_{t_n | 0}(x_{t_n} | x_0) = N ( \sqrt{\hat{alpha_n}} * x_0, (1 - \hat{alpha_n}) * I ).
|
|
||||||
Therefore, the notation \hat{alpha_n} is different from the notation alpha_t in DPM-Solver. In fact, we have
|
|
||||||
alpha_{t_n} = \sqrt{\hat{alpha_n}},
|
|
||||||
and
|
|
||||||
log(alpha_{t_n}) = 0.5 * log(\hat{alpha_n}).
|
|
||||||
|
|
||||||
|
|
||||||
2. For continuous-time DPMs:
|
|
||||||
|
|
||||||
We support two types of VPSDEs: linear (DDPM) and cosine (improved-DDPM). The hyperparameters for the noise
|
|
||||||
schedule are the default settings in DDPM and improved-DDPM:
|
|
||||||
|
|
||||||
Args:
|
|
||||||
beta_min: A `float` number. The smallest beta for the linear schedule.
|
|
||||||
beta_max: A `float` number. The largest beta for the linear schedule.
|
|
||||||
cosine_s: A `float` number. The hyperparameter in the cosine schedule.
|
|
||||||
cosine_beta_max: A `float` number. The hyperparameter in the cosine schedule.
|
|
||||||
T: A `float` number. The ending time of the forward process.
|
|
||||||
|
|
||||||
===============================================================
|
|
||||||
|
|
||||||
Args:
|
|
||||||
schedule: A `str`. The noise schedule of the forward SDE. 'discrete' for discrete-time DPMs,
|
|
||||||
'linear' or 'cosine' for continuous-time DPMs.
|
|
||||||
Returns:
|
|
||||||
A wrapper object of the forward SDE (VP type).
|
|
||||||
|
|
||||||
===============================================================
|
|
||||||
|
|
||||||
Example:
|
|
||||||
|
|
||||||
# For discrete-time DPMs, given betas (the beta array for n = 0, 1, ..., N - 1):
|
|
||||||
>>> ns = NoiseScheduleVP('discrete', betas=betas)
|
|
||||||
|
|
||||||
# For discrete-time DPMs, given alphas_cumprod (the \hat{alpha_n} array for n = 0, 1, ..., N - 1):
|
|
||||||
>>> ns = NoiseScheduleVP('discrete', alphas_cumprod=alphas_cumprod)
|
|
||||||
|
|
||||||
# For continuous-time DPMs (VPSDE), linear schedule:
|
|
||||||
>>> ns = NoiseScheduleVP('linear', continuous_beta_0=0.1, continuous_beta_1=20.)
|
|
||||||
|
|
||||||
"""
|
|
||||||
|
|
||||||
if schedule not in ['discrete', 'linear', 'cosine']:
|
|
||||||
raise ValueError("Unsupported noise schedule {}. The schedule needs to be 'discrete' or 'linear' or 'cosine'".format(schedule))
|
|
||||||
|
|
||||||
self.schedule = schedule
|
|
||||||
if schedule == 'discrete':
|
|
||||||
if betas is not None:
|
|
||||||
log_alphas = 0.5 * torch.log(1 - betas).cumsum(dim=0)
|
|
||||||
else:
|
|
||||||
assert alphas_cumprod is not None
|
|
||||||
log_alphas = 0.5 * torch.log(alphas_cumprod)
|
|
||||||
self.total_N = len(log_alphas)
|
|
||||||
self.T = 1.
|
|
||||||
self.t_array = torch.linspace(0., 1., self.total_N + 1)[1:].reshape((1, -1))
|
|
||||||
self.log_alpha_array = log_alphas.reshape((1, -1,))
|
|
||||||
else:
|
|
||||||
self.total_N = 1000
|
|
||||||
self.beta_0 = continuous_beta_0
|
|
||||||
self.beta_1 = continuous_beta_1
|
|
||||||
self.cosine_s = 0.008
|
|
||||||
self.cosine_beta_max = 999.
|
|
||||||
self.cosine_t_max = math.atan(self.cosine_beta_max * (1. + self.cosine_s) / math.pi) * 2. * (1. + self.cosine_s) / math.pi - self.cosine_s
|
|
||||||
self.cosine_log_alpha_0 = math.log(math.cos(self.cosine_s / (1. + self.cosine_s) * math.pi / 2.))
|
|
||||||
self.schedule = schedule
|
|
||||||
if schedule == 'cosine':
|
|
||||||
# For the cosine schedule, T = 1 will have numerical issues. So we manually set the ending time T.
|
|
||||||
# Note that T = 0.9946 may be not the optimal setting. However, we find it works well.
|
|
||||||
self.T = 0.9946
|
|
||||||
else:
|
|
||||||
self.T = 1.
|
|
||||||
|
|
||||||
def marginal_log_mean_coeff(self, t):
|
|
||||||
"""
|
|
||||||
Compute log(alpha_t) of a given continuous-time label t in [0, T].
|
|
||||||
"""
|
|
||||||
if self.schedule == 'discrete':
|
|
||||||
return interpolate_fn(t.reshape((-1, 1)), self.t_array.to(t.device), self.log_alpha_array.to(t.device)).reshape((-1))
|
|
||||||
elif self.schedule == 'linear':
|
|
||||||
return -0.25 * t ** 2 * (self.beta_1 - self.beta_0) - 0.5 * t * self.beta_0
|
|
||||||
elif self.schedule == 'cosine':
|
|
||||||
log_alpha_fn = lambda s: torch.log(torch.cos((s + self.cosine_s) / (1. + self.cosine_s) * math.pi / 2.))
|
|
||||||
log_alpha_t = log_alpha_fn(t) - self.cosine_log_alpha_0
|
|
||||||
return log_alpha_t
|
|
||||||
|
|
||||||
def marginal_alpha(self, t):
|
|
||||||
"""
|
|
||||||
Compute alpha_t of a given continuous-time label t in [0, T].
|
|
||||||
"""
|
|
||||||
return torch.exp(self.marginal_log_mean_coeff(t))
|
|
||||||
|
|
||||||
def marginal_std(self, t):
|
|
||||||
"""
|
|
||||||
Compute sigma_t of a given continuous-time label t in [0, T].
|
|
||||||
"""
|
|
||||||
return torch.sqrt(1. - torch.exp(2. * self.marginal_log_mean_coeff(t)))
|
|
||||||
|
|
||||||
def marginal_lambda(self, t):
|
|
||||||
"""
|
|
||||||
Compute lambda_t = log(alpha_t) - log(sigma_t) of a given continuous-time label t in [0, T].
|
|
||||||
"""
|
|
||||||
log_mean_coeff = self.marginal_log_mean_coeff(t)
|
|
||||||
log_std = 0.5 * torch.log(1. - torch.exp(2. * log_mean_coeff))
|
|
||||||
return log_mean_coeff - log_std
|
|
||||||
|
|
||||||
def inverse_lambda(self, lamb):
|
|
||||||
"""
|
|
||||||
Compute the continuous-time label t in [0, T] of a given half-logSNR lambda_t.
|
|
||||||
"""
|
|
||||||
if self.schedule == 'linear':
|
|
||||||
tmp = 2. * (self.beta_1 - self.beta_0) * torch.logaddexp(-2. * lamb, torch.zeros((1,)).to(lamb))
|
|
||||||
Delta = self.beta_0**2 + tmp
|
|
||||||
return tmp / (torch.sqrt(Delta) + self.beta_0) / (self.beta_1 - self.beta_0)
|
|
||||||
elif self.schedule == 'discrete':
|
|
||||||
log_alpha = -0.5 * torch.logaddexp(torch.zeros((1,)).to(lamb.device), -2. * lamb)
|
|
||||||
t = interpolate_fn(log_alpha.reshape((-1, 1)), torch.flip(self.log_alpha_array.to(lamb.device), [1]), torch.flip(self.t_array.to(lamb.device), [1]))
|
|
||||||
return t.reshape((-1,))
|
|
||||||
else:
|
|
||||||
log_alpha = -0.5 * torch.logaddexp(-2. * lamb, torch.zeros((1,)).to(lamb))
|
|
||||||
t_fn = lambda log_alpha_t: torch.arccos(torch.exp(log_alpha_t + self.cosine_log_alpha_0)) * 2. * (1. + self.cosine_s) / math.pi - self.cosine_s
|
|
||||||
t = t_fn(log_alpha)
|
|
||||||
return t
|
|
||||||
|
|
||||||
|
|
||||||
def model_wrapper(
|
|
||||||
model,
|
|
||||||
noise_schedule,
|
|
||||||
model_type="noise",
|
|
||||||
model_kwargs={},
|
|
||||||
guidance_type="uncond",
|
|
||||||
condition=None,
|
|
||||||
unconditional_condition=None,
|
|
||||||
guidance_scale=1.,
|
|
||||||
classifier_fn=None,
|
|
||||||
classifier_kwargs={},
|
|
||||||
):
|
|
||||||
"""Create a wrapper function for the noise prediction model.
|
|
||||||
|
|
||||||
DPM-Solver needs to solve the continuous-time diffusion ODEs. For DPMs trained on discrete-time labels, we need to
|
|
||||||
firstly wrap the model function to a noise prediction model that accepts the continuous time as the input.
|
|
||||||
|
|
||||||
We support four types of the diffusion model by setting `model_type`:
|
|
||||||
|
|
||||||
1. "noise": noise prediction model. (Trained by predicting noise).
|
|
||||||
|
|
||||||
2. "x_start": data prediction model. (Trained by predicting the data x_0 at time 0).
|
|
||||||
|
|
||||||
3. "v": velocity prediction model. (Trained by predicting the velocity).
|
|
||||||
The "v" prediction is derivation detailed in Appendix D of [1], and is used in Imagen-Video [2].
|
|
||||||
|
|
||||||
[1] Salimans, Tim, and Jonathan Ho. "Progressive distillation for fast sampling of diffusion models."
|
|
||||||
arXiv preprint arXiv:2202.00512 (2022).
|
|
||||||
[2] Ho, Jonathan, et al. "Imagen Video: High Definition Video Generation with Diffusion Models."
|
|
||||||
arXiv preprint arXiv:2210.02303 (2022).
|
|
||||||
|
|
||||||
4. "score": marginal score function. (Trained by denoising score matching).
|
|
||||||
Note that the score function and the noise prediction model follows a simple relationship:
|
|
||||||
```
|
|
||||||
noise(x_t, t) = -sigma_t * score(x_t, t)
|
|
||||||
```
|
|
||||||
|
|
||||||
We support three types of guided sampling by DPMs by setting `guidance_type`:
|
|
||||||
1. "uncond": unconditional sampling by DPMs.
|
|
||||||
The input `model` has the following format:
|
|
||||||
``
|
|
||||||
model(x, t_input, **model_kwargs) -> noise | x_start | v | score
|
|
||||||
``
|
|
||||||
|
|
||||||
2. "classifier": classifier guidance sampling [3] by DPMs and another classifier.
|
|
||||||
The input `model` has the following format:
|
|
||||||
``
|
|
||||||
model(x, t_input, **model_kwargs) -> noise | x_start | v | score
|
|
||||||
``
|
|
||||||
|
|
||||||
The input `classifier_fn` has the following format:
|
|
||||||
``
|
|
||||||
classifier_fn(x, t_input, cond, **classifier_kwargs) -> logits(x, t_input, cond)
|
|
||||||
``
|
|
||||||
|
|
||||||
[3] P. Dhariwal and A. Q. Nichol, "Diffusion models beat GANs on image synthesis,"
|
|
||||||
in Advances in Neural Information Processing Systems, vol. 34, 2021, pp. 8780-8794.
|
|
||||||
|
|
||||||
3. "classifier-free": classifier-free guidance sampling by conditional DPMs.
|
|
||||||
The input `model` has the following format:
|
|
||||||
``
|
|
||||||
model(x, t_input, cond, **model_kwargs) -> noise | x_start | v | score
|
|
||||||
``
|
|
||||||
And if cond == `unconditional_condition`, the model output is the unconditional DPM output.
|
|
||||||
|
|
||||||
[4] Ho, Jonathan, and Tim Salimans. "Classifier-free diffusion guidance."
|
|
||||||
arXiv preprint arXiv:2207.12598 (2022).
|
|
||||||
|
|
||||||
|
|
||||||
The `t_input` is the time label of the model, which may be discrete-time labels (i.e. 0 to 999)
|
|
||||||
or continuous-time labels (i.e. epsilon to T).
|
|
||||||
|
|
||||||
We wrap the model function to accept only `x` and `t_continuous` as inputs, and outputs the predicted noise:
|
|
||||||
``
|
|
||||||
def model_fn(x, t_continuous) -> noise:
|
|
||||||
t_input = get_model_input_time(t_continuous)
|
|
||||||
return noise_pred(model, x, t_input, **model_kwargs)
|
|
||||||
``
|
|
||||||
where `t_continuous` is the continuous time labels (i.e. epsilon to T). And we use `model_fn` for DPM-Solver.
|
|
||||||
|
|
||||||
===============================================================
|
|
||||||
|
|
||||||
Args:
|
|
||||||
model: A diffusion model with the corresponding format described above.
|
|
||||||
noise_schedule: A noise schedule object, such as NoiseScheduleVP.
|
|
||||||
model_type: A `str`. The parameterization type of the diffusion model.
|
|
||||||
"noise" or "x_start" or "v" or "score".
|
|
||||||
model_kwargs: A `dict`. A dict for the other inputs of the model function.
|
|
||||||
guidance_type: A `str`. The type of the guidance for sampling.
|
|
||||||
"uncond" or "classifier" or "classifier-free".
|
|
||||||
condition: A pytorch tensor. The condition for the guided sampling.
|
|
||||||
Only used for "classifier" or "classifier-free" guidance type.
|
|
||||||
unconditional_condition: A pytorch tensor. The condition for the unconditional sampling.
|
|
||||||
Only used for "classifier-free" guidance type.
|
|
||||||
guidance_scale: A `float`. The scale for the guided sampling.
|
|
||||||
classifier_fn: A classifier function. Only used for the classifier guidance.
|
|
||||||
classifier_kwargs: A `dict`. A dict for the other inputs of the classifier function.
|
|
||||||
Returns:
|
|
||||||
A noise prediction model that accepts the noised data and the continuous time as the inputs.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def get_model_input_time(t_continuous):
|
|
||||||
"""
|
|
||||||
Convert the continuous-time `t_continuous` (in [epsilon, T]) to the model input time.
|
|
||||||
For discrete-time DPMs, we convert `t_continuous` in [1 / N, 1] to `t_input` in [0, 1000 * (N - 1) / N].
|
|
||||||
For continuous-time DPMs, we just use `t_continuous`.
|
|
||||||
"""
|
|
||||||
if noise_schedule.schedule == 'discrete':
|
|
||||||
return (t_continuous - 1. / noise_schedule.total_N) * 1000.
|
|
||||||
else:
|
|
||||||
return t_continuous
|
|
||||||
|
|
||||||
def noise_pred_fn(x, t_continuous, cond=None):
|
|
||||||
if t_continuous.reshape((-1,)).shape[0] == 1:
|
|
||||||
t_continuous = t_continuous.expand((x.shape[0]))
|
|
||||||
t_input = get_model_input_time(t_continuous)
|
|
||||||
if cond is None:
|
|
||||||
output = model(x, t_input, None, **model_kwargs)
|
|
||||||
else:
|
|
||||||
output = model(x, t_input, cond, **model_kwargs)
|
|
||||||
if model_type == "noise":
|
|
||||||
return output
|
|
||||||
elif model_type == "x_start":
|
|
||||||
alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous)
|
|
||||||
dims = x.dim()
|
|
||||||
return (x - expand_dims(alpha_t, dims) * output) / expand_dims(sigma_t, dims)
|
|
||||||
elif model_type == "v":
|
|
||||||
alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous)
|
|
||||||
dims = x.dim()
|
|
||||||
return expand_dims(alpha_t, dims) * output + expand_dims(sigma_t, dims) * x
|
|
||||||
elif model_type == "score":
|
|
||||||
sigma_t = noise_schedule.marginal_std(t_continuous)
|
|
||||||
dims = x.dim()
|
|
||||||
return -expand_dims(sigma_t, dims) * output
|
|
||||||
|
|
||||||
def cond_grad_fn(x, t_input):
|
|
||||||
"""
|
|
||||||
Compute the gradient of the classifier, i.e. nabla_{x} log p_t(cond | x_t).
|
|
||||||
"""
|
|
||||||
with torch.enable_grad():
|
|
||||||
x_in = x.detach().requires_grad_(True)
|
|
||||||
log_prob = classifier_fn(x_in, t_input, condition, **classifier_kwargs)
|
|
||||||
return torch.autograd.grad(log_prob.sum(), x_in)[0]
|
|
||||||
|
|
||||||
def model_fn(x, t_continuous):
|
|
||||||
"""
|
|
||||||
The noise predicition model function that is used for DPM-Solver.
|
|
||||||
"""
|
|
||||||
if t_continuous.reshape((-1,)).shape[0] == 1:
|
|
||||||
t_continuous = t_continuous.expand((x.shape[0]))
|
|
||||||
if guidance_type == "uncond":
|
|
||||||
return noise_pred_fn(x, t_continuous)
|
|
||||||
elif guidance_type == "classifier":
|
|
||||||
assert classifier_fn is not None
|
|
||||||
t_input = get_model_input_time(t_continuous)
|
|
||||||
cond_grad = cond_grad_fn(x, t_input)
|
|
||||||
sigma_t = noise_schedule.marginal_std(t_continuous)
|
|
||||||
noise = noise_pred_fn(x, t_continuous)
|
|
||||||
return noise - guidance_scale * expand_dims(sigma_t, dims=cond_grad.dim()) * cond_grad
|
|
||||||
elif guidance_type == "classifier-free":
|
|
||||||
if guidance_scale == 1. or unconditional_condition is None:
|
|
||||||
return noise_pred_fn(x, t_continuous, cond=condition)
|
|
||||||
else:
|
|
||||||
x_in = torch.cat([x] * 2)
|
|
||||||
t_in = torch.cat([t_continuous] * 2)
|
|
||||||
c_in = torch.cat([unconditional_condition, condition])
|
|
||||||
noise_uncond, noise = noise_pred_fn(x_in, t_in, cond=c_in).chunk(2)
|
|
||||||
return noise_uncond + guidance_scale * (noise - noise_uncond)
|
|
||||||
|
|
||||||
assert model_type in ["noise", "x_start", "v"]
|
|
||||||
assert guidance_type in ["uncond", "classifier", "classifier-free"]
|
|
||||||
return model_fn
|
|
||||||
|
|
||||||
|
|
||||||
class UniPC:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
model_fn,
|
|
||||||
noise_schedule,
|
|
||||||
predict_x0=True,
|
|
||||||
thresholding=False,
|
|
||||||
max_val=1.,
|
|
||||||
variant='bh1'
|
|
||||||
):
|
|
||||||
"""Construct a UniPC.
|
|
||||||
|
|
||||||
We support both data_prediction and noise_prediction.
|
|
||||||
"""
|
|
||||||
self.model = model_fn
|
|
||||||
self.noise_schedule = noise_schedule
|
|
||||||
self.variant = variant
|
|
||||||
self.predict_x0 = predict_x0
|
|
||||||
self.thresholding = thresholding
|
|
||||||
self.max_val = max_val
|
|
||||||
|
|
||||||
def dynamic_thresholding_fn(self, x0, t=None):
|
|
||||||
"""
|
|
||||||
The dynamic thresholding method.
|
|
||||||
"""
|
|
||||||
dims = x0.dim()
|
|
||||||
p = self.dynamic_thresholding_ratio
|
|
||||||
s = torch.quantile(torch.abs(x0).reshape((x0.shape[0], -1)), p, dim=1)
|
|
||||||
s = expand_dims(torch.maximum(s, self.thresholding_max_val * torch.ones_like(s).to(s.device)), dims)
|
|
||||||
x0 = torch.clamp(x0, -s, s) / s
|
|
||||||
return x0
|
|
||||||
|
|
||||||
def noise_prediction_fn(self, x, t):
|
|
||||||
"""
|
|
||||||
Return the noise prediction model.
|
|
||||||
"""
|
|
||||||
return self.model(x, t)
|
|
||||||
|
|
||||||
def data_prediction_fn(self, x, t):
|
|
||||||
"""
|
|
||||||
Return the data prediction model (with thresholding).
|
|
||||||
"""
|
|
||||||
noise = self.noise_prediction_fn(x, t)
|
|
||||||
dims = x.dim()
|
|
||||||
alpha_t, sigma_t = self.noise_schedule.marginal_alpha(t), self.noise_schedule.marginal_std(t)
|
|
||||||
x0 = (x - expand_dims(sigma_t, dims) * noise) / expand_dims(alpha_t, dims)
|
|
||||||
if self.thresholding:
|
|
||||||
p = 0.995 # A hyperparameter in the paper of "Imagen" [1].
|
|
||||||
s = torch.quantile(torch.abs(x0).reshape((x0.shape[0], -1)), p, dim=1)
|
|
||||||
s = expand_dims(torch.maximum(s, self.max_val * torch.ones_like(s).to(s.device)), dims)
|
|
||||||
x0 = torch.clamp(x0, -s, s) / s
|
|
||||||
return x0
|
|
||||||
|
|
||||||
def model_fn(self, x, t):
|
|
||||||
"""
|
|
||||||
Convert the model to the noise prediction model or the data prediction model.
|
|
||||||
"""
|
|
||||||
if self.predict_x0:
|
|
||||||
return self.data_prediction_fn(x, t)
|
|
||||||
else:
|
|
||||||
return self.noise_prediction_fn(x, t)
|
|
||||||
|
|
||||||
def get_time_steps(self, skip_type, t_T, t_0, N, device):
|
|
||||||
"""Compute the intermediate time steps for sampling.
|
|
||||||
"""
|
|
||||||
if skip_type == 'logSNR':
|
|
||||||
lambda_T = self.noise_schedule.marginal_lambda(torch.tensor(t_T).to(device))
|
|
||||||
lambda_0 = self.noise_schedule.marginal_lambda(torch.tensor(t_0).to(device))
|
|
||||||
logSNR_steps = torch.linspace(lambda_T.cpu().item(), lambda_0.cpu().item(), N + 1).to(device)
|
|
||||||
return self.noise_schedule.inverse_lambda(logSNR_steps)
|
|
||||||
elif skip_type == 'time_uniform':
|
|
||||||
return torch.linspace(t_T, t_0, N + 1).to(device)
|
|
||||||
elif skip_type == 'time_quadratic':
|
|
||||||
t_order = 2
|
|
||||||
t = torch.linspace(t_T**(1. / t_order), t_0**(1. / t_order), N + 1).pow(t_order).to(device)
|
|
||||||
return t
|
|
||||||
else:
|
|
||||||
raise ValueError("Unsupported skip_type {}, need to be 'logSNR' or 'time_uniform' or 'time_quadratic'".format(skip_type))
|
|
||||||
|
|
||||||
def get_orders_and_timesteps_for_singlestep_solver(self, steps, order, skip_type, t_T, t_0, device):
|
|
||||||
"""
|
|
||||||
Get the order of each step for sampling by the singlestep DPM-Solver.
|
|
||||||
"""
|
|
||||||
if order == 3:
|
|
||||||
K = steps // 3 + 1
|
|
||||||
if steps % 3 == 0:
|
|
||||||
orders = [3,] * (K - 2) + [2, 1]
|
|
||||||
elif steps % 3 == 1:
|
|
||||||
orders = [3,] * (K - 1) + [1]
|
|
||||||
else:
|
|
||||||
orders = [3,] * (K - 1) + [2]
|
|
||||||
elif order == 2:
|
|
||||||
if steps % 2 == 0:
|
|
||||||
K = steps // 2
|
|
||||||
orders = [2,] * K
|
|
||||||
else:
|
|
||||||
K = steps // 2 + 1
|
|
||||||
orders = [2,] * (K - 1) + [1]
|
|
||||||
elif order == 1:
|
|
||||||
K = steps
|
|
||||||
orders = [1,] * steps
|
|
||||||
else:
|
|
||||||
raise ValueError("'order' must be '1' or '2' or '3'.")
|
|
||||||
if skip_type == 'logSNR':
|
|
||||||
# To reproduce the results in DPM-Solver paper
|
|
||||||
timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, K, device)
|
|
||||||
else:
|
|
||||||
timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, steps, device)[torch.cumsum(torch.tensor([0,] + orders), 0).to(device)]
|
|
||||||
return timesteps_outer, orders
|
|
||||||
|
|
||||||
def denoise_to_zero_fn(self, x, s):
|
|
||||||
"""
|
|
||||||
Denoise at the final step, which is equivalent to solve the ODE from lambda_s to infty by first-order discretization.
|
|
||||||
"""
|
|
||||||
return self.data_prediction_fn(x, s)
|
|
||||||
|
|
||||||
def multistep_uni_pc_update(self, x, model_prev_list, t_prev_list, t, order, **kwargs):
|
|
||||||
if len(t.shape) == 0:
|
|
||||||
t = t.view(-1)
|
|
||||||
if 'bh' in self.variant:
|
|
||||||
return self.multistep_uni_pc_bh_update(x, model_prev_list, t_prev_list, t, order, **kwargs)
|
|
||||||
else:
|
|
||||||
assert self.variant == 'vary_coeff'
|
|
||||||
return self.multistep_uni_pc_vary_update(x, model_prev_list, t_prev_list, t, order, **kwargs)
|
|
||||||
|
|
||||||
def multistep_uni_pc_vary_update(self, x, model_prev_list, t_prev_list, t, order, use_corrector=True):
|
|
||||||
print(f'using unified predictor-corrector with order {order} (solver type: vary coeff)')
|
|
||||||
ns = self.noise_schedule
|
|
||||||
assert order <= len(model_prev_list)
|
|
||||||
|
|
||||||
# first compute rks
|
|
||||||
t_prev_0 = t_prev_list[-1]
|
|
||||||
lambda_prev_0 = ns.marginal_lambda(t_prev_0)
|
|
||||||
lambda_t = ns.marginal_lambda(t)
|
|
||||||
model_prev_0 = model_prev_list[-1]
|
|
||||||
sigma_prev_0, sigma_t = ns.marginal_std(t_prev_0), ns.marginal_std(t)
|
|
||||||
log_alpha_t = ns.marginal_log_mean_coeff(t)
|
|
||||||
alpha_t = torch.exp(log_alpha_t)
|
|
||||||
|
|
||||||
h = lambda_t - lambda_prev_0
|
|
||||||
|
|
||||||
rks = []
|
|
||||||
D1s = []
|
|
||||||
for i in range(1, order):
|
|
||||||
t_prev_i = t_prev_list[-(i + 1)]
|
|
||||||
model_prev_i = model_prev_list[-(i + 1)]
|
|
||||||
lambda_prev_i = ns.marginal_lambda(t_prev_i)
|
|
||||||
rk = (lambda_prev_i - lambda_prev_0) / h
|
|
||||||
rks.append(rk)
|
|
||||||
D1s.append((model_prev_i - model_prev_0) / rk)
|
|
||||||
|
|
||||||
rks.append(1.)
|
|
||||||
rks = torch.tensor(rks, device=x.device)
|
|
||||||
|
|
||||||
K = len(rks)
|
|
||||||
# build C matrix
|
|
||||||
C = []
|
|
||||||
|
|
||||||
col = torch.ones_like(rks)
|
|
||||||
for k in range(1, K + 1):
|
|
||||||
C.append(col)
|
|
||||||
col = col * rks / (k + 1)
|
|
||||||
C = torch.stack(C, dim=1)
|
|
||||||
|
|
||||||
if len(D1s) > 0:
|
|
||||||
D1s = torch.stack(D1s, dim=1) # (B, K)
|
|
||||||
C_inv_p = torch.linalg.inv(C[:-1, :-1])
|
|
||||||
A_p = C_inv_p
|
|
||||||
|
|
||||||
if use_corrector:
|
|
||||||
print('using corrector')
|
|
||||||
C_inv = torch.linalg.inv(C)
|
|
||||||
A_c = C_inv
|
|
||||||
|
|
||||||
hh = -h if self.predict_x0 else h
|
|
||||||
h_phi_1 = torch.expm1(hh)
|
|
||||||
h_phi_ks = []
|
|
||||||
factorial_k = 1
|
|
||||||
h_phi_k = h_phi_1
|
|
||||||
for k in range(1, K + 2):
|
|
||||||
h_phi_ks.append(h_phi_k)
|
|
||||||
h_phi_k = h_phi_k / hh - 1 / factorial_k
|
|
||||||
factorial_k *= (k + 1)
|
|
||||||
|
|
||||||
model_t = None
|
|
||||||
if self.predict_x0:
|
|
||||||
x_t_ = (
|
|
||||||
sigma_t / sigma_prev_0 * x
|
|
||||||
- alpha_t * h_phi_1 * model_prev_0
|
|
||||||
)
|
|
||||||
# now predictor
|
|
||||||
x_t = x_t_
|
|
||||||
if len(D1s) > 0:
|
|
||||||
# compute the residuals for predictor
|
|
||||||
for k in range(K - 1):
|
|
||||||
x_t = x_t - alpha_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_p[k])
|
|
||||||
# now corrector
|
|
||||||
if use_corrector:
|
|
||||||
model_t = self.model_fn(x_t, t)
|
|
||||||
D1_t = (model_t - model_prev_0)
|
|
||||||
x_t = x_t_
|
|
||||||
k = 0
|
|
||||||
for k in range(K - 1):
|
|
||||||
x_t = x_t - alpha_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_c[k][:-1])
|
|
||||||
x_t = x_t - alpha_t * h_phi_ks[K] * (D1_t * A_c[k][-1])
|
|
||||||
else:
|
|
||||||
log_alpha_prev_0, log_alpha_t = ns.marginal_log_mean_coeff(t_prev_0), ns.marginal_log_mean_coeff(t)
|
|
||||||
x_t_ = (
|
|
||||||
(torch.exp(log_alpha_t - log_alpha_prev_0)) * x
|
|
||||||
- (sigma_t * h_phi_1) * model_prev_0
|
|
||||||
)
|
|
||||||
# now predictor
|
|
||||||
x_t = x_t_
|
|
||||||
if len(D1s) > 0:
|
|
||||||
# compute the residuals for predictor
|
|
||||||
for k in range(K - 1):
|
|
||||||
x_t = x_t - sigma_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_p[k])
|
|
||||||
# now corrector
|
|
||||||
if use_corrector:
|
|
||||||
model_t = self.model_fn(x_t, t)
|
|
||||||
D1_t = (model_t - model_prev_0)
|
|
||||||
x_t = x_t_
|
|
||||||
k = 0
|
|
||||||
for k in range(K - 1):
|
|
||||||
x_t = x_t - sigma_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_c[k][:-1])
|
|
||||||
x_t = x_t - sigma_t * h_phi_ks[K] * (D1_t * A_c[k][-1])
|
|
||||||
return x_t, model_t
|
|
||||||
|
|
||||||
def multistep_uni_pc_bh_update(self, x, model_prev_list, t_prev_list, t, order, x_t=None, use_corrector=True):
|
|
||||||
print(f'using unified predictor-corrector with order {order} (solver type: B(h))')
|
|
||||||
ns = self.noise_schedule
|
|
||||||
assert order <= len(model_prev_list)
|
|
||||||
dims = x.dim()
|
|
||||||
|
|
||||||
# first compute rks
|
|
||||||
t_prev_0 = t_prev_list[-1]
|
|
||||||
lambda_prev_0 = ns.marginal_lambda(t_prev_0)
|
|
||||||
lambda_t = ns.marginal_lambda(t)
|
|
||||||
model_prev_0 = model_prev_list[-1]
|
|
||||||
sigma_prev_0, sigma_t = ns.marginal_std(t_prev_0), ns.marginal_std(t)
|
|
||||||
log_alpha_prev_0, log_alpha_t = ns.marginal_log_mean_coeff(t_prev_0), ns.marginal_log_mean_coeff(t)
|
|
||||||
alpha_t = torch.exp(log_alpha_t)
|
|
||||||
|
|
||||||
h = lambda_t - lambda_prev_0
|
|
||||||
|
|
||||||
rks = []
|
|
||||||
D1s = []
|
|
||||||
for i in range(1, order):
|
|
||||||
t_prev_i = t_prev_list[-(i + 1)]
|
|
||||||
model_prev_i = model_prev_list[-(i + 1)]
|
|
||||||
lambda_prev_i = ns.marginal_lambda(t_prev_i)
|
|
||||||
rk = ((lambda_prev_i - lambda_prev_0) / h)[0]
|
|
||||||
rks.append(rk)
|
|
||||||
D1s.append((model_prev_i - model_prev_0) / rk)
|
|
||||||
|
|
||||||
rks.append(1.)
|
|
||||||
rks = torch.tensor(rks, device=x.device)
|
|
||||||
|
|
||||||
R = []
|
|
||||||
b = []
|
|
||||||
|
|
||||||
hh = -h[0] if self.predict_x0 else h[0]
|
|
||||||
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
|
|
||||||
h_phi_k = h_phi_1 / hh - 1
|
|
||||||
|
|
||||||
factorial_i = 1
|
|
||||||
|
|
||||||
if self.variant == 'bh1':
|
|
||||||
B_h = hh
|
|
||||||
elif self.variant == 'bh2':
|
|
||||||
B_h = torch.expm1(hh)
|
|
||||||
else:
|
|
||||||
raise NotImplementedError()
|
|
||||||
|
|
||||||
for i in range(1, order + 1):
|
|
||||||
R.append(torch.pow(rks, i - 1))
|
|
||||||
b.append(h_phi_k * factorial_i / B_h)
|
|
||||||
factorial_i *= (i + 1)
|
|
||||||
h_phi_k = h_phi_k / hh - 1 / factorial_i
|
|
||||||
|
|
||||||
R = torch.stack(R)
|
|
||||||
b = torch.tensor(b, device=x.device)
|
|
||||||
|
|
||||||
# now predictor
|
|
||||||
use_predictor = len(D1s) > 0 and x_t is None
|
|
||||||
if len(D1s) > 0:
|
|
||||||
D1s = torch.stack(D1s, dim=1) # (B, K)
|
|
||||||
if x_t is None:
|
|
||||||
# for order 2, we use a simplified version
|
|
||||||
if order == 2:
|
|
||||||
rhos_p = torch.tensor([0.5], device=b.device)
|
|
||||||
else:
|
|
||||||
rhos_p = torch.linalg.solve(R[:-1, :-1], b[:-1])
|
|
||||||
else:
|
|
||||||
D1s = None
|
|
||||||
|
|
||||||
if use_corrector:
|
|
||||||
print('using corrector')
|
|
||||||
# for order 1, we use a simplified version
|
|
||||||
if order == 1:
|
|
||||||
rhos_c = torch.tensor([0.5], device=b.device)
|
|
||||||
else:
|
|
||||||
rhos_c = torch.linalg.solve(R, b)
|
|
||||||
|
|
||||||
model_t = None
|
|
||||||
if self.predict_x0:
|
|
||||||
x_t_ = (
|
|
||||||
expand_dims(sigma_t / sigma_prev_0, dims) * x
|
|
||||||
- expand_dims(alpha_t * h_phi_1, dims)* model_prev_0
|
|
||||||
)
|
|
||||||
|
|
||||||
if x_t is None:
|
|
||||||
if use_predictor:
|
|
||||||
pred_res = torch.einsum('k,bkchw->bchw', rhos_p, D1s)
|
|
||||||
else:
|
|
||||||
pred_res = 0
|
|
||||||
x_t = x_t_ - expand_dims(alpha_t * B_h, dims) * pred_res
|
|
||||||
|
|
||||||
if use_corrector:
|
|
||||||
model_t = self.model_fn(x_t, t)
|
|
||||||
if D1s is not None:
|
|
||||||
corr_res = torch.einsum('k,bkchw->bchw', rhos_c[:-1], D1s)
|
|
||||||
else:
|
|
||||||
corr_res = 0
|
|
||||||
D1_t = (model_t - model_prev_0)
|
|
||||||
x_t = x_t_ - expand_dims(alpha_t * B_h, dims) * (corr_res + rhos_c[-1] * D1_t)
|
|
||||||
else:
|
|
||||||
x_t_ = (
|
|
||||||
expand_dims(torch.exp(log_alpha_t - log_alpha_prev_0), dims) * x
|
|
||||||
- expand_dims(sigma_t * h_phi_1, dims) * model_prev_0
|
|
||||||
)
|
|
||||||
if x_t is None:
|
|
||||||
if use_predictor:
|
|
||||||
pred_res = torch.einsum('k,bkchw->bchw', rhos_p, D1s)
|
|
||||||
else:
|
|
||||||
pred_res = 0
|
|
||||||
x_t = x_t_ - expand_dims(sigma_t * B_h, dims) * pred_res
|
|
||||||
|
|
||||||
if use_corrector:
|
|
||||||
model_t = self.model_fn(x_t, t)
|
|
||||||
if D1s is not None:
|
|
||||||
corr_res = torch.einsum('k,bkchw->bchw', rhos_c[:-1], D1s)
|
|
||||||
else:
|
|
||||||
corr_res = 0
|
|
||||||
D1_t = (model_t - model_prev_0)
|
|
||||||
x_t = x_t_ - expand_dims(sigma_t * B_h, dims) * (corr_res + rhos_c[-1] * D1_t)
|
|
||||||
return x_t, model_t
|
|
||||||
|
|
||||||
|
|
||||||
def sample(self, x, steps=20, t_start=None, t_end=None, order=3, skip_type='time_uniform',
|
|
||||||
method='singlestep', lower_order_final=True, denoise_to_zero=False, solver_type='dpm_solver',
|
|
||||||
atol=0.0078, rtol=0.05, corrector=False,
|
|
||||||
):
|
|
||||||
t_0 = 1. / self.noise_schedule.total_N if t_end is None else t_end
|
|
||||||
t_T = self.noise_schedule.T if t_start is None else t_start
|
|
||||||
device = x.device
|
|
||||||
if method == 'multistep':
|
|
||||||
assert steps >= order
|
|
||||||
timesteps = self.get_time_steps(skip_type=skip_type, t_T=t_T, t_0=t_0, N=steps, device=device)
|
|
||||||
assert timesteps.shape[0] - 1 == steps
|
|
||||||
with torch.no_grad():
|
|
||||||
vec_t = timesteps[0].expand((x.shape[0]))
|
|
||||||
model_prev_list = [self.model_fn(x, vec_t)]
|
|
||||||
t_prev_list = [vec_t]
|
|
||||||
# Init the first `order` values by lower order multistep DPM-Solver.
|
|
||||||
for init_order in range(1, order):
|
|
||||||
vec_t = timesteps[init_order].expand(x.shape[0])
|
|
||||||
x, model_x = self.multistep_uni_pc_update(x, model_prev_list, t_prev_list, vec_t, init_order, use_corrector=True)
|
|
||||||
if model_x is None:
|
|
||||||
model_x = self.model_fn(x, vec_t)
|
|
||||||
model_prev_list.append(model_x)
|
|
||||||
t_prev_list.append(vec_t)
|
|
||||||
for step in range(order, steps + 1):
|
|
||||||
vec_t = timesteps[step].expand(x.shape[0])
|
|
||||||
if lower_order_final:
|
|
||||||
step_order = min(order, steps + 1 - step)
|
|
||||||
else:
|
|
||||||
step_order = order
|
|
||||||
print('this step order:', step_order)
|
|
||||||
if step == steps:
|
|
||||||
print('do not run corrector at the last step')
|
|
||||||
use_corrector = False
|
|
||||||
else:
|
|
||||||
use_corrector = True
|
|
||||||
x, model_x = self.multistep_uni_pc_update(x, model_prev_list, t_prev_list, vec_t, step_order, use_corrector=use_corrector)
|
|
||||||
for i in range(order - 1):
|
|
||||||
t_prev_list[i] = t_prev_list[i + 1]
|
|
||||||
model_prev_list[i] = model_prev_list[i + 1]
|
|
||||||
t_prev_list[-1] = vec_t
|
|
||||||
# We do not need to evaluate the final model value.
|
|
||||||
if step < steps:
|
|
||||||
if model_x is None:
|
|
||||||
model_x = self.model_fn(x, vec_t)
|
|
||||||
model_prev_list[-1] = model_x
|
|
||||||
else:
|
|
||||||
raise NotImplementedError()
|
|
||||||
if denoise_to_zero:
|
|
||||||
x = self.denoise_to_zero_fn(x, torch.ones((x.shape[0],)).to(device) * t_0)
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
#############################################################
|
|
||||||
# other utility functions
|
|
||||||
#############################################################
|
|
||||||
|
|
||||||
def interpolate_fn(x, xp, yp):
|
|
||||||
"""
|
|
||||||
A piecewise linear function y = f(x), using xp and yp as keypoints.
|
|
||||||
We implement f(x) in a differentiable way (i.e. applicable for autograd).
|
|
||||||
The function f(x) is well-defined for all x-axis. (For x beyond the bounds of xp, we use the outmost points of xp to define the linear function.)
|
|
||||||
|
|
||||||
Args:
|
|
||||||
x: PyTorch tensor with shape [N, C], where N is the batch size, C is the number of channels (we use C = 1 for DPM-Solver).
|
|
||||||
xp: PyTorch tensor with shape [C, K], where K is the number of keypoints.
|
|
||||||
yp: PyTorch tensor with shape [C, K].
|
|
||||||
Returns:
|
|
||||||
The function values f(x), with shape [N, C].
|
|
||||||
"""
|
|
||||||
N, K = x.shape[0], xp.shape[1]
|
|
||||||
all_x = torch.cat([x.unsqueeze(2), xp.unsqueeze(0).repeat((N, 1, 1))], dim=2)
|
|
||||||
sorted_all_x, x_indices = torch.sort(all_x, dim=2)
|
|
||||||
x_idx = torch.argmin(x_indices, dim=2)
|
|
||||||
cand_start_idx = x_idx - 1
|
|
||||||
start_idx = torch.where(
|
|
||||||
torch.eq(x_idx, 0),
|
|
||||||
torch.tensor(1, device=x.device),
|
|
||||||
torch.where(
|
|
||||||
torch.eq(x_idx, K), torch.tensor(K - 2, device=x.device), cand_start_idx,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
end_idx = torch.where(torch.eq(start_idx, cand_start_idx), start_idx + 2, start_idx + 1)
|
|
||||||
start_x = torch.gather(sorted_all_x, dim=2, index=start_idx.unsqueeze(2)).squeeze(2)
|
|
||||||
end_x = torch.gather(sorted_all_x, dim=2, index=end_idx.unsqueeze(2)).squeeze(2)
|
|
||||||
start_idx2 = torch.where(
|
|
||||||
torch.eq(x_idx, 0),
|
|
||||||
torch.tensor(0, device=x.device),
|
|
||||||
torch.where(
|
|
||||||
torch.eq(x_idx, K), torch.tensor(K - 2, device=x.device), cand_start_idx,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
y_positions_expanded = yp.unsqueeze(0).expand(N, -1, -1)
|
|
||||||
start_y = torch.gather(y_positions_expanded, dim=2, index=start_idx2.unsqueeze(2)).squeeze(2)
|
|
||||||
end_y = torch.gather(y_positions_expanded, dim=2, index=(start_idx2 + 1).unsqueeze(2)).squeeze(2)
|
|
||||||
cand = start_y + (x - start_x) * (end_y - start_y) / (end_x - start_x)
|
|
||||||
return cand
|
|
||||||
|
|
||||||
|
|
||||||
def expand_dims(v, dims):
|
|
||||||
"""
|
|
||||||
Expand the tensor `v` to the dim `dims`.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
`v`: a PyTorch tensor with shape [N].
|
|
||||||
`dim`: a `int`.
|
|
||||||
Returns:
|
|
||||||
a PyTorch tensor with shape [N, 1, 1, ..., 1] and the total dimension is `dims`.
|
|
||||||
"""
|
|
||||||
return v[(...,) + (None,)*(dims - 1)]
|
|
||||||
@@ -1,158 +0,0 @@
|
|||||||
import math
|
|
||||||
import numpy as np
|
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from einops import repeat
|
|
||||||
|
|
||||||
|
|
||||||
def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False, dtype=None):
|
|
||||||
"""
|
|
||||||
Create sinusoidal timestep embeddings.
|
|
||||||
:param timesteps: a 1-D Tensor of N indices, one per batch element.
|
|
||||||
These may be fractional.
|
|
||||||
:param dim: the dimension of the output.
|
|
||||||
:param max_period: controls the minimum frequency of the embeddings.
|
|
||||||
:return: an [N x dim] Tensor of positional embeddings.
|
|
||||||
"""
|
|
||||||
if not repeat_only:
|
|
||||||
half = dim // 2
|
|
||||||
freqs = torch.exp(
|
|
||||||
-math.log(max_period) * torch.arange(start=0, end=half, dtype=dtype) / half
|
|
||||||
).to(device=timesteps.device)
|
|
||||||
args = timesteps[:, None].float() * freqs[None]
|
|
||||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
|
||||||
if dim % 2:
|
|
||||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
|
||||||
else:
|
|
||||||
embedding = repeat(timesteps, 'b -> b d', d=dim)
|
|
||||||
return embedding.to(dtype)
|
|
||||||
|
|
||||||
|
|
||||||
def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
|
|
||||||
if schedule == "linear":
|
|
||||||
betas = (
|
|
||||||
torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=torch.float64) ** 2
|
|
||||||
)
|
|
||||||
|
|
||||||
elif schedule == "cosine":
|
|
||||||
timesteps = (
|
|
||||||
torch.arange(n_timestep + 1, dtype=torch.float64) / n_timestep + cosine_s
|
|
||||||
)
|
|
||||||
alphas = timesteps / (1 + cosine_s) * np.pi / 2
|
|
||||||
alphas = torch.cos(alphas).pow(2)
|
|
||||||
alphas = alphas / alphas[0]
|
|
||||||
betas = 1 - alphas[1:] / alphas[:-1]
|
|
||||||
betas = np.clip(betas, a_min=0, a_max=0.999)
|
|
||||||
|
|
||||||
elif schedule == "sqrt_linear":
|
|
||||||
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64)
|
|
||||||
elif schedule == "sqrt":
|
|
||||||
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) ** 0.5
|
|
||||||
else:
|
|
||||||
raise ValueError(f"schedule '{schedule}' unknown.")
|
|
||||||
return betas.numpy()
|
|
||||||
|
|
||||||
|
|
||||||
def make_ddim_timesteps(ddim_discr_method, num_ddim_timesteps, num_ddpm_timesteps, verbose=True):
|
|
||||||
if ddim_discr_method == 'uniform':
|
|
||||||
c = num_ddpm_timesteps // num_ddim_timesteps
|
|
||||||
ddim_timesteps = np.asarray(list(range(0, num_ddpm_timesteps, c)))
|
|
||||||
steps_out = ddim_timesteps + 1
|
|
||||||
elif ddim_discr_method == 'uniform_trailing':
|
|
||||||
c = num_ddpm_timesteps / num_ddim_timesteps
|
|
||||||
ddim_timesteps = np.flip(np.round(np.arange(num_ddpm_timesteps, 0, -c))).astype(np.int64)
|
|
||||||
steps_out = ddim_timesteps - 1
|
|
||||||
elif ddim_discr_method == 'quad':
|
|
||||||
ddim_timesteps = ((np.linspace(0, np.sqrt(num_ddpm_timesteps * .8), num_ddim_timesteps)) ** 2).astype(int)
|
|
||||||
steps_out = ddim_timesteps + 1
|
|
||||||
else:
|
|
||||||
raise NotImplementedError(f'There is no ddim discretization method called "{ddim_discr_method}"')
|
|
||||||
|
|
||||||
# assert ddim_timesteps.shape[0] == num_ddim_timesteps
|
|
||||||
# add one to get the final alpha values right (the ones from first scale to data during sampling)
|
|
||||||
# steps_out = ddim_timesteps + 1
|
|
||||||
if verbose:
|
|
||||||
print(f'Selected timesteps for ddim sampler: {steps_out}')
|
|
||||||
return steps_out
|
|
||||||
|
|
||||||
|
|
||||||
def make_ddim_sampling_parameters(alphacums, ddim_timesteps, eta, verbose=True):
|
|
||||||
# select alphas for computing the variance schedule
|
|
||||||
# print(f'ddim_timesteps={ddim_timesteps}, len_alphacums={len(alphacums)}')
|
|
||||||
alphas = alphacums[ddim_timesteps]
|
|
||||||
alphas_prev = np.asarray([alphacums[0]] + alphacums[ddim_timesteps[:-1]].tolist())
|
|
||||||
|
|
||||||
# according the the formula provided in https://arxiv.org/abs/2010.02502
|
|
||||||
sigmas = eta * np.sqrt((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev))
|
|
||||||
if verbose:
|
|
||||||
print(f'Selected alphas for ddim sampler: a_t: {alphas}; a_(t-1): {alphas_prev}')
|
|
||||||
print(f'For the chosen value of eta, which is {eta}, '
|
|
||||||
f'this results in the following sigma_t schedule for ddim sampler {sigmas}')
|
|
||||||
return sigmas, alphas, alphas_prev
|
|
||||||
|
|
||||||
|
|
||||||
def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
|
|
||||||
"""
|
|
||||||
Create a beta schedule that discretizes the given alpha_t_bar function,
|
|
||||||
which defines the cumulative product of (1-beta) over time from t = [0,1].
|
|
||||||
:param num_diffusion_timesteps: the number of betas to produce.
|
|
||||||
:param alpha_bar: a lambda that takes an argument t from 0 to 1 and
|
|
||||||
produces the cumulative product of (1-beta) up to that
|
|
||||||
part of the diffusion process.
|
|
||||||
:param max_beta: the maximum beta to use; use values lower than 1 to
|
|
||||||
prevent singularities.
|
|
||||||
"""
|
|
||||||
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 np.array(betas)
|
|
||||||
|
|
||||||
def rescale_zero_terminal_snr(betas):
|
|
||||||
"""
|
|
||||||
Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
|
|
||||||
|
|
||||||
Args:
|
|
||||||
betas (`numpy.ndarray`):
|
|
||||||
the betas that the scheduler is being initialized with.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
`numpy.ndarray`: rescaled betas with zero terminal SNR
|
|
||||||
"""
|
|
||||||
# Convert betas to alphas_bar_sqrt
|
|
||||||
alphas = 1.0 - betas
|
|
||||||
alphas_cumprod = np.cumprod(alphas, axis=0)
|
|
||||||
alphas_bar_sqrt = np.sqrt(alphas_cumprod)
|
|
||||||
|
|
||||||
# Store old values.
|
|
||||||
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].copy()
|
|
||||||
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].copy()
|
|
||||||
|
|
||||||
# Shift so the last timestep is zero.
|
|
||||||
alphas_bar_sqrt -= alphas_bar_sqrt_T
|
|
||||||
|
|
||||||
# Scale so the first timestep is back to the old value.
|
|
||||||
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
|
|
||||||
|
|
||||||
# Convert alphas_bar_sqrt to betas
|
|
||||||
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
|
|
||||||
alphas = alphas_bar[1:] / alphas_bar[:-1] # Revert cumprod
|
|
||||||
alphas = np.concatenate([alphas_bar[0:1], alphas])
|
|
||||||
betas = 1 - alphas
|
|
||||||
|
|
||||||
return betas
|
|
||||||
|
|
||||||
|
|
||||||
def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
|
|
||||||
"""
|
|
||||||
Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and
|
|
||||||
Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4
|
|
||||||
"""
|
|
||||||
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
|
|
||||||
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
|
|
||||||
# rescale the results from guidance (fixes overexposure)
|
|
||||||
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
|
|
||||||
# mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images
|
|
||||||
noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg
|
|
||||||
return noise_cfg
|
|
||||||
@@ -1,809 +0,0 @@
|
|||||||
import torch
|
|
||||||
from torch import nn, einsum
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from einops import rearrange, repeat
|
|
||||||
from functools import partial
|
|
||||||
from ..common import (
|
|
||||||
checkpoint,
|
|
||||||
exists,
|
|
||||||
default,
|
|
||||||
)
|
|
||||||
from ..basics import zero_module
|
|
||||||
import comfy.ops
|
|
||||||
ops = comfy.ops.disable_weight_init
|
|
||||||
from comfy import model_management
|
|
||||||
from comfy.ldm.modules.attention import optimized_attention, optimized_attention_masked
|
|
||||||
|
|
||||||
if model_management.xformers_enabled():
|
|
||||||
import xformers
|
|
||||||
import xformers.ops
|
|
||||||
XFORMERS_IS_AVAILBLE = True
|
|
||||||
else:
|
|
||||||
XFORMERS_IS_AVAILBLE = False
|
|
||||||
|
|
||||||
class RelativePosition(nn.Module):
|
|
||||||
""" https://github.com/evelinehong/Transformer_Relative_Position_PyTorch/blob/master/relative_position.py """
|
|
||||||
|
|
||||||
def __init__(self, num_units, max_relative_position):
|
|
||||||
super().__init__()
|
|
||||||
self.num_units = num_units
|
|
||||||
self.max_relative_position = max_relative_position
|
|
||||||
self.embeddings_table = nn.Parameter(torch.Tensor(max_relative_position * 2 + 1, num_units))
|
|
||||||
nn.init.xavier_uniform_(self.embeddings_table)
|
|
||||||
|
|
||||||
def forward(self, length_q, length_k):
|
|
||||||
device = self.embeddings_table.device
|
|
||||||
range_vec_q = torch.arange(length_q, device=device)
|
|
||||||
range_vec_k = torch.arange(length_k, device=device)
|
|
||||||
distance_mat = range_vec_k[None, :] - range_vec_q[:, None]
|
|
||||||
distance_mat_clipped = torch.clamp(distance_mat, -self.max_relative_position, self.max_relative_position)
|
|
||||||
final_mat = distance_mat_clipped + self.max_relative_position
|
|
||||||
final_mat = final_mat.long()
|
|
||||||
embeddings = self.embeddings_table[final_mat]
|
|
||||||
return embeddings
|
|
||||||
|
|
||||||
|
|
||||||
# TODO Add native Comfy optimized attention.
|
|
||||||
class CrossAttention(nn.Module):
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
query_dim,
|
|
||||||
context_dim=None,
|
|
||||||
heads=8,
|
|
||||||
dim_head=64,
|
|
||||||
dropout=0.,
|
|
||||||
relative_position=False,
|
|
||||||
temporal_length=None,
|
|
||||||
video_length=None,
|
|
||||||
image_cross_attention=False,
|
|
||||||
image_cross_attention_scale=1.0,
|
|
||||||
image_cross_attention_scale_learnable=False,
|
|
||||||
text_context_len=77,
|
|
||||||
device=None,
|
|
||||||
dtype=None,
|
|
||||||
operations=ops
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
inner_dim = dim_head * heads
|
|
||||||
context_dim = default(context_dim, query_dim)
|
|
||||||
self.scale = dim_head**-0.5
|
|
||||||
self.heads = heads
|
|
||||||
self.dim_head = dim_head
|
|
||||||
self.to_q = operations.Linear(query_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
|
||||||
self.to_k = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
|
||||||
self.to_v = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
|
||||||
|
|
||||||
self.to_out = nn.Sequential(
|
|
||||||
operations.Linear(inner_dim, query_dim, device=device, dtype=dtype),
|
|
||||||
nn.Dropout(dropout)
|
|
||||||
)
|
|
||||||
|
|
||||||
self.relative_position = relative_position
|
|
||||||
if self.relative_position:
|
|
||||||
assert(temporal_length is not None)
|
|
||||||
self.relative_position_k = RelativePosition(num_units=dim_head, max_relative_position=temporal_length)
|
|
||||||
self.relative_position_v = RelativePosition(num_units=dim_head, max_relative_position=temporal_length)
|
|
||||||
else:
|
|
||||||
## only used for spatial attention, while NOT for temporal attention
|
|
||||||
if XFORMERS_IS_AVAILBLE and temporal_length is None:
|
|
||||||
self.forward = self.efficient_forward
|
|
||||||
else:
|
|
||||||
self.forward = self.comfy_efficient_forward
|
|
||||||
|
|
||||||
self.video_length = video_length
|
|
||||||
self.image_cross_attention = image_cross_attention
|
|
||||||
self.image_cross_attention_scale = image_cross_attention_scale
|
|
||||||
self.text_context_len = text_context_len
|
|
||||||
self.image_cross_attention_scale_learnable = image_cross_attention_scale_learnable
|
|
||||||
if self.image_cross_attention:
|
|
||||||
self.to_k_ip = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
|
||||||
self.to_v_ip = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
|
||||||
if image_cross_attention_scale_learnable:
|
|
||||||
self.register_parameter('alpha', nn.Parameter(torch.tensor(0.)) )
|
|
||||||
|
|
||||||
def comfy_efficient_forward(self, x, context=None, mask=None, *args, **kwargs):
|
|
||||||
spatial_self_attn = (context is None)
|
|
||||||
k_ip, v_ip, out_ip = None, None, None
|
|
||||||
|
|
||||||
h = self.heads
|
|
||||||
q = self.to_q(x)
|
|
||||||
context = default(context, x)
|
|
||||||
|
|
||||||
if self.image_cross_attention and not spatial_self_attn:
|
|
||||||
context, context_image = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:]
|
|
||||||
k = self.to_k(context)
|
|
||||||
v = self.to_v(context)
|
|
||||||
k_ip = self.to_k_ip(context_image)
|
|
||||||
v_ip = self.to_v_ip(context_image)
|
|
||||||
else:
|
|
||||||
if not spatial_self_attn:
|
|
||||||
context = context[:,:self.text_context_len,:]
|
|
||||||
k = self.to_k(context)
|
|
||||||
v = self.to_v(context)
|
|
||||||
|
|
||||||
out = optimized_attention(q, k, v, h)
|
|
||||||
|
|
||||||
if exists(mask):
|
|
||||||
## feasible for causal attention mask only
|
|
||||||
out = optimized_attention_masked(q, k, v, h)
|
|
||||||
|
|
||||||
## for image cross-attention
|
|
||||||
if k_ip is not None:
|
|
||||||
q = rearrange(q, 'b n (h d) -> (b h) n d', h=h)
|
|
||||||
k_ip, v_ip = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (k_ip, v_ip))
|
|
||||||
sim_ip = torch.einsum('b i d, b j d -> b i j', q, k_ip) * self.scale
|
|
||||||
del k_ip
|
|
||||||
sim_ip = sim_ip.softmax(dim=-1)
|
|
||||||
out_ip = torch.einsum('b i j, b j d -> b i d', sim_ip, v_ip)
|
|
||||||
out_ip = rearrange(out_ip, '(b h) n d -> b n (h d)', h=h)
|
|
||||||
|
|
||||||
if out_ip is not None:
|
|
||||||
if self.image_cross_attention_scale_learnable:
|
|
||||||
out = out + self.image_cross_attention_scale * out_ip * (torch.tanh(self.alpha)+1)
|
|
||||||
else:
|
|
||||||
out = out + self.image_cross_attention_scale * out_ip
|
|
||||||
|
|
||||||
return self.to_out(out)
|
|
||||||
|
|
||||||
def forward(self, x, context=None, mask=None):
|
|
||||||
spatial_self_attn = (context is None)
|
|
||||||
k_ip, v_ip, out_ip = None, None, None
|
|
||||||
|
|
||||||
h = self.heads
|
|
||||||
q = self.to_q(x)
|
|
||||||
context = default(context, x)
|
|
||||||
|
|
||||||
if self.image_cross_attention and not spatial_self_attn:
|
|
||||||
context, context_image = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:]
|
|
||||||
k = self.to_k(context)
|
|
||||||
v = self.to_v(context)
|
|
||||||
k_ip = self.to_k_ip(context_image)
|
|
||||||
v_ip = self.to_v_ip(context_image)
|
|
||||||
else:
|
|
||||||
|
|
||||||
# Assumed Spatial Attention (b c h w)
|
|
||||||
if not spatial_self_attn:
|
|
||||||
context = context[:,:self.text_context_len,:]
|
|
||||||
k = self.to_k(context)
|
|
||||||
v = self.to_v(context)
|
|
||||||
|
|
||||||
|
|
||||||
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))
|
|
||||||
|
|
||||||
sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale
|
|
||||||
if self.relative_position:
|
|
||||||
len_q, len_k, len_v = q.shape[1], k.shape[1], v.shape[1]
|
|
||||||
k2 = self.relative_position_k(len_q, len_k)
|
|
||||||
sim2 = einsum('b t d, t s d -> b t s', q, k2) * self.scale # TODO check
|
|
||||||
sim += sim2
|
|
||||||
del k
|
|
||||||
|
|
||||||
if exists(mask):
|
|
||||||
## feasible for causal attention mask only
|
|
||||||
max_neg_value = -torch.finfo(sim.dtype).max
|
|
||||||
mask = repeat(mask, 'b i j -> (b h) i j', h=h)
|
|
||||||
sim.masked_fill_(~(mask>0.5), max_neg_value)
|
|
||||||
|
|
||||||
# attention, what we cannot get enough of
|
|
||||||
sim = sim.softmax(dim=-1)
|
|
||||||
|
|
||||||
out = torch.einsum('b i j, b j d -> b i d', sim, v)
|
|
||||||
if self.relative_position:
|
|
||||||
v2 = self.relative_position_v(len_q, len_v)
|
|
||||||
out2 = einsum('b t s, t s d -> b t d', sim, v2) # TODO check
|
|
||||||
out += out2
|
|
||||||
out = rearrange(out, '(b h) n d -> b n (h d)', h=h)
|
|
||||||
|
|
||||||
|
|
||||||
## for image cross-attention
|
|
||||||
if k_ip is not None:
|
|
||||||
k_ip, v_ip = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (k_ip, v_ip))
|
|
||||||
sim_ip = torch.einsum('b i d, b j d -> b i j', q, k_ip) * self.scale
|
|
||||||
del k_ip
|
|
||||||
sim_ip = sim_ip.softmax(dim=-1)
|
|
||||||
out_ip = torch.einsum('b i j, b j d -> b i d', sim_ip, v_ip)
|
|
||||||
out_ip = rearrange(out_ip, '(b h) n d -> b n (h d)', h=h)
|
|
||||||
|
|
||||||
|
|
||||||
if out_ip is not None:
|
|
||||||
if self.image_cross_attention_scale_learnable:
|
|
||||||
out = out + self.image_cross_attention_scale * out_ip * (torch.tanh(self.alpha)+1)
|
|
||||||
else:
|
|
||||||
out = out + self.image_cross_attention_scale * out_ip
|
|
||||||
|
|
||||||
return self.to_out(out)
|
|
||||||
|
|
||||||
def efficient_forward(self, x, context=None, mask=None):
|
|
||||||
spatial_self_attn = (context is None)
|
|
||||||
k_ip, v_ip, out_ip = None, None, None
|
|
||||||
|
|
||||||
q = self.to_q(x)
|
|
||||||
context = default(context, x)
|
|
||||||
|
|
||||||
if self.image_cross_attention and not spatial_self_attn:
|
|
||||||
context, context_image = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:]
|
|
||||||
k = self.to_k(context)
|
|
||||||
v = self.to_v(context)
|
|
||||||
k_ip = self.to_k_ip(context_image)
|
|
||||||
v_ip = self.to_v_ip(context_image)
|
|
||||||
else:
|
|
||||||
if not spatial_self_attn:
|
|
||||||
context = context[:,:self.text_context_len,:]
|
|
||||||
k = self.to_k(context)
|
|
||||||
v = self.to_v(context)
|
|
||||||
|
|
||||||
b, _, _ = q.shape
|
|
||||||
q, k, v = map(
|
|
||||||
lambda t: t.unsqueeze(3)
|
|
||||||
.reshape(b, t.shape[1], self.heads, self.dim_head)
|
|
||||||
.permute(0, 2, 1, 3)
|
|
||||||
.reshape(b * self.heads, t.shape[1], self.dim_head)
|
|
||||||
.contiguous(),
|
|
||||||
(q, k, v),
|
|
||||||
)
|
|
||||||
# actually compute the attention, what we cannot get enough of
|
|
||||||
out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=None, op=None)
|
|
||||||
|
|
||||||
## for image cross-attention
|
|
||||||
if k_ip is not None:
|
|
||||||
k_ip, v_ip = map(
|
|
||||||
lambda t: t.unsqueeze(3)
|
|
||||||
.reshape(b, t.shape[1], self.heads, self.dim_head)
|
|
||||||
.permute(0, 2, 1, 3)
|
|
||||||
.reshape(b * self.heads, t.shape[1], self.dim_head)
|
|
||||||
.contiguous(),
|
|
||||||
(k_ip, v_ip),
|
|
||||||
)
|
|
||||||
out_ip = xformers.ops.memory_efficient_attention(q, k_ip, v_ip, attn_bias=None, op=None)
|
|
||||||
out_ip = (
|
|
||||||
out_ip.unsqueeze(0)
|
|
||||||
.reshape(b, self.heads, out.shape[1], self.dim_head)
|
|
||||||
.permute(0, 2, 1, 3)
|
|
||||||
.reshape(b, out.shape[1], self.heads * self.dim_head)
|
|
||||||
)
|
|
||||||
|
|
||||||
if exists(mask):
|
|
||||||
raise NotImplementedError
|
|
||||||
out = (
|
|
||||||
out.unsqueeze(0)
|
|
||||||
.reshape(b, self.heads, out.shape[1], self.dim_head)
|
|
||||||
.permute(0, 2, 1, 3)
|
|
||||||
.reshape(b, out.shape[1], self.heads * self.dim_head)
|
|
||||||
)
|
|
||||||
if out_ip is not None:
|
|
||||||
if self.image_cross_attention_scale_learnable:
|
|
||||||
out = out + self.image_cross_attention_scale * out_ip * (torch.tanh(self.alpha)+1)
|
|
||||||
else:
|
|
||||||
out = out + self.image_cross_attention_scale * out_ip
|
|
||||||
|
|
||||||
return self.to_out(out)
|
|
||||||
|
|
||||||
|
|
||||||
class BasicTransformerBlock(nn.Module):
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim,
|
|
||||||
n_heads,
|
|
||||||
d_head,
|
|
||||||
dropout=0.,
|
|
||||||
context_dim=None,
|
|
||||||
gated_ff=True,
|
|
||||||
checkpoint=True,
|
|
||||||
disable_self_attn=False,
|
|
||||||
attention_cls=None,
|
|
||||||
video_length=None,
|
|
||||||
inner_dim=None,
|
|
||||||
image_cross_attention=False,
|
|
||||||
image_cross_attention_scale=1.0,
|
|
||||||
image_cross_attention_scale_learnable=False,
|
|
||||||
switch_temporal_ca_to_sa=False,
|
|
||||||
text_context_len=77,
|
|
||||||
ff_in=None,
|
|
||||||
device=None,
|
|
||||||
dtype=None,
|
|
||||||
operations=ops
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
attn_cls = CrossAttention if attention_cls is None else attention_cls
|
|
||||||
|
|
||||||
self.ff_in = ff_in or inner_dim is not None
|
|
||||||
if self.ff_in:
|
|
||||||
self.norm_in = operations.LayerNorm(dim, dtype=dtype, device=device)
|
|
||||||
self.ff_in = FeedForward(
|
|
||||||
dim,
|
|
||||||
dim_out=inner_dim,
|
|
||||||
dropout=dropout,
|
|
||||||
glu=gated_ff,
|
|
||||||
dtype=dtype,
|
|
||||||
device=device,
|
|
||||||
operations=operations
|
|
||||||
)
|
|
||||||
if inner_dim is None:
|
|
||||||
inner_dim = dim
|
|
||||||
|
|
||||||
self.is_res = inner_dim == dim
|
|
||||||
self.disable_self_attn = disable_self_attn
|
|
||||||
self.attn1 = attn_cls(query_dim=dim, heads=n_heads, dim_head=d_head, dropout=dropout,
|
|
||||||
context_dim=None, device=device, dtype=dtype if self.disable_self_attn else None)
|
|
||||||
self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff, device=device, dtype=dtype)
|
|
||||||
self.attn2 = attn_cls(
|
|
||||||
query_dim=dim,
|
|
||||||
context_dim=context_dim,
|
|
||||||
heads=n_heads,
|
|
||||||
dim_head=d_head,
|
|
||||||
dropout=dropout,
|
|
||||||
video_length=video_length,
|
|
||||||
image_cross_attention=image_cross_attention,
|
|
||||||
image_cross_attention_scale=image_cross_attention_scale,
|
|
||||||
image_cross_attention_scale_learnable=image_cross_attention_scale_learnable,
|
|
||||||
text_context_len=text_context_len,
|
|
||||||
device=device,
|
|
||||||
dtype=dtype
|
|
||||||
)
|
|
||||||
self.image_cross_attention = image_cross_attention
|
|
||||||
|
|
||||||
self.norm1 = operations.LayerNorm(dim, device=device, dtype=dtype)
|
|
||||||
self.norm2 = operations.LayerNorm(dim, device=device, dtype=dtype)
|
|
||||||
self.norm3 = operations.LayerNorm(dim, device=device, dtype=dtype)
|
|
||||||
|
|
||||||
self.n_heads = n_heads
|
|
||||||
self.d_head = d_head
|
|
||||||
self.checkpoint = checkpoint
|
|
||||||
self.switch_temporal_ca_to_sa = switch_temporal_ca_to_sa
|
|
||||||
|
|
||||||
def forward(self, x, context=None, mask=None, **kwargs):
|
|
||||||
## implementation tricks: because checkpointing doesn't support non-tensor (e.g. None or scalar) arguments
|
|
||||||
input_tuple = (x,) ## should not be (x), otherwise *input_tuple will decouple x into multiple arguments
|
|
||||||
if context is not None:
|
|
||||||
input_tuple = (x, context)
|
|
||||||
if mask is not None:
|
|
||||||
forward_mask = partial(self._forward, mask=mask)
|
|
||||||
return checkpoint(forward_mask, (x,), self.parameters(), self.checkpoint)
|
|
||||||
return checkpoint(self._forward, input_tuple, self.parameters(), self.checkpoint)
|
|
||||||
|
|
||||||
|
|
||||||
def _forward(self, x, context=None, mask=None, transformer_options={}):
|
|
||||||
extra_options = {}
|
|
||||||
block = transformer_options.get("block", None)
|
|
||||||
block_index = transformer_options.get("block_index", 0)
|
|
||||||
transformer_patches = {}
|
|
||||||
transformer_patches_replace = {}
|
|
||||||
|
|
||||||
for k in transformer_options:
|
|
||||||
if k == "patches":
|
|
||||||
transformer_patches = transformer_options[k]
|
|
||||||
elif k == "patches_replace":
|
|
||||||
transformer_patches_replace = transformer_options[k]
|
|
||||||
else:
|
|
||||||
extra_options[k] = transformer_options[k]
|
|
||||||
|
|
||||||
extra_options["n_heads"] = self.n_heads
|
|
||||||
extra_options["dim_head"] = self.d_head
|
|
||||||
|
|
||||||
if self.ff_in:
|
|
||||||
x_skip = x
|
|
||||||
x = self.ff_in(self.norm_in(x))
|
|
||||||
if self.is_res:
|
|
||||||
x += x_skip
|
|
||||||
|
|
||||||
n = self.norm1(x)
|
|
||||||
if self.disable_self_attn:
|
|
||||||
context_attn1 = context
|
|
||||||
else:
|
|
||||||
context_attn1 = None
|
|
||||||
value_attn1 = None
|
|
||||||
|
|
||||||
if "attn1_patch" in transformer_patches:
|
|
||||||
patch = transformer_patches["attn1_patch"]
|
|
||||||
if context_attn1 is None:
|
|
||||||
context_attn1 = n
|
|
||||||
value_attn1 = context_attn1
|
|
||||||
for p in patch:
|
|
||||||
n, context_attn1, value_attn1 = p(n, context_attn1, value_attn1, extra_options)
|
|
||||||
|
|
||||||
if block is not None:
|
|
||||||
transformer_block = (block[0], block[1], block_index)
|
|
||||||
else:
|
|
||||||
transformer_block = None
|
|
||||||
attn1_replace_patch = transformer_patches_replace.get("attn1", {})
|
|
||||||
block_attn1 = transformer_block
|
|
||||||
if block_attn1 not in attn1_replace_patch:
|
|
||||||
block_attn1 = block
|
|
||||||
|
|
||||||
if block_attn1 in attn1_replace_patch:
|
|
||||||
if context_attn1 is None:
|
|
||||||
context_attn1 = n
|
|
||||||
value_attn1 = n
|
|
||||||
n = self.attn1.to_q(n)
|
|
||||||
context_attn1 = self.attn1.to_k(context_attn1)
|
|
||||||
value_attn1 = self.attn1.to_v(value_attn1)
|
|
||||||
n = attn1_replace_patch[block_attn1](n, context_attn1, value_attn1, extra_options)
|
|
||||||
n = self.attn1.to_out(n)
|
|
||||||
else:
|
|
||||||
n = self.attn1(n, context=context_attn1, value=value_attn1)
|
|
||||||
|
|
||||||
if "attn1_output_patch" in transformer_patches:
|
|
||||||
patch = transformer_patches["attn1_output_patch"]
|
|
||||||
for p in patch:
|
|
||||||
n = p(n, extra_options)
|
|
||||||
|
|
||||||
x += n
|
|
||||||
if "middle_patch" in transformer_patches:
|
|
||||||
patch = transformer_patches["middle_patch"]
|
|
||||||
for p in patch:
|
|
||||||
x = p(x, extra_options)
|
|
||||||
|
|
||||||
if self.attn2 is not None:
|
|
||||||
n = self.norm2(x)
|
|
||||||
if self.switch_temporal_ca_to_sa:
|
|
||||||
context_attn2 = n
|
|
||||||
else:
|
|
||||||
context_attn2 = context
|
|
||||||
value_attn2 = None
|
|
||||||
if "attn2_patch" in transformer_patches:
|
|
||||||
patch = transformer_patches["attn2_patch"]
|
|
||||||
value_attn2 = context_attn2
|
|
||||||
for p in patch:
|
|
||||||
n, context_attn2, value_attn2 = p(n, context_attn2, value_attn2, extra_options)
|
|
||||||
|
|
||||||
attn2_replace_patch = transformer_patches_replace.get("attn2", {})
|
|
||||||
block_attn2 = transformer_block
|
|
||||||
if block_attn2 not in attn2_replace_patch:
|
|
||||||
block_attn2 = block
|
|
||||||
|
|
||||||
if block_attn2 in attn2_replace_patch:
|
|
||||||
if value_attn2 is None:
|
|
||||||
value_attn2 = context_attn2
|
|
||||||
n = self.attn2.to_q(n)
|
|
||||||
context_attn2 = self.attn2.to_k(context_attn2)
|
|
||||||
value_attn2 = self.attn2.to_v(value_attn2)
|
|
||||||
n = attn2_replace_patch[block_attn2](n, context_attn2, value_attn2, extra_options)
|
|
||||||
n = self.attn2.to_out(n)
|
|
||||||
else:
|
|
||||||
n = self.attn2(n, context=context_attn2, value=value_attn2)
|
|
||||||
|
|
||||||
if "attn2_output_patch" in transformer_patches:
|
|
||||||
patch = transformer_patches["attn2_output_patch"]
|
|
||||||
for p in patch:
|
|
||||||
n = p(n, extra_options)
|
|
||||||
|
|
||||||
x += n
|
|
||||||
if self.is_res:
|
|
||||||
x_skip = x
|
|
||||||
x = self.ff(self.norm3(x))
|
|
||||||
if self.is_res:
|
|
||||||
x += x_skip
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class SpatialTransformer(nn.Module):
|
|
||||||
"""
|
|
||||||
Transformer block for image-like data in spatial axis.
|
|
||||||
First, project the input (aka embedding)
|
|
||||||
and reshape to b, t, d.
|
|
||||||
Then apply standard transformer action.
|
|
||||||
Finally, reshape to image
|
|
||||||
NEW: use_linear for more efficiency instead of the 1x1 convs
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
in_channels,
|
|
||||||
n_heads,
|
|
||||||
d_head,
|
|
||||||
depth=1,
|
|
||||||
dropout=0.,
|
|
||||||
context_dim=None,
|
|
||||||
use_checkpoint=True,
|
|
||||||
disable_self_attn=False,
|
|
||||||
use_linear=False,
|
|
||||||
video_length=None,
|
|
||||||
image_cross_attention=False,
|
|
||||||
image_cross_attention_scale_learnable=False,
|
|
||||||
device=None,
|
|
||||||
dtype=None,
|
|
||||||
operations=ops
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.in_channels = in_channels
|
|
||||||
inner_dim = n_heads * d_head
|
|
||||||
self.norm = operations.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, device=device, dtype=dtype)
|
|
||||||
if not use_linear:
|
|
||||||
self.proj_in = opeations.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0, device=device, dtype=dtype)
|
|
||||||
else:
|
|
||||||
self.proj_in = operations.Linear(in_channels, inner_dim, device=device, dtype=dtype)
|
|
||||||
|
|
||||||
attention_cls = None
|
|
||||||
self.transformer_blocks = nn.ModuleList([
|
|
||||||
BasicTransformerBlock(
|
|
||||||
inner_dim,
|
|
||||||
n_heads,
|
|
||||||
d_head,
|
|
||||||
dropout=dropout,
|
|
||||||
context_dim=context_dim,
|
|
||||||
disable_self_attn=disable_self_attn,
|
|
||||||
checkpoint=use_checkpoint,
|
|
||||||
attention_cls=attention_cls,
|
|
||||||
video_length=video_length,
|
|
||||||
image_cross_attention=image_cross_attention,
|
|
||||||
image_cross_attention_scale_learnable=image_cross_attention_scale_learnable,
|
|
||||||
device=device,
|
|
||||||
dtype=dtype
|
|
||||||
) for d in range(depth)
|
|
||||||
])
|
|
||||||
if not use_linear:
|
|
||||||
self.proj_out = zero_module(operations.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0, device=device, dtype=dtype))
|
|
||||||
else:
|
|
||||||
self.proj_out = zero_module(operations.Linear(inner_dim, in_channels, device=device, dtype=dtype))
|
|
||||||
self.use_linear = use_linear
|
|
||||||
|
|
||||||
def forward(self, x, context=None, transformer_options={}, **kwargs):
|
|
||||||
b, c, h, w = x.shape
|
|
||||||
x_in = x
|
|
||||||
x = self.norm(x)
|
|
||||||
if not self.use_linear:
|
|
||||||
x = self.proj_in(x)
|
|
||||||
x = rearrange(x, 'b c h w -> b (h w) c').contiguous()
|
|
||||||
if self.use_linear:
|
|
||||||
x = self.proj_in(x)
|
|
||||||
for i, block in enumerate(self.transformer_blocks):
|
|
||||||
transformer_options['block_index'] = i
|
|
||||||
x = block(x, context=context, **kwargs)
|
|
||||||
if self.use_linear:
|
|
||||||
x = self.proj_out(x)
|
|
||||||
x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous()
|
|
||||||
if not self.use_linear:
|
|
||||||
x = self.proj_out(x)
|
|
||||||
return x + x_in
|
|
||||||
|
|
||||||
|
|
||||||
class TemporalTransformer(nn.Module):
|
|
||||||
"""
|
|
||||||
Transformer block for image-like data in temporal axis.
|
|
||||||
First, reshape to b, t, d.
|
|
||||||
Then apply standard transformer action.
|
|
||||||
Finally, reshape to image
|
|
||||||
"""
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
in_channels,
|
|
||||||
n_heads,
|
|
||||||
d_head,
|
|
||||||
depth=1,
|
|
||||||
dropout=0.,
|
|
||||||
context_dim=None,
|
|
||||||
use_checkpoint=True,
|
|
||||||
use_linear=False,
|
|
||||||
only_self_att=True,
|
|
||||||
causal_attention=False,
|
|
||||||
causal_block_size=1,
|
|
||||||
relative_position=False,
|
|
||||||
temporal_length=None,
|
|
||||||
device=None,
|
|
||||||
dtype=None,
|
|
||||||
operations=ops
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.only_self_att = only_self_att
|
|
||||||
self.relative_position = relative_position
|
|
||||||
self.causal_attention = causal_attention
|
|
||||||
self.causal_block_size = causal_block_size
|
|
||||||
|
|
||||||
if only_self_att:
|
|
||||||
context_dim = None
|
|
||||||
|
|
||||||
self.in_channels = in_channels
|
|
||||||
inner_dim = n_heads * d_head
|
|
||||||
self.norm = operations.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, device=device, dtype=dtype)
|
|
||||||
self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0).to(device, dtype)
|
|
||||||
if not use_linear:
|
|
||||||
self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0).to(device, dtype)
|
|
||||||
else:
|
|
||||||
self.proj_in = operations.Linear(in_channels, inner_dim, device=device, dtype=dtype)
|
|
||||||
|
|
||||||
if relative_position:
|
|
||||||
assert(temporal_length is not None)
|
|
||||||
attention_cls = partial(CrossAttention, relative_position=True, temporal_length=temporal_length, device=device, dtype=dtype)
|
|
||||||
else:
|
|
||||||
attention_cls = partial(CrossAttention, temporal_length=temporal_length, device=device, dtype=dtype)
|
|
||||||
if self.causal_attention:
|
|
||||||
assert(temporal_length is not None)
|
|
||||||
self.mask = torch.tril(torch.ones([1, temporal_length, temporal_length]))
|
|
||||||
|
|
||||||
if self.only_self_att:
|
|
||||||
context_dim = None
|
|
||||||
self.transformer_blocks = nn.ModuleList([
|
|
||||||
BasicTransformerBlock(
|
|
||||||
inner_dim,
|
|
||||||
n_heads,
|
|
||||||
d_head,
|
|
||||||
dropout=dropout,
|
|
||||||
context_dim=context_dim,
|
|
||||||
attention_cls=attention_cls,
|
|
||||||
checkpoint=use_checkpoint,
|
|
||||||
device=device,
|
|
||||||
dtype=dtype
|
|
||||||
) for d in range(depth)
|
|
||||||
])
|
|
||||||
if not use_linear:
|
|
||||||
self.proj_out = zero_module(nn.Conv1d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0).to(device, dtype))
|
|
||||||
else:
|
|
||||||
self.proj_out = zero_module(operations.Linear(inner_dim, in_channels, device=device, dtype=dtype))
|
|
||||||
self.use_linear = use_linear
|
|
||||||
|
|
||||||
def forward(self, x, context=None):
|
|
||||||
b, c, t, h, w = x.shape
|
|
||||||
x_in = x
|
|
||||||
x = self.norm(x)
|
|
||||||
x = rearrange(x, 'b c t h w -> (b h w) c t').contiguous()
|
|
||||||
if not self.use_linear:
|
|
||||||
x = self.proj_in(x)
|
|
||||||
x = rearrange(x, 'bhw c t -> bhw t c').contiguous()
|
|
||||||
if self.use_linear:
|
|
||||||
x = self.proj_in(x)
|
|
||||||
|
|
||||||
temp_mask = None
|
|
||||||
if self.causal_attention:
|
|
||||||
# slice the from mask map
|
|
||||||
temp_mask = self.mask[:,:t,:t].to(x.device)
|
|
||||||
|
|
||||||
if temp_mask is not None:
|
|
||||||
mask = temp_mask.to(x.device)
|
|
||||||
mask = repeat(mask, 'l i j -> (l bhw) i j', bhw=b*h*w)
|
|
||||||
else:
|
|
||||||
mask = None
|
|
||||||
|
|
||||||
if self.only_self_att:
|
|
||||||
## note: if no context is given, cross-attention defaults to self-attention
|
|
||||||
for i, block in enumerate(self.transformer_blocks):
|
|
||||||
x = block(x, mask=mask)
|
|
||||||
x = rearrange(x, '(b hw) t c -> b hw t c', b=b).contiguous()
|
|
||||||
else:
|
|
||||||
x = rearrange(x, '(b hw) t c -> b hw t c', b=b).contiguous()
|
|
||||||
context = rearrange(context, '(b t) l con -> b t l con', t=t).contiguous()
|
|
||||||
for i, block in enumerate(self.transformer_blocks):
|
|
||||||
# calculate each batch one by one (since number in shape could not greater then 65,535 for some package)
|
|
||||||
for j in range(b):
|
|
||||||
context_j = repeat(
|
|
||||||
context[j],
|
|
||||||
't l con -> (t r) l con', r=(h * w) // t, t=t).contiguous()
|
|
||||||
## note: causal mask will not applied in cross-attention case
|
|
||||||
x[j] = block(x[j], context=context_j)
|
|
||||||
|
|
||||||
if self.use_linear:
|
|
||||||
x = self.proj_out(x)
|
|
||||||
x = rearrange(x, 'b (h w) t c -> b c t h w', h=h, w=w).contiguous()
|
|
||||||
if not self.use_linear:
|
|
||||||
x = rearrange(x, 'b hw t c -> (b hw) c t').contiguous()
|
|
||||||
x = self.proj_out(x)
|
|
||||||
x = rearrange(x, '(b h w) c t -> b c t h w', b=b, h=h, w=w).contiguous()
|
|
||||||
|
|
||||||
return x + x_in
|
|
||||||
|
|
||||||
|
|
||||||
class GEGLU(nn.Module):
|
|
||||||
def __init__(self, dim_in, dim_out, device=None, dtype=None, operations=ops):
|
|
||||||
super().__init__()
|
|
||||||
self.proj = operations.Linear(dim_in, dim_out * 2, device=device, dtype=dtype)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
x, gate = self.proj(x).chunk(2, dim=-1)
|
|
||||||
return x * F.gelu(gate)
|
|
||||||
|
|
||||||
|
|
||||||
class FeedForward(nn.Module):
|
|
||||||
def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0., device=None, dtype=None, operations=ops):
|
|
||||||
super().__init__()
|
|
||||||
inner_dim = int(dim * mult)
|
|
||||||
dim_out = default(dim_out, dim)
|
|
||||||
project_in = nn.Sequential(
|
|
||||||
operations.Linear(dim, inner_dim, device=device, dtype=dtype),
|
|
||||||
nn.GELU()
|
|
||||||
) if not glu else GEGLU(dim, inner_dim)
|
|
||||||
|
|
||||||
self.net = nn.Sequential(
|
|
||||||
project_in,
|
|
||||||
nn.Dropout(dropout),
|
|
||||||
operations.Linear(inner_dim, dim_out, device=device, dtype=dtype)
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
return self.net(x)
|
|
||||||
|
|
||||||
|
|
||||||
class LinearAttention(nn.Module):
|
|
||||||
def __init__(self, dim, heads=4, dim_head=32, device=None, dtype=None, operations=ops):
|
|
||||||
super().__init__()
|
|
||||||
self.heads = heads
|
|
||||||
hidden_dim = dim_head * heads
|
|
||||||
self.to_qkv = operations.Conv2d(dim, hidden_dim * 3, 1, bias = False, device=device, dtype=dtype)
|
|
||||||
self.to_out = operations.Conv2d(hidden_dim, dim, 1, device=device, dtype=dtype)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
b, c, h, w = x.shape
|
|
||||||
qkv = self.to_qkv(x)
|
|
||||||
q, k, v = rearrange(qkv, 'b (qkv heads c) h w -> qkv b heads c (h w)', heads = self.heads, qkv=3)
|
|
||||||
k = k.softmax(dim=-1)
|
|
||||||
context = torch.einsum('bhdn,bhen->bhde', k, v)
|
|
||||||
out = torch.einsum('bhde,bhdn->bhen', context, q)
|
|
||||||
out = rearrange(out, 'b heads c (h w) -> b (heads c) h w', heads=self.heads, h=h, w=w)
|
|
||||||
return self.to_out(out)
|
|
||||||
|
|
||||||
|
|
||||||
class SpatialSelfAttention(nn.Module):
|
|
||||||
def __init__(self, in_channels, device=None, dtype=None, operations=ops):
|
|
||||||
super().__init__()
|
|
||||||
self.in_channels = in_channels
|
|
||||||
|
|
||||||
self.norm = operations.GroupNorm(
|
|
||||||
num_groups=32,
|
|
||||||
num_channels=in_channels,
|
|
||||||
eps=1e-6,
|
|
||||||
affine=True,
|
|
||||||
device=device,
|
|
||||||
dtype=dtype
|
|
||||||
)
|
|
||||||
self.q = operations.Conv2d(
|
|
||||||
in_channels,
|
|
||||||
in_channels,
|
|
||||||
kernel_size=1,
|
|
||||||
stride=1,
|
|
||||||
padding=0,
|
|
||||||
device=device,
|
|
||||||
dtype=dtype
|
|
||||||
)
|
|
||||||
self.k = operations.Conv2d(
|
|
||||||
in_channels,
|
|
||||||
in_channels,
|
|
||||||
kernel_size=1,
|
|
||||||
stride=1,
|
|
||||||
padding=0,
|
|
||||||
device=device,
|
|
||||||
dtype=dtype
|
|
||||||
)
|
|
||||||
self.v = operations.Conv2d(
|
|
||||||
in_channels,
|
|
||||||
in_channels,
|
|
||||||
kernel_size=1,
|
|
||||||
stride=1,
|
|
||||||
padding=0,
|
|
||||||
device=device,
|
|
||||||
dtype=dtype
|
|
||||||
)
|
|
||||||
self.proj_out = operations.Conv2d(
|
|
||||||
in_channels,
|
|
||||||
in_channels,
|
|
||||||
kernel_size=1,
|
|
||||||
stride=1,
|
|
||||||
padding=0,
|
|
||||||
device=device,
|
|
||||||
dtype=dtype
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
h_ = x
|
|
||||||
h_ = self.norm(h_)
|
|
||||||
q = self.q(h_)
|
|
||||||
k = self.k(h_)
|
|
||||||
v = self.v(h_)
|
|
||||||
|
|
||||||
# compute attention
|
|
||||||
b,c,h,w = q.shape
|
|
||||||
q = rearrange(q, 'b c h w -> b (h w) c')
|
|
||||||
k = rearrange(k, 'b c h w -> b c (h w)')
|
|
||||||
w_ = torch.einsum('bij,bjk->bik', q, k)
|
|
||||||
|
|
||||||
w_ = w_ * (int(c)**(-0.5))
|
|
||||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
|
||||||
|
|
||||||
# attend to values
|
|
||||||
v = rearrange(v, 'b c h w -> b c (h w)')
|
|
||||||
w_ = rearrange(w_, 'b i j -> b j i')
|
|
||||||
h_ = torch.einsum('bij,bjk->bik', v, w_)
|
|
||||||
h_ = rearrange(h_, 'b c (h w) -> b c h w', h=h)
|
|
||||||
h_ = self.proj_out(h_)
|
|
||||||
|
|
||||||
return x+h_
|
|
||||||
@@ -1,389 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import kornia
|
|
||||||
import open_clip
|
|
||||||
from torch.utils.checkpoint import checkpoint
|
|
||||||
from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextModel
|
|
||||||
from ..common import autocast
|
|
||||||
from utils.utils import count_params
|
|
||||||
|
|
||||||
|
|
||||||
class AbstractEncoder(nn.Module):
|
|
||||||
def __init__(self):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
def encode(self, *args, **kwargs):
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
|
|
||||||
class IdentityEncoder(AbstractEncoder):
|
|
||||||
def encode(self, x):
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class ClassEmbedder(nn.Module):
|
|
||||||
def __init__(self, embed_dim, n_classes=1000, key='class', ucg_rate=0.1):
|
|
||||||
super().__init__()
|
|
||||||
self.key = key
|
|
||||||
self.embedding = nn.Embedding(n_classes, embed_dim)
|
|
||||||
self.n_classes = n_classes
|
|
||||||
self.ucg_rate = ucg_rate
|
|
||||||
|
|
||||||
def forward(self, batch, key=None, disable_dropout=False):
|
|
||||||
if key is None:
|
|
||||||
key = self.key
|
|
||||||
# this is for use in crossattn
|
|
||||||
c = batch[key][:, None]
|
|
||||||
if self.ucg_rate > 0. and not disable_dropout:
|
|
||||||
mask = 1. - torch.bernoulli(torch.ones_like(c) * self.ucg_rate)
|
|
||||||
c = mask * c + (1 - mask) * torch.ones_like(c) * (self.n_classes - 1)
|
|
||||||
c = c.long()
|
|
||||||
c = self.embedding(c)
|
|
||||||
return c
|
|
||||||
|
|
||||||
def get_unconditional_conditioning(self, bs, device="cuda"):
|
|
||||||
uc_class = self.n_classes - 1 # 1000 classes --> 0 ... 999, one extra class for ucg (class 1000)
|
|
||||||
uc = torch.ones((bs,), device=device) * uc_class
|
|
||||||
uc = {self.key: uc}
|
|
||||||
return uc
|
|
||||||
|
|
||||||
|
|
||||||
def disabled_train(self, mode=True):
|
|
||||||
"""Overwrite model.train with this function to make sure train/eval mode
|
|
||||||
does not change anymore."""
|
|
||||||
return self
|
|
||||||
|
|
||||||
|
|
||||||
class FrozenT5Embedder(AbstractEncoder):
|
|
||||||
"""Uses the T5 transformer encoder for text"""
|
|
||||||
|
|
||||||
def __init__(self, version="google/t5-v1_1-large", device="cuda", max_length=77,
|
|
||||||
freeze=True): # others are google/t5-v1_1-xl and google/t5-v1_1-xxl
|
|
||||||
super().__init__()
|
|
||||||
self.tokenizer = T5Tokenizer.from_pretrained(version)
|
|
||||||
self.transformer = T5EncoderModel.from_pretrained(version)
|
|
||||||
self.device = device
|
|
||||||
self.max_length = max_length # TODO: typical value?
|
|
||||||
if freeze:
|
|
||||||
self.freeze()
|
|
||||||
|
|
||||||
def freeze(self):
|
|
||||||
self.transformer = self.transformer.eval()
|
|
||||||
# self.train = disabled_train
|
|
||||||
for param in self.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
|
|
||||||
def forward(self, text):
|
|
||||||
batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True,
|
|
||||||
return_overflowing_tokens=False, padding="max_length", return_tensors="pt")
|
|
||||||
tokens = batch_encoding["input_ids"].to(self.device)
|
|
||||||
outputs = self.transformer(input_ids=tokens)
|
|
||||||
|
|
||||||
z = outputs.last_hidden_state
|
|
||||||
return z
|
|
||||||
|
|
||||||
def encode(self, text):
|
|
||||||
return self(text)
|
|
||||||
|
|
||||||
|
|
||||||
class FrozenCLIPEmbedder(AbstractEncoder):
|
|
||||||
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
|
||||||
LAYERS = [
|
|
||||||
"last",
|
|
||||||
"pooled",
|
|
||||||
"hidden"
|
|
||||||
]
|
|
||||||
|
|
||||||
def __init__(self, version="openai/clip-vit-large-patch14", device="cuda", max_length=77,
|
|
||||||
freeze=True, layer="last", layer_idx=None): # clip-vit-base-patch32
|
|
||||||
super().__init__()
|
|
||||||
assert layer in self.LAYERS
|
|
||||||
self.tokenizer = CLIPTokenizer.from_pretrained(version)
|
|
||||||
self.transformer = CLIPTextModel.from_pretrained(version)
|
|
||||||
self.device = device
|
|
||||||
self.max_length = max_length
|
|
||||||
if freeze:
|
|
||||||
self.freeze()
|
|
||||||
self.layer = layer
|
|
||||||
self.layer_idx = layer_idx
|
|
||||||
if layer == "hidden":
|
|
||||||
assert layer_idx is not None
|
|
||||||
assert 0 <= abs(layer_idx) <= 12
|
|
||||||
|
|
||||||
def freeze(self):
|
|
||||||
self.transformer = self.transformer.eval()
|
|
||||||
# self.train = disabled_train
|
|
||||||
for param in self.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
|
|
||||||
def forward(self, text):
|
|
||||||
batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True,
|
|
||||||
return_overflowing_tokens=False, padding="max_length", return_tensors="pt")
|
|
||||||
tokens = batch_encoding["input_ids"].to(self.device)
|
|
||||||
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden")
|
|
||||||
if self.layer == "last":
|
|
||||||
z = outputs.last_hidden_state
|
|
||||||
elif self.layer == "pooled":
|
|
||||||
z = outputs.pooler_output[:, None, :]
|
|
||||||
else:
|
|
||||||
z = outputs.hidden_states[self.layer_idx]
|
|
||||||
return z
|
|
||||||
|
|
||||||
def encode(self, text):
|
|
||||||
return self(text)
|
|
||||||
|
|
||||||
|
|
||||||
class ClipImageEmbedder(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
model,
|
|
||||||
jit=False,
|
|
||||||
device='cuda' if torch.cuda.is_available() else 'cpu',
|
|
||||||
antialias=True,
|
|
||||||
ucg_rate=0.
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
from clip import load as load_clip
|
|
||||||
self.model, _ = load_clip(name=model, device=device, jit=jit)
|
|
||||||
|
|
||||||
self.antialias = antialias
|
|
||||||
|
|
||||||
self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False)
|
|
||||||
self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False)
|
|
||||||
self.ucg_rate = ucg_rate
|
|
||||||
|
|
||||||
def preprocess(self, x):
|
|
||||||
# normalize to [0,1]
|
|
||||||
x = kornia.geometry.resize(x, (224, 224),
|
|
||||||
interpolation='bicubic', align_corners=True,
|
|
||||||
antialias=self.antialias)
|
|
||||||
x = (x + 1.) / 2.
|
|
||||||
# re-normalize according to clip
|
|
||||||
x = kornia.enhance.normalize(x, self.mean, self.std)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def forward(self, x, no_dropout=False):
|
|
||||||
# x is assumed to be in range [-1,1]
|
|
||||||
out = self.model.encode_image(self.preprocess(x))
|
|
||||||
out = out.to(x.dtype)
|
|
||||||
if self.ucg_rate > 0. and not no_dropout:
|
|
||||||
out = torch.bernoulli((1. - self.ucg_rate) * torch.ones(out.shape[0], device=out.device))[:, None] * out
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class FrozenOpenCLIPEmbedder(AbstractEncoder):
|
|
||||||
"""
|
|
||||||
Uses the OpenCLIP transformer encoder for text
|
|
||||||
"""
|
|
||||||
LAYERS = [
|
|
||||||
# "pooled",
|
|
||||||
"last",
|
|
||||||
"penultimate"
|
|
||||||
]
|
|
||||||
|
|
||||||
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77,
|
|
||||||
freeze=True, layer="last"):
|
|
||||||
super().__init__()
|
|
||||||
assert layer in self.LAYERS
|
|
||||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=version)
|
|
||||||
del model.visual
|
|
||||||
self.model = model
|
|
||||||
|
|
||||||
self.device = device
|
|
||||||
self.max_length = max_length
|
|
||||||
if freeze:
|
|
||||||
self.freeze()
|
|
||||||
self.layer = layer
|
|
||||||
if self.layer == "last":
|
|
||||||
self.layer_idx = 0
|
|
||||||
elif self.layer == "penultimate":
|
|
||||||
self.layer_idx = 1
|
|
||||||
else:
|
|
||||||
raise NotImplementedError()
|
|
||||||
|
|
||||||
def freeze(self):
|
|
||||||
self.model = self.model.eval()
|
|
||||||
for param in self.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
|
|
||||||
def forward(self, text):
|
|
||||||
tokens = open_clip.tokenize(text) ## all clip models use 77 as context length
|
|
||||||
z = self.encode_with_transformer(tokens.to(self.device))
|
|
||||||
return z
|
|
||||||
|
|
||||||
def encode_with_transformer(self, text):
|
|
||||||
x = self.model.token_embedding(text) # [batch_size, n_ctx, d_model]
|
|
||||||
x = x + self.model.positional_embedding
|
|
||||||
x = x.permute(1, 0, 2) # NLD -> LND
|
|
||||||
x = self.text_transformer_forward(x, attn_mask=self.model.attn_mask)
|
|
||||||
x = x.permute(1, 0, 2) # LND -> NLD
|
|
||||||
x = self.model.ln_final(x)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def text_transformer_forward(self, x: torch.Tensor, attn_mask=None):
|
|
||||||
for i, r in enumerate(self.model.transformer.resblocks):
|
|
||||||
if i == len(self.model.transformer.resblocks) - self.layer_idx:
|
|
||||||
break
|
|
||||||
if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting():
|
|
||||||
x = checkpoint(r, x, attn_mask)
|
|
||||||
else:
|
|
||||||
x = r(x, attn_mask=attn_mask)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def encode(self, text):
|
|
||||||
return self(text)
|
|
||||||
|
|
||||||
|
|
||||||
class FrozenOpenCLIPImageEmbedder(AbstractEncoder):
|
|
||||||
"""
|
|
||||||
Uses the OpenCLIP vision transformer encoder for images
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77,
|
|
||||||
freeze=True, layer="pooled", antialias=True, ucg_rate=0.):
|
|
||||||
super().__init__()
|
|
||||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'),
|
|
||||||
pretrained=version, )
|
|
||||||
del model.transformer
|
|
||||||
self.model = model
|
|
||||||
# self.mapper = torch.nn.Linear(1280, 1024)
|
|
||||||
self.device = device
|
|
||||||
self.max_length = max_length
|
|
||||||
if freeze:
|
|
||||||
self.freeze()
|
|
||||||
self.layer = layer
|
|
||||||
if self.layer == "penultimate":
|
|
||||||
raise NotImplementedError()
|
|
||||||
self.layer_idx = 1
|
|
||||||
|
|
||||||
self.antialias = antialias
|
|
||||||
|
|
||||||
self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False)
|
|
||||||
self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False)
|
|
||||||
self.ucg_rate = ucg_rate
|
|
||||||
|
|
||||||
def preprocess(self, x):
|
|
||||||
# normalize to [0,1]
|
|
||||||
x = kornia.geometry.resize(x, (224, 224),
|
|
||||||
interpolation='bicubic', align_corners=True,
|
|
||||||
antialias=self.antialias)
|
|
||||||
x = (x + 1.) / 2.
|
|
||||||
# renormalize according to clip
|
|
||||||
x = kornia.enhance.normalize(x, self.mean, self.std)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def freeze(self):
|
|
||||||
self.model = self.model.eval()
|
|
||||||
for param in self.model.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
|
|
||||||
@autocast
|
|
||||||
def forward(self, image, no_dropout=False):
|
|
||||||
z = self.encode_with_vision_transformer(image)
|
|
||||||
if self.ucg_rate > 0. and not no_dropout:
|
|
||||||
z = torch.bernoulli((1. - self.ucg_rate) * torch.ones(z.shape[0], device=z.device))[:, None] * z
|
|
||||||
return z
|
|
||||||
|
|
||||||
def encode_with_vision_transformer(self, img):
|
|
||||||
img = self.preprocess(img)
|
|
||||||
x = self.model.visual(img)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def encode(self, text):
|
|
||||||
return self(text)
|
|
||||||
|
|
||||||
class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder):
|
|
||||||
"""
|
|
||||||
Uses the OpenCLIP vision transformer encoder for images
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda",
|
|
||||||
freeze=True, layer="pooled", antialias=True):
|
|
||||||
super().__init__()
|
|
||||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'),
|
|
||||||
pretrained=version, )
|
|
||||||
del model.transformer
|
|
||||||
self.model = model
|
|
||||||
self.device = device
|
|
||||||
|
|
||||||
if freeze:
|
|
||||||
self.freeze()
|
|
||||||
self.layer = layer
|
|
||||||
if self.layer == "penultimate":
|
|
||||||
raise NotImplementedError()
|
|
||||||
self.layer_idx = 1
|
|
||||||
|
|
||||||
self.antialias = antialias
|
|
||||||
|
|
||||||
self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False)
|
|
||||||
self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False)
|
|
||||||
|
|
||||||
|
|
||||||
def preprocess(self, x):
|
|
||||||
# normalize to [0,1]
|
|
||||||
x = kornia.geometry.resize(x, (224, 224),
|
|
||||||
interpolation='bicubic', align_corners=True,
|
|
||||||
antialias=self.antialias)
|
|
||||||
x = (x + 1.) / 2.
|
|
||||||
# renormalize according to clip
|
|
||||||
x = kornia.enhance.normalize(x, self.mean, self.std)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def freeze(self):
|
|
||||||
self.model = self.model.eval()
|
|
||||||
for param in self.model.parameters():
|
|
||||||
param.requires_grad = False
|
|
||||||
|
|
||||||
def forward(self, image, no_dropout=False):
|
|
||||||
## image: b c h w
|
|
||||||
z = self.encode_with_vision_transformer(image)
|
|
||||||
return z
|
|
||||||
|
|
||||||
def encode_with_vision_transformer(self, x):
|
|
||||||
x = self.preprocess(x)
|
|
||||||
|
|
||||||
# to patches - whether to use dual patchnorm - https://arxiv.org/abs/2302.01327v1
|
|
||||||
if self.model.visual.input_patchnorm:
|
|
||||||
# einops - rearrange(x, 'b c (h p1) (w p2) -> b (h w) (c p1 p2)')
|
|
||||||
x = x.reshape(x.shape[0], x.shape[1], self.model.visual.grid_size[0], self.model.visual.patch_size[0], self.model.visual.grid_size[1], self.model.visual.patch_size[1])
|
|
||||||
x = x.permute(0, 2, 4, 1, 3, 5)
|
|
||||||
x = x.reshape(x.shape[0], self.model.visual.grid_size[0] * self.model.visual.grid_size[1], -1)
|
|
||||||
x = self.model.visual.patchnorm_pre_ln(x)
|
|
||||||
x = self.model.visual.conv1(x)
|
|
||||||
else:
|
|
||||||
x = self.model.visual.conv1(x) # shape = [*, width, grid, grid]
|
|
||||||
x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2]
|
|
||||||
x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width]
|
|
||||||
|
|
||||||
# class embeddings and positional embeddings
|
|
||||||
x = torch.cat(
|
|
||||||
[self.model.visual.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device),
|
|
||||||
x], dim=1) # shape = [*, grid ** 2 + 1, width]
|
|
||||||
x = x + self.model.visual.positional_embedding.to(x.dtype)
|
|
||||||
|
|
||||||
# a patch_dropout of 0. would mean it is disabled and this function would do nothing but return what was passed in
|
|
||||||
x = self.model.visual.patch_dropout(x)
|
|
||||||
x = self.model.visual.ln_pre(x)
|
|
||||||
|
|
||||||
x = x.permute(1, 0, 2) # NLD -> LND
|
|
||||||
x = self.model.visual.transformer(x)
|
|
||||||
x = x.permute(1, 0, 2) # LND -> NLD
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
class FrozenCLIPT5Encoder(AbstractEncoder):
|
|
||||||
def __init__(self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", device="cuda",
|
|
||||||
clip_max_length=77, t5_max_length=77):
|
|
||||||
super().__init__()
|
|
||||||
self.clip_encoder = FrozenCLIPEmbedder(clip_version, device, max_length=clip_max_length)
|
|
||||||
self.t5_encoder = FrozenT5Embedder(t5_version, device, max_length=t5_max_length)
|
|
||||||
print(f"{self.clip_encoder.__class__.__name__} has {count_params(self.clip_encoder) * 1.e-6:.2f} M parameters, "
|
|
||||||
f"{self.t5_encoder.__class__.__name__} comes with {count_params(self.t5_encoder) * 1.e-6:.2f} M params.")
|
|
||||||
|
|
||||||
def encode(self, text):
|
|
||||||
return self(text)
|
|
||||||
|
|
||||||
def forward(self, text):
|
|
||||||
clip_z = self.clip_encoder.encode(text)
|
|
||||||
t5_z = self.t5_encoder.encode(text)
|
|
||||||
return [clip_z, t5_z]
|
|
||||||
@@ -1,145 +0,0 @@
|
|||||||
# modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py
|
|
||||||
# and https://github.com/lucidrains/imagen-pytorch/blob/main/imagen_pytorch/imagen_pytorch.py
|
|
||||||
# and https://github.com/tencent-ailab/IP-Adapter/blob/main/ip_adapter/resampler.py
|
|
||||||
import math
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
|
|
||||||
|
|
||||||
class ImageProjModel(nn.Module):
|
|
||||||
"""Projection Model"""
|
|
||||||
def __init__(self, cross_attention_dim=1024, clip_embeddings_dim=1024, clip_extra_context_tokens=4):
|
|
||||||
super().__init__()
|
|
||||||
self.cross_attention_dim = cross_attention_dim
|
|
||||||
self.clip_extra_context_tokens = clip_extra_context_tokens
|
|
||||||
self.proj = nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim)
|
|
||||||
self.norm = nn.LayerNorm(cross_attention_dim)
|
|
||||||
|
|
||||||
def forward(self, image_embeds):
|
|
||||||
#embeds = image_embeds
|
|
||||||
embeds = image_embeds.type(list(self.proj.parameters())[0].dtype)
|
|
||||||
clip_extra_context_tokens = self.proj(embeds).reshape(-1, self.clip_extra_context_tokens, self.cross_attention_dim)
|
|
||||||
clip_extra_context_tokens = self.norm(clip_extra_context_tokens)
|
|
||||||
return clip_extra_context_tokens
|
|
||||||
|
|
||||||
|
|
||||||
# FFN
|
|
||||||
def FeedForward(dim, mult=4):
|
|
||||||
inner_dim = int(dim * mult)
|
|
||||||
return nn.Sequential(
|
|
||||||
nn.LayerNorm(dim),
|
|
||||||
nn.Linear(dim, inner_dim, bias=False),
|
|
||||||
nn.GELU(),
|
|
||||||
nn.Linear(inner_dim, dim, bias=False),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def reshape_tensor(x, heads):
|
|
||||||
bs, length, width = x.shape
|
|
||||||
#(bs, length, width) --> (bs, length, n_heads, dim_per_head)
|
|
||||||
x = x.view(bs, length, heads, -1)
|
|
||||||
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
|
|
||||||
x = x.transpose(1, 2)
|
|
||||||
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
|
|
||||||
x = x.reshape(bs, heads, length, -1)
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class PerceiverAttention(nn.Module):
|
|
||||||
def __init__(self, *, dim, dim_head=64, heads=8):
|
|
||||||
super().__init__()
|
|
||||||
self.scale = dim_head**-0.5
|
|
||||||
self.dim_head = dim_head
|
|
||||||
self.heads = heads
|
|
||||||
inner_dim = dim_head * heads
|
|
||||||
|
|
||||||
self.norm1 = nn.LayerNorm(dim)
|
|
||||||
self.norm2 = nn.LayerNorm(dim)
|
|
||||||
|
|
||||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
|
||||||
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
|
|
||||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
|
||||||
|
|
||||||
|
|
||||||
def forward(self, x, latents):
|
|
||||||
"""
|
|
||||||
Args:
|
|
||||||
x (torch.Tensor): image features
|
|
||||||
shape (b, n1, D)
|
|
||||||
latent (torch.Tensor): latent features
|
|
||||||
shape (b, n2, D)
|
|
||||||
"""
|
|
||||||
x = self.norm1(x)
|
|
||||||
latents = self.norm2(latents)
|
|
||||||
|
|
||||||
b, l, _ = latents.shape
|
|
||||||
|
|
||||||
q = self.to_q(latents)
|
|
||||||
kv_input = torch.cat((x, latents), dim=-2)
|
|
||||||
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
|
||||||
|
|
||||||
q = reshape_tensor(q, self.heads)
|
|
||||||
k = reshape_tensor(k, self.heads)
|
|
||||||
v = reshape_tensor(v, self.heads)
|
|
||||||
|
|
||||||
# attention
|
|
||||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
|
||||||
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
|
|
||||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
|
||||||
out = weight @ v
|
|
||||||
|
|
||||||
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
|
|
||||||
|
|
||||||
return self.to_out(out)
|
|
||||||
|
|
||||||
|
|
||||||
class Resampler(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim=1024,
|
|
||||||
depth=8,
|
|
||||||
dim_head=64,
|
|
||||||
heads=16,
|
|
||||||
num_queries=8,
|
|
||||||
embedding_dim=768,
|
|
||||||
output_dim=1024,
|
|
||||||
ff_mult=4,
|
|
||||||
video_length=None, # using frame-wise version or not
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
## queries for a single frame / image
|
|
||||||
self.num_queries = num_queries
|
|
||||||
self.video_length = video_length
|
|
||||||
|
|
||||||
## <num_queries> queries for each frame
|
|
||||||
if video_length is not None:
|
|
||||||
num_queries = num_queries * video_length
|
|
||||||
|
|
||||||
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5)
|
|
||||||
self.proj_in = nn.Linear(embedding_dim, dim)
|
|
||||||
self.proj_out = nn.Linear(dim, output_dim)
|
|
||||||
self.norm_out = nn.LayerNorm(output_dim)
|
|
||||||
|
|
||||||
self.layers = nn.ModuleList([])
|
|
||||||
for _ in range(depth):
|
|
||||||
self.layers.append(
|
|
||||||
nn.ModuleList(
|
|
||||||
[
|
|
||||||
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
|
|
||||||
FeedForward(dim=dim, mult=ff_mult),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
latents = self.latents.repeat(x.size(0), 1, 1) ## B (T L) C
|
|
||||||
x = self.proj_in(x)
|
|
||||||
|
|
||||||
for attn, ff in self.layers:
|
|
||||||
latents = attn(x, latents) + latents
|
|
||||||
latents = ff(latents) + latents
|
|
||||||
|
|
||||||
latents = self.proj_out(latents)
|
|
||||||
latents = self.norm_out(latents) # B L C or B (T L) C
|
|
||||||
|
|
||||||
return latents
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,822 +0,0 @@
|
|||||||
from functools import partial
|
|
||||||
from abc import abstractmethod
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
from einops import rearrange
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from ...models.utils_diffusion import timestep_embedding
|
|
||||||
from ...common import checkpoint
|
|
||||||
from ...basics import (
|
|
||||||
zero_module,
|
|
||||||
conv_nd,
|
|
||||||
linear,
|
|
||||||
avg_pool_nd,
|
|
||||||
normalization
|
|
||||||
)
|
|
||||||
from ...modules.attention import SpatialTransformer, TemporalTransformer
|
|
||||||
import comfy.ops
|
|
||||||
import logging
|
|
||||||
|
|
||||||
ops = comfy.ops.disable_weight_init
|
|
||||||
|
|
||||||
class TimestepBlock(nn.Module):
|
|
||||||
"""
|
|
||||||
Any module where forward() takes timestep embeddings as a second argument.
|
|
||||||
"""
|
|
||||||
@abstractmethod
|
|
||||||
def forward(self, x, emb):
|
|
||||||
"""
|
|
||||||
Apply the module to `x` given `emb` timestep embeddings.
|
|
||||||
"""
|
|
||||||
|
|
||||||
#This is needed because accelerate makes a copy of transformer_options which breaks "transformer_index"
|
|
||||||
def forward_timestep_embed(ts, x, emb, context=None, batch_size=None, transformer_options={}):
|
|
||||||
for layer in ts:
|
|
||||||
if isinstance(layer, TimestepBlock):
|
|
||||||
x = layer(x, emb, batch_size=batch_size)
|
|
||||||
elif isinstance(layer, SpatialTransformer):
|
|
||||||
x = layer(x, context)
|
|
||||||
if "transformer_index" in transformer_options:
|
|
||||||
transformer_options["transformer_index"] += 1
|
|
||||||
elif isinstance(layer, TemporalTransformer):
|
|
||||||
x = rearrange(x, '(b f) c h w -> b c f h w', b=batch_size)
|
|
||||||
x = layer(x, context)
|
|
||||||
if "transformer_index" in transformer_options:
|
|
||||||
transformer_options["transformer_index"] += 1
|
|
||||||
x = rearrange(x, 'b c f h w -> (b f) c h w')
|
|
||||||
else:
|
|
||||||
x = layer(x)
|
|
||||||
return x
|
|
||||||
|
|
||||||
class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
|
|
||||||
"""
|
|
||||||
A sequential module that passes timestep embeddings to the children that
|
|
||||||
support it as an extra input.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def forward(self, *args, **kwargs):
|
|
||||||
return forward_timestep_embed(self, *args, **kwargs)
|
|
||||||
|
|
||||||
class Downsample(nn.Module):
|
|
||||||
"""
|
|
||||||
A downsampling layer with an optional convolution.
|
|
||||||
:param channels: channels in the inputs and outputs.
|
|
||||||
:param use_conv: a bool determining if a convolution is applied.
|
|
||||||
:param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then
|
|
||||||
downsampling occurs in the inner-two dimensions.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, channels, use_conv, dims=2, out_channels=None, padding=1, dtype=None, device=None, operations=ops):
|
|
||||||
super().__init__()
|
|
||||||
self.channels = channels
|
|
||||||
self.out_channels = out_channels or channels
|
|
||||||
self.use_conv = use_conv
|
|
||||||
self.dims = dims
|
|
||||||
stride = 2 if dims != 3 else (1, 2, 2)
|
|
||||||
if use_conv:
|
|
||||||
self.op = operations.conv_nd(
|
|
||||||
dims, self.channels, self.out_channels, 3, stride=stride, padding=padding
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
assert self.channels == self.out_channels
|
|
||||||
self.op = avg_pool_nd(dims, kernel_size=stride, stride=stride)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
assert x.shape[1] == self.channels
|
|
||||||
return self.op(x)
|
|
||||||
|
|
||||||
class Upsample(nn.Module):
|
|
||||||
"""
|
|
||||||
An upsampling layer with an optional convolution.
|
|
||||||
:param channels: channels in the inputs and outputs.
|
|
||||||
:param use_conv: a bool determining if a convolution is applied.
|
|
||||||
:param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then
|
|
||||||
upsampling occurs in the inner-two dimensions.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, channels, use_conv, dims=2, out_channels=None, padding=1, dtype=None, device=None, operations=ops):
|
|
||||||
super().__init__()
|
|
||||||
self.channels = channels
|
|
||||||
self.out_channels = out_channels or channels
|
|
||||||
self.use_conv = use_conv
|
|
||||||
self.dims = dims
|
|
||||||
if use_conv:
|
|
||||||
self.conv = operations.conv_nd(dims, self.channels, self.out_channels, 3, padding=padding, dtype=dtype, device=device)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
assert x.shape[1] == self.channels
|
|
||||||
if self.dims == 3:
|
|
||||||
x = F.interpolate(x, (x.shape[2], x.shape[3] * 2, x.shape[4] * 2), mode='nearest')
|
|
||||||
else:
|
|
||||||
x = F.interpolate(x, scale_factor=2, mode='nearest')
|
|
||||||
if self.use_conv:
|
|
||||||
x = self.conv(x)
|
|
||||||
return x
|
|
||||||
|
|
||||||
class ResBlock(TimestepBlock):
|
|
||||||
"""
|
|
||||||
A residual block that can optionally change the number of channels.
|
|
||||||
:param channels: the number of input channels.
|
|
||||||
:param emb_channels: the number of timestep embedding channels.
|
|
||||||
:param dropout: the rate of dropout.
|
|
||||||
:param out_channels: if specified, the number of out channels.
|
|
||||||
:param use_conv: if True and out_channels is specified, use a spatial
|
|
||||||
convolution instead of a smaller 1x1 convolution to change the
|
|
||||||
channels in the skip connection.
|
|
||||||
:param dims: determines if the signal is 1D, 2D, or 3D.
|
|
||||||
:param up: if True, use this block for upsampling.
|
|
||||||
:param down: if True, use this block for downsampling.
|
|
||||||
:param use_temporal_conv: if True, use the temporal convolution.
|
|
||||||
:param use_image_dataset: if True, the temporal parameters will not be optimized.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
channels,
|
|
||||||
emb_channels,
|
|
||||||
dropout,
|
|
||||||
out_channels=None,
|
|
||||||
use_scale_shift_norm=False,
|
|
||||||
dims=2,
|
|
||||||
use_checkpoint=False,
|
|
||||||
use_conv=False,
|
|
||||||
up=False,
|
|
||||||
down=False,
|
|
||||||
kernel_size=3,
|
|
||||||
use_temporal_conv=False,
|
|
||||||
tempspatial_aware=False,
|
|
||||||
dtype=None,
|
|
||||||
device=None,
|
|
||||||
operations=ops
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.channels = channels
|
|
||||||
self.emb_channels = emb_channels
|
|
||||||
self.dropout = dropout
|
|
||||||
self.out_channels = out_channels or channels
|
|
||||||
self.use_conv = use_conv
|
|
||||||
self.use_checkpoint = use_checkpoint
|
|
||||||
self.use_scale_shift_norm = use_scale_shift_norm
|
|
||||||
self.use_temporal_conv = use_temporal_conv
|
|
||||||
|
|
||||||
if isinstance(kernel_size, list):
|
|
||||||
padding =[k // 2 for k in kernel_size]
|
|
||||||
else:
|
|
||||||
padding = kernel_size // 2
|
|
||||||
|
|
||||||
# operations used in normalization function
|
|
||||||
self.in_layers = nn.Sequential(
|
|
||||||
normalization(channels, dtype=dtype, device=device),
|
|
||||||
nn.SiLU(),
|
|
||||||
operations.conv_nd(dims, channels, self.out_channels, 3, padding=1, dtype=dtype, device=device),
|
|
||||||
)
|
|
||||||
|
|
||||||
self.updown = up or down
|
|
||||||
|
|
||||||
if up:
|
|
||||||
self.h_upd = Upsample(channels, False, dims, dtype=dtype, device=device)
|
|
||||||
self.x_upd = Upsample(channels, False, dims, dtype=dtype, device=device)
|
|
||||||
elif down:
|
|
||||||
self.h_upd = Downsample(channels, False, dims, dtype=dtype, device=device)
|
|
||||||
self.x_upd = Downsample(channels, False, dims, dtype=dtype, device=device)
|
|
||||||
else:
|
|
||||||
self.h_upd = self.x_upd = nn.Identity()
|
|
||||||
|
|
||||||
self.emb_layers = nn.Sequential(
|
|
||||||
nn.SiLU(),
|
|
||||||
operations.Linear(
|
|
||||||
emb_channels,
|
|
||||||
2 * self.out_channels if use_scale_shift_norm else self.out_channels,
|
|
||||||
dtype=dtype,
|
|
||||||
device=device
|
|
||||||
),
|
|
||||||
)
|
|
||||||
self.out_layers = nn.Sequential(
|
|
||||||
normalization(self.out_channels, dtype=dtype, device=device),
|
|
||||||
nn.SiLU(),
|
|
||||||
nn.Dropout(p=dropout),
|
|
||||||
zero_module(operations.Conv2d(self.out_channels, self.out_channels, 3, padding=1, dtype=dtype, device=device)),
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.out_channels == channels:
|
|
||||||
self.skip_connection = nn.Identity()
|
|
||||||
elif use_conv:
|
|
||||||
self.skip_connection = operations.conv_nd(dims, channels, self.out_channels, 3, padding=1, dtype=dtype, device=device)
|
|
||||||
else:
|
|
||||||
self.skip_connection = operations.conv_nd(dims, channels, self.out_channels, 1, dtype=dtype, device=device)
|
|
||||||
|
|
||||||
if self.use_temporal_conv:
|
|
||||||
self.temopral_conv = TemporalConvBlock(
|
|
||||||
self.out_channels,
|
|
||||||
self.out_channels,
|
|
||||||
dropout=0.1,
|
|
||||||
spatial_aware=tempspatial_aware,
|
|
||||||
dtype=dtype,
|
|
||||||
device=device
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(self, x, emb, batch_size=None):
|
|
||||||
"""
|
|
||||||
Apply the block to a Tensor, conditioned on a timestep embedding.
|
|
||||||
:param x: an [N x C x ...] Tensor of features.
|
|
||||||
:param emb: an [N x emb_channels] Tensor of timestep embeddings.
|
|
||||||
:return: an [N x C x ...] Tensor of outputs.
|
|
||||||
"""
|
|
||||||
input_tuple = (x, emb)
|
|
||||||
if batch_size:
|
|
||||||
forward_batchsize = partial(self._forward, batch_size=batch_size)
|
|
||||||
return checkpoint(forward_batchsize, input_tuple, self.parameters(), self.use_checkpoint)
|
|
||||||
return checkpoint(self._forward, input_tuple, self.parameters(), self.use_checkpoint)
|
|
||||||
|
|
||||||
def _forward(self, x, emb, batch_size=None):
|
|
||||||
if self.updown:
|
|
||||||
in_rest, in_conv = self.in_layers[:-1], self.in_layers[-1]
|
|
||||||
h = in_rest(x)
|
|
||||||
h = self.h_upd(h)
|
|
||||||
x = self.x_upd(x)
|
|
||||||
h = in_conv(h)
|
|
||||||
else:
|
|
||||||
h = self.in_layers(x)
|
|
||||||
emb_out = self.emb_layers(emb).type(h.dtype)
|
|
||||||
while len(emb_out.shape) < len(h.shape):
|
|
||||||
emb_out = emb_out[..., None]
|
|
||||||
if self.use_scale_shift_norm:
|
|
||||||
out_norm, out_rest = self.out_layers[0], self.out_layers[1:]
|
|
||||||
scale, shift = torch.chunk(emb_out, 2, dim=1)
|
|
||||||
h = out_norm(h) * (1 + scale) + shift
|
|
||||||
h = out_rest(h)
|
|
||||||
else:
|
|
||||||
h = h + emb_out
|
|
||||||
h = self.out_layers(h)
|
|
||||||
h = self.skip_connection(x) + h
|
|
||||||
|
|
||||||
if self.use_temporal_conv and batch_size:
|
|
||||||
h = rearrange(h, '(b t) c h w -> b c t h w', b=batch_size)
|
|
||||||
h = self.temopral_conv(h)
|
|
||||||
h = rearrange(h, 'b c t h w -> (b t) c h w')
|
|
||||||
return h
|
|
||||||
|
|
||||||
class TemporalConvBlock(nn.Module):
|
|
||||||
"""
|
|
||||||
Adapted from modelscope: https://github.com/modelscope/modelscope/blob/master/modelscope/models/multi_modal/video_synthesis/unet_sd.py
|
|
||||||
"""
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
in_channels,
|
|
||||||
out_channels=None,
|
|
||||||
dropout=0.0,
|
|
||||||
spatial_aware=False,
|
|
||||||
dtype=None,
|
|
||||||
device=None,
|
|
||||||
operations=ops
|
|
||||||
):
|
|
||||||
super(TemporalConvBlock, self).__init__()
|
|
||||||
if out_channels is None:
|
|
||||||
out_channels = in_channels
|
|
||||||
self.in_channels = in_channels
|
|
||||||
self.out_channels = out_channels
|
|
||||||
th_kernel_shape = (3, 1, 1) if not spatial_aware else (3, 3, 1)
|
|
||||||
th_padding_shape = (1, 0, 0) if not spatial_aware else (1, 1, 0)
|
|
||||||
tw_kernel_shape = (3, 1, 1) if not spatial_aware else (3, 1, 3)
|
|
||||||
tw_padding_shape = (1, 0, 0) if not spatial_aware else (1, 0, 1)
|
|
||||||
|
|
||||||
# conv layers
|
|
||||||
self.conv1 = nn.Sequential(
|
|
||||||
operations.GroupNorm(32, in_channels, device=device, dtype=dtype), nn.SiLU(),
|
|
||||||
operations.Conv3d(in_channels, out_channels, th_kernel_shape, padding=th_padding_shape, device=device, dtype=dtype))
|
|
||||||
self.conv2 = nn.Sequential(
|
|
||||||
operations.GroupNorm(32, out_channels, device=device, dtype=dtype), nn.SiLU(), nn.Dropout(dropout),
|
|
||||||
operations.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape, device=device, dtype=dtype))
|
|
||||||
self.conv3 = nn.Sequential(
|
|
||||||
operations.GroupNorm(32, out_channels, device=device, dtype=dtype), nn.SiLU(), nn.Dropout(dropout),
|
|
||||||
operations.Conv3d(out_channels, in_channels, th_kernel_shape, padding=th_padding_shape, device=device, dtype=dtype))
|
|
||||||
self.conv4 = nn.Sequential(
|
|
||||||
operations.GroupNorm(32, out_channels, device=device, dtype=dtype), nn.SiLU(), nn.Dropout(dropout),
|
|
||||||
operations.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape, device=device, dtype=dtype))
|
|
||||||
|
|
||||||
# zero out the last layer params,so the conv block is identity
|
|
||||||
nn.init.zeros_(self.conv4[-1].weight)
|
|
||||||
nn.init.zeros_(self.conv4[-1].bias)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
identity = x
|
|
||||||
x = self.conv1(x)
|
|
||||||
x = self.conv2(x)
|
|
||||||
x = self.conv3(x)
|
|
||||||
x = self.conv4(x)
|
|
||||||
|
|
||||||
return identity + x
|
|
||||||
|
|
||||||
def context_processor(context, t, img_emb=None, temporal_size=16, concat_only=False, disable_concat=False):
|
|
||||||
if disable_concat:
|
|
||||||
return context
|
|
||||||
|
|
||||||
## repeat t times for context [(b t) 77 768] & time embedding
|
|
||||||
## check if we use per-frame image conditioning
|
|
||||||
|
|
||||||
if img_emb is not None:
|
|
||||||
context = torch.cat([context, img_emb.to(context.device, context.dtype)], dim=1)
|
|
||||||
|
|
||||||
if concat_only:
|
|
||||||
return context
|
|
||||||
|
|
||||||
b, l_context, _ = context.shape
|
|
||||||
if l_context == 77 + t * temporal_size:
|
|
||||||
context_text, context_img = context[:,:77,:], context[:,77:,:]
|
|
||||||
context_text = context_text.repeat_interleave(repeats=t, dim=0)
|
|
||||||
context_img = rearrange(context_img, 'b (t l) c -> (b t) l c', t=t)
|
|
||||||
context = torch.cat([context_text, context_img], dim=1)
|
|
||||||
else:
|
|
||||||
context = context.repeat_interleave(repeats=t, dim=0)
|
|
||||||
|
|
||||||
return context
|
|
||||||
|
|
||||||
def apply_control(h, control, name, cond_idx=None):
|
|
||||||
if control is not None and name in control and len(control[name]) > 0:
|
|
||||||
frames = h.shape[0]
|
|
||||||
ctrl = control[name].pop()
|
|
||||||
if ctrl is not None:
|
|
||||||
try:
|
|
||||||
if cond_idx is not None and ctrl.shape[0] > frames:
|
|
||||||
ctrl_frames_list = list(range(ctrl.shape[0]))
|
|
||||||
ctrl_frames = len(ctrl_frames_list)
|
|
||||||
|
|
||||||
idxs = (
|
|
||||||
ctrl_frames_list[ctrl_frames // 2:] if cond_idx == 0 else \
|
|
||||||
ctrl_frames_list[:ctrl_frames // 2]
|
|
||||||
)
|
|
||||||
|
|
||||||
ctrl = ctrl[idxs]
|
|
||||||
|
|
||||||
h += ctrl
|
|
||||||
except Exception as e:
|
|
||||||
if h.shape != ctrl.shape:
|
|
||||||
logging.warning(
|
|
||||||
"warning control could not be applied {} {}".format(h.shape, ctrl.shape)
|
|
||||||
)
|
|
||||||
logging.warning(e)
|
|
||||||
return h
|
|
||||||
|
|
||||||
class UNetModel(nn.Module):
|
|
||||||
"""
|
|
||||||
The full UNet model with attention and timestep embedding.
|
|
||||||
:param in_channels: in_channels in the input Tensor.
|
|
||||||
:param model_channels: base channel count for the model.
|
|
||||||
:param out_channels: channels in the output Tensor.
|
|
||||||
:param num_res_blocks: number of residual blocks per downsample.
|
|
||||||
:param attention_resolutions: a collection of downsample rates at which
|
|
||||||
attention will take place. May be a set, list, or tuple.
|
|
||||||
For example, if this contains 4, then at 4x downsampling, attention
|
|
||||||
will be used.
|
|
||||||
:param dropout: the dropout probability.
|
|
||||||
:param channel_mult: channel multiplier for each level of the UNet.
|
|
||||||
:param conv_resample: if True, use learned convolutions for upsampling and
|
|
||||||
downsampling.
|
|
||||||
:param dims: determines if the signal is 1D, 2D, or 3D.
|
|
||||||
:param num_classes: if specified (as an int), then this model will be
|
|
||||||
class-conditional with `num_classes` classes.
|
|
||||||
:param use_checkpoint: use gradient checkpointing to reduce memory usage.
|
|
||||||
:param num_heads: the number of attention heads in each attention layer.
|
|
||||||
:param num_heads_channels: if specified, ignore num_heads and instead use
|
|
||||||
a fixed channel width per attention head.
|
|
||||||
:param num_heads_upsample: works with num_heads to set a different number
|
|
||||||
of heads for upsampling. Deprecated.
|
|
||||||
:param use_scale_shift_norm: use a FiLM-like conditioning mechanism.
|
|
||||||
:param resblock_updown: use residual blocks for up/downsampling.
|
|
||||||
:param use_new_attention_order: use a different attention pattern for potentially
|
|
||||||
increased efficiency.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self,
|
|
||||||
in_channels,
|
|
||||||
model_channels,
|
|
||||||
out_channels,
|
|
||||||
num_res_blocks,
|
|
||||||
attention_resolutions,
|
|
||||||
dropout=0.0,
|
|
||||||
channel_mult=(1, 2, 4, 8),
|
|
||||||
conv_resample=True,
|
|
||||||
dims=2,
|
|
||||||
context_dim=None,
|
|
||||||
use_scale_shift_norm=False,
|
|
||||||
resblock_updown=False,
|
|
||||||
num_heads=-1,
|
|
||||||
num_head_channels=-1,
|
|
||||||
transformer_depth=1,
|
|
||||||
use_linear=False,
|
|
||||||
use_checkpoint=False,
|
|
||||||
temporal_conv=False,
|
|
||||||
tempspatial_aware=False,
|
|
||||||
temporal_attention=True,
|
|
||||||
use_relative_position=True,
|
|
||||||
use_causal_attention=False,
|
|
||||||
temporal_length=None,
|
|
||||||
use_fp16=False,
|
|
||||||
addition_attention=False,
|
|
||||||
temporal_selfatt_only=True,
|
|
||||||
image_cross_attention=False,
|
|
||||||
image_cross_attention_scale_learnable=False,
|
|
||||||
default_fs=4,
|
|
||||||
fs_condition=False,
|
|
||||||
device=None,
|
|
||||||
dtype=torch.float16,
|
|
||||||
operations=ops
|
|
||||||
):
|
|
||||||
super(UNetModel, self).__init__()
|
|
||||||
if num_heads == -1:
|
|
||||||
assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set'
|
|
||||||
if num_head_channels == -1:
|
|
||||||
assert num_heads != -1, 'Either num_heads or num_head_channels has to be set'
|
|
||||||
|
|
||||||
self.in_channels = in_channels
|
|
||||||
self.model_channels = model_channels
|
|
||||||
self.out_channels = out_channels
|
|
||||||
self.num_res_blocks = num_res_blocks
|
|
||||||
self.attention_resolutions = attention_resolutions
|
|
||||||
self.dropout = dropout
|
|
||||||
self.channel_mult = channel_mult
|
|
||||||
self.conv_resample = conv_resample
|
|
||||||
self.temporal_attention = temporal_attention
|
|
||||||
time_embed_dim = model_channels * 4
|
|
||||||
self.use_checkpoint = use_checkpoint
|
|
||||||
temporal_self_att_only = True
|
|
||||||
self.addition_attention = addition_attention
|
|
||||||
self.temporal_length = temporal_length
|
|
||||||
self.image_cross_attention = image_cross_attention
|
|
||||||
self.image_cross_attention_scale_learnable = image_cross_attention_scale_learnable
|
|
||||||
self.default_fs = default_fs
|
|
||||||
self.fs_condition = fs_condition
|
|
||||||
self.device = device
|
|
||||||
#self.dtype = dtype
|
|
||||||
self.dtype = torch.float32
|
|
||||||
|
|
||||||
## Time embedding blocks
|
|
||||||
self.time_embed = nn.Sequential(
|
|
||||||
linear(model_channels, time_embed_dim, device=device, dtype=self.dtype),
|
|
||||||
nn.SiLU(),
|
|
||||||
linear(time_embed_dim, time_embed_dim, device=device, dtype=self.dtype),
|
|
||||||
)
|
|
||||||
if fs_condition:
|
|
||||||
self.fps_embedding = nn.Sequential(
|
|
||||||
linear(model_channels, time_embed_dim, device=device, dtype=self.dtype),
|
|
||||||
nn.SiLU(),
|
|
||||||
linear(time_embed_dim, time_embed_dim, device=device, dtype=self.dtype),
|
|
||||||
)
|
|
||||||
nn.init.zeros_(self.fps_embedding[-1].weight)
|
|
||||||
nn.init.zeros_(self.fps_embedding[-1].bias)
|
|
||||||
## Input Block
|
|
||||||
self.input_blocks = nn.ModuleList(
|
|
||||||
[
|
|
||||||
TimestepEmbedSequential(
|
|
||||||
operations.conv_nd(
|
|
||||||
dims,
|
|
||||||
in_channels,
|
|
||||||
model_channels,
|
|
||||||
3,
|
|
||||||
padding=1,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
))
|
|
||||||
]
|
|
||||||
)
|
|
||||||
if self.addition_attention:
|
|
||||||
self.init_attn=TimestepEmbedSequential(
|
|
||||||
TemporalTransformer(
|
|
||||||
model_channels,
|
|
||||||
n_heads=8,
|
|
||||||
d_head=num_head_channels,
|
|
||||||
depth=transformer_depth,
|
|
||||||
context_dim=context_dim,
|
|
||||||
use_checkpoint=use_checkpoint, only_self_att=temporal_selfatt_only,
|
|
||||||
causal_attention=False, relative_position=use_relative_position,
|
|
||||||
temporal_length=temporal_length,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
))
|
|
||||||
|
|
||||||
input_block_chans = [model_channels]
|
|
||||||
ch = model_channels
|
|
||||||
ds = 1
|
|
||||||
for level, mult in enumerate(channel_mult):
|
|
||||||
for _ in range(num_res_blocks):
|
|
||||||
layers = [
|
|
||||||
ResBlock(ch, time_embed_dim, dropout,
|
|
||||||
out_channels=mult * model_channels, dims=dims, use_checkpoint=use_checkpoint,
|
|
||||||
use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware,
|
|
||||||
use_temporal_conv=temporal_conv,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
]
|
|
||||||
ch = mult * model_channels
|
|
||||||
if ds in attention_resolutions:
|
|
||||||
if num_head_channels == -1:
|
|
||||||
dim_head = ch // num_heads
|
|
||||||
else:
|
|
||||||
num_heads = ch // num_head_channels
|
|
||||||
dim_head = num_head_channels
|
|
||||||
layers.append(
|
|
||||||
SpatialTransformer(ch, num_heads, dim_head,
|
|
||||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
|
||||||
use_checkpoint=use_checkpoint, disable_self_attn=False,
|
|
||||||
video_length=temporal_length, image_cross_attention=self.image_cross_attention,
|
|
||||||
image_cross_attention_scale_learnable=self.image_cross_attention_scale_learnable,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
)
|
|
||||||
if self.temporal_attention:
|
|
||||||
layers.append(
|
|
||||||
TemporalTransformer(ch, num_heads, dim_head,
|
|
||||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
|
||||||
use_checkpoint=use_checkpoint, only_self_att=temporal_self_att_only,
|
|
||||||
causal_attention=use_causal_attention, relative_position=use_relative_position,
|
|
||||||
temporal_length=temporal_length,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
)
|
|
||||||
self.input_blocks.append(TimestepEmbedSequential(*layers))
|
|
||||||
input_block_chans.append(ch)
|
|
||||||
if level != len(channel_mult) - 1:
|
|
||||||
out_ch = ch
|
|
||||||
self.input_blocks.append(
|
|
||||||
TimestepEmbedSequential(
|
|
||||||
ResBlock(ch, time_embed_dim, dropout,
|
|
||||||
out_channels=out_ch, dims=dims, use_checkpoint=use_checkpoint,
|
|
||||||
use_scale_shift_norm=use_scale_shift_norm,
|
|
||||||
down=True,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
if resblock_updown
|
|
||||||
else Downsample(
|
|
||||||
ch,
|
|
||||||
conv_resample,
|
|
||||||
dims=dims,
|
|
||||||
out_channels=out_ch,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
)
|
|
||||||
)
|
|
||||||
ch = out_ch
|
|
||||||
input_block_chans.append(ch)
|
|
||||||
ds *= 2
|
|
||||||
|
|
||||||
if num_head_channels == -1:
|
|
||||||
dim_head = ch // num_heads
|
|
||||||
else:
|
|
||||||
num_heads = ch // num_head_channels
|
|
||||||
dim_head = num_head_channels
|
|
||||||
layers = [
|
|
||||||
ResBlock(ch, time_embed_dim, dropout,
|
|
||||||
dims=dims, use_checkpoint=use_checkpoint,
|
|
||||||
use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware,
|
|
||||||
use_temporal_conv=temporal_conv,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
),
|
|
||||||
SpatialTransformer(ch, num_heads, dim_head,
|
|
||||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
|
||||||
use_checkpoint=use_checkpoint, disable_self_attn=False, video_length=temporal_length,
|
|
||||||
image_cross_attention=self.image_cross_attention,image_cross_attention_scale_learnable=self.image_cross_attention_scale_learnable,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
]
|
|
||||||
if self.temporal_attention:
|
|
||||||
layers.append(
|
|
||||||
TemporalTransformer(ch, num_heads, dim_head,
|
|
||||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
|
||||||
use_checkpoint=use_checkpoint, only_self_att=temporal_self_att_only,
|
|
||||||
causal_attention=use_causal_attention, relative_position=use_relative_position,
|
|
||||||
temporal_length=temporal_length,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
)
|
|
||||||
layers.append(
|
|
||||||
ResBlock(ch, time_embed_dim, dropout,
|
|
||||||
dims=dims, use_checkpoint=use_checkpoint,
|
|
||||||
use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware,
|
|
||||||
use_temporal_conv=temporal_conv,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
## Middle Block
|
|
||||||
self.middle_block = TimestepEmbedSequential(*layers)
|
|
||||||
|
|
||||||
## Output Block
|
|
||||||
self.output_blocks = nn.ModuleList([])
|
|
||||||
for level, mult in list(enumerate(channel_mult))[::-1]:
|
|
||||||
for i in range(num_res_blocks + 1):
|
|
||||||
ich = input_block_chans.pop()
|
|
||||||
layers = [
|
|
||||||
ResBlock(ch + ich, time_embed_dim, dropout,
|
|
||||||
out_channels=mult * model_channels, dims=dims, use_checkpoint=use_checkpoint,
|
|
||||||
use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware,
|
|
||||||
use_temporal_conv=temporal_conv,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
]
|
|
||||||
ch = model_channels * mult
|
|
||||||
if ds in attention_resolutions:
|
|
||||||
if num_head_channels == -1:
|
|
||||||
dim_head = ch // num_heads
|
|
||||||
else:
|
|
||||||
num_heads = ch // num_head_channels
|
|
||||||
dim_head = num_head_channels
|
|
||||||
layers.append(
|
|
||||||
SpatialTransformer(ch, num_heads, dim_head,
|
|
||||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
|
||||||
use_checkpoint=use_checkpoint, disable_self_attn=False, video_length=temporal_length,
|
|
||||||
image_cross_attention=self.image_cross_attention,image_cross_attention_scale_learnable=self.image_cross_attention_scale_learnable,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
)
|
|
||||||
if self.temporal_attention:
|
|
||||||
layers.append(
|
|
||||||
TemporalTransformer(ch, num_heads, dim_head,
|
|
||||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
|
||||||
use_checkpoint=use_checkpoint, only_self_att=temporal_self_att_only,
|
|
||||||
causal_attention=use_causal_attention, relative_position=use_relative_position,
|
|
||||||
temporal_length=temporal_length,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
)
|
|
||||||
if level and i == num_res_blocks:
|
|
||||||
out_ch = ch
|
|
||||||
layers.append(
|
|
||||||
ResBlock(ch, time_embed_dim, dropout,
|
|
||||||
out_channels=out_ch, dims=dims, use_checkpoint=use_checkpoint,
|
|
||||||
use_scale_shift_norm=use_scale_shift_norm,
|
|
||||||
up=True,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
if resblock_updown
|
|
||||||
else Upsample(ch, conv_resample, dims=dims, out_channels=out_ch)
|
|
||||||
)
|
|
||||||
ds //= 2
|
|
||||||
self.output_blocks.append(TimestepEmbedSequential(*layers))
|
|
||||||
|
|
||||||
self.out = nn.Sequential(
|
|
||||||
normalization(ch, device=device, dtype=self.dtype),
|
|
||||||
nn.SiLU(),
|
|
||||||
zero_module(
|
|
||||||
operations.conv_nd(
|
|
||||||
dims,
|
|
||||||
model_channels,
|
|
||||||
out_channels,
|
|
||||||
3,
|
|
||||||
padding=1,
|
|
||||||
device=device,
|
|
||||||
dtype=self.dtype
|
|
||||||
)
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
# TODO Add Transformer options to leverage the usage of patches.
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x,
|
|
||||||
timesteps,
|
|
||||||
context=None,
|
|
||||||
context_in=None,
|
|
||||||
cc_concat=None,
|
|
||||||
num_video_frames=16,
|
|
||||||
features_adapter=None,
|
|
||||||
fs=None,
|
|
||||||
img_emb=None,
|
|
||||||
control=None,
|
|
||||||
transformer_options={},
|
|
||||||
cond_idx=None,
|
|
||||||
**kwargs
|
|
||||||
):
|
|
||||||
|
|
||||||
if any([fs is None, img_emb is None, cc_concat is None]):
|
|
||||||
raise ValueError("One or more of the required inputs for UNet Forward is None.")
|
|
||||||
|
|
||||||
cond_idx = transformer_options.get("cond_idx", None)
|
|
||||||
transformer_options['original_shape'] = list(x.shape)
|
|
||||||
transformer_options['transformer_index'] = 0
|
|
||||||
transformer_patches = transformer_options.get("patches", {})
|
|
||||||
|
|
||||||
# In ComfyUI, the frames are always with the batch, so we deconstruct it here.
|
|
||||||
# This is mandatory as this is a video based model.
|
|
||||||
# We usually denote "f" as frames, but will use "t" (time) to be consistent with DynamiCrafter.
|
|
||||||
b,_,t,_,_ = x.shape
|
|
||||||
|
|
||||||
context = context_in
|
|
||||||
cc_concat = cc_concat.to(x.device, x.dtype)
|
|
||||||
x = torch.cat([x, cc_concat], dim=1)
|
|
||||||
|
|
||||||
fs = fs.to(x.device, x.dtype)
|
|
||||||
|
|
||||||
timestep = timesteps
|
|
||||||
context = context_processor(context, num_video_frames, img_emb=img_emb)
|
|
||||||
|
|
||||||
t_emb = timestep_embedding(timestep, self.model_channels, repeat_only=False, dtype=self.dtype)
|
|
||||||
emb = self.time_embed(t_emb)
|
|
||||||
emb = emb.repeat_interleave(repeats=t, dim=0)
|
|
||||||
|
|
||||||
## always in shape (b t) c h w, except for temporal layer
|
|
||||||
x = rearrange(x, 'b c t h w -> (b t) c h w')
|
|
||||||
|
|
||||||
## combine emb
|
|
||||||
if self.fs_condition:
|
|
||||||
if fs is None:
|
|
||||||
fs = torch.tensor(
|
|
||||||
[self.default_fs] * b, dtype=torch.long, device=x.device)
|
|
||||||
fs_emb = timestep_embedding(fs, self.model_channels, repeat_only=False, dtype=self.dtype).type(x.dtype)
|
|
||||||
|
|
||||||
fs_embed = self.fps_embedding(fs_emb)
|
|
||||||
fs_embed = fs_embed.repeat_interleave(repeats=t, dim=0)
|
|
||||||
|
|
||||||
emb = emb + fs_embed
|
|
||||||
|
|
||||||
h = x.type(self.dtype)
|
|
||||||
adapter_idx = 0
|
|
||||||
hs = []
|
|
||||||
|
|
||||||
for id, module in enumerate(self.input_blocks):
|
|
||||||
transformer_options["block"] = ("input", id)
|
|
||||||
#h = module(h, emb, context=context, batch_size=b)
|
|
||||||
h = forward_timestep_embed(
|
|
||||||
module,
|
|
||||||
h,
|
|
||||||
emb,
|
|
||||||
context=context,
|
|
||||||
batch_size=b,
|
|
||||||
transformer_options=transformer_options
|
|
||||||
)
|
|
||||||
h = apply_control(h, control, 'input', cond_idx)
|
|
||||||
|
|
||||||
if "input_block_patch" in transformer_patches:
|
|
||||||
patch = transformer_patches["input_block_patch"]
|
|
||||||
for p in patch:
|
|
||||||
h = p(h, transformer_options)
|
|
||||||
|
|
||||||
if id ==0 and self.addition_attention:
|
|
||||||
h = forward_timestep_embed(
|
|
||||||
self.init_attn,
|
|
||||||
h,
|
|
||||||
emb,
|
|
||||||
context=context,
|
|
||||||
batch_size=b,
|
|
||||||
transformer_options=transformer_options
|
|
||||||
)
|
|
||||||
## plug-in adapter features
|
|
||||||
if ((id+1)%3 == 0) and features_adapter is not None:
|
|
||||||
h = h + features_adapter[adapter_idx]
|
|
||||||
adapter_idx += 1
|
|
||||||
hs.append(h)
|
|
||||||
if "input_block_patch_after_skip" in transformer_patches:
|
|
||||||
patch = transformer_patches["input_block_patch_after_skip"]
|
|
||||||
for p in patch:
|
|
||||||
h = p(h, transformer_options)
|
|
||||||
if features_adapter is not None:
|
|
||||||
assert len(features_adapter)==adapter_idx, 'Wrong features_adapter'
|
|
||||||
transformer_options["block"] = ("middle", 0)
|
|
||||||
h = forward_timestep_embed(
|
|
||||||
self.middle_block,
|
|
||||||
h,
|
|
||||||
emb,
|
|
||||||
context=context,
|
|
||||||
batch_size=b,
|
|
||||||
transformer_options=transformer_options
|
|
||||||
)
|
|
||||||
h = apply_control(h, control, 'middle', cond_idx)
|
|
||||||
for id, module in enumerate(self.output_blocks):
|
|
||||||
transformer_options["block"] = ("output", id)
|
|
||||||
hsp = hs.pop()
|
|
||||||
hsp = apply_control(hsp, control, 'output', cond_idx)
|
|
||||||
|
|
||||||
if "output_block_patch" in transformer_patches:
|
|
||||||
patch = transformer_patches["output_block_patch"]
|
|
||||||
for p in patch:
|
|
||||||
h, hsp = p(h, hsp, transformer_options)
|
|
||||||
|
|
||||||
h = torch.cat([h, hsp], dim=1)
|
|
||||||
del hsp
|
|
||||||
h = forward_timestep_embed(
|
|
||||||
module,
|
|
||||||
h,
|
|
||||||
emb,
|
|
||||||
context=context,
|
|
||||||
batch_size=b,
|
|
||||||
transformer_options=transformer_options
|
|
||||||
)
|
|
||||||
h = h.type(x.dtype)
|
|
||||||
h = self.out(h)
|
|
||||||
|
|
||||||
# We output with the tensor unfolded framewise, then reshape them to batched using ComfyUI nodes.
|
|
||||||
h = rearrange(h, '(b t) c h w -> b c t h w', t=num_video_frames)
|
|
||||||
|
|
||||||
return h
|
|
||||||
@@ -1,639 +0,0 @@
|
|||||||
"""shout-out to https://github.com/lucidrains/x-transformers/tree/main/x_transformers"""
|
|
||||||
from functools import partial
|
|
||||||
from inspect import isfunction
|
|
||||||
from collections import namedtuple
|
|
||||||
from einops import rearrange, repeat
|
|
||||||
import torch
|
|
||||||
from torch import nn, einsum
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
# constants
|
|
||||||
DEFAULT_DIM_HEAD = 64
|
|
||||||
|
|
||||||
Intermediates = namedtuple('Intermediates', [
|
|
||||||
'pre_softmax_attn',
|
|
||||||
'post_softmax_attn'
|
|
||||||
])
|
|
||||||
|
|
||||||
LayerIntermediates = namedtuple('Intermediates', [
|
|
||||||
'hiddens',
|
|
||||||
'attn_intermediates'
|
|
||||||
])
|
|
||||||
|
|
||||||
|
|
||||||
class AbsolutePositionalEmbedding(nn.Module):
|
|
||||||
def __init__(self, dim, max_seq_len):
|
|
||||||
super().__init__()
|
|
||||||
self.emb = nn.Embedding(max_seq_len, dim)
|
|
||||||
self.init_()
|
|
||||||
|
|
||||||
def init_(self):
|
|
||||||
nn.init.normal_(self.emb.weight, std=0.02)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
n = torch.arange(x.shape[1], device=x.device)
|
|
||||||
return self.emb(n)[None, :, :]
|
|
||||||
|
|
||||||
|
|
||||||
class FixedPositionalEmbedding(nn.Module):
|
|
||||||
def __init__(self, dim):
|
|
||||||
super().__init__()
|
|
||||||
inv_freq = 1. / (10000 ** (torch.arange(0, dim, 2).float() / dim))
|
|
||||||
self.register_buffer('inv_freq', inv_freq)
|
|
||||||
|
|
||||||
def forward(self, x, seq_dim=1, offset=0):
|
|
||||||
t = torch.arange(x.shape[seq_dim], device=x.device).type_as(self.inv_freq) + offset
|
|
||||||
sinusoid_inp = torch.einsum('i , j -> i j', t, self.inv_freq)
|
|
||||||
emb = torch.cat((sinusoid_inp.sin(), sinusoid_inp.cos()), dim=-1)
|
|
||||||
return emb[None, :, :]
|
|
||||||
|
|
||||||
|
|
||||||
# helpers
|
|
||||||
|
|
||||||
def exists(val):
|
|
||||||
return val is not None
|
|
||||||
|
|
||||||
|
|
||||||
def default(val, d):
|
|
||||||
if exists(val):
|
|
||||||
return val
|
|
||||||
return d() if isfunction(d) else d
|
|
||||||
|
|
||||||
|
|
||||||
def always(val):
|
|
||||||
def inner(*args, **kwargs):
|
|
||||||
return val
|
|
||||||
return inner
|
|
||||||
|
|
||||||
|
|
||||||
def not_equals(val):
|
|
||||||
def inner(x):
|
|
||||||
return x != val
|
|
||||||
return inner
|
|
||||||
|
|
||||||
|
|
||||||
def equals(val):
|
|
||||||
def inner(x):
|
|
||||||
return x == val
|
|
||||||
return inner
|
|
||||||
|
|
||||||
|
|
||||||
def max_neg_value(tensor):
|
|
||||||
return -torch.finfo(tensor.dtype).max
|
|
||||||
|
|
||||||
|
|
||||||
# keyword argument helpers
|
|
||||||
|
|
||||||
def pick_and_pop(keys, d):
|
|
||||||
values = list(map(lambda key: d.pop(key), keys))
|
|
||||||
return dict(zip(keys, values))
|
|
||||||
|
|
||||||
|
|
||||||
def group_dict_by_key(cond, d):
|
|
||||||
return_val = [dict(), dict()]
|
|
||||||
for key in d.keys():
|
|
||||||
match = bool(cond(key))
|
|
||||||
ind = int(not match)
|
|
||||||
return_val[ind][key] = d[key]
|
|
||||||
return (*return_val,)
|
|
||||||
|
|
||||||
|
|
||||||
def string_begins_with(prefix, str):
|
|
||||||
return str.startswith(prefix)
|
|
||||||
|
|
||||||
|
|
||||||
def group_by_key_prefix(prefix, d):
|
|
||||||
return group_dict_by_key(partial(string_begins_with, prefix), d)
|
|
||||||
|
|
||||||
|
|
||||||
def groupby_prefix_and_trim(prefix, d):
|
|
||||||
kwargs_with_prefix, kwargs = group_dict_by_key(partial(string_begins_with, prefix), d)
|
|
||||||
kwargs_without_prefix = dict(map(lambda x: (x[0][len(prefix):], x[1]), tuple(kwargs_with_prefix.items())))
|
|
||||||
return kwargs_without_prefix, kwargs
|
|
||||||
|
|
||||||
|
|
||||||
# classes
|
|
||||||
class Scale(nn.Module):
|
|
||||||
def __init__(self, value, fn):
|
|
||||||
super().__init__()
|
|
||||||
self.value = value
|
|
||||||
self.fn = fn
|
|
||||||
|
|
||||||
def forward(self, x, **kwargs):
|
|
||||||
x, *rest = self.fn(x, **kwargs)
|
|
||||||
return (x * self.value, *rest)
|
|
||||||
|
|
||||||
|
|
||||||
class Rezero(nn.Module):
|
|
||||||
def __init__(self, fn):
|
|
||||||
super().__init__()
|
|
||||||
self.fn = fn
|
|
||||||
self.g = nn.Parameter(torch.zeros(1))
|
|
||||||
|
|
||||||
def forward(self, x, **kwargs):
|
|
||||||
x, *rest = self.fn(x, **kwargs)
|
|
||||||
return (x * self.g, *rest)
|
|
||||||
|
|
||||||
|
|
||||||
class ScaleNorm(nn.Module):
|
|
||||||
def __init__(self, dim, eps=1e-5):
|
|
||||||
super().__init__()
|
|
||||||
self.scale = dim ** -0.5
|
|
||||||
self.eps = eps
|
|
||||||
self.g = nn.Parameter(torch.ones(1))
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
norm = torch.norm(x, dim=-1, keepdim=True) * self.scale
|
|
||||||
return x / norm.clamp(min=self.eps) * self.g
|
|
||||||
|
|
||||||
|
|
||||||
class RMSNorm(nn.Module):
|
|
||||||
def __init__(self, dim, eps=1e-8):
|
|
||||||
super().__init__()
|
|
||||||
self.scale = dim ** -0.5
|
|
||||||
self.eps = eps
|
|
||||||
self.g = nn.Parameter(torch.ones(dim))
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
norm = torch.norm(x, dim=-1, keepdim=True) * self.scale
|
|
||||||
return x / norm.clamp(min=self.eps) * self.g
|
|
||||||
|
|
||||||
|
|
||||||
class Residual(nn.Module):
|
|
||||||
def forward(self, x, residual):
|
|
||||||
return x + residual
|
|
||||||
|
|
||||||
|
|
||||||
class GRUGating(nn.Module):
|
|
||||||
def __init__(self, dim):
|
|
||||||
super().__init__()
|
|
||||||
self.gru = nn.GRUCell(dim, dim)
|
|
||||||
|
|
||||||
def forward(self, x, residual):
|
|
||||||
gated_output = self.gru(
|
|
||||||
rearrange(x, 'b n d -> (b n) d'),
|
|
||||||
rearrange(residual, 'b n d -> (b n) d')
|
|
||||||
)
|
|
||||||
|
|
||||||
return gated_output.reshape_as(x)
|
|
||||||
|
|
||||||
|
|
||||||
# feedforward
|
|
||||||
|
|
||||||
class GEGLU(nn.Module):
|
|
||||||
def __init__(self, dim_in, dim_out):
|
|
||||||
super().__init__()
|
|
||||||
self.proj = nn.Linear(dim_in, dim_out * 2)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
x, gate = self.proj(x).chunk(2, dim=-1)
|
|
||||||
return x * F.gelu(gate)
|
|
||||||
|
|
||||||
|
|
||||||
class FeedForward(nn.Module):
|
|
||||||
def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.):
|
|
||||||
super().__init__()
|
|
||||||
inner_dim = int(dim * mult)
|
|
||||||
dim_out = default(dim_out, dim)
|
|
||||||
project_in = nn.Sequential(
|
|
||||||
nn.Linear(dim, inner_dim),
|
|
||||||
nn.GELU()
|
|
||||||
) if not glu else GEGLU(dim, inner_dim)
|
|
||||||
|
|
||||||
self.net = nn.Sequential(
|
|
||||||
project_in,
|
|
||||||
nn.Dropout(dropout),
|
|
||||||
nn.Linear(inner_dim, dim_out)
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
return self.net(x)
|
|
||||||
|
|
||||||
|
|
||||||
# attention.
|
|
||||||
class Attention(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim,
|
|
||||||
dim_head=DEFAULT_DIM_HEAD,
|
|
||||||
heads=8,
|
|
||||||
causal=False,
|
|
||||||
mask=None,
|
|
||||||
talking_heads=False,
|
|
||||||
sparse_topk=None,
|
|
||||||
use_entmax15=False,
|
|
||||||
num_mem_kv=0,
|
|
||||||
dropout=0.,
|
|
||||||
on_attn=False
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
if use_entmax15:
|
|
||||||
raise NotImplementedError("Check out entmax activation instead of softmax activation!")
|
|
||||||
self.scale = dim_head ** -0.5
|
|
||||||
self.heads = heads
|
|
||||||
self.causal = causal
|
|
||||||
self.mask = mask
|
|
||||||
|
|
||||||
inner_dim = dim_head * heads
|
|
||||||
|
|
||||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
|
||||||
self.to_k = nn.Linear(dim, inner_dim, bias=False)
|
|
||||||
self.to_v = nn.Linear(dim, inner_dim, bias=False)
|
|
||||||
self.dropout = nn.Dropout(dropout)
|
|
||||||
|
|
||||||
# talking heads
|
|
||||||
self.talking_heads = talking_heads
|
|
||||||
if talking_heads:
|
|
||||||
self.pre_softmax_proj = nn.Parameter(torch.randn(heads, heads))
|
|
||||||
self.post_softmax_proj = nn.Parameter(torch.randn(heads, heads))
|
|
||||||
|
|
||||||
# explicit topk sparse attention
|
|
||||||
self.sparse_topk = sparse_topk
|
|
||||||
|
|
||||||
# entmax
|
|
||||||
#self.attn_fn = entmax15 if use_entmax15 else F.softmax
|
|
||||||
self.attn_fn = F.softmax
|
|
||||||
|
|
||||||
# add memory key / values
|
|
||||||
self.num_mem_kv = num_mem_kv
|
|
||||||
if num_mem_kv > 0:
|
|
||||||
self.mem_k = nn.Parameter(torch.randn(heads, num_mem_kv, dim_head))
|
|
||||||
self.mem_v = nn.Parameter(torch.randn(heads, num_mem_kv, dim_head))
|
|
||||||
|
|
||||||
# attention on attention
|
|
||||||
self.attn_on_attn = on_attn
|
|
||||||
self.to_out = nn.Sequential(nn.Linear(inner_dim, dim * 2), nn.GLU()) if on_attn else nn.Linear(inner_dim, dim)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x,
|
|
||||||
context=None,
|
|
||||||
mask=None,
|
|
||||||
context_mask=None,
|
|
||||||
rel_pos=None,
|
|
||||||
sinusoidal_emb=None,
|
|
||||||
prev_attn=None,
|
|
||||||
mem=None
|
|
||||||
):
|
|
||||||
b, n, _, h, talking_heads, device = *x.shape, self.heads, self.talking_heads, x.device
|
|
||||||
kv_input = default(context, x)
|
|
||||||
|
|
||||||
q_input = x
|
|
||||||
k_input = kv_input
|
|
||||||
v_input = kv_input
|
|
||||||
|
|
||||||
if exists(mem):
|
|
||||||
k_input = torch.cat((mem, k_input), dim=-2)
|
|
||||||
v_input = torch.cat((mem, v_input), dim=-2)
|
|
||||||
|
|
||||||
if exists(sinusoidal_emb):
|
|
||||||
# in shortformer, the query would start at a position offset depending on the past cached memory
|
|
||||||
offset = k_input.shape[-2] - q_input.shape[-2]
|
|
||||||
q_input = q_input + sinusoidal_emb(q_input, offset=offset)
|
|
||||||
k_input = k_input + sinusoidal_emb(k_input)
|
|
||||||
|
|
||||||
q = self.to_q(q_input)
|
|
||||||
k = self.to_k(k_input)
|
|
||||||
v = self.to_v(v_input)
|
|
||||||
|
|
||||||
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h=h), (q, k, v))
|
|
||||||
|
|
||||||
input_mask = None
|
|
||||||
if any(map(exists, (mask, context_mask))):
|
|
||||||
q_mask = default(mask, lambda: torch.ones((b, n), device=device).bool())
|
|
||||||
k_mask = q_mask if not exists(context) else context_mask
|
|
||||||
k_mask = default(k_mask, lambda: torch.ones((b, k.shape[-2]), device=device).bool())
|
|
||||||
q_mask = rearrange(q_mask, 'b i -> b () i ()')
|
|
||||||
k_mask = rearrange(k_mask, 'b j -> b () () j')
|
|
||||||
input_mask = q_mask * k_mask
|
|
||||||
|
|
||||||
if self.num_mem_kv > 0:
|
|
||||||
mem_k, mem_v = map(lambda t: repeat(t, 'h n d -> b h n d', b=b), (self.mem_k, self.mem_v))
|
|
||||||
k = torch.cat((mem_k, k), dim=-2)
|
|
||||||
v = torch.cat((mem_v, v), dim=-2)
|
|
||||||
if exists(input_mask):
|
|
||||||
input_mask = F.pad(input_mask, (self.num_mem_kv, 0), value=True)
|
|
||||||
|
|
||||||
dots = einsum('b h i d, b h j d -> b h i j', q, k) * self.scale
|
|
||||||
mask_value = max_neg_value(dots)
|
|
||||||
|
|
||||||
if exists(prev_attn):
|
|
||||||
dots = dots + prev_attn
|
|
||||||
|
|
||||||
pre_softmax_attn = dots
|
|
||||||
|
|
||||||
if talking_heads:
|
|
||||||
dots = einsum('b h i j, h k -> b k i j', dots, self.pre_softmax_proj).contiguous()
|
|
||||||
|
|
||||||
if exists(rel_pos):
|
|
||||||
dots = rel_pos(dots)
|
|
||||||
|
|
||||||
if exists(input_mask):
|
|
||||||
dots.masked_fill_(~input_mask, mask_value)
|
|
||||||
del input_mask
|
|
||||||
|
|
||||||
if self.causal:
|
|
||||||
i, j = dots.shape[-2:]
|
|
||||||
r = torch.arange(i, device=device)
|
|
||||||
mask = rearrange(r, 'i -> () () i ()') < rearrange(r, 'j -> () () () j')
|
|
||||||
mask = F.pad(mask, (j - i, 0), value=False)
|
|
||||||
dots.masked_fill_(mask, mask_value)
|
|
||||||
del mask
|
|
||||||
|
|
||||||
if exists(self.sparse_topk) and self.sparse_topk < dots.shape[-1]:
|
|
||||||
top, _ = dots.topk(self.sparse_topk, dim=-1)
|
|
||||||
vk = top[..., -1].unsqueeze(-1).expand_as(dots)
|
|
||||||
mask = dots < vk
|
|
||||||
dots.masked_fill_(mask, mask_value)
|
|
||||||
del mask
|
|
||||||
|
|
||||||
attn = self.attn_fn(dots, dim=-1)
|
|
||||||
post_softmax_attn = attn
|
|
||||||
|
|
||||||
attn = self.dropout(attn)
|
|
||||||
|
|
||||||
if talking_heads:
|
|
||||||
attn = einsum('b h i j, h k -> b k i j', attn, self.post_softmax_proj).contiguous()
|
|
||||||
|
|
||||||
out = einsum('b h i j, b h j d -> b h i d', attn, v)
|
|
||||||
out = rearrange(out, 'b h n d -> b n (h d)')
|
|
||||||
|
|
||||||
intermediates = Intermediates(
|
|
||||||
pre_softmax_attn=pre_softmax_attn,
|
|
||||||
post_softmax_attn=post_softmax_attn
|
|
||||||
)
|
|
||||||
|
|
||||||
return self.to_out(out), intermediates
|
|
||||||
|
|
||||||
|
|
||||||
class AttentionLayers(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim,
|
|
||||||
depth,
|
|
||||||
heads=8,
|
|
||||||
causal=False,
|
|
||||||
cross_attend=False,
|
|
||||||
only_cross=False,
|
|
||||||
use_scalenorm=False,
|
|
||||||
use_rmsnorm=False,
|
|
||||||
use_rezero=False,
|
|
||||||
rel_pos_num_buckets=32,
|
|
||||||
rel_pos_max_distance=128,
|
|
||||||
position_infused_attn=False,
|
|
||||||
custom_layers=None,
|
|
||||||
sandwich_coef=None,
|
|
||||||
par_ratio=None,
|
|
||||||
residual_attn=False,
|
|
||||||
cross_residual_attn=False,
|
|
||||||
macaron=False,
|
|
||||||
pre_norm=True,
|
|
||||||
gate_residual=False,
|
|
||||||
**kwargs
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
ff_kwargs, kwargs = groupby_prefix_and_trim('ff_', kwargs)
|
|
||||||
attn_kwargs, _ = groupby_prefix_and_trim('attn_', kwargs)
|
|
||||||
|
|
||||||
dim_head = attn_kwargs.get('dim_head', DEFAULT_DIM_HEAD)
|
|
||||||
|
|
||||||
self.dim = dim
|
|
||||||
self.depth = depth
|
|
||||||
self.layers = nn.ModuleList([])
|
|
||||||
|
|
||||||
self.has_pos_emb = position_infused_attn
|
|
||||||
self.pia_pos_emb = FixedPositionalEmbedding(dim) if position_infused_attn else None
|
|
||||||
self.rotary_pos_emb = always(None)
|
|
||||||
|
|
||||||
assert rel_pos_num_buckets <= rel_pos_max_distance, 'number of relative position buckets must be less than the relative position max distance'
|
|
||||||
self.rel_pos = None
|
|
||||||
|
|
||||||
self.pre_norm = pre_norm
|
|
||||||
|
|
||||||
self.residual_attn = residual_attn
|
|
||||||
self.cross_residual_attn = cross_residual_attn
|
|
||||||
|
|
||||||
norm_class = ScaleNorm if use_scalenorm else nn.LayerNorm
|
|
||||||
norm_class = RMSNorm if use_rmsnorm else norm_class
|
|
||||||
norm_fn = partial(norm_class, dim)
|
|
||||||
|
|
||||||
norm_fn = nn.Identity if use_rezero else norm_fn
|
|
||||||
branch_fn = Rezero if use_rezero else None
|
|
||||||
|
|
||||||
if cross_attend and not only_cross:
|
|
||||||
default_block = ('a', 'c', 'f')
|
|
||||||
elif cross_attend and only_cross:
|
|
||||||
default_block = ('c', 'f')
|
|
||||||
else:
|
|
||||||
default_block = ('a', 'f')
|
|
||||||
|
|
||||||
if macaron:
|
|
||||||
default_block = ('f',) + default_block
|
|
||||||
|
|
||||||
if exists(custom_layers):
|
|
||||||
layer_types = custom_layers
|
|
||||||
elif exists(par_ratio):
|
|
||||||
par_depth = depth * len(default_block)
|
|
||||||
assert 1 < par_ratio <= par_depth, 'par ratio out of range'
|
|
||||||
default_block = tuple(filter(not_equals('f'), default_block))
|
|
||||||
par_attn = par_depth // par_ratio
|
|
||||||
depth_cut = par_depth * 2 // 3 # 2 / 3 attention layer cutoff suggested by PAR paper
|
|
||||||
par_width = (depth_cut + depth_cut // par_attn) // par_attn
|
|
||||||
assert len(default_block) <= par_width, 'default block is too large for par_ratio'
|
|
||||||
par_block = default_block + ('f',) * (par_width - len(default_block))
|
|
||||||
par_head = par_block * par_attn
|
|
||||||
layer_types = par_head + ('f',) * (par_depth - len(par_head))
|
|
||||||
elif exists(sandwich_coef):
|
|
||||||
assert sandwich_coef > 0 and sandwich_coef <= depth, 'sandwich coefficient should be less than the depth'
|
|
||||||
layer_types = ('a',) * sandwich_coef + default_block * (depth - sandwich_coef) + ('f',) * sandwich_coef
|
|
||||||
else:
|
|
||||||
layer_types = default_block * depth
|
|
||||||
|
|
||||||
self.layer_types = layer_types
|
|
||||||
self.num_attn_layers = len(list(filter(equals('a'), layer_types)))
|
|
||||||
|
|
||||||
for layer_type in self.layer_types:
|
|
||||||
if layer_type == 'a':
|
|
||||||
layer = Attention(dim, heads=heads, causal=causal, **attn_kwargs)
|
|
||||||
elif layer_type == 'c':
|
|
||||||
layer = Attention(dim, heads=heads, **attn_kwargs)
|
|
||||||
elif layer_type == 'f':
|
|
||||||
layer = FeedForward(dim, **ff_kwargs)
|
|
||||||
layer = layer if not macaron else Scale(0.5, layer)
|
|
||||||
else:
|
|
||||||
raise Exception(f'invalid layer type {layer_type}')
|
|
||||||
|
|
||||||
if isinstance(layer, Attention) and exists(branch_fn):
|
|
||||||
layer = branch_fn(layer)
|
|
||||||
|
|
||||||
if gate_residual:
|
|
||||||
residual_fn = GRUGating(dim)
|
|
||||||
else:
|
|
||||||
residual_fn = Residual()
|
|
||||||
|
|
||||||
self.layers.append(nn.ModuleList([
|
|
||||||
norm_fn(),
|
|
||||||
layer,
|
|
||||||
residual_fn
|
|
||||||
]))
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x,
|
|
||||||
context=None,
|
|
||||||
mask=None,
|
|
||||||
context_mask=None,
|
|
||||||
mems=None,
|
|
||||||
return_hiddens=False
|
|
||||||
):
|
|
||||||
hiddens = []
|
|
||||||
intermediates = []
|
|
||||||
prev_attn = None
|
|
||||||
prev_cross_attn = None
|
|
||||||
|
|
||||||
mems = mems.copy() if exists(mems) else [None] * self.num_attn_layers
|
|
||||||
|
|
||||||
for ind, (layer_type, (norm, block, residual_fn)) in enumerate(zip(self.layer_types, self.layers)):
|
|
||||||
is_last = ind == (len(self.layers) - 1)
|
|
||||||
|
|
||||||
if layer_type == 'a':
|
|
||||||
hiddens.append(x)
|
|
||||||
layer_mem = mems.pop(0)
|
|
||||||
|
|
||||||
residual = x
|
|
||||||
|
|
||||||
if self.pre_norm:
|
|
||||||
x = norm(x)
|
|
||||||
|
|
||||||
if layer_type == 'a':
|
|
||||||
out, inter = block(x, mask=mask, sinusoidal_emb=self.pia_pos_emb, rel_pos=self.rel_pos,
|
|
||||||
prev_attn=prev_attn, mem=layer_mem)
|
|
||||||
elif layer_type == 'c':
|
|
||||||
out, inter = block(x, context=context, mask=mask, context_mask=context_mask, prev_attn=prev_cross_attn)
|
|
||||||
elif layer_type == 'f':
|
|
||||||
out = block(x)
|
|
||||||
|
|
||||||
x = residual_fn(out, residual)
|
|
||||||
|
|
||||||
if layer_type in ('a', 'c'):
|
|
||||||
intermediates.append(inter)
|
|
||||||
|
|
||||||
if layer_type == 'a' and self.residual_attn:
|
|
||||||
prev_attn = inter.pre_softmax_attn
|
|
||||||
elif layer_type == 'c' and self.cross_residual_attn:
|
|
||||||
prev_cross_attn = inter.pre_softmax_attn
|
|
||||||
|
|
||||||
if not self.pre_norm and not is_last:
|
|
||||||
x = norm(x)
|
|
||||||
|
|
||||||
if return_hiddens:
|
|
||||||
intermediates = LayerIntermediates(
|
|
||||||
hiddens=hiddens,
|
|
||||||
attn_intermediates=intermediates
|
|
||||||
)
|
|
||||||
|
|
||||||
return x, intermediates
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class Encoder(AttentionLayers):
|
|
||||||
def __init__(self, **kwargs):
|
|
||||||
assert 'causal' not in kwargs, 'cannot set causality on encoder'
|
|
||||||
super().__init__(causal=False, **kwargs)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class TransformerWrapper(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
num_tokens,
|
|
||||||
max_seq_len,
|
|
||||||
attn_layers,
|
|
||||||
emb_dim=None,
|
|
||||||
max_mem_len=0.,
|
|
||||||
emb_dropout=0.,
|
|
||||||
num_memory_tokens=None,
|
|
||||||
tie_embedding=False,
|
|
||||||
use_pos_emb=True
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
assert isinstance(attn_layers, AttentionLayers), 'attention layers must be one of Encoder or Decoder'
|
|
||||||
|
|
||||||
dim = attn_layers.dim
|
|
||||||
emb_dim = default(emb_dim, dim)
|
|
||||||
|
|
||||||
self.max_seq_len = max_seq_len
|
|
||||||
self.max_mem_len = max_mem_len
|
|
||||||
self.num_tokens = num_tokens
|
|
||||||
|
|
||||||
self.token_emb = nn.Embedding(num_tokens, emb_dim)
|
|
||||||
self.pos_emb = AbsolutePositionalEmbedding(emb_dim, max_seq_len) if (
|
|
||||||
use_pos_emb and not attn_layers.has_pos_emb) else always(0)
|
|
||||||
self.emb_dropout = nn.Dropout(emb_dropout)
|
|
||||||
|
|
||||||
self.project_emb = nn.Linear(emb_dim, dim) if emb_dim != dim else nn.Identity()
|
|
||||||
self.attn_layers = attn_layers
|
|
||||||
self.norm = nn.LayerNorm(dim)
|
|
||||||
|
|
||||||
self.init_()
|
|
||||||
|
|
||||||
self.to_logits = nn.Linear(dim, num_tokens) if not tie_embedding else lambda t: t @ self.token_emb.weight.t()
|
|
||||||
|
|
||||||
# memory tokens (like [cls]) from Memory Transformers paper
|
|
||||||
num_memory_tokens = default(num_memory_tokens, 0)
|
|
||||||
self.num_memory_tokens = num_memory_tokens
|
|
||||||
if num_memory_tokens > 0:
|
|
||||||
self.memory_tokens = nn.Parameter(torch.randn(num_memory_tokens, dim))
|
|
||||||
|
|
||||||
# let funnel encoder know number of memory tokens, if specified
|
|
||||||
if hasattr(attn_layers, 'num_memory_tokens'):
|
|
||||||
attn_layers.num_memory_tokens = num_memory_tokens
|
|
||||||
|
|
||||||
def init_(self):
|
|
||||||
nn.init.normal_(self.token_emb.weight, std=0.02)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x,
|
|
||||||
return_embeddings=False,
|
|
||||||
mask=None,
|
|
||||||
return_mems=False,
|
|
||||||
return_attn=False,
|
|
||||||
mems=None,
|
|
||||||
**kwargs
|
|
||||||
):
|
|
||||||
b, n, device, num_mem = *x.shape, x.device, self.num_memory_tokens
|
|
||||||
x = self.token_emb(x)
|
|
||||||
x += self.pos_emb(x)
|
|
||||||
x = self.emb_dropout(x)
|
|
||||||
|
|
||||||
x = self.project_emb(x)
|
|
||||||
|
|
||||||
if num_mem > 0:
|
|
||||||
mem = repeat(self.memory_tokens, 'n d -> b n d', b=b)
|
|
||||||
x = torch.cat((mem, x), dim=1)
|
|
||||||
|
|
||||||
# auto-handle masking after appending memory tokens
|
|
||||||
if exists(mask):
|
|
||||||
mask = F.pad(mask, (num_mem, 0), value=True)
|
|
||||||
|
|
||||||
x, intermediates = self.attn_layers(x, mask=mask, mems=mems, return_hiddens=True, **kwargs)
|
|
||||||
x = self.norm(x)
|
|
||||||
|
|
||||||
mem, x = x[:, :num_mem], x[:, num_mem:]
|
|
||||||
|
|
||||||
out = self.to_logits(x) if not return_embeddings else x
|
|
||||||
|
|
||||||
if return_mems:
|
|
||||||
hiddens = intermediates.hiddens
|
|
||||||
new_mems = list(map(lambda pair: torch.cat(pair, dim=-2), zip(mems, hiddens))) if exists(mems) else hiddens
|
|
||||||
new_mems = list(map(lambda t: t[..., -self.max_mem_len:, :].detach(), new_mems))
|
|
||||||
return out, new_mems
|
|
||||||
|
|
||||||
if return_attn:
|
|
||||||
attn_maps = list(map(lambda t: t.post_softmax_attn, intermediates.attn_intermediates))
|
|
||||||
return out, attn_maps
|
|
||||||
|
|
||||||
return out
|
|
||||||
@@ -1,143 +0,0 @@
|
|||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from collections import OrderedDict
|
|
||||||
|
|
||||||
from comfy import model_base
|
|
||||||
from comfy import utils
|
|
||||||
from comfy import diffusers_convert
|
|
||||||
|
|
||||||
from comfy import sd2_clip
|
|
||||||
|
|
||||||
from comfy import supported_models_base
|
|
||||||
from comfy import latent_formats
|
|
||||||
|
|
||||||
from ..lvdm.modules.encoders.resampler import Resampler
|
|
||||||
|
|
||||||
DYNAMICRAFTER_CONFIG = {
|
|
||||||
'in_channels': 8,
|
|
||||||
'out_channels': 4,
|
|
||||||
'model_channels': 320,
|
|
||||||
'attention_resolutions': [4, 2, 1],
|
|
||||||
'num_res_blocks': 2,
|
|
||||||
'channel_mult': [1, 2, 4, 4],
|
|
||||||
'num_head_channels': 64,
|
|
||||||
'transformer_depth': 1,
|
|
||||||
'context_dim': 1024,
|
|
||||||
'use_linear': True,
|
|
||||||
'use_checkpoint': False,
|
|
||||||
'temporal_conv': True,
|
|
||||||
'temporal_attention': True,
|
|
||||||
'temporal_selfatt_only': True,
|
|
||||||
'use_relative_position': False,
|
|
||||||
'use_causal_attention': False,
|
|
||||||
'temporal_length': 16,
|
|
||||||
'addition_attention': True,
|
|
||||||
'image_cross_attention': True,
|
|
||||||
'image_cross_attention_scale_learnable': True,
|
|
||||||
'default_fs': 3,
|
|
||||||
'fs_condition': True
|
|
||||||
}
|
|
||||||
|
|
||||||
IMAGE_PROJ_CONFIG = {
|
|
||||||
"dim": 1024,
|
|
||||||
"depth": 4,
|
|
||||||
"dim_head": 64,
|
|
||||||
"heads": 12,
|
|
||||||
"num_queries": 16,
|
|
||||||
"embedding_dim": 1280,
|
|
||||||
"output_dim": 1024,
|
|
||||||
"ff_mult": 4,
|
|
||||||
"video_length": 16
|
|
||||||
}
|
|
||||||
|
|
||||||
def process_list_or_str(target_key_or_keys, k):
|
|
||||||
if isinstance(target_key_or_keys, list):
|
|
||||||
return any([list_k in k for list_k in target_key_or_keys])
|
|
||||||
else:
|
|
||||||
return target_key_or_keys in k
|
|
||||||
|
|
||||||
def simple_state_dict_loader(state_dict: dict, target_key: str, target_dict: dict = None):
|
|
||||||
out_dict = {}
|
|
||||||
|
|
||||||
if target_dict is None:
|
|
||||||
for k, v in state_dict.items():
|
|
||||||
if process_list_or_str(target_key, k):
|
|
||||||
out_dict[k] = v
|
|
||||||
else:
|
|
||||||
for k, v in target_dict.items():
|
|
||||||
out_dict[k] = state_dict[k]
|
|
||||||
|
|
||||||
return out_dict
|
|
||||||
|
|
||||||
def load_image_proj_dict(state_dict: dict):
|
|
||||||
return simple_state_dict_loader(state_dict, 'image_proj')
|
|
||||||
|
|
||||||
def load_dynamicrafter_dict(state_dict: dict):
|
|
||||||
return simple_state_dict_loader(state_dict, 'model.diffusion_model')
|
|
||||||
|
|
||||||
def load_vae_dict(state_dict: dict):
|
|
||||||
return simple_state_dict_loader(state_dict, 'first_stage_model')
|
|
||||||
|
|
||||||
def get_base_model(state_dict: dict, version_checker=False):
|
|
||||||
|
|
||||||
is_256_model = False
|
|
||||||
|
|
||||||
for k in state_dict.keys():
|
|
||||||
if "framestride_embed" in k:
|
|
||||||
is_256_model = True
|
|
||||||
break
|
|
||||||
|
|
||||||
def get_image_proj_model(state_dict: dict):
|
|
||||||
|
|
||||||
state_dict = {k.replace('image_proj_model.', ''): v for k, v in state_dict.items()}
|
|
||||||
#target_dict = Resampler().state_dict()
|
|
||||||
|
|
||||||
ImageProjModel = Resampler(**IMAGE_PROJ_CONFIG)
|
|
||||||
ImageProjModel.load_state_dict(state_dict)
|
|
||||||
|
|
||||||
print("Image Projection Model loaded successfully")
|
|
||||||
#del target_dict
|
|
||||||
return ImageProjModel
|
|
||||||
|
|
||||||
class DynamiCrafterBase(supported_models_base.BASE):
|
|
||||||
unet_config = {}
|
|
||||||
unet_extra_config = {}
|
|
||||||
|
|
||||||
latent_format = latent_formats.SD15
|
|
||||||
|
|
||||||
def process_clip_state_dict(self, state_dict):
|
|
||||||
replace_prefix = {}
|
|
||||||
replace_prefix["conditioner.embedders.0.model."] = "clip_h." #SD2 in sgm format
|
|
||||||
replace_prefix["cond_stage_model.model."] = "clip_h."
|
|
||||||
state_dict = utils.state_dict_prefix_replace(state_dict, replace_prefix, filter_keys=True)
|
|
||||||
state_dict = utils.clip_text_transformers_convert(state_dict, "clip_h.", "clip_h.transformer.")
|
|
||||||
return state_dict
|
|
||||||
|
|
||||||
def process_clip_state_dict_for_saving(self, state_dict):
|
|
||||||
replace_prefix = {}
|
|
||||||
replace_prefix["clip_h"] = "cond_stage_model.model"
|
|
||||||
state_dict = utils.state_dict_prefix_replace(state_dict, replace_prefix)
|
|
||||||
state_dict = diffusers_convert.convert_text_enc_state_dict_v20(state_dict)
|
|
||||||
return state_dict
|
|
||||||
|
|
||||||
def clip_target(self):
|
|
||||||
return supported_models_base.ClipTarget(sd2_clip.SD2Tokenizer, sd2_clip.SD2ClipModel)
|
|
||||||
|
|
||||||
def process_dict_version(self, state_dict: dict):
|
|
||||||
processed_dict = OrderedDict()
|
|
||||||
is_eps = False
|
|
||||||
|
|
||||||
for k in list(state_dict.keys()):
|
|
||||||
if "framestride_embed" in k:
|
|
||||||
new_key = k.replace("framestride_embed", "fps_embedding")
|
|
||||||
processed_dict[new_key] = state_dict[k]
|
|
||||||
is_eps = True
|
|
||||||
continue
|
|
||||||
|
|
||||||
processed_dict[k] = state_dict[k]
|
|
||||||
|
|
||||||
return processed_dict, is_eps
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@@ -1,82 +0,0 @@
|
|||||||
import importlib
|
|
||||||
import numpy as np
|
|
||||||
import cv2
|
|
||||||
import torch
|
|
||||||
import torch.distributed as dist
|
|
||||||
|
|
||||||
MODEL_EXTS = ['ckpt', 'safetensors', 'bin']
|
|
||||||
|
|
||||||
def get_models_directory(directory: list):
|
|
||||||
files_list = list(filter(lambda f: f.split(".")[-1] in MODEL_EXTS, directory))
|
|
||||||
return files_list
|
|
||||||
|
|
||||||
def count_params(model, verbose=False):
|
|
||||||
total_params = sum(p.numel() for p in model.parameters())
|
|
||||||
if verbose:
|
|
||||||
print(f"{model.__class__.__name__} has {total_params*1.e-6:.2f} M params.")
|
|
||||||
return total_params
|
|
||||||
|
|
||||||
|
|
||||||
def check_istarget(name, para_list):
|
|
||||||
"""
|
|
||||||
name: full name of source para
|
|
||||||
para_list: partial name of target para
|
|
||||||
"""
|
|
||||||
istarget=False
|
|
||||||
for para in para_list:
|
|
||||||
if para in name:
|
|
||||||
return True
|
|
||||||
return istarget
|
|
||||||
|
|
||||||
|
|
||||||
def instantiate_from_config(config):
|
|
||||||
if not "target" in config:
|
|
||||||
if config == '__is_first_stage__':
|
|
||||||
return None
|
|
||||||
elif config == "__is_unconditional__":
|
|
||||||
return None
|
|
||||||
raise KeyError("Expected key `target` to instantiate.")
|
|
||||||
return get_obj_from_str(config["target"])(**config.get("params", dict()))
|
|
||||||
|
|
||||||
|
|
||||||
def get_obj_from_str(string, reload=False):
|
|
||||||
module, cls = string.rsplit(".", 1)
|
|
||||||
if reload:
|
|
||||||
module_imp = importlib.import_module(module)
|
|
||||||
importlib.reload(module_imp)
|
|
||||||
return getattr(importlib.import_module(module, package=None), cls)
|
|
||||||
|
|
||||||
|
|
||||||
def load_npz_from_dir(data_dir):
|
|
||||||
data = [np.load(os.path.join(data_dir, data_name))['arr_0'] for data_name in os.listdir(data_dir)]
|
|
||||||
data = np.concatenate(data, axis=0)
|
|
||||||
return data
|
|
||||||
|
|
||||||
|
|
||||||
def load_npz_from_paths(data_paths):
|
|
||||||
data = [np.load(data_path)['arr_0'] for data_path in data_paths]
|
|
||||||
data = np.concatenate(data, axis=0)
|
|
||||||
return data
|
|
||||||
|
|
||||||
|
|
||||||
def resize_numpy_image(image, max_resolution=512 * 512, resize_short_edge=None):
|
|
||||||
h, w = image.shape[:2]
|
|
||||||
if resize_short_edge is not None:
|
|
||||||
k = resize_short_edge / min(h, w)
|
|
||||||
else:
|
|
||||||
k = max_resolution / (h * w)
|
|
||||||
k = k**0.5
|
|
||||||
h = int(np.round(h * k / 64)) * 64
|
|
||||||
w = int(np.round(w * k / 64)) * 64
|
|
||||||
image = cv2.resize(image, (w, h), interpolation=cv2.INTER_LANCZOS4)
|
|
||||||
return image
|
|
||||||
|
|
||||||
|
|
||||||
def setup_dist(args):
|
|
||||||
if dist.is_initialized():
|
|
||||||
return
|
|
||||||
torch.cuda.set_device(args.local_rank)
|
|
||||||
torch.distributed.init_process_group(
|
|
||||||
'nccl',
|
|
||||||
init_method='env://'
|
|
||||||
)
|
|
||||||
-7539
File diff suppressed because it is too large
Load Diff
@@ -1,23 +0,0 @@
|
|||||||
from .parsing_api import onnx_inference
|
|
||||||
from ..libs.utils import install_package
|
|
||||||
|
|
||||||
class HumanParsing:
|
|
||||||
def __init__(self, model_path):
|
|
||||||
self.model_path = model_path
|
|
||||||
self.session = None
|
|
||||||
|
|
||||||
def __call__(self, input_image, mask_components):
|
|
||||||
if self.session is None:
|
|
||||||
install_package('onnxruntime')
|
|
||||||
import onnxruntime as ort
|
|
||||||
|
|
||||||
session_options = ort.SessionOptions()
|
|
||||||
session_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
|
|
||||||
session_options.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
|
|
||||||
# session_options.add_session_config_entry('gpu_id', str(gpu_id))
|
|
||||||
self.session = ort.InferenceSession(self.model_path, sess_options=session_options,
|
|
||||||
providers=['CPUExecutionProvider'])
|
|
||||||
|
|
||||||
parsed_image, mask = onnx_inference(self.session, input_image, mask_components)
|
|
||||||
return parsed_image, mask
|
|
||||||
|
|
||||||
@@ -1,187 +0,0 @@
|
|||||||
#credit to huchenlei for this module
|
|
||||||
#from https://github.com/huchenlei/ComfyUI-IC-Light-Native
|
|
||||||
import torch
|
|
||||||
import numpy as np
|
|
||||||
from typing import Tuple, TypedDict, Callable
|
|
||||||
|
|
||||||
import comfy.model_management
|
|
||||||
from comfy.sd import load_unet
|
|
||||||
from comfy.ldm.models.autoencoder import AutoencoderKL
|
|
||||||
from comfy.model_base import BaseModel
|
|
||||||
from PIL import Image
|
|
||||||
from nodes import VAEEncode
|
|
||||||
|
|
||||||
from ..layer_diffuse.model import ModelPatcher, calculate_weight_adjust_channel
|
|
||||||
from ..libs.image import np2tensor, pil2tensor
|
|
||||||
|
|
||||||
class UnetParams(TypedDict):
|
|
||||||
input: torch.Tensor
|
|
||||||
timestep: torch.Tensor
|
|
||||||
c: dict
|
|
||||||
cond_or_uncond: torch.Tensor
|
|
||||||
|
|
||||||
|
|
||||||
class VAEEncodeArgMax(VAEEncode):
|
|
||||||
def encode(self, vae, pixels):
|
|
||||||
assert isinstance(
|
|
||||||
vae.first_stage_model, AutoencoderKL
|
|
||||||
), "ArgMax only supported for AutoencoderKL"
|
|
||||||
original_sample_mode = vae.first_stage_model.regularization.sample
|
|
||||||
vae.first_stage_model.regularization.sample = False
|
|
||||||
ret = super().encode(vae, pixels)
|
|
||||||
vae.first_stage_model.regularization.sample = original_sample_mode
|
|
||||||
return ret
|
|
||||||
|
|
||||||
class ICLight:
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def apply_c_concat(params: UnetParams, concat_conds) -> UnetParams:
|
|
||||||
"""Apply c_concat on unet call."""
|
|
||||||
sample = params["input"]
|
|
||||||
params["c"]["c_concat"] = torch.cat(
|
|
||||||
(
|
|
||||||
[concat_conds.to(sample.device)]
|
|
||||||
* (sample.shape[0] // concat_conds.shape[0])
|
|
||||||
),
|
|
||||||
dim=0,
|
|
||||||
)
|
|
||||||
return params
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def create_custom_conv(
|
|
||||||
original_conv: torch.nn.Module,
|
|
||||||
dtype: torch.dtype,
|
|
||||||
device=torch.device,
|
|
||||||
) -> torch.nn.Module:
|
|
||||||
with torch.no_grad():
|
|
||||||
new_conv_in = torch.nn.Conv2d(
|
|
||||||
8,
|
|
||||||
original_conv.out_channels,
|
|
||||||
original_conv.kernel_size,
|
|
||||||
original_conv.stride,
|
|
||||||
original_conv.padding,
|
|
||||||
)
|
|
||||||
new_conv_in.weight.zero_()
|
|
||||||
new_conv_in.weight[:, :4, :, :].copy_(original_conv.weight)
|
|
||||||
new_conv_in.bias = original_conv.bias
|
|
||||||
return new_conv_in.to(dtype=dtype, device=device)
|
|
||||||
|
|
||||||
def generate_lighting_image(self, original_image, direction):
|
|
||||||
_, image_height, image_width, _ = original_image.shape
|
|
||||||
match direction:
|
|
||||||
case 'Left Light':
|
|
||||||
gradient = np.linspace(255, 0, image_width)
|
|
||||||
image = np.tile(gradient, (image_height, 1))
|
|
||||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
|
||||||
return np2tensor(input_bg)
|
|
||||||
case 'Right Light':
|
|
||||||
gradient = np.linspace(0, 255, image_width)
|
|
||||||
image = np.tile(gradient, (image_height, 1))
|
|
||||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
|
||||||
return np2tensor(input_bg)
|
|
||||||
case 'Top Light':
|
|
||||||
gradient = np.linspace(255, 0, image_height)[:, None]
|
|
||||||
image = np.tile(gradient, (1, image_width))
|
|
||||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
|
||||||
return np2tensor(input_bg)
|
|
||||||
case 'Bottom Light':
|
|
||||||
gradient = np.linspace(0, 255, image_height)[:, None]
|
|
||||||
image = np.tile(gradient, (1, image_width))
|
|
||||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
|
||||||
return np2tensor(input_bg)
|
|
||||||
case 'Circle Light':
|
|
||||||
x = np.linspace(-1, 1, image_width)
|
|
||||||
y = np.linspace(-1, 1, image_height)
|
|
||||||
x, y = np.meshgrid(x, y)
|
|
||||||
r = np.sqrt(x ** 2 + y ** 2)
|
|
||||||
r = r / r.max()
|
|
||||||
color1 = np.array([0, 0, 0])[np.newaxis, np.newaxis, :]
|
|
||||||
color2 = np.array([255, 255, 255])[np.newaxis, np.newaxis, :]
|
|
||||||
gradient = (color1 * r[..., np.newaxis] + color2 * (1 - r)[..., np.newaxis]).astype(np.uint8)
|
|
||||||
image = pil2tensor(Image.fromarray(gradient))
|
|
||||||
return image
|
|
||||||
case _:
|
|
||||||
image = pil2tensor(Image.new('RGB', (1, 1), (0, 0, 0)))
|
|
||||||
return image
|
|
||||||
|
|
||||||
def generate_source_image(self, original_image, source):
|
|
||||||
batch_size, image_height, image_width, _ = original_image.shape
|
|
||||||
match source:
|
|
||||||
case 'Use Flipped Background Image':
|
|
||||||
if batch_size < 2:
|
|
||||||
raise ValueError('Must be at least 2 image to use flipped background image.')
|
|
||||||
original_image = [img.unsqueeze(0) for img in original_image]
|
|
||||||
image = torch.flip(original_image[1], [2])
|
|
||||||
return image
|
|
||||||
case 'Ambient':
|
|
||||||
input_bg = np.zeros(shape=(image_height, image_width, 3), dtype=np.uint8) + 64
|
|
||||||
return np2tensor(input_bg)
|
|
||||||
case 'Left Light':
|
|
||||||
gradient = np.linspace(224, 32, image_width)
|
|
||||||
image = np.tile(gradient, (image_height, 1))
|
|
||||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
|
||||||
return np2tensor(input_bg)
|
|
||||||
case 'Right Light':
|
|
||||||
gradient = np.linspace(32, 224, image_width)
|
|
||||||
image = np.tile(gradient, (image_height, 1))
|
|
||||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
|
||||||
return np2tensor(input_bg)
|
|
||||||
case 'Top Light':
|
|
||||||
gradient = np.linspace(224, 32, image_height)[:, None]
|
|
||||||
image = np.tile(gradient, (1, image_width))
|
|
||||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
|
||||||
return np2tensor(input_bg)
|
|
||||||
case 'Bottom Light':
|
|
||||||
gradient = np.linspace(32, 224, image_height)[:, None]
|
|
||||||
image = np.tile(gradient, (1, image_width))
|
|
||||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
|
||||||
return np2tensor(input_bg)
|
|
||||||
case _:
|
|
||||||
image = pil2tensor(Image.new('RGB', (1, 1), (0, 0, 0)))
|
|
||||||
return image
|
|
||||||
|
|
||||||
|
|
||||||
def apply(self, ic_model_path, model: ModelPatcher, c_concat: dict, ic_model=None) -> Tuple[ModelPatcher]:
|
|
||||||
try:
|
|
||||||
ModelPatcher.calculate_weight = calculate_weight_adjust_channel(ModelPatcher.calculate_weight)
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
|
|
||||||
device = comfy.model_management.get_torch_device()
|
|
||||||
dtype = comfy.model_management.unet_dtype()
|
|
||||||
work_model = model.clone()
|
|
||||||
|
|
||||||
# Apply scale factor.
|
|
||||||
base_model: BaseModel = work_model.model
|
|
||||||
scale_factor = base_model.model_config.latent_format.scale_factor
|
|
||||||
|
|
||||||
# [B, 4, H, W]
|
|
||||||
concat_conds: torch.Tensor = c_concat["samples"] * scale_factor
|
|
||||||
# [1, 4 * B, H, W]
|
|
||||||
concat_conds = torch.cat([c[None, ...] for c in concat_conds], dim=1)
|
|
||||||
|
|
||||||
def unet_dummy_apply(unet_apply: Callable, params: UnetParams):
|
|
||||||
"""A dummy unet apply wrapper serving as the endpoint of wrapper
|
|
||||||
chain."""
|
|
||||||
return unet_apply(x=params["input"], t=params["timestep"], **params["c"])
|
|
||||||
|
|
||||||
existing_wrapper = work_model.model_options.get(
|
|
||||||
"model_function_wrapper", unet_dummy_apply
|
|
||||||
)
|
|
||||||
|
|
||||||
def wrapper_func(unet_apply: Callable, params: UnetParams):
|
|
||||||
return existing_wrapper(unet_apply, params=self.apply_c_concat(params, concat_conds))
|
|
||||||
|
|
||||||
work_model.set_model_unet_function_wrapper(wrapper_func)
|
|
||||||
if not ic_model:
|
|
||||||
ic_model = load_unet(ic_model_path)
|
|
||||||
ic_model_state_dict = ic_model.model.diffusion_model.state_dict()
|
|
||||||
|
|
||||||
work_model.add_patches(
|
|
||||||
patches={
|
|
||||||
("diffusion_model." + key): (value.to(dtype=dtype, device=device),)
|
|
||||||
for key, value in ic_model_state_dict.items()
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
return (work_model, ic_model)
|
|
||||||
@@ -6,10 +6,10 @@ import itertools
|
|||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
from comfy.sdxl_clip import SDXLClipModel, SDXLRefinerClipModel, SDXLClipG
|
from comfy.sdxl_clip import SDXLClipModel, SDXLRefinerClipModel, SDXLClipG
|
||||||
try:
|
try:
|
||||||
|
from comfy.text_encoders.sd3_clip import SD3ClipModel, T5XXLModel
|
||||||
|
except ImportError:
|
||||||
from comfy.sd3_clip import SD3ClipModel, T5XXLModel
|
from comfy.sd3_clip import SD3ClipModel, T5XXLModel
|
||||||
except:
|
|
||||||
SD3ClipModel, T5XXLModel = None, None
|
|
||||||
pass
|
|
||||||
from nodes import NODE_CLASS_MAPPINGS, ConditioningConcat, ConditioningZeroOut, ConditioningSetTimestepRange, ConditioningCombine
|
from nodes import NODE_CLASS_MAPPINGS, ConditioningConcat, ConditioningZeroOut, ConditioningSetTimestepRange, ConditioningCombine
|
||||||
|
|
||||||
def _grouper(n, iterable):
|
def _grouper(n, iterable):
|
||||||
@@ -325,7 +325,7 @@ def advanced_encode(clip, text, token_normalization, weight_interpretation, w_ma
|
|||||||
out = None
|
out = None
|
||||||
|
|
||||||
if len(tokenized['l']) > 0 or len(tokenized['g']) > 0:
|
if len(tokenized['l']) > 0 or len(tokenized['g']) > 0:
|
||||||
if 'l' in tokenized:
|
if clip.cond_stage_model.clip_l is not None:
|
||||||
lg_out, l_pooled = advanced_encode_from_tokens(tokenized['l'],
|
lg_out, l_pooled = advanced_encode_from_tokens(tokenized['l'],
|
||||||
token_normalization,
|
token_normalization,
|
||||||
weight_interpretation,
|
weight_interpretation,
|
||||||
@@ -334,7 +334,7 @@ def advanced_encode(clip, text, token_normalization, weight_interpretation, w_ma
|
|||||||
else:
|
else:
|
||||||
l_pooled = torch.zeros((1, 768), device=model_management.intermediate_device())
|
l_pooled = torch.zeros((1, 768), device=model_management.intermediate_device())
|
||||||
|
|
||||||
if 'g' in tokenized:
|
if clip.cond_stage_model.clip_g is not None:
|
||||||
g_out, g_pooled = advanced_encode_from_tokens(tokenized['g'],
|
g_out, g_pooled = advanced_encode_from_tokens(tokenized['g'],
|
||||||
token_normalization,
|
token_normalization,
|
||||||
weight_interpretation,
|
weight_interpretation,
|
||||||
@@ -354,7 +354,7 @@ def advanced_encode(clip, text, token_normalization, weight_interpretation, w_ma
|
|||||||
pooled = torch.cat((l_pooled, g_pooled), dim=-1)
|
pooled = torch.cat((l_pooled, g_pooled), dim=-1)
|
||||||
|
|
||||||
# t5xxl
|
# t5xxl
|
||||||
if 't5xxl' in tokenized and clip.cond_stage_model.t5xxl is not None:
|
if 't5xxl' in tokenized:
|
||||||
t5_out, t5_pooled = advanced_encode_from_tokens(tokenized['t5xxl'],
|
t5_out, t5_pooled = advanced_encode_from_tokens(tokenized['t5xxl'],
|
||||||
token_normalization,
|
token_normalization,
|
||||||
weight_interpretation,
|
weight_interpretation,
|
||||||
|
|||||||
@@ -0,0 +1,372 @@
|
|||||||
|
import yaml
|
||||||
|
import pathlib
|
||||||
|
import base64
|
||||||
|
import io
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import pickle
|
||||||
|
import zlib
|
||||||
|
import urllib.parse
|
||||||
|
import urllib.request
|
||||||
|
import urllib.error
|
||||||
|
from enum import Enum
|
||||||
|
from functools import singledispatch
|
||||||
|
from typing import Any, List, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
root_path = pathlib.Path(__file__).parent.parent.parent.parent
|
||||||
|
config_path = os.path.join(root_path, 'config.yaml')
|
||||||
|
|
||||||
|
class BizyAIRAPI:
|
||||||
|
def __init__(self):
|
||||||
|
self.base_url = 'https://bizyair-api.siliconflow.cn/x/v1'
|
||||||
|
self.api_key = None
|
||||||
|
|
||||||
|
|
||||||
|
def getAPIKey(self):
|
||||||
|
if self.api_key is None:
|
||||||
|
if os.path.isfile(config_path):
|
||||||
|
with open(config_path, 'r') as f:
|
||||||
|
data = yaml.load(f, Loader=yaml.FullLoader)
|
||||||
|
if 'BIZYAIR_API_KEY' not in data:
|
||||||
|
raise Exception("Please add BIZYAIR_API_KEY to config.yaml")
|
||||||
|
self.api_key = data['BIZYAIR_API_KEY']
|
||||||
|
else:
|
||||||
|
raise Exception("Please add config.yaml to root path")
|
||||||
|
return self.api_key
|
||||||
|
|
||||||
|
def send_post_request(self, url, payload, headers):
|
||||||
|
try:
|
||||||
|
data = json.dumps(payload).encode("utf-8")
|
||||||
|
req = urllib.request.Request(url, data=data, headers=headers, method="POST")
|
||||||
|
with urllib.request.urlopen(req) as response:
|
||||||
|
response_data = response.read().decode("utf-8")
|
||||||
|
return response_data
|
||||||
|
except urllib.error.URLError as e:
|
||||||
|
if "Unauthorized" in str(e):
|
||||||
|
raise Exception(
|
||||||
|
"Key is invalid, please refer to https://cloud.siliconflow.cn to get the API key.\n"
|
||||||
|
"If you have the key, please click the 'BizyAir Key' button at the bottom right to set the key."
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise Exception(
|
||||||
|
f"Failed to connect to the server: {e}, if you have no key, "
|
||||||
|
)
|
||||||
|
|
||||||
|
# joycaption
|
||||||
|
def joyCaption(self, payload, image, apikey_override=None, API_URL='/supernode/joycaption2'):
|
||||||
|
if apikey_override is not None:
|
||||||
|
api_key = apikey_override
|
||||||
|
else:
|
||||||
|
api_key = self.getAPIKey()
|
||||||
|
url = f"{self.base_url}{API_URL}"
|
||||||
|
print('Sending request to:', url)
|
||||||
|
auth = f"Bearer {api_key}"
|
||||||
|
headers = {
|
||||||
|
"accept": "application/json",
|
||||||
|
"content-type": "application/json",
|
||||||
|
"authorization": auth,
|
||||||
|
}
|
||||||
|
input_image = encode_data(image, disable_image_marker=True)
|
||||||
|
payload["image"] = input_image
|
||||||
|
|
||||||
|
ret: str = self.send_post_request(url=url, payload=payload, headers=headers)
|
||||||
|
ret = json.loads(ret)
|
||||||
|
|
||||||
|
try:
|
||||||
|
if "result" in ret:
|
||||||
|
ret = json.loads(ret["result"])
|
||||||
|
except Exception as e:
|
||||||
|
raise Exception(f"Unexpected response: {ret} {e=}")
|
||||||
|
|
||||||
|
if ret["type"] == "error":
|
||||||
|
raise Exception(ret["message"])
|
||||||
|
|
||||||
|
msg = ret["data"]
|
||||||
|
if msg["type"] not in ("comfyair", "bizyair",):
|
||||||
|
raise Exception(f"Unexpected response type: {msg}")
|
||||||
|
|
||||||
|
caption = msg["data"]
|
||||||
|
|
||||||
|
return caption
|
||||||
|
|
||||||
|
bizyairAPI = BizyAIRAPI()
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
BIZYAIR_DEBUG = True
|
||||||
|
# Marker to identify base64-encoded tensors
|
||||||
|
TENSOR_MARKER = "TENSOR:"
|
||||||
|
IMAGE_MARKER = "IMAGE:"
|
||||||
|
|
||||||
|
|
||||||
|
class TaskStatus(Enum):
|
||||||
|
PENDING = "pending"
|
||||||
|
PROCESSING = "processing"
|
||||||
|
COMPLETED = "completed"
|
||||||
|
|
||||||
|
|
||||||
|
def convert_image_to_rgb(image: Image.Image) -> Image.Image:
|
||||||
|
if image.mode != "RGB":
|
||||||
|
return image.convert("RGB")
|
||||||
|
return image
|
||||||
|
|
||||||
|
|
||||||
|
def encode_image_to_base64(
|
||||||
|
image: Image.Image, format: str = "png", quality: int = 100, lossless=False
|
||||||
|
) -> str:
|
||||||
|
image = convert_image_to_rgb(image)
|
||||||
|
with io.BytesIO() as output:
|
||||||
|
image.save(output, format=format, quality=quality, lossless=lossless)
|
||||||
|
output.seek(0)
|
||||||
|
img_bytes = output.getvalue()
|
||||||
|
if BIZYAIR_DEBUG:
|
||||||
|
print(f"encode_image_to_base64: {format_bytes(len(img_bytes))}")
|
||||||
|
return base64.b64encode(img_bytes).decode("utf-8")
|
||||||
|
|
||||||
|
|
||||||
|
def decode_base64_to_np(img_data: str, format: str = "png") -> np.ndarray:
|
||||||
|
img_bytes = base64.b64decode(img_data)
|
||||||
|
if BIZYAIR_DEBUG:
|
||||||
|
print(f"decode_base64_to_np: {format_bytes(len(img_bytes))}")
|
||||||
|
with io.BytesIO(img_bytes) as input_buffer:
|
||||||
|
img = Image.open(input_buffer)
|
||||||
|
# https://github.com/comfyanonymous/ComfyUI/blob/a178e25912b01abf436eba1cfaab316ba02d272d/nodes.py#L1511
|
||||||
|
img = img.convert("RGB")
|
||||||
|
return np.array(img)
|
||||||
|
|
||||||
|
|
||||||
|
def decode_base64_to_image(img_data: str) -> Image.Image:
|
||||||
|
img_bytes = base64.b64decode(img_data)
|
||||||
|
with io.BytesIO(img_bytes) as input_buffer:
|
||||||
|
img = Image.open(input_buffer)
|
||||||
|
if BIZYAIR_DEBUG:
|
||||||
|
format_info = img.format.upper() if img.format else "Unknown"
|
||||||
|
print(f"decode image format: {format_info}")
|
||||||
|
return img
|
||||||
|
|
||||||
|
|
||||||
|
def format_bytes(num_bytes: int) -> str:
|
||||||
|
"""
|
||||||
|
Converts a number of bytes to a human-readable string with units (B, KB, or MB).
|
||||||
|
|
||||||
|
:param num_bytes: The number of bytes to convert.
|
||||||
|
:return: A string representing the number of bytes in a human-readable format.
|
||||||
|
"""
|
||||||
|
if num_bytes < 1024:
|
||||||
|
return f"{num_bytes} B"
|
||||||
|
elif num_bytes < 1024 * 1024:
|
||||||
|
return f"{num_bytes / 1024:.2f} KB"
|
||||||
|
else:
|
||||||
|
return f"{num_bytes / (1024 * 1024):.2f} MB"
|
||||||
|
|
||||||
|
|
||||||
|
def _legacy_encode_comfy_image(image: torch.Tensor, image_format="png") -> str:
|
||||||
|
input_image = image.cpu().detach().numpy()
|
||||||
|
i = 255.0 * input_image[0]
|
||||||
|
input_image = np.clip(i, 0, 255).astype(np.uint8)
|
||||||
|
base64ed_image = encode_image_to_base64(
|
||||||
|
Image.fromarray(input_image), format=image_format
|
||||||
|
)
|
||||||
|
return base64ed_image
|
||||||
|
|
||||||
|
|
||||||
|
def _legacy_decode_comfy_image(
|
||||||
|
img_data: Union[List, str], image_format="png"
|
||||||
|
) -> torch.tensor:
|
||||||
|
if isinstance(img_data, List):
|
||||||
|
decoded_imgs = [decode_comfy_image(x, old_version=True) for x in img_data]
|
||||||
|
|
||||||
|
combined_imgs = torch.cat(decoded_imgs, dim=0)
|
||||||
|
return combined_imgs
|
||||||
|
|
||||||
|
out = decode_base64_to_np(img_data, format=image_format)
|
||||||
|
out = np.array(out).astype(np.float32) / 255.0
|
||||||
|
output = torch.from_numpy(out)[None,]
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def _new_encode_comfy_image(images: torch.Tensor, image_format="WEBP", **kwargs) -> str:
|
||||||
|
"""https://docs.comfy.org/essentials/custom_node_snippets#save-an-image-batch
|
||||||
|
Encode a batch of images to base64 strings.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
images (torch.Tensor): A batch of images.
|
||||||
|
image_format (str, optional): The format of the images. Defaults to "WEBP".
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: A JSON string containing the base64-encoded images.
|
||||||
|
"""
|
||||||
|
results = {}
|
||||||
|
for batch_number, image in enumerate(images):
|
||||||
|
i = 255.0 * image.cpu().numpy()
|
||||||
|
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
|
||||||
|
base64ed_image = encode_image_to_base64(img, format=image_format, **kwargs)
|
||||||
|
results[batch_number] = base64ed_image
|
||||||
|
|
||||||
|
return json.dumps(results)
|
||||||
|
|
||||||
|
|
||||||
|
def _new_decode_comfy_image(img_datas: str, image_format="WEBP") -> torch.tensor:
|
||||||
|
"""
|
||||||
|
Decode a batch of base64-encoded images.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
img_datas (str): A JSON string containing the base64-encoded images.
|
||||||
|
image_format (str, optional): The format of the images. Defaults to "WEBP".
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: A tensor containing the decoded images.
|
||||||
|
"""
|
||||||
|
img_datas = json.loads(img_datas)
|
||||||
|
|
||||||
|
decoded_imgs = []
|
||||||
|
for img_data in img_datas.values():
|
||||||
|
decoded_image = decode_base64_to_np(img_data, format=image_format)
|
||||||
|
decoded_image = np.array(decoded_image).astype(np.float32) / 255.0
|
||||||
|
decoded_imgs.append(torch.from_numpy(decoded_image)[None,])
|
||||||
|
|
||||||
|
return torch.cat(decoded_imgs, dim=0)
|
||||||
|
|
||||||
|
|
||||||
|
def encode_comfy_image(
|
||||||
|
image: torch.Tensor, image_format="WEBP", old_version=False, lossless=False
|
||||||
|
) -> str:
|
||||||
|
if old_version:
|
||||||
|
return _legacy_encode_comfy_image(image, image_format)
|
||||||
|
return _new_encode_comfy_image(image, image_format, lossless=lossless)
|
||||||
|
|
||||||
|
|
||||||
|
def decode_comfy_image(
|
||||||
|
img_data: Union[List, str], image_format="WEBP", old_version=False
|
||||||
|
) -> torch.tensor:
|
||||||
|
if old_version:
|
||||||
|
return _legacy_decode_comfy_image(img_data, image_format)
|
||||||
|
return _new_decode_comfy_image(img_data, image_format)
|
||||||
|
|
||||||
|
|
||||||
|
def tensor_to_base64(tensor: torch.Tensor, compress=True) -> str:
|
||||||
|
tensor_np = tensor.cpu().detach().numpy()
|
||||||
|
|
||||||
|
tensor_bytes = pickle.dumps(tensor_np)
|
||||||
|
if compress:
|
||||||
|
tensor_bytes = zlib.compress(tensor_bytes)
|
||||||
|
|
||||||
|
tensor_b64 = base64.b64encode(tensor_bytes).decode("utf-8")
|
||||||
|
return tensor_b64
|
||||||
|
|
||||||
|
|
||||||
|
def base64_to_tensor(tensor_b64: str, compress=True) -> torch.Tensor:
|
||||||
|
tensor_bytes = base64.b64decode(tensor_b64)
|
||||||
|
|
||||||
|
if compress:
|
||||||
|
tensor_bytes = zlib.decompress(tensor_bytes)
|
||||||
|
|
||||||
|
tensor_np = pickle.loads(tensor_bytes)
|
||||||
|
|
||||||
|
tensor = torch.from_numpy(tensor_np)
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
|
||||||
|
@singledispatch
|
||||||
|
def decode_data(input, old_version=False):
|
||||||
|
raise NotImplementedError(f"Unsupported type: {type(input)}")
|
||||||
|
|
||||||
|
|
||||||
|
@decode_data.register(int)
|
||||||
|
@decode_data.register(float)
|
||||||
|
@decode_data.register(bool)
|
||||||
|
@decode_data.register(type(None))
|
||||||
|
def _(input, **kwargs):
|
||||||
|
return input
|
||||||
|
|
||||||
|
|
||||||
|
@decode_data.register(dict)
|
||||||
|
def _(input, **kwargs):
|
||||||
|
return {k: decode_data(v, **kwargs) for k, v in input.items()}
|
||||||
|
|
||||||
|
|
||||||
|
@decode_data.register(list)
|
||||||
|
def _(input, **kwargs):
|
||||||
|
return [decode_data(x, **kwargs) for x in input]
|
||||||
|
|
||||||
|
|
||||||
|
@decode_data.register(str)
|
||||||
|
def _(input: str, **kwargs):
|
||||||
|
if input.startswith(TENSOR_MARKER):
|
||||||
|
tensor_b64 = input[len(TENSOR_MARKER) :]
|
||||||
|
return base64_to_tensor(tensor_b64)
|
||||||
|
elif input.startswith(IMAGE_MARKER):
|
||||||
|
tensor_b64 = input[len(IMAGE_MARKER) :]
|
||||||
|
old_version = kwargs.get("old_version", False)
|
||||||
|
return decode_comfy_image(tensor_b64, old_version=old_version)
|
||||||
|
return input
|
||||||
|
|
||||||
|
|
||||||
|
@singledispatch
|
||||||
|
def encode_data(output, disable_image_marker=False, old_version=False):
|
||||||
|
raise NotImplementedError(f"Unsupported type: {type(output)}")
|
||||||
|
|
||||||
|
|
||||||
|
@encode_data.register(dict)
|
||||||
|
def _(output, **kwargs):
|
||||||
|
return {k: encode_data(v, **kwargs) for k, v in output.items()}
|
||||||
|
|
||||||
|
|
||||||
|
@encode_data.register(list)
|
||||||
|
def _(output, **kwargs):
|
||||||
|
return [encode_data(x, **kwargs) for x in output]
|
||||||
|
|
||||||
|
|
||||||
|
def is_image_tensor(tensor) -> bool:
|
||||||
|
"""https://docs.comfy.org/essentials/custom_node_datatypes#image
|
||||||
|
|
||||||
|
Check if the given tensor is in the format of an IMAGE (shape [B, H, W, C] where C=3).
|
||||||
|
|
||||||
|
`Args`:
|
||||||
|
tensor (torch.Tensor): The tensor to check.
|
||||||
|
|
||||||
|
`Returns`:
|
||||||
|
bool: True if the tensor is in the IMAGE format, False otherwise.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
if not isinstance(tensor, torch.Tensor):
|
||||||
|
return False
|
||||||
|
|
||||||
|
if len(tensor.shape) != 4:
|
||||||
|
return False
|
||||||
|
|
||||||
|
B, H, W, C = tensor.shape
|
||||||
|
if C != 3:
|
||||||
|
return False
|
||||||
|
|
||||||
|
return True
|
||||||
|
except:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
@encode_data.register(torch.Tensor)
|
||||||
|
def _(output, **kwargs):
|
||||||
|
if is_image_tensor(output) and not kwargs.get("disable_image_marker", False):
|
||||||
|
old_version = kwargs.get("old_version", False)
|
||||||
|
lossless = kwargs.get("lossless", True)
|
||||||
|
return IMAGE_MARKER + encode_comfy_image(
|
||||||
|
output, image_format="WEBP", old_version=old_version, lossless=lossless
|
||||||
|
)
|
||||||
|
return TENSOR_MARKER + tensor_to_base64(output)
|
||||||
|
|
||||||
|
|
||||||
|
@encode_data.register(int)
|
||||||
|
@encode_data.register(float)
|
||||||
|
@encode_data.register(bool)
|
||||||
|
@encode_data.register(type(None))
|
||||||
|
def _(output, **kwargs):
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
@encode_data.register(str)
|
||||||
|
def _(output, **kwargs):
|
||||||
|
return output
|
||||||
@@ -0,0 +1,51 @@
|
|||||||
|
import json
|
||||||
|
import os
|
||||||
|
import yaml
|
||||||
|
import requests
|
||||||
|
import pathlib
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
root_path = pathlib.Path(__file__).parent.parent.parent.parent
|
||||||
|
config_path = os.path.join(root_path,'config.yaml')
|
||||||
|
class FluxAIAPI:
|
||||||
|
def __init__(self):
|
||||||
|
self.api_url = "https://fluxaiimagegenerator.com/api"
|
||||||
|
self.origin = "https://fluxaiimagegenerator.com"
|
||||||
|
self.user_agent = None
|
||||||
|
self.cookie = None
|
||||||
|
|
||||||
|
def promptGenerate(self, text, cookies=None):
|
||||||
|
cookie = self.cookie if cookies is None else cookies
|
||||||
|
if cookie is None:
|
||||||
|
if os.path.isfile(config_path):
|
||||||
|
with open(config_path, 'r') as f:
|
||||||
|
data = yaml.load(f, Loader=yaml.FullLoader)
|
||||||
|
if 'FLUXAI_COOKIE' not in data:
|
||||||
|
raise Exception("Please add FLUXAI_COOKIE to config.yaml")
|
||||||
|
if "FLUXAI_USER_AGENT" in data:
|
||||||
|
self.user_agent = data["FLUXAI_USER_AGENT"]
|
||||||
|
self.cookie = cookie = data['FLUXAI_COOKIE']
|
||||||
|
|
||||||
|
headers = {
|
||||||
|
"Cookie": cookie,
|
||||||
|
"Referer": "https://fluxaiimagegenerator.com/flux-prompt-generator",
|
||||||
|
"Origin": self.origin,
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
}
|
||||||
|
if self.user_agent is not None:
|
||||||
|
headers['User-Agent'] = self.user_agent
|
||||||
|
|
||||||
|
url = self.api_url + '/prompt'
|
||||||
|
json = {
|
||||||
|
"prompt": text
|
||||||
|
}
|
||||||
|
|
||||||
|
response = requests.post(url, json=json, headers=headers)
|
||||||
|
res = response.json()
|
||||||
|
if "error" in res:
|
||||||
|
return res['error']
|
||||||
|
elif "data" in res and "prompt" in res['data']:
|
||||||
|
return res['data']['prompt']
|
||||||
|
|
||||||
|
fluxaiAPI = FluxAIAPI()
|
||||||
|
|
||||||
@@ -5,21 +5,21 @@ import requests
|
|||||||
import pathlib
|
import pathlib
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
from server import PromptServer
|
from server import PromptServer
|
||||||
from .image import tensor2pil, pil2tensor, image2base64, pil2byte
|
from ..image import tensor2pil, pil2tensor, image2base64, pil2byte
|
||||||
from .log import log_node_error
|
from ..log import log_node_error
|
||||||
|
|
||||||
|
|
||||||
root_path = pathlib.Path(__file__).parent.parent.parent
|
root_path = pathlib.Path(__file__).parent.parent.parent.parent
|
||||||
config_path = os.path.join(root_path,'config.yaml')
|
config_path = os.path.join(root_path,'config.yaml')
|
||||||
default_key = [{'name':'Default', 'key':''}]
|
default_key = [{'name':'Default', 'key':''}]
|
||||||
|
|
||||||
|
|
||||||
class StabilityAPI:
|
class StabilityAPI:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.api_url = "https://api.stability.ai"
|
self.api_url = "https://api.stability.ai"
|
||||||
self.api_keys = None
|
self.api_keys = None
|
||||||
self.api_current = 0
|
self.api_current = 0
|
||||||
self.user_info = {}
|
self.user_info = {}
|
||||||
self.getAPIKeys()
|
|
||||||
|
|
||||||
def getErrors(self, code):
|
def getErrors(self, code):
|
||||||
errors = {
|
errors = {
|
||||||
@@ -154,7 +154,6 @@ class StabilityAPI:
|
|||||||
|
|
||||||
stableAPI = StabilityAPI()
|
stableAPI = StabilityAPI()
|
||||||
|
|
||||||
|
|
||||||
@PromptServer.instance.routes.get("/easyuse/stability/api_keys")
|
@PromptServer.instance.routes.get("/easyuse/stability/api_keys")
|
||||||
async def get_stability_api_keys(request):
|
async def get_stability_api_keys(request):
|
||||||
stableAPI.getAPIKeys()
|
stableAPI.getAPIKeys()
|
||||||
+4
-3
@@ -77,10 +77,11 @@ def update_cache(k, tag, v):
|
|||||||
else:
|
else:
|
||||||
cache_count[k] += 1
|
cache_count[k] += 1
|
||||||
def remove_cache(key):
|
def remove_cache(key):
|
||||||
global cache
|
|
||||||
if key == '*':
|
if key == '*':
|
||||||
cache = TaggedCache(cache_settings)
|
cache.clear()
|
||||||
|
cache_count.clear()
|
||||||
elif key in cache:
|
elif key in cache:
|
||||||
del cache[key]
|
del cache[key]
|
||||||
|
cache_count.pop(key, None)
|
||||||
else:
|
else:
|
||||||
print(f"invalid {key}")
|
print(f"invalid {key}")
|
||||||
|
|||||||
+139
-38
@@ -1,52 +1,153 @@
|
|||||||
|
from threading import Event
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
from server import PromptServer
|
from server import PromptServer
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
from comfy import model_management as mm
|
||||||
|
from comfy_execution.graph import ExecutionBlocker
|
||||||
import time
|
import time
|
||||||
|
|
||||||
class ChooserCancelled(Exception):
|
class ChooserCancelled(Exception):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
class ChooserMessage:
|
def get_chooser_cache():
|
||||||
stash = {}
|
"""获取选择器缓存"""
|
||||||
messages = {}
|
if not hasattr(PromptServer.instance, '_easyuse_chooser_node'):
|
||||||
cancelled = False
|
PromptServer.instance._easyuse_chooser_node = {}
|
||||||
|
return PromptServer.instance._easyuse_chooser_node
|
||||||
|
|
||||||
@classmethod
|
def cleanup_session_data(node_id):
|
||||||
def addMessage(cls, id, message):
|
"""清理会话数据"""
|
||||||
if message == '__cancel__':
|
node_data = get_chooser_cache()
|
||||||
cls.messages = {}
|
if node_id in node_data:
|
||||||
cls.cancelled = True
|
session_keys = ["event", "selected", "images", "total_count", "cancelled"]
|
||||||
elif message == '__start__':
|
for key in session_keys:
|
||||||
cls.messages = {}
|
if key in node_data[node_id]:
|
||||||
cls.stash = {}
|
del node_data[node_id][key]
|
||||||
cls.cancelled = False
|
|
||||||
else:
|
def wait_for_chooser(id, images, mode, period=0.1):
|
||||||
cls.messages[str(id)] = message
|
try:
|
||||||
|
node_data = get_chooser_cache()
|
||||||
|
images = [images[i:i + 1, ...] for i in range(images.shape[0])]
|
||||||
|
if mode == "Keep Last Selection":
|
||||||
|
if id in node_data and "last_selection" in node_data[id]:
|
||||||
|
last_selection = node_data[id]["last_selection"]
|
||||||
|
if last_selection and len(last_selection) > 0:
|
||||||
|
valid_indices = [idx for idx in last_selection if 0 <= idx < len(images)]
|
||||||
|
if valid_indices:
|
||||||
|
try:
|
||||||
|
PromptServer.instance.send_sync("easyuse-image-keep-selection", {
|
||||||
|
"id": id,
|
||||||
|
"selected": valid_indices
|
||||||
|
})
|
||||||
|
except Exception as e:
|
||||||
|
pass
|
||||||
|
cleanup_session_data(id)
|
||||||
|
indices_str = ','.join(str(i) for i in valid_indices)
|
||||||
|
images = [images[idx] for idx in valid_indices]
|
||||||
|
images = torch.cat(images, dim=0)
|
||||||
|
return {"result": (images,)}
|
||||||
|
|
||||||
|
if id in node_data:
|
||||||
|
del node_data[id]
|
||||||
|
|
||||||
|
event = Event()
|
||||||
|
node_data[id] = {
|
||||||
|
"event": event,
|
||||||
|
"images": images,
|
||||||
|
"selected": None,
|
||||||
|
"total_count": len(images),
|
||||||
|
"cancelled": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
while id in node_data:
|
||||||
|
node_info = node_data[id]
|
||||||
|
if node_info.get("cancelled", False):
|
||||||
|
cleanup_session_data(id)
|
||||||
|
raise ChooserCancelled("Manual selection cancelled")
|
||||||
|
|
||||||
|
if "selected" in node_info and node_info["selected"] is not None:
|
||||||
|
break
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def waitForMessage(cls, id, period=0.1, asList=False):
|
|
||||||
sid = str(id)
|
|
||||||
while not (sid in cls.messages) and not ("-1" in cls.messages):
|
|
||||||
if cls.cancelled:
|
|
||||||
cls.cancelled = False
|
|
||||||
raise ChooserCancelled()
|
|
||||||
time.sleep(period)
|
time.sleep(period)
|
||||||
if cls.cancelled:
|
|
||||||
cls.cancelled = False
|
if id in node_data:
|
||||||
raise ChooserCancelled()
|
node_info = node_data[id]
|
||||||
message = cls.messages.pop(str(id), None) or cls.messages.pop("-1")
|
selected_indices = node_info.get("selected")
|
||||||
try:
|
|
||||||
if asList:
|
if selected_indices is not None and len(selected_indices) > 0:
|
||||||
return [int(x.strip()) for x in message.split(",")]
|
valid_indices = [idx for idx in selected_indices if 0 <= idx < len(images)]
|
||||||
|
if valid_indices:
|
||||||
|
selected_images = [images[idx] for idx in valid_indices]
|
||||||
|
|
||||||
|
if id not in node_data:
|
||||||
|
node_data[id] = {}
|
||||||
|
node_data[id]["last_selection"] = valid_indices
|
||||||
|
cleanup_session_data(id)
|
||||||
|
selected_images = torch.cat(selected_images, dim=0)
|
||||||
|
return {"result": (selected_images,)}
|
||||||
|
else:
|
||||||
|
cleanup_session_data(id)
|
||||||
|
return {"result": (images[0] if len(images) > 0 else ExecutionBlocker(None),)}
|
||||||
else:
|
else:
|
||||||
return int(message.strip())
|
cleanup_session_data(id)
|
||||||
except ValueError:
|
return {
|
||||||
print(
|
"result": (images[0] if len(images) > 0 else ExecutionBlocker(None),)}
|
||||||
f"ERROR IN IMAGE_CHOOSER - failed to parse '${message}' as ${'comma separated list of ints' if asList else 'int'}")
|
else:
|
||||||
return [1] if asList else 1
|
return {"result": (images[0] if len(images) > 0 else ExecutionBlocker(None),)}
|
||||||
|
|
||||||
|
except ChooserCancelled:
|
||||||
|
raise mm.InterruptProcessingException()
|
||||||
|
except Exception as e:
|
||||||
|
node_data = get_chooser_cache()
|
||||||
|
if id in node_data:
|
||||||
|
cleanup_session_data(id)
|
||||||
|
if 'image_list' in locals() and len(images) > 0:
|
||||||
|
return {"result": (images[0])}
|
||||||
|
else:
|
||||||
|
return {"result": (ExecutionBlocker(None),)}
|
||||||
|
|
||||||
|
|
||||||
@PromptServer.instance.routes.post('/easyuse/image_chooser_message')
|
@PromptServer.instance.routes.post('/easyuse/image_chooser_message')
|
||||||
async def make_image_selection(request):
|
async def handle_image_selection(request):
|
||||||
post = await request.post()
|
try:
|
||||||
ChooserMessage.addMessage(post.get("id"), post.get("message"))
|
data = await request.json()
|
||||||
return web.json_response({})
|
node_id = data.get("node_id")
|
||||||
|
selected = data.get("selected", [])
|
||||||
|
action = data.get("action")
|
||||||
|
|
||||||
|
node_data = get_chooser_cache()
|
||||||
|
|
||||||
|
if node_id not in node_data:
|
||||||
|
return web.json_response({"code": -1, "error": "Node data does not exist"})
|
||||||
|
|
||||||
|
try:
|
||||||
|
node_info = node_data[node_id]
|
||||||
|
|
||||||
|
if "total_count" not in node_info:
|
||||||
|
return web.json_response({"code": -1, "error": "The node has been processed"})
|
||||||
|
|
||||||
|
if action == "cancel":
|
||||||
|
node_info["cancelled"] = True
|
||||||
|
node_info["selected"] = []
|
||||||
|
elif action == "select" and isinstance(selected, list):
|
||||||
|
valid_indices = [idx for idx in selected if isinstance(idx, int) and 0 <= idx < node_info["total_count"]]
|
||||||
|
if valid_indices:
|
||||||
|
node_info["selected"] = valid_indices
|
||||||
|
node_info["cancelled"] = False
|
||||||
|
else:
|
||||||
|
return web.json_response({"code": -1, "error": "Invalid Selection Index"})
|
||||||
|
else:
|
||||||
|
return web.json_response({"code": -1, "error": "Invalid operation"})
|
||||||
|
|
||||||
|
node_info["event"].set()
|
||||||
|
return web.json_response({"code": 1})
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
if node_id in node_data and "event" in node_data[node_id]:
|
||||||
|
node_data[node_id]["event"].set()
|
||||||
|
return web.json_response({"code": -1, "message": "Processing Failed"})
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
return web.json_response({"code": -1, "message": "Request Failed"})
|
||||||
|
|||||||
+17
-9
@@ -4,32 +4,40 @@ from .translate import zh_to_en, has_chinese
|
|||||||
from .wildcards import process_with_loras
|
from .wildcards import process_with_loras
|
||||||
from .adv_encode import advanced_encode
|
from .adv_encode import advanced_encode
|
||||||
|
|
||||||
from nodes import ConditioningConcat, ConditioningCombine, ConditioningAverage, ConditioningSetTimestepRange
|
from nodes import ConditioningConcat, ConditioningCombine, ConditioningAverage, ConditioningSetTimestepRange, CLIPTextEncode
|
||||||
|
|
||||||
def prompt_to_cond(type, model, clip, clip_skip, lora_stack, text, prompt_token_normalization, prompt_weight_interpretation, a1111_prompt_style ,my_unique_id, prompt, easyCache, can_load_lora=True, steps=None):
|
def prompt_to_cond(type, model, clip, clip_skip, lora_stack, text, prompt_token_normalization, prompt_weight_interpretation, a1111_prompt_style ,my_unique_id, prompt, easyCache, can_load_lora=True, steps=None, model_type=None):
|
||||||
styles_selector = is_linked_styles_selector(prompt, my_unique_id, type)
|
styles_selector = is_linked_styles_selector(prompt, my_unique_id, type)
|
||||||
title = "正面提示词" if type == 'positive' else "负面提示词"
|
title = "Positive encoding" if type == 'positive' else "Negative encoding"
|
||||||
log_node_warn("正在进行" + title + "...")
|
|
||||||
|
|
||||||
# Translate cn to en
|
# Translate cn to en
|
||||||
if has_chinese(text):
|
if model_type not in ['hydit'] and text is not None and has_chinese(text):
|
||||||
text = zh_to_en([text])[0]
|
text = zh_to_en([text])[0]
|
||||||
|
|
||||||
|
if model_type in ['hydit', 'flux', 'mochi', 'anima', 'krea2']:
|
||||||
|
log_node_warn(title + "...")
|
||||||
|
embeddings_final, = CLIPTextEncode().encode(clip, text) if text is not None else (None,)
|
||||||
|
|
||||||
|
return (embeddings_final, "", model, clip)
|
||||||
|
|
||||||
|
log_node_warn(title + "...")
|
||||||
|
|
||||||
positive_seed = find_wildcards_seed(my_unique_id, text, prompt)
|
positive_seed = find_wildcards_seed(my_unique_id, text, prompt)
|
||||||
model, clip, text, cond_decode, show_prompt, pipe_lora_stack = process_with_loras(
|
model, clip, text, cond_decode, show_prompt, pipe_lora_stack = process_with_loras(
|
||||||
text, model, clip, type, positive_seed, can_load_lora, lora_stack, easyCache)
|
text, model, clip, type, positive_seed, can_load_lora, lora_stack, easyCache)
|
||||||
wildcard_prompt = cond_decode if show_prompt or styles_selector else ""
|
wildcard_prompt = cond_decode if show_prompt or styles_selector else ""
|
||||||
|
|
||||||
clipped = clip.clone()
|
clipped = clip.clone()
|
||||||
if clip_skip != 0:
|
# 当clip模型不存在t5xxl时,可执行跳过层
|
||||||
clipped.clip_layer(clip_skip)
|
if not hasattr(clip.cond_stage_model, 't5xxl'):
|
||||||
|
if clip_skip != 0:
|
||||||
|
clipped.clip_layer(clip_skip)
|
||||||
|
|
||||||
log_node_warn("正在进行" + title + "编码...")
|
|
||||||
steps = steps if steps is not None else find_nearest_steps(my_unique_id, prompt)
|
steps = steps if steps is not None else find_nearest_steps(my_unique_id, prompt)
|
||||||
return (advanced_encode(clipped, text, prompt_token_normalization,
|
return (advanced_encode(clipped, text, prompt_token_normalization,
|
||||||
prompt_weight_interpretation, w_max=1.0,
|
prompt_weight_interpretation, w_max=1.0,
|
||||||
apply_to_pooled='enable',
|
apply_to_pooled='enable',
|
||||||
a1111_prompt_style=a1111_prompt_style, steps=steps), wildcard_prompt, model, clipped)
|
a1111_prompt_style=a1111_prompt_style, steps=steps) if text is not None else None, wildcard_prompt, model, clipped)
|
||||||
|
|
||||||
def set_cond(old_cond, new_cond, mode, average_strength, old_cond_start, old_cond_end, new_cond_start, new_cond_end):
|
def set_cond(old_cond, new_cond, mode, average_strength, old_cond_start, old_cond_end, new_cond_start, new_cond_end):
|
||||||
if not old_cond:
|
if not old_cond:
|
||||||
|
|||||||
+28
-4
@@ -3,16 +3,40 @@ import comfy.controlnet
|
|||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
from nodes import NODE_CLASS_MAPPINGS
|
from nodes import NODE_CLASS_MAPPINGS
|
||||||
|
|
||||||
|
union_controlnet_types = {"auto": -1, "openpose": 0, "depth": 1, "hed/pidi/scribble/ted": 2, "canny/lineart/anime_lineart/mlsd": 3, "normal": 4, "segment": 5, "tile": 6, "repaint": 7}
|
||||||
|
|
||||||
class easyControlnet:
|
class easyControlnet:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def apply(self, control_net_name, image, positive, negative, strength, start_percent=0, end_percent=1, control_net=None, scale_soft_weights=1, mask=None, easyCache=None, use_cache=True):
|
def apply(self, control_net_name, image, positive, negative, strength, start_percent=0, end_percent=1, control_net=None, scale_soft_weights=1, mask=None, union_type=None, easyCache=None, use_cache=True, model=None, vae=None):
|
||||||
if strength == 0:
|
if strength == 0:
|
||||||
return (positive, negative)
|
return (positive, negative)
|
||||||
|
|
||||||
if control_net is None:
|
# kolors controlnet patch
|
||||||
control_net = easyCache.load_controlnet(control_net_name, scale_soft_weights, use_cache)
|
from ..modules.kolors.loader import is_kolors_model, applyKolorsUnet
|
||||||
|
if is_kolors_model(model):
|
||||||
|
from ..modules.kolors.model_patch import patch_controlnet
|
||||||
|
if control_net is None:
|
||||||
|
with applyKolorsUnet():
|
||||||
|
control_net = easyCache.load_controlnet(control_net_name, scale_soft_weights, use_cache)
|
||||||
|
control_net = patch_controlnet(model, control_net)
|
||||||
|
else:
|
||||||
|
if control_net is None:
|
||||||
|
if easyCache is not None:
|
||||||
|
control_net = easyCache.load_controlnet(control_net_name, scale_soft_weights, use_cache)
|
||||||
|
else:
|
||||||
|
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||||
|
control_net = comfy.controlnet.load_controlnet(controlnet_path)
|
||||||
|
|
||||||
|
# union controlnet
|
||||||
|
if union_type is not None:
|
||||||
|
control_net = control_net.copy()
|
||||||
|
type_number = union_controlnet_types[union_type]
|
||||||
|
if type_number >= 0:
|
||||||
|
control_net.set_extra_arg("control_type", [type_number])
|
||||||
|
else:
|
||||||
|
control_net.set_extra_arg("control_type", [])
|
||||||
|
|
||||||
if mask is not None:
|
if mask is not None:
|
||||||
mask = mask.to(self.device)
|
mask = mask.to(self.device)
|
||||||
@@ -49,7 +73,7 @@ class easyControlnet:
|
|||||||
if prev_cnet in cnets:
|
if prev_cnet in cnets:
|
||||||
c_net = cnets[prev_cnet]
|
c_net = cnets[prev_cnet]
|
||||||
else:
|
else:
|
||||||
c_net = control_net.copy().set_cond_hint(control_hint, strength, (start_percent, end_percent))
|
c_net = control_net.copy().set_cond_hint(control_hint, strength, (start_percent, end_percent), vae)
|
||||||
c_net.set_previous_controlnet(prev_cnet)
|
c_net.set_previous_controlnet(prev_cnet)
|
||||||
cnets[prev_cnet] = c_net
|
cnets[prev_cnet] = c_net
|
||||||
|
|
||||||
|
|||||||
@@ -1,113 +0,0 @@
|
|||||||
#credit to Acly for this module
|
|
||||||
#from https://github.com/Acly/comfyui-inpaint-nodes
|
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
import comfy
|
|
||||||
from comfy.model_base import BaseModel
|
|
||||||
from comfy.model_patcher import ModelPatcher
|
|
||||||
from comfy.model_management import cast_to_device
|
|
||||||
|
|
||||||
from .log import log_node_warn, log_node_error, log_node_info
|
|
||||||
|
|
||||||
# Inpaint
|
|
||||||
original_calculate_weight = ModelPatcher.calculate_weight
|
|
||||||
injected_model_patcher_calculate_weight = False
|
|
||||||
|
|
||||||
class InpaintHead(torch.nn.Module):
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self.head = torch.nn.Parameter(torch.empty(size=(320, 5, 3, 3), device="cpu"))
|
|
||||||
|
|
||||||
def __call__(self, x):
|
|
||||||
x = F.pad(x, (1, 1, 1, 1), "replicate")
|
|
||||||
return F.conv2d(x, weight=self.head)
|
|
||||||
|
|
||||||
def calculate_weight_patched(self: ModelPatcher, patches, weight, key):
|
|
||||||
remaining = []
|
|
||||||
|
|
||||||
for p in patches:
|
|
||||||
alpha = p[0]
|
|
||||||
v = p[1]
|
|
||||||
|
|
||||||
is_fooocus_patch = isinstance(v, tuple) and len(v) == 2 and v[0] == "fooocus"
|
|
||||||
if not is_fooocus_patch:
|
|
||||||
remaining.append(p)
|
|
||||||
continue
|
|
||||||
|
|
||||||
if alpha != 0.0:
|
|
||||||
v = v[1]
|
|
||||||
w1 = cast_to_device(v[0], weight.device, torch.float32)
|
|
||||||
if w1.shape == weight.shape:
|
|
||||||
w_min = cast_to_device(v[1], weight.device, torch.float32)
|
|
||||||
w_max = cast_to_device(v[2], weight.device, torch.float32)
|
|
||||||
w1 = (w1 / 255.0) * (w_max - w_min) + w_min
|
|
||||||
weight += alpha * cast_to_device(w1, weight.device, weight.dtype)
|
|
||||||
else:
|
|
||||||
pass
|
|
||||||
# log_node_warn(self.node_name,
|
|
||||||
# f"Shape mismatch {key}, weight not merged ({w1.shape} != {weight.shape})"
|
|
||||||
# )
|
|
||||||
|
|
||||||
if len(remaining) > 0:
|
|
||||||
return original_calculate_weight(self, remaining, weight, key)
|
|
||||||
return weight
|
|
||||||
|
|
||||||
def inject_patched_calculate_weight():
|
|
||||||
global injected_model_patcher_calculate_weight
|
|
||||||
if not injected_model_patcher_calculate_weight:
|
|
||||||
print(
|
|
||||||
"[comfyui-inpaint-nodes] Injecting patched comfy.model_patcher.ModelPatcher.calculate_weight"
|
|
||||||
)
|
|
||||||
ModelPatcher.calculate_weight = calculate_weight_patched
|
|
||||||
injected_model_patcher_calculate_weight = True
|
|
||||||
|
|
||||||
class InpaintWorker:
|
|
||||||
def __init__(self, node_name):
|
|
||||||
self.node_name = node_name if node_name is not None else ""
|
|
||||||
|
|
||||||
def load_fooocus_patch(self, lora: dict, to_load: dict):
|
|
||||||
patch_dict = {}
|
|
||||||
loaded_keys = set()
|
|
||||||
for key in to_load.values():
|
|
||||||
if value := lora.get(key, None):
|
|
||||||
patch_dict[key] = ("fooocus", value)
|
|
||||||
loaded_keys.add(key)
|
|
||||||
|
|
||||||
not_loaded = sum(1 for x in lora if x not in loaded_keys)
|
|
||||||
if not_loaded > 0:
|
|
||||||
log_node_info(self.node_name,
|
|
||||||
f"{len(loaded_keys)} Lora keys loaded, {not_loaded} remaining keys not found in model."
|
|
||||||
)
|
|
||||||
return patch_dict
|
|
||||||
|
|
||||||
|
|
||||||
def patch(self, model, latent, patch):
|
|
||||||
base_model: BaseModel = model.model
|
|
||||||
latent_pixels = base_model.process_latent_in(latent["samples"])
|
|
||||||
noise_mask = latent["noise_mask"].round()
|
|
||||||
latent_mask = F.max_pool2d(noise_mask, (8, 8)).round().to(latent_pixels)
|
|
||||||
|
|
||||||
inpaint_head_model, inpaint_lora = patch
|
|
||||||
feed = torch.cat([latent_mask, latent_pixels], dim=1)
|
|
||||||
inpaint_head_model.to(device=feed.device, dtype=feed.dtype)
|
|
||||||
inpaint_head_feature = inpaint_head_model(feed)
|
|
||||||
|
|
||||||
def input_block_patch(h, transformer_options):
|
|
||||||
if transformer_options["block"][1] == 0:
|
|
||||||
h = h + inpaint_head_feature.to(h)
|
|
||||||
return h
|
|
||||||
|
|
||||||
lora_keys = comfy.lora.model_lora_keys_unet(model.model, {})
|
|
||||||
lora_keys.update({x: x for x in base_model.state_dict().keys()})
|
|
||||||
loaded_lora = self.load_fooocus_patch(inpaint_lora, lora_keys)
|
|
||||||
|
|
||||||
m = model.clone()
|
|
||||||
m.set_model_input_block_patch(input_block_patch)
|
|
||||||
patched = m.add_patches(loaded_lora, 1.0)
|
|
||||||
|
|
||||||
not_patched_count = sum(1 for x in loaded_lora if x not in patched)
|
|
||||||
if not_patched_count > 0:
|
|
||||||
log_node_error(self.node_name, f"Failed to patch {not_patched_count} keys")
|
|
||||||
|
|
||||||
inject_patched_calculate_weight()
|
|
||||||
return (m,)
|
|
||||||
+34
-1
@@ -105,6 +105,11 @@ class blendImage:
|
|||||||
return blended_image
|
return blended_image
|
||||||
|
|
||||||
|
|
||||||
|
def empty_image(width, height, batch_size=1, color=0):
|
||||||
|
r = torch.full([batch_size, height, width, 1], ((color >> 16) & 0xFF) / 0xFF)
|
||||||
|
g = torch.full([batch_size, height, width, 1], ((color >> 8) & 0xFF) / 0xFF)
|
||||||
|
b = torch.full([batch_size, height, width, 1], ((color) & 0xFF) / 0xFF)
|
||||||
|
return torch.cat((r, g, b), dim=-1)
|
||||||
|
|
||||||
|
|
||||||
class ResizeMode(Enum):
|
class ResizeMode(Enum):
|
||||||
@@ -120,7 +125,35 @@ class ResizeMode(Enum):
|
|||||||
return 2
|
return 2
|
||||||
assert False, "NOTREACHED"
|
assert False, "NOTREACHED"
|
||||||
|
|
||||||
|
# credit by https://github.com/chflame163/ComfyUI_LayerStyle/blob/main/py/imagefunc.py#L591C1-L617C22
|
||||||
|
def fit_resize_image(image: Image, target_width: int, target_height: int, fit: str, resize_sampler: str,
|
||||||
|
background_color: str = '#000000') -> Image:
|
||||||
|
image = image.convert('RGB')
|
||||||
|
orig_width, orig_height = image.size
|
||||||
|
if image is not None:
|
||||||
|
if fit == 'letterbox':
|
||||||
|
if orig_width / orig_height > target_width / target_height: # 更宽,上下留黑
|
||||||
|
fit_width = target_width
|
||||||
|
fit_height = int(target_width / orig_width * orig_height)
|
||||||
|
else: # 更瘦,左右留黑
|
||||||
|
fit_height = target_height
|
||||||
|
fit_width = int(target_height / orig_height * orig_width)
|
||||||
|
fit_image = image.resize((fit_width, fit_height), resize_sampler)
|
||||||
|
ret_image = Image.new('RGB', size=(target_width, target_height), color=background_color)
|
||||||
|
ret_image.paste(fit_image, box=((target_width - fit_width) // 2, (target_height - fit_height) // 2))
|
||||||
|
elif fit == 'crop':
|
||||||
|
if orig_width / orig_height > target_width / target_height: # 更宽,裁左右
|
||||||
|
fit_width = int(orig_height * target_width / target_height)
|
||||||
|
fit_image = image.crop(
|
||||||
|
((orig_width - fit_width) // 2, 0, (orig_width - fit_width) // 2 + fit_width, orig_height))
|
||||||
|
else: # 更瘦,裁上下
|
||||||
|
fit_height = int(orig_width * target_height / target_width)
|
||||||
|
fit_image = image.crop(
|
||||||
|
(0, (orig_height - fit_height) // 2, orig_width, (orig_height - fit_height) // 2 + fit_height))
|
||||||
|
ret_image = fit_image.resize((target_width, target_height), resize_sampler)
|
||||||
|
else:
|
||||||
|
ret_image = image.resize((target_width, target_height), resize_sampler)
|
||||||
|
return ret_image
|
||||||
|
|
||||||
# CLIP反推
|
# CLIP反推
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
|
|||||||
+170
-86
@@ -1,4 +1,4 @@
|
|||||||
import time, os, psutil
|
import re, time, os, psutil
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import comfy.sd
|
import comfy.sd
|
||||||
@@ -8,17 +8,18 @@ from comfy.model_patcher import ModelPatcher
|
|||||||
from nodes import NODE_CLASS_MAPPINGS
|
from nodes import NODE_CLASS_MAPPINGS
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from .log import log_node_info, log_node_error
|
from .log import log_node_info, log_node_error
|
||||||
from ..dit.hunyuanDiT.loader import EXM_HyDiT_Tenc_Temp, load_hydit
|
from .utils import get_sd_version
|
||||||
from ..dit.pixArt.loader import load_pixart
|
from ..config import DIFFUSION_MODEL_XY_DEFAULTS, DIFFUSION_MODEL_CLIP_TYPES
|
||||||
|
from ..modules.dit.pixArt.loader import load_pixart
|
||||||
|
|
||||||
stable_diffusion_loaders = ["easy fullLoader", "easy a1111Loader", "easy comfyLoader", "easy zero123Loader", "easy svdLoader"]
|
diffusion_loaders = ["easy fullLoader", "easy a1111Loader", "easy fluxLoader", "easy comfyLoader", "easy hunyuanDiTLoader", "easy zero123Loader", "easy svdLoader"]
|
||||||
stable_cascade_loaders = ["easy cascadeLoader"]
|
stable_cascade_loaders = ["easy cascadeLoader"]
|
||||||
dit_loaders = ['easy hunyuanDiTLoader', 'easy pixArtLoader']
|
dit_loaders = ['easy pixArtLoader']
|
||||||
controlnet_loaders = ["easy controlnetLoader", "easy controlnetLoaderADV"]
|
controlnet_loaders = ["easy controlnetLoader", "easy controlnetLoaderADV", "easy controlnetLoader++"]
|
||||||
instant_loaders = ["easy instantIDApply", "easy instantIDApplyADV"]
|
instant_loaders = ["easy instantIDApply", "easy instantIDApplyADV"]
|
||||||
cascade_vae_node = ["easy preSamplingCascade", "easy fullCascadeKSampler"]
|
cascade_vae_node = ["easy preSamplingCascade", "easy fullCascadeKSampler"]
|
||||||
model_merge_node = ["easy XYInputs: ModelMergeBlocks"]
|
model_merge_node = ["easy XYInputs: ModelMergeBlocks"]
|
||||||
lora_widget = ["easy fullLoader", "easy a1111Loader", "easy comfyLoader"]
|
lora_widget = ["easy fullLoader", "easy a1111Loader", "easy comfyLoader", "easy fluxLoader"]
|
||||||
|
|
||||||
class easyLoader:
|
class easyLoader:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
@@ -32,8 +33,9 @@ class easyLoader:
|
|||||||
"lora": defaultdict(dict), # {lora_name: {UID: (model_lora, clip_lora)}}
|
"lora": defaultdict(dict), # {lora_name: {UID: (model_lora, clip_lora)}}
|
||||||
"controlnet": defaultdict(dict),
|
"controlnet": defaultdict(dict),
|
||||||
"t5": defaultdict(tuple),
|
"t5": defaultdict(tuple),
|
||||||
|
"chatglm3": defaultdict(tuple),
|
||||||
}
|
}
|
||||||
self.memory_threshold = self.determine_memory_threshold(0.7)
|
self.memory_threshold = self.determine_memory_threshold(1)
|
||||||
self.lora_name_cache = []
|
self.lora_name_cache = []
|
||||||
|
|
||||||
def clean_values(self, values: str):
|
def clean_values(self, values: str):
|
||||||
@@ -91,6 +93,7 @@ class easyLoader:
|
|||||||
desired_lora_settings = set()
|
desired_lora_settings = set()
|
||||||
desired_controlnet_names = set()
|
desired_controlnet_names = set()
|
||||||
desired_t5_names = set()
|
desired_t5_names = set()
|
||||||
|
desired_glm3_names = set()
|
||||||
|
|
||||||
for entry in prompt.values():
|
for entry in prompt.values():
|
||||||
class_type = entry["class_type"]
|
class_type = entry["class_type"]
|
||||||
@@ -100,10 +103,15 @@ class easyLoader:
|
|||||||
setting = f'{lora_name};{entry["inputs"]["lora_model_strength"]};{entry["inputs"]["lora_clip_strength"]}'
|
setting = f'{lora_name};{entry["inputs"]["lora_model_strength"]};{entry["inputs"]["lora_clip_strength"]}'
|
||||||
desired_lora_settings.add(setting)
|
desired_lora_settings.add(setting)
|
||||||
|
|
||||||
if class_type in stable_diffusion_loaders:
|
if class_type in diffusion_loaders:
|
||||||
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name", prompt))
|
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name", prompt))
|
||||||
desired_vae_names.add(self.get_input_value(entry, "vae_name"))
|
desired_vae_names.add(self.get_input_value(entry, "vae_name"))
|
||||||
|
|
||||||
|
elif class_type in ['easy kolorsLoader']:
|
||||||
|
desired_unet_names.add(self.get_input_value(entry, "unet_name"))
|
||||||
|
desired_vae_names.add(self.get_input_value(entry, "vae_name"))
|
||||||
|
desired_glm3_names.add(self.get_input_value(entry, "chatglm3_name"))
|
||||||
|
|
||||||
elif class_type in dit_loaders:
|
elif class_type in dit_loaders:
|
||||||
t5_name = self.get_input_value(entry, "mt5_name") if "mt5_name" in entry["inputs"] else None
|
t5_name = self.get_input_value(entry, "mt5_name") if "mt5_name" in entry["inputs"] else None
|
||||||
clip_name = self.get_input_value(entry, "clip_name") if "clip_name" in entry["inputs"] else None
|
clip_name = self.get_input_value(entry, "clip_name") if "clip_name" in entry["inputs"] else None
|
||||||
@@ -139,6 +147,28 @@ class easyLoader:
|
|||||||
scale_soft_weights = self.get_input_value(entry, "cn_soft_weights")
|
scale_soft_weights = self.get_input_value(entry, "cn_soft_weights")
|
||||||
desired_controlnet_names.add(f'{control_net_name};{scale_soft_weights}')
|
desired_controlnet_names.add(f'{control_net_name};{scale_soft_weights}')
|
||||||
|
|
||||||
|
elif class_type == "easy diffusionModelLoader":
|
||||||
|
desired_unet_names.add(self.get_input_value(entry, "model_name", prompt))
|
||||||
|
clip_name = self.get_input_value(entry, "clip_name", prompt)
|
||||||
|
vae_name = self.get_input_value(entry, "vae_name", prompt)
|
||||||
|
if clip_name not in ("None", "Auto"):
|
||||||
|
desired_clip_names.add(clip_name)
|
||||||
|
if vae_name not in ("None", "Auto"):
|
||||||
|
desired_vae_names.add(vae_name)
|
||||||
|
|
||||||
|
elif class_type == "easy XYInputs: DiffusionModel":
|
||||||
|
model_count = int(self.get_input_value(entry, "model_count", prompt) or 0)
|
||||||
|
for i in range(1, model_count + 1):
|
||||||
|
model_name = self.get_input_value(entry, f"model_name_{i}", prompt)
|
||||||
|
if model_name and model_name != "None":
|
||||||
|
desired_unet_names.add(model_name)
|
||||||
|
clip_name = self.get_input_value(entry, f"clip_name_{i}", prompt)
|
||||||
|
if clip_name not in ("None", "Auto"):
|
||||||
|
desired_clip_names.add(clip_name)
|
||||||
|
vae_name = self.get_input_value(entry, f"vae_name_{i}", prompt)
|
||||||
|
if vae_name not in ("None", "Auto"):
|
||||||
|
desired_vae_names.add(vae_name)
|
||||||
|
|
||||||
elif class_type in model_merge_node:
|
elif class_type in model_merge_node:
|
||||||
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name_1"))
|
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name_1"))
|
||||||
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name_2"))
|
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name_2"))
|
||||||
@@ -161,6 +191,8 @@ class easyLoader:
|
|||||||
desired_names = desired_controlnet_names
|
desired_names = desired_controlnet_names
|
||||||
elif object_type == "t5":
|
elif object_type == "t5":
|
||||||
desired_names = desired_t5_names
|
desired_names = desired_t5_names
|
||||||
|
elif object_type == "chatglm3":
|
||||||
|
desired_names = desired_glm3_names
|
||||||
else:
|
else:
|
||||||
desired_names = desired_lora_names
|
desired_names = desired_lora_names
|
||||||
self.clear_unused_objects(desired_names, object_type)
|
self.clear_unused_objects(desired_names, object_type)
|
||||||
@@ -199,7 +231,7 @@ class easyLoader:
|
|||||||
current_memory = self.get_memory_usage()
|
current_memory = self.get_memory_usage()
|
||||||
if current_memory < self.memory_threshold:
|
if current_memory < self.memory_threshold:
|
||||||
return
|
return
|
||||||
eviction_order = ["vae", "lora", "bvae", "clip", "ckpt", "controlnet"]
|
eviction_order = ["vae", "lora", "bvae", "clip", "ckpt", "controlnet", "unet", "t5", "chatglm3"]
|
||||||
for obj_type in eviction_order:
|
for obj_type in eviction_order:
|
||||||
if current_memory < self.memory_threshold:
|
if current_memory < self.memory_threshold:
|
||||||
break
|
break
|
||||||
@@ -230,7 +262,11 @@ class easyLoader:
|
|||||||
config_path = folder_paths.get_full_path("configs", config_name)
|
config_path = folder_paths.get_full_path("configs", config_name)
|
||||||
loaded_ckpt = comfy.sd.load_checkpoint(config_path, ckpt_path, output_vae=True, output_clip=output_clip, embedding_directory=folder_paths.get_folder_paths("embeddings"))
|
loaded_ckpt = comfy.sd.load_checkpoint(config_path, ckpt_path, output_vae=True, output_clip=output_clip, embedding_directory=folder_paths.get_folder_paths("embeddings"))
|
||||||
else:
|
else:
|
||||||
loaded_ckpt = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=output_clip, output_clipvision=output_clipvision, embedding_directory=folder_paths.get_folder_paths("embeddings"))
|
model_options = {}
|
||||||
|
if re.search("nf4", ckpt_name):
|
||||||
|
from ..modules.bitsandbytes_NF4 import OPS
|
||||||
|
model_options = {"custom_operations": OPS}
|
||||||
|
loaded_ckpt = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=output_clip, output_clipvision=output_clipvision, embedding_directory=folder_paths.get_folder_paths("embeddings"), model_options=model_options)
|
||||||
|
|
||||||
self.add_to_cache("ckpt", cache_name, loaded_ckpt[0])
|
self.add_to_cache("ckpt", cache_name, loaded_ckpt[0])
|
||||||
self.add_to_cache("bvae", cache_name, loaded_ckpt[2])
|
self.add_to_cache("bvae", cache_name, loaded_ckpt[2])
|
||||||
@@ -260,6 +296,7 @@ class easyLoader:
|
|||||||
|
|
||||||
def load_unet(self, unet_name):
|
def load_unet(self, unet_name):
|
||||||
if unet_name in self.loaded_objects["unet"]:
|
if unet_name in self.loaded_objects["unet"]:
|
||||||
|
log_node_info("Load UNet", f"{unet_name} cached")
|
||||||
return self.loaded_objects["unet"][unet_name][0]
|
return self.loaded_objects["unet"][unet_name][0]
|
||||||
|
|
||||||
unet_path = folder_paths.get_full_path("unet", unet_name)
|
unet_path = folder_paths.get_full_path("unet", unet_name)
|
||||||
@@ -269,6 +306,57 @@ class easyLoader:
|
|||||||
|
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
def load_diffusion_model(self, model_name):
|
||||||
|
if model_name in self.loaded_objects["unet"]:
|
||||||
|
log_node_info("Load Diffusion Model", f"{model_name} cached")
|
||||||
|
return self.loaded_objects["unet"][model_name][0]
|
||||||
|
|
||||||
|
model_path = folder_paths.get_full_path("diffusion_models", model_name)
|
||||||
|
if not model_path:
|
||||||
|
raise FileNotFoundError(f"[EasyUse] diffusion model not found: {model_name}")
|
||||||
|
|
||||||
|
model = comfy.sd.load_diffusion_model(model_path)
|
||||||
|
self.add_to_cache("unet", model_name, model)
|
||||||
|
self.eviction_based_on_memory()
|
||||||
|
|
||||||
|
return model
|
||||||
|
|
||||||
|
def load_diffusion_xy_model(self, model_name, clip_name, vae_name):
|
||||||
|
model = self.load_diffusion_model(model_name)
|
||||||
|
family = get_sd_version(model)
|
||||||
|
|
||||||
|
defaults = DIFFUSION_MODEL_XY_DEFAULTS.get(family)
|
||||||
|
if defaults is None:
|
||||||
|
raise RuntimeError(f"[EasyUse] unsupported diffusion model family: {family}")
|
||||||
|
|
||||||
|
if clip_name in ("Auto", None):
|
||||||
|
clip_name = defaults["clip_name"]
|
||||||
|
if vae_name in ("Auto", None):
|
||||||
|
vae_name = defaults["vae_name"]
|
||||||
|
|
||||||
|
clip = self.load_clip(clip_name, type=defaults["clip_type"])
|
||||||
|
vae = self.load_vae(vae_name)
|
||||||
|
|
||||||
|
return model, clip, vae, family
|
||||||
|
|
||||||
|
def load_diffusion_model_required(self, model_name, clip_name, vae_name):
|
||||||
|
if clip_name in ("None", None):
|
||||||
|
raise RuntimeError("[EasyUse] clip_name is required: please select a text encoder")
|
||||||
|
if vae_name in ("None", None):
|
||||||
|
raise RuntimeError("[EasyUse] vae_name is required: please select a VAE")
|
||||||
|
|
||||||
|
model = self.load_diffusion_model(model_name)
|
||||||
|
family = get_sd_version(model)
|
||||||
|
|
||||||
|
clip_type = DIFFUSION_MODEL_CLIP_TYPES.get(family)
|
||||||
|
if clip_type is None:
|
||||||
|
raise RuntimeError(f"[EasyUse] unsupported diffusion model family: {family}")
|
||||||
|
|
||||||
|
clip = self.load_clip(clip_name, type=clip_type)
|
||||||
|
vae = self.load_vae(vae_name)
|
||||||
|
|
||||||
|
return model, clip, vae, family
|
||||||
|
|
||||||
def load_controlnet(self, control_net_name, scale_soft_weights=1, use_cache=True):
|
def load_controlnet(self, control_net_name, scale_soft_weights=1, use_cache=True):
|
||||||
unique_id = f'{control_net_name};{str(scale_soft_weights)}'
|
unique_id = f'{control_net_name};{str(scale_soft_weights)}'
|
||||||
if use_cache and unique_id in self.loaded_objects["controlnet"]:
|
if use_cache and unique_id in self.loaded_objects["controlnet"]:
|
||||||
@@ -280,18 +368,19 @@ class easyLoader:
|
|||||||
cn_adv_cls = NODE_CLASS_MAPPINGS['ControlNetLoaderAdvanced']
|
cn_adv_cls = NODE_CLASS_MAPPINGS['ControlNetLoaderAdvanced']
|
||||||
control_net, = cn_adv_cls().load_controlnet(control_net_name, timestep_keyframe)
|
control_net, = cn_adv_cls().load_controlnet(control_net_name, timestep_keyframe)
|
||||||
else:
|
else:
|
||||||
raise Exception(
|
raise Exception(f"[Advanced-ControlNet Not Found] you need to install 'COMFYUI-Advanced-ControlNet'")
|
||||||
f"[Advanced-ControlNet Not Found] you need to install 'COMFYUI-Advanced-ControlNet'")
|
|
||||||
else:
|
else:
|
||||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||||
control_net = comfy.controlnet.load_controlnet(controlnet_path)
|
control_net = comfy.controlnet.load_controlnet(controlnet_path)
|
||||||
if use_cache:
|
if use_cache:
|
||||||
self.add_to_cache("controlnet", unique_id, control_net)
|
self.add_to_cache("controlnet", unique_id, control_net)
|
||||||
self.eviction_based_on_memory()
|
self.eviction_based_on_memory()
|
||||||
|
|
||||||
return control_net
|
return control_net
|
||||||
def load_clip(self, clip_name, type='stable_diffusion', load_clip=None):
|
def load_clip(self, clip_name, type='stable_diffusion', load_clip=None):
|
||||||
if clip_name in self.loaded_objects["clip"]:
|
cache_key = f"{clip_name}::{type}"
|
||||||
return self.loaded_objects["clip"][clip_name][0]
|
if cache_key in self.loaded_objects["clip"]:
|
||||||
|
return self.loaded_objects["clip"][cache_key][0]
|
||||||
|
|
||||||
if type == 'stable_diffusion':
|
if type == 'stable_diffusion':
|
||||||
clip_type = comfy.sd.CLIPType.STABLE_DIFFUSION
|
clip_type = comfy.sd.CLIPType.STABLE_DIFFUSION
|
||||||
@@ -299,16 +388,22 @@ class easyLoader:
|
|||||||
clip_type = comfy.sd.CLIPType.STABLE_CASCADE
|
clip_type = comfy.sd.CLIPType.STABLE_CASCADE
|
||||||
elif type == 'sd3':
|
elif type == 'sd3':
|
||||||
clip_type = comfy.sd.CLIPType.SD3
|
clip_type = comfy.sd.CLIPType.SD3
|
||||||
|
elif type == 'flux':
|
||||||
|
clip_type = comfy.sd.CLIPType.FLUX
|
||||||
elif type == 'stable_audio':
|
elif type == 'stable_audio':
|
||||||
clip_type = comfy.sd.CLIPType.STABLE_AUDIO
|
clip_type = comfy.sd.CLIPType.STABLE_AUDIO
|
||||||
|
elif type == 'krea2':
|
||||||
|
clip_type = comfy.sd.CLIPType.KREA2
|
||||||
|
elif type == 'anima':
|
||||||
|
clip_type = comfy.sd.CLIPType.STABLE_DIFFUSION
|
||||||
clip_path = folder_paths.get_full_path("clip", clip_name)
|
clip_path = folder_paths.get_full_path("clip", clip_name)
|
||||||
load_clip = comfy.sd.load_clip(ckpt_paths=[clip_path], embedding_directory=folder_paths.get_folder_paths("embeddings"), clip_type=clip_type)
|
load_clip = comfy.sd.load_clip(ckpt_paths=[clip_path], embedding_directory=folder_paths.get_folder_paths("embeddings"), clip_type=clip_type)
|
||||||
self.add_to_cache("clip", clip_name, load_clip)
|
self.add_to_cache("clip", cache_key, load_clip)
|
||||||
self.eviction_based_on_memory()
|
self.eviction_based_on_memory()
|
||||||
|
|
||||||
return load_clip
|
return load_clip
|
||||||
|
|
||||||
def load_lora(self, lora, model=None, clip=None, type=None):
|
def load_lora(self, lora, model=None, clip=None, type=None , use_cache=True):
|
||||||
lora_name = lora["lora_name"]
|
lora_name = lora["lora_name"]
|
||||||
model = model if model is not None else lora["model"]
|
model = model if model is not None else lora["model"]
|
||||||
clip = clip if clip is not None else lora["clip"]
|
clip = clip if clip is not None else lora["clip"]
|
||||||
@@ -319,11 +414,12 @@ class easyLoader:
|
|||||||
lbw_b = lora["lbw_b"] if "lbw_b" in lora else None
|
lbw_b = lora["lbw_b"] if "lbw_b" in lora else None
|
||||||
|
|
||||||
model_hash = str(model)[44:-1]
|
model_hash = str(model)[44:-1]
|
||||||
clip_hash = str(clip)[25:-1]
|
clip_hash = str(clip)[25:-1] if clip else ''
|
||||||
|
|
||||||
unique_id = f'{model_hash};{clip_hash};{lora_name};{model_strength};{clip_strength}'
|
unique_id = f'{model_hash};{clip_hash};{lora_name};{model_strength};{clip_strength}'
|
||||||
|
|
||||||
if unique_id in self.loaded_objects["lora"] and unique_id in self.loaded_objects["lora"][lora_name]:
|
if use_cache and unique_id in self.loaded_objects["lora"]:
|
||||||
|
log_node_info("Load LORA",f"{lora_name} cached")
|
||||||
return self.loaded_objects["lora"][unique_id][0]
|
return self.loaded_objects["lora"][unique_id][0]
|
||||||
|
|
||||||
orig_lora_name = lora_name
|
orig_lora_name = lora_name
|
||||||
@@ -335,7 +431,7 @@ class easyLoader:
|
|||||||
lora_path = None
|
lora_path = None
|
||||||
|
|
||||||
if lora_path is not None:
|
if lora_path is not None:
|
||||||
log_node_info("Load LORA",f"{lora_name}: {model_strength}, {clip_strength}, LBW={lbw}, A={lbw_a}, B={lbw_b}")
|
log_node_info("Load LORA",f"{lora_name}: model={model_strength:.3f}, clip={clip_strength:.3f}, LBW={lbw}, A={lbw_a}, B={lbw_b}")
|
||||||
if lbw:
|
if lbw:
|
||||||
lbw = lora["lbw"]
|
lbw = lora["lbw"]
|
||||||
lbw_a = lora["lbw_a"]
|
lbw_a = lora["lbw_a"]
|
||||||
@@ -375,13 +471,14 @@ class easyLoader:
|
|||||||
|
|
||||||
# PixArt
|
# PixArt
|
||||||
if type is not None and type == 'PixArt':
|
if type is not None and type == 'PixArt':
|
||||||
from ..dit.pixArt.loader import load_pixart_lora
|
from ..modules.dit.pixArt.loader import load_pixart_lora
|
||||||
model = load_pixart_lora(model, _lora, lora_path, model_strength)
|
model = load_pixart_lora(model, _lora, lora_path, model_strength)
|
||||||
else:
|
else:
|
||||||
model, clip = comfy.sd.load_lora_for_models(model, clip, _lora, model_strength, clip_strength)
|
model, clip = comfy.sd.load_lora_for_models(model, clip, _lora, model_strength, clip_strength)
|
||||||
|
|
||||||
self.add_to_cache("lora", unique_id, (model, clip))
|
if use_cache:
|
||||||
self.eviction_based_on_memory()
|
self.add_to_cache("lora", unique_id, (model, clip))
|
||||||
|
self.eviction_based_on_memory()
|
||||||
else:
|
else:
|
||||||
log_node_error(f"LORA NOT FOUND", orig_lora_name)
|
log_node_error(f"LORA NOT FOUND", orig_lora_name)
|
||||||
|
|
||||||
@@ -408,17 +505,20 @@ class easyLoader:
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def load_main(self, ckpt_name, config_name, vae_name, lora_name, lora_model_strength, lora_clip_strength, optional_lora_stack, model_override, clip_override, vae_override, prompt):
|
def load_main(self, ckpt_name, config_name, vae_name, lora_name, lora_model_strength, lora_clip_strength, optional_lora_stack, model_override, clip_override, vae_override, prompt, nf4=False):
|
||||||
model: ModelPatcher | None = None
|
model: ModelPatcher | None = None
|
||||||
clip: comfy.sd.CLIP | None = None
|
clip: comfy.sd.CLIP | None = None
|
||||||
vae: comfy.sd.VAE | None = None
|
vae: comfy.sd.VAE | None = None
|
||||||
clip_vision = None
|
clip_vision = None
|
||||||
lora_stack = []
|
lora_stack = []
|
||||||
|
|
||||||
|
# Check for model override
|
||||||
can_load_lora = True
|
can_load_lora = True
|
||||||
# 判断是否存在 模型或Lora叠加xyplot, 若存在优先缓存第一个模型
|
# 判断是否存在 模型或Lora叠加xyplot, 若存在优先缓存第一个模型
|
||||||
|
# Determine whether there is a model or Lora overlapping xyplot, and if there is, prioritize caching the first model.
|
||||||
xy_model_id = next((x for x in prompt if str(prompt[x]["class_type"]) in ["easy XYInputs: ModelMergeBlocks",
|
xy_model_id = next((x for x in prompt if str(prompt[x]["class_type"]) in ["easy XYInputs: ModelMergeBlocks",
|
||||||
"easy XYInputs: Checkpoint"]), None)
|
"easy XYInputs: Checkpoint"]), None)
|
||||||
|
# This will find nodes that aren't actively connected to anything, and skip loading lora's for them.
|
||||||
xy_lora_id = next((x for x in prompt if str(prompt[x]["class_type"]) == "easy XYInputs: Lora"), None)
|
xy_lora_id = next((x for x in prompt if str(prompt[x]["class_type"]) == "easy XYInputs: Lora"), None)
|
||||||
if xy_lora_id is not None:
|
if xy_lora_id is not None:
|
||||||
can_load_lora = False
|
can_load_lora = False
|
||||||
@@ -428,22 +528,23 @@ class easyLoader:
|
|||||||
ckpt_name_1 = node["inputs"]["ckpt_name_1"]
|
ckpt_name_1 = node["inputs"]["ckpt_name_1"]
|
||||||
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name_1)
|
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name_1)
|
||||||
can_load_lora = False
|
can_load_lora = False
|
||||||
# Load models
|
|
||||||
elif model_override is not None and clip_override is not None and vae_override is not None:
|
elif model_override is not None and clip_override is not None and vae_override is not None:
|
||||||
model = model_override
|
model = model_override
|
||||||
clip = clip_override
|
clip = clip_override
|
||||||
vae = vae_override
|
vae = vae_override
|
||||||
elif model_override is not None:
|
|
||||||
raise Exception(f"[ERROR] clip or vae is missing")
|
|
||||||
elif vae_override is not None:
|
|
||||||
raise Exception(f"[ERROR] model or clip is missing")
|
|
||||||
elif clip_override is not None:
|
|
||||||
raise Exception(f"[ERROR] model or vae is missing")
|
|
||||||
else:
|
else:
|
||||||
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name, config_name)
|
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name, config_name)
|
||||||
|
if model_override is not None:
|
||||||
|
model = model_override
|
||||||
|
if vae_override is not None:
|
||||||
|
vae = vae_override
|
||||||
|
elif clip_override is not None:
|
||||||
|
clip = clip_override
|
||||||
|
|
||||||
|
|
||||||
if optional_lora_stack is not None and can_load_lora:
|
if optional_lora_stack is not None and can_load_lora:
|
||||||
for lora in optional_lora_stack:
|
for lora in optional_lora_stack:
|
||||||
|
# This is a subtle bit of code because it uses the model created by the last call, and passes it to the next call.
|
||||||
lora = {"lora_name": lora[0], "model": model, "clip": clip, "model_strength": lora[1],
|
lora = {"lora_name": lora[0], "model": model, "clip": clip, "model_strength": lora[1],
|
||||||
"clip_strength": lora[2]}
|
"clip_strength": lora[2]}
|
||||||
model, clip = self.load_lora(lora)
|
model, clip = self.load_lora(lora)
|
||||||
@@ -466,18 +567,46 @@ class easyLoader:
|
|||||||
|
|
||||||
return model, clip, vae, clip_vision, lora_stack
|
return model, clip, vae, clip_vision, lora_stack
|
||||||
|
|
||||||
|
# Kolors
|
||||||
|
def load_kolors_unet(self, unet_name):
|
||||||
|
if unet_name in self.loaded_objects["unet"]:
|
||||||
|
log_node_info("Load Kolors UNet", f"{unet_name} cached")
|
||||||
|
return self.loaded_objects["unet"][unet_name][0]
|
||||||
|
else:
|
||||||
|
from ..modules.kolors.loader import applyKolorsUnet
|
||||||
|
with applyKolorsUnet():
|
||||||
|
unet_path = folder_paths.get_full_path("unet", unet_name)
|
||||||
|
sd = comfy.utils.load_torch_file(unet_path)
|
||||||
|
model = comfy.sd.load_unet_state_dict(sd)
|
||||||
|
if model is None:
|
||||||
|
raise RuntimeError("ERROR: Could not detect model type of: {}".format(unet_path))
|
||||||
|
|
||||||
|
self.add_to_cache("unet", unet_name, model)
|
||||||
|
self.eviction_based_on_memory()
|
||||||
|
|
||||||
|
return model
|
||||||
|
|
||||||
|
def load_chatglm3(self, chatglm3_name):
|
||||||
|
from ..modules.kolors.loader import load_chatglm3
|
||||||
|
if chatglm3_name in self.loaded_objects["chatglm3"]:
|
||||||
|
log_node_info("Load ChatGLM3", f"{chatglm3_name} cached")
|
||||||
|
return self.loaded_objects["chatglm3"][chatglm3_name][0]
|
||||||
|
|
||||||
|
chatglm_model = load_chatglm3(model_path=folder_paths.get_full_path("llm", chatglm3_name))
|
||||||
|
self.add_to_cache("chatglm3", chatglm3_name, chatglm_model)
|
||||||
|
self.eviction_based_on_memory()
|
||||||
|
|
||||||
|
return chatglm_model
|
||||||
|
|
||||||
|
|
||||||
# DiT
|
# DiT
|
||||||
def load_dit_ckpt(self, ckpt_name, model_name, **kwargs):
|
def load_dit_ckpt(self, ckpt_name, model_name, **kwargs):
|
||||||
if (ckpt_name+'_'+model_name) in self.loaded_objects["ckpt"]:
|
if (ckpt_name+'_'+model_name) in self.loaded_objects["ckpt"]:
|
||||||
return self.loaded_objects["ckpt"][ckpt_name+'_'+model_name][0]
|
return self.loaded_objects["ckpt"][ckpt_name+'_'+model_name][0]
|
||||||
model = None
|
model = None
|
||||||
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||||
model_type = kwargs['model_type'] if "model_type" in kwargs else 'HyDiT'
|
model_type = kwargs['model_type'] if "model_type" in kwargs else 'PixArt'
|
||||||
if model_type == 'HyDiT':
|
if model_type == 'PixArt':
|
||||||
hydit_conf = kwargs['hydit_conf']
|
|
||||||
model_conf = hydit_conf[model_name]
|
|
||||||
model = load_hydit(ckpt_path, model_conf)
|
|
||||||
elif model_type == 'PixArt':
|
|
||||||
pixart_conf = kwargs['pixart_conf']
|
pixart_conf = kwargs['pixart_conf']
|
||||||
model_conf = pixart_conf[model_name]
|
model_conf = pixart_conf[model_name]
|
||||||
model = load_pixart(ckpt_path, model_conf)
|
model = load_pixart(ckpt_path, model_conf)
|
||||||
@@ -486,56 +615,11 @@ class easyLoader:
|
|||||||
self.eviction_based_on_memory()
|
self.eviction_based_on_memory()
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
def load_dit_clip(self, clip_name, **kwargs):
|
|
||||||
if clip_name in self.loaded_objects["clip"]:
|
|
||||||
return self.loaded_objects["clip"][clip_name][0]
|
|
||||||
|
|
||||||
model_type = kwargs['model_type'] if "model_type" in kwargs else 'HyDiT'
|
|
||||||
if model_type == 'HyDiT':
|
|
||||||
del kwargs['model_type']
|
|
||||||
model = EXM_HyDiT_Tenc_Temp(model_class="clip", **kwargs)
|
|
||||||
clip_path = folder_paths.get_full_path("clip", clip_name)
|
|
||||||
sd = comfy.utils.load_torch_file(clip_path)
|
|
||||||
|
|
||||||
prefix = "bert."
|
|
||||||
state_dict = {}
|
|
||||||
for key in sd:
|
|
||||||
nkey = key
|
|
||||||
if key.startswith(prefix):
|
|
||||||
nkey = key[len(prefix):]
|
|
||||||
state_dict[nkey] = sd[key]
|
|
||||||
|
|
||||||
m, e = model.load_sd(state_dict)
|
|
||||||
if len(m) > 0 or len(e) > 0:
|
|
||||||
print(f"{clip_name}: clip missing {len(m)} keys ({len(e)} extra)")
|
|
||||||
|
|
||||||
self.add_to_cache("clip", clip_name, model)
|
|
||||||
self.eviction_based_on_memory()
|
|
||||||
|
|
||||||
return model
|
|
||||||
|
|
||||||
def load_dit_t5(self, t5_name, **kwargs):
|
|
||||||
if t5_name in self.loaded_objects["t5"]:
|
|
||||||
return self.loaded_objects["t5"][t5_name][0]
|
|
||||||
|
|
||||||
model_type = kwargs['model_type'] if "model_type" in kwargs else 'HyDiT'
|
|
||||||
if model_type == 'HyDiT':
|
|
||||||
del kwargs['model_type']
|
|
||||||
model = EXM_HyDiT_Tenc_Temp(model_class="mT5", **kwargs)
|
|
||||||
t5_path = folder_paths.get_full_path("t5", t5_name)
|
|
||||||
sd = comfy.utils.load_torch_file(t5_path)
|
|
||||||
m, e = model.load_sd(sd)
|
|
||||||
if len(m) > 0 or len(e) > 0:
|
|
||||||
print(f"{t5_name}: mT5 missing {len(m)} keys ({len(e)} extra)")
|
|
||||||
|
|
||||||
self.add_to_cache("t5", t5_name, model)
|
|
||||||
self.eviction_based_on_memory()
|
|
||||||
|
|
||||||
return model
|
|
||||||
|
|
||||||
def load_t5_from_sd3_clip(self, sd3_clip, padding):
|
def load_t5_from_sd3_clip(self, sd3_clip, padding):
|
||||||
from comfy.sd3_clip import SD3Tokenizer, SD3ClipModel
|
try:
|
||||||
|
from comfy.text_encoders.sd3_clip import SD3Tokenizer, SD3ClipModel
|
||||||
|
except:
|
||||||
|
from comfy.sd3_clip import SD3Tokenizer, SD3ClipModel
|
||||||
import copy
|
import copy
|
||||||
|
|
||||||
clip = sd3_clip.clone()
|
clip = sd3_clip.clone()
|
||||||
|
|||||||
+148
@@ -0,0 +1,148 @@
|
|||||||
|
"""
|
||||||
|
Math utility functions for formula evaluation
|
||||||
|
"""
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
|
||||||
|
def evaluate_formula(formula: str, a=0, b=0, c=0, d=0):
|
||||||
|
"""
|
||||||
|
计算字符串数学公式
|
||||||
|
|
||||||
|
支持的运算符和函数:
|
||||||
|
- 基本运算:+, -, *, /, //, %, **
|
||||||
|
- 比较运算:>, <, >=, <=, ==, !=
|
||||||
|
- 数学函数:abs, pow, round, ceil, floor, sqrt, exp, log, log10
|
||||||
|
- 三角函数:sin, cos, tan, asin, acos, atan
|
||||||
|
- 常量:pi, e
|
||||||
|
|
||||||
|
Args:
|
||||||
|
formula: 数学公式字符串,可以使用变量a、b、c、d
|
||||||
|
a: 变量a的值
|
||||||
|
b: 变量b的值
|
||||||
|
c: 变量c的值
|
||||||
|
d: 变量d的值
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
如果任意输入为list则返回list[float],否则返回float
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
>>> evaluate_formula("a + b", 1, 2)
|
||||||
|
3.0
|
||||||
|
>>> evaluate_formula("pow(a, 2)", 5)
|
||||||
|
25.0
|
||||||
|
>>> evaluate_formula("ceil(a / b)", 5, 2)
|
||||||
|
3.0
|
||||||
|
>>> evaluate_formula("(a>b)*b+(a<=b)*a", 5, 3)
|
||||||
|
3.0
|
||||||
|
>>> evaluate_formula("(a>b)*b+(a<=b)*a", 2, 3)
|
||||||
|
2.0
|
||||||
|
"""
|
||||||
|
# 安全的数学函数白名单
|
||||||
|
safe_dict = {
|
||||||
|
# 基本运算
|
||||||
|
'abs': abs,
|
||||||
|
'pow': pow,
|
||||||
|
'round': round,
|
||||||
|
# 数学函数
|
||||||
|
'ceil': math.ceil,
|
||||||
|
'floor': math.floor,
|
||||||
|
'sqrt': math.sqrt,
|
||||||
|
'exp': math.exp,
|
||||||
|
'log': math.log,
|
||||||
|
'log10': math.log10,
|
||||||
|
# 三角函数
|
||||||
|
'sin': math.sin,
|
||||||
|
'cos': math.cos,
|
||||||
|
'tan': math.tan,
|
||||||
|
'asin': math.asin,
|
||||||
|
'acos': math.acos,
|
||||||
|
'atan': math.atan,
|
||||||
|
# 常量
|
||||||
|
'pi': math.pi,
|
||||||
|
'e': math.e,
|
||||||
|
}
|
||||||
|
|
||||||
|
# 判断是否有 list 输入
|
||||||
|
list_inputs = {k: v for k, v in {'a': a, 'b': b, 'c': c, 'd': d}.items() if isinstance(v, (list, tuple))}
|
||||||
|
scalar_inputs = {k: v for k, v in {'a': a, 'b': b, 'c': c, 'd': d}.items() if not isinstance(v, (list, tuple))}
|
||||||
|
|
||||||
|
def _eval_single(vals: dict) -> float:
|
||||||
|
env = dict(safe_dict)
|
||||||
|
env.update({k: float(v) for k, v in vals.items()})
|
||||||
|
try:
|
||||||
|
result = eval(formula, {"__builtins__": {}}, env)
|
||||||
|
return float(result)
|
||||||
|
except Exception as e:
|
||||||
|
raise ValueError(f"公式计算错误: {str(e)}")
|
||||||
|
|
||||||
|
if not list_inputs:
|
||||||
|
# 全是标量
|
||||||
|
return _eval_single({k: v for k, v in {'a': a, 'b': b, 'c': c, 'd': d}.items()})
|
||||||
|
|
||||||
|
# 有 list 输入,逐元素计算
|
||||||
|
max_len = max(len(v) for v in list_inputs.values())
|
||||||
|
results = []
|
||||||
|
for i in range(max_len):
|
||||||
|
vals = {k: float(v) for k, v in scalar_inputs.items()}
|
||||||
|
for k, v in list_inputs.items():
|
||||||
|
vals[k] = float(v[i] if i < len(v) else v[-1])
|
||||||
|
results.append(_eval_single(vals))
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
def ceil_value(value: float) -> int:
|
||||||
|
"""向上取整"""
|
||||||
|
return math.ceil(value)
|
||||||
|
|
||||||
|
|
||||||
|
def floor_value(value: float) -> int:
|
||||||
|
"""向下取整"""
|
||||||
|
return math.floor(value)
|
||||||
|
|
||||||
|
|
||||||
|
def round_value(value: float, decimals: int = 0) -> float:
|
||||||
|
"""
|
||||||
|
四舍五入
|
||||||
|
|
||||||
|
Args:
|
||||||
|
value: 要取整的值
|
||||||
|
decimals: 保留小数位数
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
四舍五入后的值
|
||||||
|
"""
|
||||||
|
return round(value, decimals)
|
||||||
|
|
||||||
|
|
||||||
|
def power(base: float, exponent: float) -> float:
|
||||||
|
"""计算幂运算"""
|
||||||
|
return math.pow(base, exponent)
|
||||||
|
|
||||||
|
|
||||||
|
def sqrt_value(value: float) -> float:
|
||||||
|
"""计算平方根"""
|
||||||
|
if value < 0:
|
||||||
|
raise ValueError("不能对负数求平方根")
|
||||||
|
return math.sqrt(value)
|
||||||
|
|
||||||
|
|
||||||
|
def add(a: float, b: float) -> float:
|
||||||
|
"""加法"""
|
||||||
|
return a + b
|
||||||
|
|
||||||
|
|
||||||
|
def subtract(a: float, b: float) -> float:
|
||||||
|
"""减法"""
|
||||||
|
return a - b
|
||||||
|
|
||||||
|
|
||||||
|
def multiply(a: float, b: float) -> float:
|
||||||
|
"""乘法"""
|
||||||
|
return a * b
|
||||||
|
|
||||||
|
|
||||||
|
def divide(a: float, b: float) -> float:
|
||||||
|
"""除法"""
|
||||||
|
if b == 0:
|
||||||
|
raise ValueError("除数不能为零")
|
||||||
|
return a / b
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
from server import PromptServer
|
||||||
|
from aiohttp import web
|
||||||
|
import time
|
||||||
|
import json
|
||||||
|
|
||||||
|
class MessageCancelled(Exception):
|
||||||
|
pass
|
||||||
|
|
||||||
|
class Message:
|
||||||
|
stash = {}
|
||||||
|
messages = {}
|
||||||
|
cancelled = False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def addMessage(cls, id, message):
|
||||||
|
if message == '__cancel__':
|
||||||
|
cls.messages = {}
|
||||||
|
cls.cancelled = True
|
||||||
|
elif message == '__start__':
|
||||||
|
cls.messages = {}
|
||||||
|
cls.stash = {}
|
||||||
|
cls.cancelled = False
|
||||||
|
else:
|
||||||
|
cls.messages[str(id)] = message
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def waitForMessage(cls, id, period=0.1, asList=False):
|
||||||
|
sid = str(id)
|
||||||
|
while not (sid in cls.messages) and not ("-1" in cls.messages):
|
||||||
|
if cls.cancelled:
|
||||||
|
cls.cancelled = False
|
||||||
|
raise MessageCancelled()
|
||||||
|
time.sleep(period)
|
||||||
|
if cls.cancelled:
|
||||||
|
cls.cancelled = False
|
||||||
|
raise MessageCancelled()
|
||||||
|
message = cls.messages.pop(str(id), None) or cls.messages.pop("-1")
|
||||||
|
try:
|
||||||
|
if asList:
|
||||||
|
return [str(x.strip()) for x in message.split(",")]
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
return json.loads(message)
|
||||||
|
except ValueError:
|
||||||
|
return message
|
||||||
|
except ValueError:
|
||||||
|
print( f"ERROR IN MESSAGE - failed to parse '${message}' as ${'comma separated list of strings' if asList else 'string'}")
|
||||||
|
return [message] if asList else message
|
||||||
|
|
||||||
|
|
||||||
|
@PromptServer.instance.routes.post('/easyuse/message_callback')
|
||||||
|
async def message_callback(request):
|
||||||
|
post = await request.post()
|
||||||
|
Message.addMessage(post.get("id"), post.get("message"))
|
||||||
|
return web.json_response({})
|
||||||
@@ -0,0 +1,29 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_output_file_path(output_root, output_file_path, file_name, file_extension):
|
||||||
|
"""Resolve a workflow-provided output path beneath ``output_root``.
|
||||||
|
|
||||||
|
Relative output directories remain supported, but are interpreted relative
|
||||||
|
to ComfyUI's configured output directory rather than the process working
|
||||||
|
directory. Resolving both paths prevents ``..`` components and existing
|
||||||
|
symlinks from escaping the allowed root.
|
||||||
|
"""
|
||||||
|
output_root = os.path.realpath(output_root)
|
||||||
|
requested_directory = output_file_path
|
||||||
|
if not os.path.isabs(requested_directory):
|
||||||
|
requested_directory = os.path.join(output_root, requested_directory)
|
||||||
|
|
||||||
|
candidate = os.path.realpath(
|
||||||
|
os.path.join(requested_directory, f"{file_name}.{file_extension}")
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
is_within_output = os.path.commonpath((output_root, candidate)) == output_root
|
||||||
|
except ValueError:
|
||||||
|
# Different Windows drives and paths containing null bytes are unsafe.
|
||||||
|
is_within_output = False
|
||||||
|
|
||||||
|
if not is_within_output:
|
||||||
|
raise ValueError("Saving outside the ComfyUI output directory is not allowed")
|
||||||
|
|
||||||
|
return candidate
|
||||||
+1055
-909
File diff suppressed because it is too large
Load Diff
@@ -12,8 +12,8 @@ from .utils import install_package
|
|||||||
try:
|
try:
|
||||||
from lark import Lark, Transformer, v_args
|
from lark import Lark, Transformer, v_args
|
||||||
except:
|
except:
|
||||||
print('install lark-parser...')
|
print('install lark...')
|
||||||
install_package('lark-parser')
|
install_package('lark')
|
||||||
from lark import Lark, Transformer, v_args
|
from lark import Lark, Transformer, v_args
|
||||||
|
|
||||||
model_path = os.path.join(folder_paths.models_dir, 'prompt_generator')
|
model_path = os.path.join(folder_paths.models_dir, 'prompt_generator')
|
||||||
@@ -80,7 +80,7 @@ def has_chinese(text):
|
|||||||
_text = text
|
_text = text
|
||||||
_text = re.sub(r'<.*?>', '', _text)
|
_text = re.sub(r'<.*?>', '', _text)
|
||||||
_text = re.sub(r'__.*?__', '', _text)
|
_text = re.sub(r'__.*?__', '', _text)
|
||||||
_text = re.sub(r'embedding:.*?(\d+)?', '', _text)
|
_text = re.sub(r'embedding:.*?$', '', _text)
|
||||||
for char in _text:
|
for char in _text:
|
||||||
if '\u4e00' <= char <= '\u9fff':
|
if '\u4e00' <= char <= '\u9fff':
|
||||||
has_cn = True
|
has_cn = True
|
||||||
@@ -95,7 +95,6 @@ def translate(text):
|
|||||||
if not os.path.exists(zh_en_model_path):
|
if not os.path.exists(zh_en_model_path):
|
||||||
zh_en_model_path = 'Helsinki-NLP/opus-mt-zh-en'
|
zh_en_model_path = 'Helsinki-NLP/opus-mt-zh-en'
|
||||||
|
|
||||||
print(zh_en_model_path)
|
|
||||||
if zh_en_model is None:
|
if zh_en_model is None:
|
||||||
|
|
||||||
zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval()
|
zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval()
|
||||||
@@ -186,7 +185,7 @@ class ChinesePromptTranslate(Transformer):
|
|||||||
|
|
||||||
|
|
||||||
#定义Prompt文法
|
#定义Prompt文法
|
||||||
grammar = """
|
grammar = r"""
|
||||||
start: sentence
|
start: sentence
|
||||||
sentence: phrase ("," phrase)*
|
sentence: phrase ("," phrase)*
|
||||||
phrase: emphasis | weight | word | lora | embedding | schedule
|
phrase: emphasis | weight | word | lora | embedding | schedule
|
||||||
|
|||||||
+55
-9
@@ -5,6 +5,19 @@ class AlwaysEqualProxy(str):
|
|||||||
def __ne__(self, _):
|
def __ne__(self, _):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
class TautologyStr(str):
|
||||||
|
def __ne__(self, other):
|
||||||
|
return False
|
||||||
|
|
||||||
|
class ByPassTypeTuple(tuple):
|
||||||
|
def __getitem__(self, index):
|
||||||
|
if index>0:
|
||||||
|
index=0
|
||||||
|
item = super().__getitem__(index)
|
||||||
|
if isinstance(item, str):
|
||||||
|
return TautologyStr(item)
|
||||||
|
return item
|
||||||
|
|
||||||
comfy_ui_revision = None
|
comfy_ui_revision = None
|
||||||
def get_comfyui_revision():
|
def get_comfyui_revision():
|
||||||
try:
|
try:
|
||||||
@@ -22,9 +35,13 @@ import sys
|
|||||||
import importlib.util
|
import importlib.util
|
||||||
import importlib.metadata
|
import importlib.metadata
|
||||||
import comfy.model_management as mm
|
import comfy.model_management as mm
|
||||||
|
import logging
|
||||||
import gc
|
import gc
|
||||||
from packaging import version
|
from packaging import version
|
||||||
from server import PromptServer
|
from server import PromptServer
|
||||||
|
|
||||||
|
LOG = logging.getLogger(__name__)
|
||||||
|
|
||||||
def is_package_installed(package):
|
def is_package_installed(package):
|
||||||
try:
|
try:
|
||||||
module = importlib.util.find_spec(package)
|
module = importlib.util.find_spec(package)
|
||||||
@@ -69,6 +86,7 @@ def compare_revision(num):
|
|||||||
if not comfy_ui_revision:
|
if not comfy_ui_revision:
|
||||||
comfy_ui_revision = get_comfyui_revision()
|
comfy_ui_revision = get_comfyui_revision()
|
||||||
return True if comfy_ui_revision == 'Unknown' or int(comfy_ui_revision) >= num else False
|
return True if comfy_ui_revision == 'Unknown' or int(comfy_ui_revision) >= num else False
|
||||||
|
|
||||||
def find_tags(string: str, sep="/") -> list[str]:
|
def find_tags(string: str, sep="/") -> list[str]:
|
||||||
"""
|
"""
|
||||||
find tags from string use the sep for split
|
find tags from string use the sep for split
|
||||||
@@ -92,6 +110,8 @@ def get_sd_version(model):
|
|||||||
model_config: comfy.supported_models.supported_models_base.BASE = base.model_config
|
model_config: comfy.supported_models.supported_models_base.BASE = base.model_config
|
||||||
if isinstance(model_config, comfy.supported_models.SDXL):
|
if isinstance(model_config, comfy.supported_models.SDXL):
|
||||||
return 'sdxl'
|
return 'sdxl'
|
||||||
|
elif isinstance(model_config, comfy.supported_models.SDXLRefiner):
|
||||||
|
return 'sdxl_refiner'
|
||||||
elif isinstance(
|
elif isinstance(
|
||||||
model_config, (comfy.supported_models.SD15, comfy.supported_models.SD20)
|
model_config, (comfy.supported_models.SD15, comfy.supported_models.SD20)
|
||||||
):
|
):
|
||||||
@@ -102,6 +122,16 @@ def get_sd_version(model):
|
|||||||
return 'svd'
|
return 'svd'
|
||||||
elif isinstance(model_config, comfy.supported_models.SD3):
|
elif isinstance(model_config, comfy.supported_models.SD3):
|
||||||
return 'sd3'
|
return 'sd3'
|
||||||
|
elif isinstance(model_config, comfy.supported_models.HunyuanDiT):
|
||||||
|
return 'hydit'
|
||||||
|
elif isinstance(model_config, comfy.supported_models.Flux):
|
||||||
|
return 'flux'
|
||||||
|
elif isinstance(model_config, comfy.supported_models.GenmoMochi):
|
||||||
|
return 'mochi'
|
||||||
|
elif isinstance(model_config, comfy.supported_models.Anima):
|
||||||
|
return 'anima'
|
||||||
|
elif isinstance(model_config, comfy.supported_models.Krea2):
|
||||||
|
return 'krea2'
|
||||||
else:
|
else:
|
||||||
return 'unknown'
|
return 'unknown'
|
||||||
|
|
||||||
@@ -163,8 +193,9 @@ def find_wildcards_seed(clip_id, text, prompt):
|
|||||||
else:
|
else:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def is_linked_styles_selector(prompt, my_unique_id, prompt_type='positive'):
|
def is_linked_styles_selector(prompt, unique_id, prompt_type='positive'):
|
||||||
inputs_values = prompt[my_unique_id]['inputs'][prompt_type] if prompt_type in prompt[my_unique_id][
|
unique_id = unique_id.split('.')[len(unique_id.split('.')) - 1] if "." in unique_id else unique_id
|
||||||
|
inputs_values = prompt[unique_id]['inputs'][prompt_type] if prompt_type in prompt[unique_id][
|
||||||
'inputs'] else None
|
'inputs'] else None
|
||||||
if type(inputs_values) == list and inputs_values != 'undefined' and inputs_values[0]:
|
if type(inputs_values) == list and inputs_values != 'undefined' and inputs_values[0]:
|
||||||
return True if prompt[inputs_values[0]] and prompt[inputs_values[0]]['class_type'] == 'easy stylesSelector' else False
|
return True if prompt[inputs_values[0]] and prompt[inputs_values[0]]['class_type'] == 'easy stylesSelector' else False
|
||||||
@@ -195,14 +226,15 @@ def get_local_filepath(url, dirname, local_file_name=None):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
use_mirror = True
|
use_mirror = True
|
||||||
url = url.replace('huggingface.co', 'hf-mirror.com')
|
url = url.replace('huggingface.co', 'hf-mirror.com')
|
||||||
print(f'无法从huggingface下载,正在尝试从 {url} 下载...')
|
print(f'Unable to download from huggingface, trying mirror: {url}')
|
||||||
PromptServer.instance.send_sync("easyuse-toast", {'content': f'无法连接huggingface,正在尝试从 {url} 下载...', 'duration': 10000})
|
PromptServer.instance.send_sync("easyuse-toast", {'content': f'Unable to connect to huggingface, trying mirror: {url}', 'duration': 10000})
|
||||||
try:
|
try:
|
||||||
download_url_to_file(url, destination)
|
download_url_to_file(url, destination)
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
|
error_msg = str(err.args[0]) if err.args else str(err)
|
||||||
PromptServer.instance.send_sync("easyuse-toast",
|
PromptServer.instance.send_sync("easyuse-toast",
|
||||||
{'content': f'无法从 {url} 下载模型', 'type':'error'})
|
{'content': f'Unable to download model from {url}', 'type':'error'})
|
||||||
raise Exception(f'无法从 {url} 下载,错误信息:{str(err.args[0])}')
|
raise Exception(f'Download failed. Original URL and mirror both failed.\nError: {error_msg}')
|
||||||
return destination
|
return destination
|
||||||
|
|
||||||
def to_lora_patch_dict(state_dict: dict) -> dict:
|
def to_lora_patch_dict(state_dict: dict) -> dict:
|
||||||
@@ -227,9 +259,9 @@ def to_lora_patch_dict(state_dict: dict) -> dict:
|
|||||||
def easySave(images, filename_prefix, output_type, prompt=None, extra_pnginfo=None):
|
def easySave(images, filename_prefix, output_type, prompt=None, extra_pnginfo=None):
|
||||||
"""Save or Preview Image"""
|
"""Save or Preview Image"""
|
||||||
from nodes import PreviewImage, SaveImage
|
from nodes import PreviewImage, SaveImage
|
||||||
if output_type == "Hide":
|
if output_type in ["Hide", "None"]:
|
||||||
return list()
|
return list()
|
||||||
if output_type in ["Preview", "Preview&Choose"]:
|
elif output_type in ["Preview", "Preview&Choose"]:
|
||||||
filename_prefix = 'easyPreview'
|
filename_prefix = 'easyPreview'
|
||||||
results = PreviewImage().save_images(images, filename_prefix, prompt, extra_pnginfo)
|
results = PreviewImage().save_images(images, filename_prefix, prompt, extra_pnginfo)
|
||||||
return results['ui']['images']
|
return results['ui']['images']
|
||||||
@@ -253,6 +285,20 @@ def getMetadata(filepath):
|
|||||||
return header
|
return header
|
||||||
|
|
||||||
def cleanGPUUsedForce():
|
def cleanGPUUsedForce():
|
||||||
|
from .cache import remove_cache
|
||||||
|
|
||||||
|
remove_cache("*")
|
||||||
gc.collect()
|
gc.collect()
|
||||||
|
try:
|
||||||
|
import torch
|
||||||
|
except (ImportError, OSError, RuntimeError) as exc:
|
||||||
|
LOG.debug("Skipping CUDA synchronize during cleanGPUUsedForce: torch import failed: %s", exc)
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
except (AttributeError, OSError, RuntimeError) as exc:
|
||||||
|
LOG.debug("Skipping CUDA synchronize during cleanGPUUsedForce: %s", exc)
|
||||||
|
|
||||||
mm.unload_all_models()
|
mm.unload_all_models()
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
|
|||||||
+182
-10
@@ -1,9 +1,13 @@
|
|||||||
import re
|
|
||||||
import random
|
|
||||||
import os
|
|
||||||
import folder_paths
|
|
||||||
import yaml
|
|
||||||
import json
|
import json
|
||||||
|
import os
|
||||||
|
import random
|
||||||
|
import re
|
||||||
|
from math import prod
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
import folder_paths
|
||||||
|
|
||||||
from .log import log_node_info
|
from .log import log_node_info
|
||||||
|
|
||||||
easy_wildcard_dict = {}
|
easy_wildcard_dict = {}
|
||||||
@@ -34,16 +38,16 @@ def read_wildcard_dict(wildcard_path):
|
|||||||
key = os.path.splitext(rel_path)[0].replace('\\', '/').lower()
|
key = os.path.splitext(rel_path)[0].replace('\\', '/').lower()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
with open(file_path, 'r', encoding="ISO-8859-1") as f:
|
with open(file_path, 'r', encoding="UTF-8", errors="ignore") as f:
|
||||||
lines = f.read().splitlines()
|
lines = f.read().splitlines()
|
||||||
easy_wildcard_dict[key] = lines
|
easy_wildcard_dict[key] = lines
|
||||||
except UnicodeDecodeError:
|
except UnicodeDecodeError:
|
||||||
with open(file_path, 'r', encoding="UTF-8", errors="ignore") as f:
|
with open(file_path, 'r', encoding="ISO-8859-1") as f:
|
||||||
lines = f.read().splitlines()
|
lines = f.read().splitlines()
|
||||||
easy_wildcard_dict[key] = lines
|
easy_wildcard_dict[key] = lines
|
||||||
elif file.endswith('.yaml'):
|
elif file.endswith('.yaml'):
|
||||||
file_path = os.path.join(root, file)
|
file_path = os.path.join(root, file)
|
||||||
with open(file_path, 'r') as f:
|
with open(file_path, 'r', encoding="utf-8") as f:
|
||||||
yaml_data = yaml.load(f, Loader=yaml.FullLoader)
|
yaml_data = yaml.load(f, Loader=yaml.FullLoader)
|
||||||
|
|
||||||
for k, v in yaml_data.items():
|
for k, v in yaml_data.items():
|
||||||
@@ -51,7 +55,7 @@ def read_wildcard_dict(wildcard_path):
|
|||||||
elif file.endswith('.json'):
|
elif file.endswith('.json'):
|
||||||
file_path = os.path.join(root, file)
|
file_path = os.path.join(root, file)
|
||||||
try:
|
try:
|
||||||
with open(file_path, 'r') as f:
|
with open(file_path, 'r', encoding="utf-8") as f:
|
||||||
json_data = json.load(f)
|
json_data = json.load(f)
|
||||||
for key, value in json_data.items():
|
for key, value in json_data.items():
|
||||||
key = wildcard_normalize(key)
|
key = wildcard_normalize(key)
|
||||||
@@ -168,7 +172,7 @@ def process(text, seed=None):
|
|||||||
replacements_found = True
|
replacements_found = True
|
||||||
string = string.replace(f"__{match}__", replacement, 1)
|
string = string.replace(f"__{match}__", replacement, 1)
|
||||||
elif '*' in keyword:
|
elif '*' in keyword:
|
||||||
subpattern = keyword.replace('*', '.*').replace('+','\+')
|
subpattern = keyword.replace('*', '.*').replace('+', r'\+')
|
||||||
total_patterns = []
|
total_patterns = []
|
||||||
found = False
|
found = False
|
||||||
for k, v in easy_wildcard_dict.items():
|
for k, v in easy_wildcard_dict.items():
|
||||||
@@ -302,3 +306,171 @@ def process_with_loras(wildcard_opt, model, clip, title="Positive", seed=None, c
|
|||||||
log_node_info("easy wildcards",f'{title}_decode: {pass1}')
|
log_node_info("easy wildcards",f'{title}_decode: {pass1}')
|
||||||
|
|
||||||
return model, clip, pass2, pass1, show_wildcard_prompt, pipe_lora_stack
|
return model, clip, pass2, pass1, show_wildcard_prompt, pipe_lora_stack
|
||||||
|
|
||||||
|
|
||||||
|
def expand_wildcard(keyword: str) -> tuple[str]:
|
||||||
|
"""传入文件通配符的关键词,从 easy_wildcard_dict 中获取通配符的所有选项。"""
|
||||||
|
global easy_wildcard_dict
|
||||||
|
if keyword in easy_wildcard_dict:
|
||||||
|
return tuple(easy_wildcard_dict[keyword])
|
||||||
|
elif '*' in keyword:
|
||||||
|
subpattern = keyword.replace('*', '.*').replace('+', r"\+")
|
||||||
|
total_pattern = []
|
||||||
|
for k, v in easy_wildcard_dict.items():
|
||||||
|
if re.match(subpattern, k) is not None:
|
||||||
|
total_pattern.extend(v)
|
||||||
|
if total_pattern:
|
||||||
|
return tuple(total_pattern)
|
||||||
|
elif '/' not in keyword:
|
||||||
|
return expand_wildcard(f"*/{keyword}")
|
||||||
|
|
||||||
|
def expand_options(options: str) -> tuple[str]:
|
||||||
|
"""传入去掉 {} 的选项。
|
||||||
|
展开选项通配符,返回该选项中的每一项,这里的每一项都是一个替换项。
|
||||||
|
不会对选项内容进行任何处理,即便存在空格或特殊符号,也会原样返回。"""
|
||||||
|
return tuple(options.split("|"))
|
||||||
|
|
||||||
|
|
||||||
|
def decimal_to_irregular(n, bases):
|
||||||
|
"""
|
||||||
|
将十进制数转换为不规则进制
|
||||||
|
|
||||||
|
:param n: 十进制数
|
||||||
|
:param bases: 各位置的基数列表,从低位到高位
|
||||||
|
:return: 不规则进制表示的列表,从低位到高位
|
||||||
|
"""
|
||||||
|
if n == 0:
|
||||||
|
return [0] * len(bases) if bases else [0]
|
||||||
|
|
||||||
|
digits = []
|
||||||
|
remaining = n
|
||||||
|
|
||||||
|
# 从低位到高位处理
|
||||||
|
for base in bases:
|
||||||
|
digit = remaining % base
|
||||||
|
digits.append(digit)
|
||||||
|
remaining = remaining // base
|
||||||
|
|
||||||
|
return digits
|
||||||
|
|
||||||
|
|
||||||
|
class WildcardProcessor:
|
||||||
|
"""通配符处理器
|
||||||
|
|
||||||
|
通配符格式:
|
||||||
|
+ option : {a|b}
|
||||||
|
+ wildcard: __keyword__ 通配符内容将从 Easy-Use 插件提供的 easy_wildcard_dict 中获取
|
||||||
|
"""
|
||||||
|
|
||||||
|
RE_OPTIONS = re.compile(r"{([^{}]*?)}")
|
||||||
|
RE_WILDCARD = re.compile(r"__([\w\s.\-+/*\\]+?)__")
|
||||||
|
RE_REPLACER = re.compile(r"{([^{}]*?)}|__([\w\s.\-+/*\\]+?)__")
|
||||||
|
|
||||||
|
# 将输入的提示词转化成符合 python str.format 要求格式的模板,并将 option 和 wildcard 按照顺序在模板中留下 {0}, {1} 等占位符
|
||||||
|
template: str
|
||||||
|
# option、wildcard 的替换项列表,按照在模板中出现的顺序排列,相同的替换项列表只保留第一份
|
||||||
|
replacers: dict[int, tuple[str]]
|
||||||
|
# 占位符的编号和替换项列表的索引的映射,占位符编号按照在模板中出现的顺序排列,方便减少替换项的存储占用
|
||||||
|
placeholder_mapping: dict[str, int] # placeholder_id => replacer_id
|
||||||
|
# 各替换项列表的项数,按照在模板中出现的顺序排列,提前计算,方便后续使用
|
||||||
|
placeholder_choices: dict[str, int] # placeholder_id => len(replacer)
|
||||||
|
|
||||||
|
def __init__(self, text: str):
|
||||||
|
self.__make_template(text)
|
||||||
|
self.__total = None
|
||||||
|
|
||||||
|
def random(self, seed=None) -> str:
|
||||||
|
"从所有可能性中随机获取一个"
|
||||||
|
if seed is not None:
|
||||||
|
random.seed(seed)
|
||||||
|
return self.getn(random.randint(0, self.total() - 1))
|
||||||
|
|
||||||
|
def getn(self, n: int) -> str:
|
||||||
|
"从所有可能性中获取第 n 个,以 self.total() 为周期循环"
|
||||||
|
n = n % self.total()
|
||||||
|
indice = decimal_to_irregular(n, self.placeholder_choices.values())
|
||||||
|
replacements = {
|
||||||
|
placeholder_id: self.replacers[self.placeholder_mapping[placeholder_id]][i]
|
||||||
|
for placeholder_id, i in zip(self.placeholder_mapping.keys(), indice)
|
||||||
|
}
|
||||||
|
return self.template.format(**replacements)
|
||||||
|
|
||||||
|
def getmany(self, limit: int, offset: int = 0) -> list[str]:
|
||||||
|
"""返回一组可能性组成的列表,为了避免结果太长导致内存占用超限,使用 limit 限制列表的长度,使用 offset 调整偏移。
|
||||||
|
若 limit 和 offset 的设置导致预期的结果长度超过剩下的实际长度,则会回到开头。
|
||||||
|
"""
|
||||||
|
return [self.getn(n) for n in range(offset, offset + limit)]
|
||||||
|
|
||||||
|
def total(self) -> int:
|
||||||
|
"计算可能性的数目"
|
||||||
|
if self.__total is None:
|
||||||
|
self.__total = prod(self.placeholder_choices.values())
|
||||||
|
return self.__total
|
||||||
|
|
||||||
|
def __make_template(self, text: str):
|
||||||
|
"""将输入的提示词转化成符合 python str.format 要求格式的模板,
|
||||||
|
并将 option 和 wildcard 按照顺序在模板中留下 {r0}, {r1} 等占位符,
|
||||||
|
即使遇到相同的 option 或 wildcard,留下的占位符编号也不同,从而使每项都独立变化。
|
||||||
|
"""
|
||||||
|
self.placeholder_mapping = {}
|
||||||
|
placeholder_id = 0
|
||||||
|
replacer_id = 0
|
||||||
|
replacers_rev = {} # replacers => id
|
||||||
|
blocks = []
|
||||||
|
# 记录所处理过的通配符末尾在文本中的位置,用于拼接完整的模板
|
||||||
|
tail = 0
|
||||||
|
for match in self.RE_REPLACER.finditer(text):
|
||||||
|
# 提取并展开通配符内容
|
||||||
|
m = match.group(0)
|
||||||
|
if m.startswith("{"):
|
||||||
|
choices = expand_options(m[1:-1])
|
||||||
|
elif m.startswith("__"):
|
||||||
|
keyword = m[2:-2].lower()
|
||||||
|
keyword = wildcard_normalize(keyword)
|
||||||
|
choices = expand_wildcard(keyword)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"{m!r} is not a wildcard or option")
|
||||||
|
|
||||||
|
# 记录通配符的替换项列表和ID,相同的通配符只保留第一个
|
||||||
|
if choices not in replacers_rev:
|
||||||
|
replacers_rev[choices] = replacer_id
|
||||||
|
replacer_id += 1
|
||||||
|
|
||||||
|
# 拼接通配符前方文本
|
||||||
|
start, end = match.span()
|
||||||
|
blocks.append(text[tail:start])
|
||||||
|
tail = end
|
||||||
|
# 将通配符替换为占位符,并记录占位符和替换项列表的索引的映射
|
||||||
|
blocks.append(f"{{r{placeholder_id}}}")
|
||||||
|
self.placeholder_mapping[f"r{placeholder_id}"] = replacers_rev[choices]
|
||||||
|
placeholder_id += 1
|
||||||
|
|
||||||
|
if tail < len(text):
|
||||||
|
blocks.append(text[tail:])
|
||||||
|
self.template = "".join(blocks)
|
||||||
|
self.replacers = {v: k for k, v in replacers_rev.items()}
|
||||||
|
self.placeholder_choices = {
|
||||||
|
placeholder_id: len(self.replacers[replacer_id])
|
||||||
|
for placeholder_id, replacer_id in self.placeholder_mapping.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_option():
|
||||||
|
text = "{|a|b|c}"
|
||||||
|
answer = ["", "a", "b", "c"]
|
||||||
|
p = WildcardProcessor(text)
|
||||||
|
assert p.total() == len(answer)
|
||||||
|
assert p.getn(0) == answer[0]
|
||||||
|
assert p.getmany(4) == answer
|
||||||
|
assert p.getmany(4, 1) == answer[1:]
|
||||||
|
|
||||||
|
|
||||||
|
def test_same():
|
||||||
|
text = "{a|b},{a|b}"
|
||||||
|
answer = ["a,a", "b,a", "a,b", "b,b"]
|
||||||
|
p = WildcardProcessor(text)
|
||||||
|
assert p.total() == len(answer)
|
||||||
|
assert p.getn(0) == answer[0]
|
||||||
|
assert p.getmany(4) == answer
|
||||||
|
assert p.getmany(4, 1) == answer[1:]
|
||||||
|
|
||||||
|
|||||||
+261
-62
@@ -1,12 +1,18 @@
|
|||||||
import os, torch
|
import os, torch
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from PIL import Image, ImageDraw, ImageFont
|
from PIL import Image, ImageDraw, ImageFont
|
||||||
from .utils import easySave
|
from .utils import easySave, get_sd_version
|
||||||
from .adv_encode import advanced_encode
|
from .adv_encode import advanced_encode
|
||||||
from .controlnet import easyControlnet
|
from .controlnet import easyControlnet
|
||||||
from .log import log_node_warn
|
from .log import log_node_warn
|
||||||
from ..layer_diffuse import LayerDiffuse
|
from ..modules.layer_diffuse import LayerDiffuse
|
||||||
from ..config import RESOURCES_DIR
|
from ..config import RESOURCES_DIR
|
||||||
|
from nodes import CLIPTextEncode
|
||||||
|
import pprint
|
||||||
|
try:
|
||||||
|
from comfy_extras.nodes_flux import FluxGuidance
|
||||||
|
except:
|
||||||
|
FluxGuidance = None
|
||||||
|
|
||||||
class easyXYPlot():
|
class easyXYPlot():
|
||||||
|
|
||||||
@@ -15,6 +21,7 @@ class easyXYPlot():
|
|||||||
self.y_node_type, self.y_type = sampler.safe_split(xyPlotData.get("y_axis"), ': ')
|
self.y_node_type, self.y_type = sampler.safe_split(xyPlotData.get("y_axis"), ': ')
|
||||||
self.x_values = xyPlotData.get("x_vals") if self.x_type != "None" else []
|
self.x_values = xyPlotData.get("x_vals") if self.x_type != "None" else []
|
||||||
self.y_values = xyPlotData.get("y_vals") if self.y_type != "None" else []
|
self.y_values = xyPlotData.get("y_vals") if self.y_type != "None" else []
|
||||||
|
self.custom_font = xyPlotData.get("custom_font")
|
||||||
|
|
||||||
self.grid_spacing = xyPlotData.get("grid_spacing")
|
self.grid_spacing = xyPlotData.get("grid_spacing")
|
||||||
self.latent_id = 0
|
self.latent_id = 0
|
||||||
@@ -46,7 +53,7 @@ class easyXYPlot():
|
|||||||
|
|
||||||
plot_image_vars[value_type] = value
|
plot_image_vars[value_type] = value
|
||||||
if value_type in ["seed", "Seeds++ Batch"]:
|
if value_type in ["seed", "Seeds++ Batch"]:
|
||||||
value_label = f"{value}"
|
value_label = f"seed: {value}"
|
||||||
else:
|
else:
|
||||||
value_label = f"{value_type}: {value}"
|
value_label = f"{value_type}: {value}"
|
||||||
|
|
||||||
@@ -54,7 +61,16 @@ class easyXYPlot():
|
|||||||
value_label = f"ControlNet {index + 1}"
|
value_label = f"ControlNet {index + 1}"
|
||||||
|
|
||||||
if value_type in ['Lora', 'Checkpoint']:
|
if value_type in ['Lora', 'Checkpoint']:
|
||||||
value_label = f"{os.path.basename(os.path.splitext(value.split(',')[0])[0])}"
|
arr = value.split(',')
|
||||||
|
model_name = os.path.basename(os.path.splitext(arr[0])[0])
|
||||||
|
trigger_words = ' ' + arr[3] if value_type == 'Lora' and len(arr) > 3 and len(arr[3]) > 2 else ''
|
||||||
|
lora_weight = float(arr[1]) if value_type == 'Lora' and len(arr) > 1 else 0
|
||||||
|
lora_weight_desc = f" w:{lora_weight:.2f}" if value_type == 'Lora' and lora_weight != 1.0 else ''
|
||||||
|
value_label = f"{model_name[:25]}{lora_weight_desc}{trigger_words}"
|
||||||
|
|
||||||
|
if value_type == "DiffusionModel":
|
||||||
|
model_name = os.path.basename(os.path.splitext(value.split(",")[0])[0])
|
||||||
|
value_label = model_name[:25]
|
||||||
|
|
||||||
if value_type in ["ModelMergeBlocks"]:
|
if value_type in ["ModelMergeBlocks"]:
|
||||||
if ":" in value:
|
if ":" in value:
|
||||||
@@ -87,8 +103,36 @@ class easyXYPlot():
|
|||||||
return plot_image_vars, value_label
|
return plot_image_vars, value_label
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_font(font_size):
|
def _ensure_latent_for_model(model, vae, samples, plot_image_vars):
|
||||||
return ImageFont.truetype(str(Path(os.path.join(RESOURCES_DIR, 'OpenSans-Medium.ttf'))), font_size)
|
fmt = model.model.latent_format
|
||||||
|
x = samples["samples"]
|
||||||
|
expected_ndim = 2 + fmt.latent_dimensions
|
||||||
|
|
||||||
|
if x.ndim == expected_ndim and x.shape[1] == fmt.latent_channels:
|
||||||
|
return samples
|
||||||
|
|
||||||
|
if fmt.latent_dimensions == 3 and x.ndim == 4:
|
||||||
|
if x.count_nonzero() == 0:
|
||||||
|
x = torch.zeros(
|
||||||
|
[x.shape[0], fmt.latent_channels, 1, x.shape[2], x.shape[3]],
|
||||||
|
dtype=x.dtype, device=x.device)
|
||||||
|
elif plot_image_vars.get("images") is not None:
|
||||||
|
x = vae.encode(plot_image_vars["images"][..., :3])
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Switching to a 3D-latent model requires an input image "
|
||||||
|
"or an empty latent"
|
||||||
|
)
|
||||||
|
|
||||||
|
return {**samples, "samples": x}
|
||||||
|
|
||||||
|
return samples
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_font(font_size, font_path=None):
|
||||||
|
if font_path is None:
|
||||||
|
font_path = str(Path(os.path.join(RESOURCES_DIR, 'OpenSans-Medium.ttf')))
|
||||||
|
return ImageFont.truetype(font_path, font_size)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def update_label(label, value, num_items):
|
def update_label(label, value, num_items):
|
||||||
@@ -107,24 +151,32 @@ class easyXYPlot():
|
|||||||
|
|
||||||
def calculate_background_dimensions(self):
|
def calculate_background_dimensions(self):
|
||||||
border_size = int((self.max_width // 8) * 1.5) if self.y_type != "None" or self.x_type != "None" else 0
|
border_size = int((self.max_width // 8) * 1.5) if self.y_type != "None" or self.x_type != "None" else 0
|
||||||
|
|
||||||
bg_width = self.num_cols * (self.max_width + self.grid_spacing) - self.grid_spacing + border_size * (
|
bg_width = self.num_cols * (self.max_width + self.grid_spacing) - self.grid_spacing + border_size * (
|
||||||
self.y_type != "None")
|
self.y_type != "None")
|
||||||
bg_height = self.num_rows * (self.max_height + self.grid_spacing) - self.grid_spacing + border_size * (
|
bg_height = self.num_rows * (self.max_height + self.grid_spacing) - self.grid_spacing + border_size * (
|
||||||
self.x_type != "None")
|
self.x_type != "None")
|
||||||
|
|
||||||
|
# Add space at the bottom of the image for common informaiton about the image
|
||||||
|
bg_height = bg_height + (border_size*2)
|
||||||
|
# print(f"Grid Size: width = {bg_width} height = {bg_height} border_size = {border_size}")
|
||||||
|
|
||||||
x_offset_initial = border_size if self.y_type != "None" else 0
|
x_offset_initial = border_size if self.y_type != "None" else 0
|
||||||
y_offset = border_size if self.x_type != "None" else 0
|
y_offset = border_size if self.x_type != "None" else 0
|
||||||
|
|
||||||
return bg_width, bg_height, x_offset_initial, y_offset
|
return bg_width, bg_height, x_offset_initial, y_offset
|
||||||
|
|
||||||
|
|
||||||
def adjust_font_size(self, text, initial_font_size, label_width):
|
def adjust_font_size(self, text, initial_font_size, label_width):
|
||||||
font = self.get_font(initial_font_size)
|
font = self.get_font(initial_font_size, self.custom_font)
|
||||||
text_width = font.getbbox(text)
|
text_width = font.getbbox(text)
|
||||||
|
# pprint.pp(f"Initial font size: {initial_font_size}, text: {text}, text_width: {text_width}")
|
||||||
if text_width and text_width[2]:
|
if text_width and text_width[2]:
|
||||||
text_width = text_width[2]
|
text_width = text_width[2]
|
||||||
|
|
||||||
scaling_factor = 0.9
|
scaling_factor = 0.9
|
||||||
if text_width > (label_width * scaling_factor):
|
if text_width > (label_width * scaling_factor):
|
||||||
|
# print(f"Adjusting font size from {initial_font_size} to fit text width {text_width} into label width {label_width} scaling_factor {scaling_factor}")
|
||||||
return int(initial_font_size * (label_width / text_width) * scaling_factor)
|
return int(initial_font_size * (label_width / text_width) * scaling_factor)
|
||||||
else:
|
else:
|
||||||
return initial_font_size
|
return initial_font_size
|
||||||
@@ -133,20 +185,27 @@ class easyXYPlot():
|
|||||||
_, _, width, height = d.textbbox((0, 0), text=text, font=font)
|
_, _, width, height = d.textbbox((0, 0), text=text, font=font)
|
||||||
return width, height
|
return width, height
|
||||||
|
|
||||||
def create_label(self, img, text, initial_font_size, is_x_label=True, max_font_size=70, min_font_size=10):
|
def create_label(self, img, text, initial_font_size, is_x_label=True, max_font_size=70, min_font_size=10, label_width=0, label_height=0):
|
||||||
label_width = img.width if is_x_label else img.height
|
|
||||||
|
|
||||||
|
# if the label_width is specified, leave it along. Otherwise do the old logic.
|
||||||
|
if label_width == 0:
|
||||||
|
label_width = img.width if is_x_label else img.height
|
||||||
|
|
||||||
|
text_lines = text.split('\n')
|
||||||
|
longest_line = max(text_lines, key=len)
|
||||||
|
|
||||||
# Adjust font size
|
# Adjust font size
|
||||||
font_size = self.adjust_font_size(text, initial_font_size, label_width)
|
font_size = self.adjust_font_size(longest_line, initial_font_size, label_width)
|
||||||
font_size = min(max_font_size, font_size) # Ensure font isn't too large
|
font_size = min(max_font_size, font_size) # Ensure font isn't too large
|
||||||
font_size = max(min_font_size, font_size) # Ensure font isn't too small
|
font_size = max(min_font_size, font_size) # Ensure font isn't too small
|
||||||
|
|
||||||
label_height = int(font_size * 1.5) if is_x_label else font_size
|
if label_height == 0:
|
||||||
|
label_height = int(font_size * 1.5) if is_x_label else font_size
|
||||||
|
|
||||||
label_bg = Image.new('RGBA', (label_width, label_height), color=(255, 255, 255, 0))
|
label_bg = Image.new('RGBA', (label_width, label_height), color=(255, 255, 255, 0))
|
||||||
d = ImageDraw.Draw(label_bg)
|
d = ImageDraw.Draw(label_bg)
|
||||||
|
|
||||||
font = self.get_font(font_size)
|
font = self.get_font(font_size, self.custom_font)
|
||||||
|
|
||||||
# Check if text will fit, if not insert ellipsis and reduce text
|
# Check if text will fit, if not insert ellipsis and reduce text
|
||||||
if self.textsize(d, text, font=font)[0] > label_width:
|
if self.textsize(d, text, font=font)[0] > label_width:
|
||||||
@@ -155,7 +214,7 @@ class easyXYPlot():
|
|||||||
text = text + '...'
|
text = text + '...'
|
||||||
|
|
||||||
# Compute text width and height for multi-line text
|
# Compute text width and height for multi-line text
|
||||||
text_lines = text.split('\n')
|
|
||||||
text_widths, text_heights = zip(*[self.textsize(d, line, font=font) for line in text_lines])
|
text_widths, text_heights = zip(*[self.textsize(d, line, font=font) for line in text_lines])
|
||||||
max_text_width = max(text_widths)
|
max_text_width = max(text_widths)
|
||||||
total_text_height = sum(text_heights)
|
total_text_height = sum(text_heights)
|
||||||
@@ -184,6 +243,7 @@ class easyXYPlot():
|
|||||||
clip = clip if clip is not None else plot_image_vars["clip"]
|
clip = clip if clip is not None else plot_image_vars["clip"]
|
||||||
steps = plot_image_vars['steps'] if "steps" in plot_image_vars else 1
|
steps = plot_image_vars['steps'] if "steps" in plot_image_vars else 1
|
||||||
|
|
||||||
|
sd_version = get_sd_version(plot_image_vars['model'])
|
||||||
# 高级用法
|
# 高级用法
|
||||||
if plot_image_vars["x_node_type"] == "advanced" or plot_image_vars["y_node_type"] == "advanced":
|
if plot_image_vars["x_node_type"] == "advanced" or plot_image_vars["y_node_type"] == "advanced":
|
||||||
if self.x_type == "Seeds++ Batch" or self.y_type == "Seeds++ Batch":
|
if self.x_type == "Seeds++ Batch" or self.y_type == "Seeds++ Batch":
|
||||||
@@ -332,19 +392,58 @@ class easyXYPlot():
|
|||||||
if "negative_cond" in plot_image_vars:
|
if "negative_cond" in plot_image_vars:
|
||||||
negative = negative + plot_image_vars["negative_cond"]
|
negative = negative + plot_image_vars["negative_cond"]
|
||||||
|
|
||||||
|
# DiffusionModel
|
||||||
|
if self.x_type == "DiffusionModel" or self.y_type == "DiffusionModel":
|
||||||
|
xy_values = x_value if self.x_type == "DiffusionModel" else y_value
|
||||||
|
model_name, clip_name, vae_name = xy_values.split(",")
|
||||||
|
model, clip, vae, family = self.easyCache.load_diffusion_xy_model(
|
||||||
|
model_name.replace("*", ","),
|
||||||
|
clip_name.replace("*", ","),
|
||||||
|
vae_name.replace("*", ","),
|
||||||
|
)
|
||||||
|
sd_version = family
|
||||||
|
|
||||||
|
positive = plot_image_vars["positive"]
|
||||||
|
negative = plot_image_vars["negative"]
|
||||||
|
if positive is not None:
|
||||||
|
positive, = CLIPTextEncode().encode(clip, positive)
|
||||||
|
if negative is not None:
|
||||||
|
negative, = CLIPTextEncode().encode(clip, negative)
|
||||||
|
|
||||||
|
samples = self._ensure_latent_for_model(
|
||||||
|
model, vae, samples, plot_image_vars
|
||||||
|
)
|
||||||
|
|
||||||
# Lora
|
# Lora
|
||||||
if self.x_type == "Lora" or self.y_type == "Lora":
|
if self.x_type == "Lora" or self.y_type == "Lora":
|
||||||
|
# print(f"Lora: {x_value} {y_value}")
|
||||||
model = model if model is not None else plot_image_vars["model"]
|
model = model if model is not None else plot_image_vars["model"]
|
||||||
clip = clip if clip is not None else plot_image_vars["clip"]
|
clip = clip if clip is not None else plot_image_vars["clip"]
|
||||||
|
|
||||||
|
# Build lora_stack from both X and Y axes if both are LoRA types
|
||||||
|
lora_stack = []
|
||||||
|
|
||||||
|
# Add X axis LoRA if present
|
||||||
|
if self.x_type == "Lora":
|
||||||
|
lora_name, lora_model_strength, lora_clip_strength, _ = x_value.split(",")
|
||||||
|
lora_stack.append({"lora_name": lora_name, "model": model, "clip": clip, "model_strength": float(lora_model_strength), "clip_strength": float(lora_clip_strength)})
|
||||||
|
|
||||||
|
# Add Y axis LoRA if present
|
||||||
|
if self.y_type == "Lora":
|
||||||
|
lora_name, lora_model_strength, lora_clip_strength, _ = y_value.split(",")
|
||||||
|
lora_stack.append({"lora_name": lora_name, "model": model, "clip": clip, "model_strength": float(lora_model_strength), "clip_strength": float(lora_clip_strength)})
|
||||||
|
|
||||||
|
# print(f"new_lora_stack: {new_lora_stack}")
|
||||||
|
|
||||||
xy_values = x_value if self.x_type == "Lora" else y_value
|
|
||||||
lora_name, lora_model_strength, lora_clip_strength = xy_values.split(",")
|
|
||||||
lora_stack = [{"lora_name": lora_name, "model": model, "clip" :clip, "model_strength": float(lora_model_strength), "clip_strength": float(lora_clip_strength)}]
|
|
||||||
if 'lora_stack' in plot_image_vars:
|
if 'lora_stack' in plot_image_vars:
|
||||||
lora_stack = lora_stack + plot_image_vars['lora_stack']
|
lora_stack = lora_stack + plot_image_vars['lora_stack']
|
||||||
|
|
||||||
if lora_stack is not None and lora_stack != []:
|
if lora_stack is not None and lora_stack != []:
|
||||||
for lora in lora_stack:
|
for lora in lora_stack:
|
||||||
|
# Each generation of the model, must use the reference to previously created model / clip objects.
|
||||||
|
lora['model'] = model
|
||||||
|
lora['clip'] = clip
|
||||||
model, clip = self.easyCache.load_lora(lora)
|
model, clip = self.easyCache.load_lora(lora)
|
||||||
|
|
||||||
# 提示词
|
# 提示词
|
||||||
@@ -352,11 +451,14 @@ class easyXYPlot():
|
|||||||
if self.x_type == 'Positive Prompt S/R' or self.y_type == 'Positive Prompt S/R':
|
if self.x_type == 'Positive Prompt S/R' or self.y_type == 'Positive Prompt S/R':
|
||||||
positive = x_value if self.x_type == "Positive Prompt S/R" else y_value
|
positive = x_value if self.x_type == "Positive Prompt S/R" else y_value
|
||||||
|
|
||||||
positive = advanced_encode(clip, positive,
|
if sd_version in ("flux", "anima", "krea2"):
|
||||||
plot_image_vars['positive_token_normalization'],
|
positive, = CLIPTextEncode().encode(clip, positive)
|
||||||
plot_image_vars['positive_weight_interpretation'],
|
else:
|
||||||
w_max=1.0,
|
positive = advanced_encode(clip, positive,
|
||||||
apply_to_pooled="enable", a1111_prompt_style=a1111_prompt_style, steps=steps)
|
plot_image_vars['positive_token_normalization'],
|
||||||
|
plot_image_vars['positive_weight_interpretation'],
|
||||||
|
w_max=1.0,
|
||||||
|
apply_to_pooled="enable", a1111_prompt_style=a1111_prompt_style, steps=steps)
|
||||||
|
|
||||||
# if "positive_cond" in plot_image_vars:
|
# if "positive_cond" in plot_image_vars:
|
||||||
# positive = positive + plot_image_vars["positive_cond"]
|
# positive = positive + plot_image_vars["positive_cond"]
|
||||||
@@ -365,11 +467,14 @@ class easyXYPlot():
|
|||||||
if self.x_type == 'Negative Prompt S/R' or self.y_type == 'Negative Prompt S/R':
|
if self.x_type == 'Negative Prompt S/R' or self.y_type == 'Negative Prompt S/R':
|
||||||
negative = x_value if self.x_type == "Negative Prompt S/R" else y_value
|
negative = x_value if self.x_type == "Negative Prompt S/R" else y_value
|
||||||
|
|
||||||
negative = advanced_encode(clip, negative,
|
if sd_version in ("flux", "anima", "krea2"):
|
||||||
plot_image_vars['negative_token_normalization'],
|
negative, = CLIPTextEncode().encode(clip, negative)
|
||||||
plot_image_vars['negative_weight_interpretation'],
|
else:
|
||||||
w_max=1.0,
|
negative = advanced_encode(clip, negative,
|
||||||
apply_to_pooled="enable", a1111_prompt_style=a1111_prompt_style, steps=steps)
|
plot_image_vars['negative_token_normalization'],
|
||||||
|
plot_image_vars['negative_weight_interpretation'],
|
||||||
|
w_max=1.0,
|
||||||
|
apply_to_pooled="enable", a1111_prompt_style=a1111_prompt_style, steps=steps)
|
||||||
# if "negative_cond" in plot_image_vars:
|
# if "negative_cond" in plot_image_vars:
|
||||||
# negative = negative + plot_image_vars["negative_cond"]
|
# negative = negative + plot_image_vars["negative_cond"]
|
||||||
|
|
||||||
@@ -387,19 +492,42 @@ class easyXYPlot():
|
|||||||
strength = item[2]
|
strength = item[2]
|
||||||
start_percent = item[3]
|
start_percent = item[3]
|
||||||
end_percent = item[4]
|
end_percent = item[4]
|
||||||
positive, negative = easyControlnet().apply(control_net_name, image, positive, negative, strength, start_percent, end_percent, None, 1)
|
provided_control_net = item[5] if len(item) > 5 else None
|
||||||
|
positive, negative = easyControlnet().apply(control_net_name, image, positive, negative, strength, start_percent, end_percent, provided_control_net, 1)
|
||||||
|
# Flux guidance
|
||||||
|
if self.x_type == "Flux Guidance" or self.y_type == "Flux Guidance":
|
||||||
|
positive = plot_image_vars["positive_cond"] if "positive" in plot_image_vars else None
|
||||||
|
flux_guidance = float(x_value) if self.x_type == "Flux Guidance" else float(y_value)
|
||||||
|
positive, = FluxGuidance().append(positive, flux_guidance)
|
||||||
|
|
||||||
# 简单用法
|
# 简单用法
|
||||||
if plot_image_vars["x_node_type"] == "loader" or plot_image_vars["y_node_type"] == "loader":
|
if plot_image_vars["x_node_type"] == "loader" or plot_image_vars["y_node_type"] == "loader":
|
||||||
model, clip, vae, clip_vision = self.easyCache.load_checkpoint(plot_image_vars['ckpt_name'])
|
if self.x_type == 'ckpt_name' or self.y_type == 'ckpt_name':
|
||||||
|
ckpt_name = x_value if self.x_type == "ckpt_name" else y_value
|
||||||
|
model, clip, vae, clip_vision = self.easyCache.load_checkpoint(ckpt_name)
|
||||||
|
|
||||||
if plot_image_vars['lora_name'] != "None":
|
if self.x_type == 'lora_name' or self.y_type == 'lora_name':
|
||||||
lora = {"lora_name": plot_image_vars['lora_name'], "model": model, "clip": clip, "model_strength": plot_image_vars['model_strength'], "clip_strength": plot_image_vars['lora_clip_strength']}
|
model, clip, vae, clip_vision = self.easyCache.load_checkpoint(plot_image_vars['ckpt_name'])
|
||||||
|
lora_name = x_value if self.x_type == "lora_name" else y_value
|
||||||
|
lora = {"lora_name": lora_name, "model": model, "clip": clip, "model_strength": 1, "clip_strength": 1}
|
||||||
|
model, clip = self.easyCache.load_lora(lora)
|
||||||
|
|
||||||
|
if self.x_type == 'lora_model_strength' or self.y_type == 'lora_model_strength':
|
||||||
|
model, clip, vae, clip_vision = self.easyCache.load_checkpoint(plot_image_vars['ckpt_name'])
|
||||||
|
lora_model_strength = float(x_value) if self.x_type == "lora_model_strength" else float(y_value)
|
||||||
|
lora = {"lora_name": plot_image_vars['lora_name'], "model": model, "clip": clip, "model_strength": lora_model_strength, "clip_strength": plot_image_vars['lora_clip_strength']}
|
||||||
|
model, clip = self.easyCache.load_lora(lora)
|
||||||
|
|
||||||
|
if self.x_type == 'lora_clip_strength' or self.y_type == 'lora_clip_strength':
|
||||||
|
model, clip, vae, clip_vision = self.easyCache.load_checkpoint(plot_image_vars['ckpt_name'])
|
||||||
|
lora_clip_strength = float(x_value) if self.x_type == "lora_clip_strength" else float(y_value)
|
||||||
|
lora = {"lora_name": plot_image_vars['lora_name'], "model": model, "clip": clip, "model_strength": plot_image_vars['lora_model_strength'], "clip_strength": lora_clip_strength}
|
||||||
model, clip = self.easyCache.load_lora(lora)
|
model, clip = self.easyCache.load_lora(lora)
|
||||||
|
|
||||||
# Check for custom VAE
|
# Check for custom VAE
|
||||||
if plot_image_vars['vae_name'] not in ["Baked-VAE", "Baked VAE"]:
|
if self.x_type == 'vae_name' or self.y_type == 'vae_name':
|
||||||
vae = self.easyCache.load_vae(plot_image_vars['vae_name'])
|
vae_name = x_value if self.x_type == "vae_name" else y_value
|
||||||
|
vae = self.easyCache.load_vae(vae_name)
|
||||||
|
|
||||||
# CLIP skip
|
# CLIP skip
|
||||||
if not clip:
|
if not clip:
|
||||||
@@ -407,15 +535,22 @@ class easyXYPlot():
|
|||||||
clip = clip.clone()
|
clip = clip.clone()
|
||||||
clip.clip_layer(plot_image_vars['clip_skip'])
|
clip.clip_layer(plot_image_vars['clip_skip'])
|
||||||
|
|
||||||
positive = advanced_encode(clip, plot_image_vars['positive'],
|
if sd_version in ("flux", "anima", "krea2"):
|
||||||
plot_image_vars['positive_token_normalization'],
|
positive, = CLIPTextEncode().encode(clip, positive)
|
||||||
plot_image_vars['positive_weight_interpretation'], w_max=1.0,
|
else:
|
||||||
apply_to_pooled="enable",a1111_prompt_style=a1111_prompt_style, steps=steps)
|
positive = advanced_encode(clip, plot_image_vars['positive'],
|
||||||
|
plot_image_vars['positive_token_normalization'],
|
||||||
|
plot_image_vars['positive_weight_interpretation'], w_max=1.0,
|
||||||
|
apply_to_pooled="enable",a1111_prompt_style=a1111_prompt_style, steps=steps)
|
||||||
|
|
||||||
|
if sd_version in ("flux", "anima", "krea2"):
|
||||||
|
negative, = CLIPTextEncode().encode(clip, negative)
|
||||||
|
else:
|
||||||
|
negative = advanced_encode(clip, plot_image_vars['negative'],
|
||||||
|
plot_image_vars['negative_token_normalization'],
|
||||||
|
plot_image_vars['negative_weight_interpretation'], w_max=1.0,
|
||||||
|
apply_to_pooled="enable", a1111_prompt_style=a1111_prompt_style, steps=steps)
|
||||||
|
|
||||||
negative = advanced_encode(clip, plot_image_vars['negative'],
|
|
||||||
plot_image_vars['negative_token_normalization'],
|
|
||||||
plot_image_vars['negative_weight_interpretation'], w_max=1.0,
|
|
||||||
apply_to_pooled="enable", a1111_prompt_style=a1111_prompt_style, steps=steps)
|
|
||||||
|
|
||||||
model = model if model is not None else plot_image_vars["model"]
|
model = model if model is not None else plot_image_vars["model"]
|
||||||
vae = vae if vae is not None else plot_image_vars["vae"]
|
vae = vae if vae is not None else plot_image_vars["vae"]
|
||||||
@@ -429,6 +564,8 @@ class easyXYPlot():
|
|||||||
scheduler = scheduler if scheduler is not None else plot_image_vars["scheduler"]
|
scheduler = scheduler if scheduler is not None else plot_image_vars["scheduler"]
|
||||||
denoise = denoise if denoise is not None else plot_image_vars["denoise"]
|
denoise = denoise if denoise is not None else plot_image_vars["denoise"]
|
||||||
|
|
||||||
|
noise_device = plot_image_vars["noise_device"] if "noise_device" in plot_image_vars else 'cpu'
|
||||||
|
|
||||||
# LayerDiffuse
|
# LayerDiffuse
|
||||||
layer_diffusion_method = plot_image_vars["layer_diffusion_method"] if "layer_diffusion_method" in plot_image_vars else None
|
layer_diffusion_method = plot_image_vars["layer_diffusion_method"] if "layer_diffusion_method" in plot_image_vars else None
|
||||||
empty_samples = plot_image_vars["empty_samples"] if "empty_samples" in plot_image_vars else None
|
empty_samples = plot_image_vars["empty_samples"] if "empty_samples" in plot_image_vars else None
|
||||||
@@ -448,7 +585,7 @@ class easyXYPlot():
|
|||||||
samples = self.sampler.common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, samples,
|
samples = self.sampler.common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, samples,
|
||||||
denoise=denoise, disable_noise=disable_noise, preview_latent=preview_latent,
|
denoise=denoise, disable_noise=disable_noise, preview_latent=preview_latent,
|
||||||
start_step=start_step, last_step=last_step,
|
start_step=start_step, last_step=last_step,
|
||||||
force_full_denoise=force_full_denoise)
|
force_full_denoise=force_full_denoise, noise_device=noise_device)
|
||||||
|
|
||||||
# Decode images and store
|
# Decode images and store
|
||||||
latent = samples["samples"]
|
latent = samples["samples"]
|
||||||
@@ -502,28 +639,40 @@ class easyXYPlot():
|
|||||||
|
|
||||||
def get_labels_and_sample(self, plot_image_vars, latent_image, preview_latent, start_step, last_step,
|
def get_labels_and_sample(self, plot_image_vars, latent_image, preview_latent, start_step, last_step,
|
||||||
force_full_denoise, disable_noise):
|
force_full_denoise, disable_noise):
|
||||||
for x_index, x_value in enumerate(self.x_values):
|
# Handle X-only variation (Y is "None")
|
||||||
plot_image_vars, x_value_label = self.define_variable(plot_image_vars, self.x_type, x_value,
|
if self.y_type == 'None':
|
||||||
x_index)
|
for x_index, x_value in enumerate(self.x_values):
|
||||||
self.x_label = self.update_label(self.x_label, x_value_label, len(self.x_values))
|
plot_image_vars, x_value_label = self.define_variable(plot_image_vars, self.x_type, x_value, x_index)
|
||||||
if self.y_type != 'None':
|
self.x_label = self.update_label(self.x_label, x_value_label, len(self.x_values))
|
||||||
|
|
||||||
|
self.image_list, self.max_width, self.max_height, self.latents_plot = self.sample_plot_image(
|
||||||
|
plot_image_vars, latent_image, preview_latent, self.latents_plot, self.image_list,
|
||||||
|
disable_noise, start_step, last_step, force_full_denoise, x_value)
|
||||||
|
self.num += 1
|
||||||
|
# Handle Y-only variation (X is "None")
|
||||||
|
elif self.x_type == 'None':
|
||||||
|
for y_index, y_value in enumerate(self.y_values):
|
||||||
|
plot_image_vars, y_value_label = self.define_variable(plot_image_vars, self.y_type, y_value, y_index)
|
||||||
|
self.y_label = self.update_label(self.y_label, y_value_label, len(self.y_values))
|
||||||
|
|
||||||
|
self.image_list, self.max_width, self.max_height, self.latents_plot = self.sample_plot_image(
|
||||||
|
plot_image_vars, latent_image, preview_latent, self.latents_plot, self.image_list,
|
||||||
|
disable_noise, start_step, last_step, force_full_denoise, y_value=y_value)
|
||||||
|
self.num += 1
|
||||||
|
# Handle both X and Y variation
|
||||||
|
else:
|
||||||
|
for x_index, x_value in enumerate(self.x_values):
|
||||||
|
plot_image_vars, x_value_label = self.define_variable(plot_image_vars, self.x_type, x_value, x_index)
|
||||||
|
self.x_label = self.update_label(self.x_label, x_value_label, len(self.x_values))
|
||||||
|
|
||||||
for y_index, y_value in enumerate(self.y_values):
|
for y_index, y_value in enumerate(self.y_values):
|
||||||
plot_image_vars, y_value_label = self.define_variable(plot_image_vars, self.y_type, y_value,
|
plot_image_vars, y_value_label = self.define_variable(plot_image_vars, self.y_type, y_value, y_index)
|
||||||
y_index)
|
|
||||||
self.y_label = self.update_label(self.y_label, y_value_label, len(self.y_values))
|
self.y_label = self.update_label(self.y_label, y_value_label, len(self.y_values))
|
||||||
# ttNl(f'{CC.GREY}X: {x_value_label}, Y: {y_value_label}').t(
|
|
||||||
# f'Plot Values {self.num}/{self.total} ->').p()
|
|
||||||
|
|
||||||
self.image_list, self.max_width, self.max_height, self.latents_plot = self.sample_plot_image(
|
self.image_list, self.max_width, self.max_height, self.latents_plot = self.sample_plot_image(
|
||||||
plot_image_vars, latent_image, preview_latent, self.latents_plot, self.image_list,
|
plot_image_vars, latent_image, preview_latent, self.latents_plot, self.image_list,
|
||||||
disable_noise, start_step, last_step, force_full_denoise, x_value, y_value)
|
disable_noise, start_step, last_step, force_full_denoise, x_value, y_value)
|
||||||
self.num += 1
|
self.num += 1
|
||||||
else:
|
|
||||||
# ttNl(f'{CC.GREY}X: {x_value_label}').t(f'Plot Values {self.num}/{self.total} ->').p()
|
|
||||||
self.image_list, self.max_width, self.max_height, self.latents_plot = self.sample_plot_image(
|
|
||||||
plot_image_vars, latent_image, preview_latent, self.latents_plot, self.image_list, disable_noise,
|
|
||||||
start_step, last_step, force_full_denoise, x_value)
|
|
||||||
self.num += 1
|
|
||||||
|
|
||||||
# Rearrange latent array to match preview image grid
|
# Rearrange latent array to match preview image grid
|
||||||
self.latents_plot = self.rearrange_tensors(self.latents_plot, self.num_cols, self.num_rows)
|
self.latents_plot = self.rearrange_tensors(self.latents_plot, self.num_cols, self.num_rows)
|
||||||
@@ -533,11 +682,10 @@ class easyXYPlot():
|
|||||||
|
|
||||||
return self.latents_plot
|
return self.latents_plot
|
||||||
|
|
||||||
def plot_images_and_labels(self):
|
def plot_images_and_labels(self, plot_image_vars):
|
||||||
# Calculate the background dimensions
|
|
||||||
bg_width, bg_height, x_offset_initial, y_offset = self.calculate_background_dimensions()
|
bg_width, bg_height, x_offset_initial, y_offset = self.calculate_background_dimensions()
|
||||||
|
|
||||||
# Create the white background image
|
|
||||||
background = Image.new('RGBA', (int(bg_width), int(bg_height)), color=(255, 255, 255, 255))
|
background = Image.new('RGBA', (int(bg_width), int(bg_height)), color=(255, 255, 255, 255))
|
||||||
|
|
||||||
output_image = []
|
output_image = []
|
||||||
@@ -569,4 +717,55 @@ class easyXYPlot():
|
|||||||
|
|
||||||
y_offset += img.height + self.grid_spacing
|
y_offset += img.height + self.grid_spacing
|
||||||
|
|
||||||
return (self.sampler.pil2tensor(background), output_image)
|
# lookup used models in the image
|
||||||
|
common_label = ""
|
||||||
|
# Update to add a function to do the heavy lifting. Parameters are plot_image_vars name, label to use, names of the axis,
|
||||||
|
|
||||||
|
# pprint.pp(plot_image_vars)
|
||||||
|
|
||||||
|
# We don't process LORAs here because there can be multiple of them.
|
||||||
|
labels = [
|
||||||
|
{"id": "ckpt_name", "id_desc": "ckpt", "axis_type" : "Checkpoint"},
|
||||||
|
{"id": "vae_name", "id_desc": '', "axis_type" : "vae_name"},
|
||||||
|
{"id": "sampler_name", "id_desc": "sampler", "axis_type" : "Sampler"},
|
||||||
|
{"id": "scheduler", "id_desc": '', "axis_type" : "Scheduler"},
|
||||||
|
{"id": "steps", "id_desc": '', "axis_type" : "Steps"},
|
||||||
|
{"id": "Flux Guidance", "id_desc": 'guidance', "axis_type" : "Flux Guidance"},
|
||||||
|
{"id": "seed", "id_desc": '', "axis_type" : "Seeds++ Batch"}
|
||||||
|
]
|
||||||
|
|
||||||
|
for item in labels:
|
||||||
|
# Only add the label if it's not one of the axis
|
||||||
|
# print(f"Checking item: {item['id']} axis_type {item['axis_type']} x_type: {self.x_type} y_type: {self.y_type}")
|
||||||
|
if self.x_type != item['axis_type'] and self.y_type != item['axis_type']:
|
||||||
|
common_label += self.add_common_label(item['id'], plot_image_vars, item['id_desc'])
|
||||||
|
common_label += f"\n"
|
||||||
|
|
||||||
|
if plot_image_vars['lora_stack'] is not None and plot_image_vars['lora_stack'] != []:
|
||||||
|
# print(f"lora_stack: {plot_image_vars['lora_stack']}")
|
||||||
|
for lora in plot_image_vars['lora_stack']:
|
||||||
|
|
||||||
|
lora_name = lora['lora_name']
|
||||||
|
lora_weight = lora['model_strength']
|
||||||
|
if lora_name is not None and len(lora_name) > 0 and lora_weight > 0:
|
||||||
|
common_label += f"LORA: {lora_name} weight: {lora_weight:.2f} \n"
|
||||||
|
|
||||||
|
common_label = common_label.strip()
|
||||||
|
|
||||||
|
if len(common_label) > 0:
|
||||||
|
label_height = background.height - y_offset
|
||||||
|
label_bg = self.create_label(background, common_label, int(48 * background.width / 512), label_width=background.width, label_height=label_height)
|
||||||
|
label_x = (background.width - label_bg.width) // 2
|
||||||
|
label_y = y_offset
|
||||||
|
# print(f"Adding common label: {common_label} x = {label_x} y = {label_y}")
|
||||||
|
background.alpha_composite(label_bg, (label_x, label_y))
|
||||||
|
|
||||||
|
return (self.sampler.pil2tensor(background), output_image)
|
||||||
|
|
||||||
|
def add_common_label(self, tag, plot_image_vars, description = ''):
|
||||||
|
label = ''
|
||||||
|
if description == '': description = tag
|
||||||
|
if tag in plot_image_vars and plot_image_vars[tag] is not None and plot_image_vars[tag] != 'None':
|
||||||
|
label += f"{description}: {plot_image_vars[tag]} "
|
||||||
|
# print(f"add_common_label: {tag} description: {description} label: {label}" )
|
||||||
|
return label
|
||||||
|
|||||||
-623
@@ -1,623 +0,0 @@
|
|||||||
from typing import Iterator, List, Tuple, Dict, Any, Union, Optional
|
|
||||||
from _decimal import Context, getcontext
|
|
||||||
from decimal import Decimal
|
|
||||||
from .libs.utils import AlwaysEqualProxy, cleanGPUUsedForce
|
|
||||||
from .libs.cache import remove_cache
|
|
||||||
import numpy as np
|
|
||||||
import json
|
|
||||||
|
|
||||||
def validate_list_args(args: Dict[str, List[Any]]) -> Tuple[bool, Optional[str], Optional[str]]:
|
|
||||||
"""
|
|
||||||
Checks that if there are multiple arguments, they are all the same length or 1
|
|
||||||
:param args:
|
|
||||||
:return: Tuple (Status, mismatched_key_1, mismatched_key_2)
|
|
||||||
"""
|
|
||||||
# Only have 1 arg
|
|
||||||
if len(args) == 1:
|
|
||||||
return True, None, None
|
|
||||||
|
|
||||||
len_to_match = None
|
|
||||||
matched_arg_name = None
|
|
||||||
for arg_name, arg in args.items():
|
|
||||||
if arg_name == 'self':
|
|
||||||
# self is in locals()
|
|
||||||
continue
|
|
||||||
|
|
||||||
if len(arg) != 1:
|
|
||||||
if len_to_match is None:
|
|
||||||
len_to_match = len(arg)
|
|
||||||
matched_arg_name = arg_name
|
|
||||||
elif len(arg) != len_to_match:
|
|
||||||
return False, arg_name, matched_arg_name
|
|
||||||
|
|
||||||
return True, None, None
|
|
||||||
def error_if_mismatched_list_args(args: Dict[str, List[Any]]) -> None:
|
|
||||||
is_valid, failed_key1, failed_key2 = validate_list_args(args)
|
|
||||||
if not is_valid:
|
|
||||||
assert failed_key1 is not None
|
|
||||||
assert failed_key2 is not None
|
|
||||||
raise ValueError(
|
|
||||||
f"Mismatched list inputs received. {failed_key1}({len(args[failed_key1])}) !== {failed_key2}({len(args[failed_key2])})"
|
|
||||||
)
|
|
||||||
|
|
||||||
def zip_with_fill(*lists: Union[List[Any], None]) -> Iterator[Tuple[Any, ...]]:
|
|
||||||
"""
|
|
||||||
Zips lists together, but if a list has 1 element, it will be repeated for each element in the other lists.
|
|
||||||
If a list is None, None will be used for that element.
|
|
||||||
(Not intended for use with lists of different lengths)
|
|
||||||
:param lists:
|
|
||||||
:return: Iterator of tuples of length len(lists)
|
|
||||||
"""
|
|
||||||
max_len = max(len(lst) if lst is not None else 0 for lst in lists)
|
|
||||||
for i in range(max_len):
|
|
||||||
yield tuple(None if lst is None else (lst[0] if len(lst) == 1 else lst[i]) for lst in lists)
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------类型 开始----------------------------------------------------------------------#
|
|
||||||
|
|
||||||
# 字符串
|
|
||||||
class String:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {
|
|
||||||
"required": {"value": ("STRING", {"default": ""})},
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("STRING",)
|
|
||||||
RETURN_NAMES = ("string",)
|
|
||||||
FUNCTION = "execute"
|
|
||||||
CATEGORY = "EasyUse/Logic/Type"
|
|
||||||
|
|
||||||
def execute(self, value):
|
|
||||||
return (value,)
|
|
||||||
|
|
||||||
# 整数
|
|
||||||
class Int:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {
|
|
||||||
"required": {"value": ("INT", {"default": 0, "min": -999999, "max": 999999,})},
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("INT",)
|
|
||||||
RETURN_NAMES = ("int",)
|
|
||||||
FUNCTION = "execute"
|
|
||||||
CATEGORY = "EasyUse/Logic/Type"
|
|
||||||
|
|
||||||
def execute(self, value):
|
|
||||||
return (value,)
|
|
||||||
|
|
||||||
# 整数范围
|
|
||||||
class RangeInt:
|
|
||||||
def __init__(self) -> None:
|
|
||||||
pass
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s) -> Dict[str, Dict[str, Any]]:
|
|
||||||
return {
|
|
||||||
"required": {
|
|
||||||
"range_mode": (["step", "num_steps"], {"default": "step"}),
|
|
||||||
"start": ("INT", {"default": 0, "min": -4096, "max": 4096, "step": 1}),
|
|
||||||
"stop": ("INT", {"default": 0, "min": -4096, "max": 4096, "step": 1}),
|
|
||||||
"step": ("INT", {"default": 0, "min": -4096, "max": 4096, "step": 1}),
|
|
||||||
"num_steps": ("INT", {"default": 0, "min": -4096, "max": 4096, "step": 1}),
|
|
||||||
"end_mode": (["Inclusive", "Exclusive"], {"default": "Inclusive"}),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("INT", "INT")
|
|
||||||
RETURN_NAMES = ("range", "range_sizes")
|
|
||||||
INPUT_IS_LIST = True
|
|
||||||
OUTPUT_IS_LIST = (True, True)
|
|
||||||
FUNCTION = "build_range"
|
|
||||||
|
|
||||||
CATEGORY = "EasyUse/Logic/Type"
|
|
||||||
|
|
||||||
def build_range(
|
|
||||||
self, range_mode, start, stop, step, num_steps, end_mode
|
|
||||||
) -> Tuple[List[int], List[int]]:
|
|
||||||
error_if_mismatched_list_args(locals())
|
|
||||||
|
|
||||||
ranges = []
|
|
||||||
range_sizes = []
|
|
||||||
for range_mode, e_start, e_stop, e_num_steps, e_step, e_end_mode in zip_with_fill(
|
|
||||||
range_mode, start, stop, num_steps, step, end_mode
|
|
||||||
):
|
|
||||||
if range_mode == 'step':
|
|
||||||
if e_end_mode == "Inclusive":
|
|
||||||
e_stop += 1
|
|
||||||
vals = list(range(e_start, e_stop, e_step))
|
|
||||||
ranges.extend(vals)
|
|
||||||
range_sizes.append(len(vals))
|
|
||||||
elif range_mode == 'num_steps':
|
|
||||||
direction = 1 if e_stop > e_start else -1
|
|
||||||
if e_end_mode == "Exclusive":
|
|
||||||
e_stop -= direction
|
|
||||||
vals = (np.rint(np.linspace(e_start, e_stop, e_num_steps)).astype(int).tolist())
|
|
||||||
ranges.extend(vals)
|
|
||||||
range_sizes.append(len(vals))
|
|
||||||
return ranges, range_sizes
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
# 浮点数
|
|
||||||
class Float:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {
|
|
||||||
"required": {"value": ("FLOAT", {"default": 0, "step": 0.01, "min": -999999, "max": 999999,})},
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("FLOAT",)
|
|
||||||
RETURN_NAMES = ("float",)
|
|
||||||
FUNCTION = "execute"
|
|
||||||
CATEGORY = "EasyUse/Logic/Type"
|
|
||||||
|
|
||||||
def execute(self, value):
|
|
||||||
return (value,)
|
|
||||||
|
|
||||||
|
|
||||||
# 浮点数范围
|
|
||||||
class RangeFloat:
|
|
||||||
def __init__(self) -> None:
|
|
||||||
pass
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s) -> Dict[str, Dict[str, Any]]:
|
|
||||||
return {
|
|
||||||
"required": {
|
|
||||||
"range_mode": (["step", "num_steps"], {"default": "step"}),
|
|
||||||
"start": ("FLOAT", {"default": 0, "min": -4096, "max": 4096, "step": 0.1}),
|
|
||||||
"stop": ("FLOAT", {"default": 0, "min": -4096, "max": 4096, "step": 0.1}),
|
|
||||||
"step": ("FLOAT", {"default": 0, "min": -4096, "max": 4096, "step": 0.1}),
|
|
||||||
"num_steps": ("INT", {"default": 0, "min": -4096, "max": 4096, "step": 1}),
|
|
||||||
"end_mode": (["Inclusive", "Exclusive"], {"default": "Inclusive"}),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("FLOAT", "INT")
|
|
||||||
RETURN_NAMES = ("range", "range_sizes")
|
|
||||||
INPUT_IS_LIST = True
|
|
||||||
OUTPUT_IS_LIST = (True, True)
|
|
||||||
FUNCTION = "build_range"
|
|
||||||
|
|
||||||
CATEGORY = "EasyUse/Logic/Type"
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _decimal_range(
|
|
||||||
range_mode: String, start: Decimal, stop: Decimal, step: Decimal, num_steps: Int, inclusive: bool
|
|
||||||
) -> Iterator[float]:
|
|
||||||
if range_mode == 'step':
|
|
||||||
ret_val = start
|
|
||||||
if inclusive:
|
|
||||||
stop = stop + step
|
|
||||||
direction = 1 if step > 0 else -1
|
|
||||||
while (ret_val - stop) * direction < 0:
|
|
||||||
yield float(ret_val)
|
|
||||||
ret_val += step
|
|
||||||
elif range_mode == 'num_steps':
|
|
||||||
step = (stop - start) / (num_steps - 1)
|
|
||||||
direction = 1 if step > 0 else -1
|
|
||||||
|
|
||||||
ret_val = start
|
|
||||||
for _ in range(num_steps):
|
|
||||||
if (ret_val - stop) * direction > 0: # Ensure we don't exceed the 'stop' value
|
|
||||||
break
|
|
||||||
yield float(ret_val)
|
|
||||||
ret_val += step
|
|
||||||
|
|
||||||
def build_range(
|
|
||||||
self,
|
|
||||||
range_mode,
|
|
||||||
start,
|
|
||||||
stop,
|
|
||||||
step,
|
|
||||||
num_steps,
|
|
||||||
end_mode,
|
|
||||||
) -> Tuple[List[float], List[int]]:
|
|
||||||
error_if_mismatched_list_args(locals())
|
|
||||||
getcontext().prec = 12
|
|
||||||
|
|
||||||
start = [Decimal(s) for s in start]
|
|
||||||
stop = [Decimal(s) for s in stop]
|
|
||||||
step = [Decimal(s) for s in step]
|
|
||||||
|
|
||||||
ranges = []
|
|
||||||
range_sizes = []
|
|
||||||
for range_mode, e_start, e_stop, e_step, e_num_steps, e_end_mode in zip_with_fill(
|
|
||||||
range_mode, start, stop, step, num_steps, end_mode
|
|
||||||
):
|
|
||||||
vals = list(
|
|
||||||
self._decimal_range(range_mode, e_start, e_stop, e_step, e_num_steps, e_end_mode == 'Inclusive')
|
|
||||||
)
|
|
||||||
ranges.extend(vals)
|
|
||||||
range_sizes.append(len(vals))
|
|
||||||
|
|
||||||
return ranges, range_sizes
|
|
||||||
|
|
||||||
|
|
||||||
# 布尔
|
|
||||||
class Boolean:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {
|
|
||||||
"required": {"value": ("BOOLEAN", {"default": False})},
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("BOOLEAN",)
|
|
||||||
RETURN_NAMES = ("boolean",)
|
|
||||||
FUNCTION = "execute"
|
|
||||||
CATEGORY = "EasyUse/Logic/Type"
|
|
||||||
|
|
||||||
def execute(self, value):
|
|
||||||
return (value,)
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------开关 开始----------------------------------------------------------------------#
|
|
||||||
class imageSwitch:
|
|
||||||
def __init__(self):
|
|
||||||
pass
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(cls):
|
|
||||||
return {
|
|
||||||
"required": {
|
|
||||||
"image_a": ("IMAGE",),
|
|
||||||
"image_b": ("IMAGE",),
|
|
||||||
"boolean": ("BOOLEAN", {"default": False}),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE",)
|
|
||||||
FUNCTION = "image_switch"
|
|
||||||
|
|
||||||
CATEGORY = "EasyUse/Logic/Switch"
|
|
||||||
|
|
||||||
def image_switch(self, image_a, image_b, boolean):
|
|
||||||
|
|
||||||
if boolean:
|
|
||||||
return (image_a, )
|
|
||||||
else:
|
|
||||||
return (image_b, )
|
|
||||||
|
|
||||||
class textSwitch:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(cls):
|
|
||||||
return {
|
|
||||||
"required": {
|
|
||||||
"input": ("INT", {"default": 1, "min": 1, "max": 2}),
|
|
||||||
},
|
|
||||||
"optional": {
|
|
||||||
"text1": ("STRING", {"forceInput": True}),
|
|
||||||
"text2": ("STRING", {"forceInput": True}),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("STRING",)
|
|
||||||
RETURN_NAMES = ("STRING",)
|
|
||||||
CATEGORY = "EasyUse/Logic/Switch"
|
|
||||||
FUNCTION = "switch"
|
|
||||||
|
|
||||||
def switch(self, input, text1=None, text2=None,):
|
|
||||||
if input == 1:
|
|
||||||
return (text1,)
|
|
||||||
else:
|
|
||||||
return (text2,)
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------运算 开始----------------------------------------------------------------------#
|
|
||||||
|
|
||||||
COMPARE_FUNCTIONS = {
|
|
||||||
"a == b": lambda a, b: a == b,
|
|
||||||
"a != b": lambda a, b: a != b,
|
|
||||||
"a < b": lambda a, b: a < b,
|
|
||||||
"a > b": lambda a, b: a > b,
|
|
||||||
"a <= b": lambda a, b: a <= b,
|
|
||||||
"a >= b": lambda a, b: a >= b,
|
|
||||||
}
|
|
||||||
|
|
||||||
# 比较
|
|
||||||
class Compare:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
s.compare_functions = list(COMPARE_FUNCTIONS.keys())
|
|
||||||
return {
|
|
||||||
"required": {
|
|
||||||
"a": (AlwaysEqualProxy("*"), {"default": 0}),
|
|
||||||
"b": (AlwaysEqualProxy("*"), {"default": 0}),
|
|
||||||
"comparison": (s.compare_functions, {"default": "a == b"}),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("BOOLEAN",)
|
|
||||||
RETURN_NAMES = ("boolean",)
|
|
||||||
FUNCTION = "compare"
|
|
||||||
CATEGORY = "EasyUse/Logic/Math"
|
|
||||||
|
|
||||||
def compare(self, a, b, comparison):
|
|
||||||
return (COMPARE_FUNCTIONS[comparison](a, b),)
|
|
||||||
|
|
||||||
# 判断
|
|
||||||
class If:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {
|
|
||||||
"required": {
|
|
||||||
"any": (AlwaysEqualProxy("*"),),
|
|
||||||
"if": (AlwaysEqualProxy("*"),),
|
|
||||||
"else": (AlwaysEqualProxy("*"),),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = (AlwaysEqualProxy("*"),)
|
|
||||||
RETURN_NAMES = ("?",)
|
|
||||||
FUNCTION = "execute"
|
|
||||||
CATEGORY = "EasyUse/Logic/Math"
|
|
||||||
|
|
||||||
def execute(self, *args, **kwargs):
|
|
||||||
return (kwargs['if'] if kwargs['any'] else kwargs['else'],)
|
|
||||||
|
|
||||||
|
|
||||||
#是否为SDXL
|
|
||||||
from comfy.sdxl_clip import SDXLClipModel, SDXLRefinerClipModel, SDXLClipG
|
|
||||||
class isSDXL:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {
|
|
||||||
"required": {},
|
|
||||||
"optional": {
|
|
||||||
"optional_pipe": ("PIPE_LINE",),
|
|
||||||
"optional_clip": ("CLIP",),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("BOOLEAN",)
|
|
||||||
RETURN_NAMES = ("boolean",)
|
|
||||||
FUNCTION = "execute"
|
|
||||||
CATEGORY = "EasyUse/Logic"
|
|
||||||
|
|
||||||
def execute(self, optional_pipe=None, optional_clip=None):
|
|
||||||
if optional_pipe is None and optional_clip is None:
|
|
||||||
raise Exception(f"[ERROR] optional_pipe or optional_clip is missing")
|
|
||||||
clip = optional_clip if optional_clip is not None else optional_pipe['clip']
|
|
||||||
if isinstance(clip.cond_stage_model, (SDXLClipModel, SDXLRefinerClipModel, SDXLClipG)):
|
|
||||||
return (True,)
|
|
||||||
else:
|
|
||||||
return (False,)
|
|
||||||
|
|
||||||
#xy矩阵
|
|
||||||
class xyAny:
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
|
|
||||||
return {
|
|
||||||
"required": {
|
|
||||||
"X": (AlwaysEqualProxy("*"), {}),
|
|
||||||
"Y": (AlwaysEqualProxy("*"), {}),
|
|
||||||
"direction": (["horizontal", "vertical"], {"default": "horizontal"})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = (AlwaysEqualProxy("*"), AlwaysEqualProxy("*"))
|
|
||||||
RETURN_NAMES = ("X", "Y")
|
|
||||||
INPUT_IS_LIST = True
|
|
||||||
OUTPUT_IS_LIST = (True, True)
|
|
||||||
CATEGORY = "EasyUse/Logic"
|
|
||||||
FUNCTION = "to_xy"
|
|
||||||
|
|
||||||
def to_xy(self, X, Y, direction):
|
|
||||||
new_x = list()
|
|
||||||
new_y = list()
|
|
||||||
if direction[0] == "horizontal":
|
|
||||||
for y in Y:
|
|
||||||
for x in X:
|
|
||||||
new_x.append(x)
|
|
||||||
new_y.append(y)
|
|
||||||
else:
|
|
||||||
for x in X:
|
|
||||||
for y in Y:
|
|
||||||
new_x.append(x)
|
|
||||||
new_y.append(y)
|
|
||||||
|
|
||||||
return (new_x, new_y)
|
|
||||||
|
|
||||||
# 转换所有类型
|
|
||||||
class ConvertAnything:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {"required": {
|
|
||||||
"anything": (AlwaysEqualProxy("*"),),
|
|
||||||
"output_type": (["string", "int", "float", "boolean"], {"default": "string"}),
|
|
||||||
}}
|
|
||||||
|
|
||||||
RETURN_TYPES = (AlwaysEqualProxy("*"),),
|
|
||||||
RETURN_NAMES = ('*',)
|
|
||||||
OUTPUT_NODE = True
|
|
||||||
FUNCTION = "convert"
|
|
||||||
CATEGORY = "EasyUse/Logic"
|
|
||||||
|
|
||||||
def convert(self, *args, **kwargs):
|
|
||||||
print(kwargs)
|
|
||||||
anything = kwargs['anything']
|
|
||||||
output_type = kwargs['output_type']
|
|
||||||
params = None
|
|
||||||
if output_type == 'string':
|
|
||||||
params = str(anything)
|
|
||||||
elif output_type == 'int':
|
|
||||||
params = int(anything)
|
|
||||||
elif output_type == 'float':
|
|
||||||
params = float(anything)
|
|
||||||
elif output_type == 'boolean':
|
|
||||||
params = bool(anything)
|
|
||||||
return (params,)
|
|
||||||
|
|
||||||
# 将所有类型的内容都转成字符串输出
|
|
||||||
class showAnything:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {"required": {}, "optional": {"anything": (AlwaysEqualProxy("*"), {}), },
|
|
||||||
"hidden": {"unique_id": "UNIQUE_ID", "extra_pnginfo": "EXTRA_PNGINFO",
|
|
||||||
}}
|
|
||||||
|
|
||||||
RETURN_TYPES = ()
|
|
||||||
INPUT_IS_LIST = True
|
|
||||||
OUTPUT_NODE = True
|
|
||||||
FUNCTION = "log_input"
|
|
||||||
CATEGORY = "EasyUse/Logic"
|
|
||||||
|
|
||||||
def log_input(self, unique_id=None, extra_pnginfo=None, **kwargs):
|
|
||||||
|
|
||||||
values = []
|
|
||||||
if "anything" in kwargs:
|
|
||||||
for val in kwargs['anything']:
|
|
||||||
try:
|
|
||||||
if type(val) is str:
|
|
||||||
values.append(val)
|
|
||||||
else:
|
|
||||||
val = json.dumps(val)
|
|
||||||
values.append(str(val))
|
|
||||||
except Exception:
|
|
||||||
values.append(str(val))
|
|
||||||
pass
|
|
||||||
|
|
||||||
if not extra_pnginfo:
|
|
||||||
print("Error: extra_pnginfo is empty")
|
|
||||||
elif (not isinstance(extra_pnginfo[0], dict) or "workflow" not in extra_pnginfo[0]):
|
|
||||||
print("Error: extra_pnginfo[0] is not a dict or missing 'workflow' key")
|
|
||||||
else:
|
|
||||||
workflow = extra_pnginfo[0]["workflow"]
|
|
||||||
node = next((x for x in workflow["nodes"] if str(x["id"]) == unique_id[0]), None)
|
|
||||||
if node:
|
|
||||||
node["widgets_values"] = [values]
|
|
||||||
|
|
||||||
return {"ui": {"text": values}}
|
|
||||||
|
|
||||||
class showTensorShape:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {"required": {"tensor": (AlwaysEqualProxy("*"),)}, "optional": {},
|
|
||||||
"hidden": {"unique_id": "UNIQUE_ID", "extra_pnginfo": "EXTRA_PNGINFO"
|
|
||||||
}}
|
|
||||||
|
|
||||||
RETURN_TYPES = ()
|
|
||||||
RETURN_NAMES = ()
|
|
||||||
OUTPUT_NODE = True
|
|
||||||
FUNCTION = "log_input"
|
|
||||||
CATEGORY = "EasyUse/Logic"
|
|
||||||
|
|
||||||
def log_input(self, tensor, unique_id=None, extra_pnginfo=None):
|
|
||||||
shapes = []
|
|
||||||
|
|
||||||
def tensorShape(tensor):
|
|
||||||
if isinstance(tensor, dict):
|
|
||||||
for k in tensor:
|
|
||||||
tensorShape(tensor[k])
|
|
||||||
elif isinstance(tensor, list):
|
|
||||||
for i in range(len(tensor)):
|
|
||||||
tensorShape(tensor[i])
|
|
||||||
elif hasattr(tensor, 'shape'):
|
|
||||||
shapes.append(list(tensor.shape))
|
|
||||||
|
|
||||||
tensorShape(tensor)
|
|
||||||
|
|
||||||
return {"ui": {"text": shapes}}
|
|
||||||
|
|
||||||
# cleanGpuUsed
|
|
||||||
class cleanGPUUsed:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {"required": {"anything": (AlwaysEqualProxy("*"), {})}, "optional": {},
|
|
||||||
"hidden": {"unique_id": "UNIQUE_ID", "extra_pnginfo": "EXTRA_PNGINFO",
|
|
||||||
}}
|
|
||||||
|
|
||||||
RETURN_TYPES = ()
|
|
||||||
RETURN_NAMES = ()
|
|
||||||
OUTPUT_NODE = True
|
|
||||||
FUNCTION = "empty_cache"
|
|
||||||
CATEGORY = "EasyUse/Logic"
|
|
||||||
|
|
||||||
def empty_cache(self, anything, unique_id=None, extra_pnginfo=None):
|
|
||||||
cleanGPUUsedForce()
|
|
||||||
remove_cache('*')
|
|
||||||
return ()
|
|
||||||
|
|
||||||
class clearCacheKey:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {"required": {
|
|
||||||
"anything": (AlwaysEqualProxy("*"), {}),
|
|
||||||
"cache_key": ("STRING", {"default": "*"}),
|
|
||||||
}, "optional": {},
|
|
||||||
"hidden": {"unique_id": "UNIQUE_ID", "extra_pnginfo": "EXTRA_PNGINFO",}
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ()
|
|
||||||
RETURN_NAMES = ()
|
|
||||||
OUTPUT_NODE = True
|
|
||||||
FUNCTION = "empty_cache"
|
|
||||||
CATEGORY = "EasyUse/Logic"
|
|
||||||
|
|
||||||
def empty_cache(self, anything, cache_name, unique_id=None, extra_pnginfo=None):
|
|
||||||
remove_cache(cache_name)
|
|
||||||
return ()
|
|
||||||
|
|
||||||
class clearCacheAll:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {"required": {
|
|
||||||
"anything": (AlwaysEqualProxy("*"), {}),
|
|
||||||
}, "optional": {},
|
|
||||||
"hidden": {"unique_id": "UNIQUE_ID", "extra_pnginfo": "EXTRA_PNGINFO",}
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ()
|
|
||||||
RETURN_NAMES = ()
|
|
||||||
OUTPUT_NODE = True
|
|
||||||
FUNCTION = "empty_cache"
|
|
||||||
CATEGORY = "EasyUse/Logic"
|
|
||||||
|
|
||||||
def empty_cache(self, anything, unique_id=None, extra_pnginfo=None):
|
|
||||||
remove_cache('*')
|
|
||||||
return ()
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
|
||||||
"easy string": String,
|
|
||||||
"easy int": Int,
|
|
||||||
"easy rangeInt": RangeInt,
|
|
||||||
"easy float": Float,
|
|
||||||
"easy rangeFloat": RangeFloat,
|
|
||||||
"easy boolean": Boolean,
|
|
||||||
"easy compare": Compare,
|
|
||||||
"easy imageSwitch": imageSwitch,
|
|
||||||
"easy textSwitch": textSwitch,
|
|
||||||
"easy if": If,
|
|
||||||
"easy isSDXL": isSDXL,
|
|
||||||
"easy xyAny": xyAny,
|
|
||||||
"easy convertAnything": ConvertAnything,
|
|
||||||
"easy showAnything": showAnything,
|
|
||||||
"easy showTensorShape": showTensorShape,
|
|
||||||
"easy clearCacheKey": clearCacheKey,
|
|
||||||
"easy clearCacheAll": clearCacheAll,
|
|
||||||
"easy cleanGpuUsed": cleanGPUUsed,
|
|
||||||
}
|
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
|
||||||
"easy string": "String",
|
|
||||||
"easy int": "Int",
|
|
||||||
"easy rangeInt": "Range(Int)",
|
|
||||||
"easy float": "Float",
|
|
||||||
"easy rangeFloat": "Range(Float)",
|
|
||||||
"easy boolean": "Boolean",
|
|
||||||
"easy compare": "Compare",
|
|
||||||
"easy imageSwitch": "Image Switch",
|
|
||||||
"easy textSwitch": "Text Switch",
|
|
||||||
"easy if": "If",
|
|
||||||
"easy isSDXL": "Is SDXL",
|
|
||||||
"easy xyAny": "XYAny",
|
|
||||||
"easy convertAnything": "Convert Any",
|
|
||||||
"easy showAnything": "Show Any",
|
|
||||||
"easy showTensorShape": "Show Tensor Shape",
|
|
||||||
"easy clearCacheKey": "Clear Cache Key",
|
|
||||||
"easy clearCacheAll": "Clear Cache All",
|
|
||||||
"easy cleanGpuUsed": "Clean GPU Used"
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,167 @@
|
|||||||
|
#credit to comfyanonymous for this module
|
||||||
|
#from https://github.com/comfyanonymous/ComfyUI_bitsandbytes_NF4
|
||||||
|
import comfy.ops
|
||||||
|
import torch
|
||||||
|
import folder_paths
|
||||||
|
from ...libs.utils import install_package
|
||||||
|
|
||||||
|
try:
|
||||||
|
from bitsandbytes.nn.modules import Params4bit, QuantState
|
||||||
|
except ImportError:
|
||||||
|
Params4bit = torch.nn.Parameter
|
||||||
|
raise ImportError("Please install bitsandbytes>=0.43.3")
|
||||||
|
|
||||||
|
def functional_linear_4bits(x, weight, bias):
|
||||||
|
try:
|
||||||
|
install_package("bitsandbytes", "0.43.3", True, "0.43.3")
|
||||||
|
import bitsandbytes as bnb
|
||||||
|
except ImportError:
|
||||||
|
raise ImportError("Please install bitsandbytes>=0.43.3")
|
||||||
|
|
||||||
|
out = bnb.matmul_4bit(x, weight.t(), bias=bias, quant_state=weight.quant_state)
|
||||||
|
out = out.to(x)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def copy_quant_state(state, device: torch.device = None):
|
||||||
|
if state is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
device = device or state.absmax.device
|
||||||
|
|
||||||
|
state2 = (
|
||||||
|
QuantState(
|
||||||
|
absmax=state.state2.absmax.to(device),
|
||||||
|
shape=state.state2.shape,
|
||||||
|
code=state.state2.code.to(device),
|
||||||
|
blocksize=state.state2.blocksize,
|
||||||
|
quant_type=state.state2.quant_type,
|
||||||
|
dtype=state.state2.dtype,
|
||||||
|
)
|
||||||
|
if state.nested
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
return QuantState(
|
||||||
|
absmax=state.absmax.to(device),
|
||||||
|
shape=state.shape,
|
||||||
|
code=state.code.to(device),
|
||||||
|
blocksize=state.blocksize,
|
||||||
|
quant_type=state.quant_type,
|
||||||
|
dtype=state.dtype,
|
||||||
|
offset=state.offset.to(device) if state.nested else None,
|
||||||
|
state2=state2,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ForgeParams4bit(Params4bit):
|
||||||
|
|
||||||
|
def to(self, *args, **kwargs):
|
||||||
|
device, dtype, non_blocking, convert_to_format = torch._C._nn._parse_to(*args, **kwargs)
|
||||||
|
if device is not None and device.type == "cuda" and not self.bnb_quantized:
|
||||||
|
return self._quantize(device)
|
||||||
|
else:
|
||||||
|
n = ForgeParams4bit(
|
||||||
|
torch.nn.Parameter.to(self, device=device, dtype=dtype, non_blocking=non_blocking),
|
||||||
|
requires_grad=self.requires_grad,
|
||||||
|
quant_state=copy_quant_state(self.quant_state, device),
|
||||||
|
blocksize=self.blocksize,
|
||||||
|
compress_statistics=self.compress_statistics,
|
||||||
|
quant_type=self.quant_type,
|
||||||
|
quant_storage=self.quant_storage,
|
||||||
|
bnb_quantized=self.bnb_quantized,
|
||||||
|
module=self.module
|
||||||
|
)
|
||||||
|
self.module.quant_state = n.quant_state
|
||||||
|
self.data = n.data
|
||||||
|
self.quant_state = n.quant_state
|
||||||
|
return n
|
||||||
|
|
||||||
|
class ForgeLoader4Bit(torch.nn.Module):
|
||||||
|
def __init__(self, *, device, dtype, quant_type, **kwargs):
|
||||||
|
super().__init__()
|
||||||
|
self.dummy = torch.nn.Parameter(torch.empty(1, device=device, dtype=dtype))
|
||||||
|
self.weight = None
|
||||||
|
self.quant_state = None
|
||||||
|
self.bias = None
|
||||||
|
self.quant_type = quant_type
|
||||||
|
|
||||||
|
def _save_to_state_dict(self, destination, prefix, keep_vars):
|
||||||
|
super()._save_to_state_dict(destination, prefix, keep_vars)
|
||||||
|
quant_state = getattr(self.weight, "quant_state", None)
|
||||||
|
if quant_state is not None:
|
||||||
|
for k, v in quant_state.as_dict(packed=True).items():
|
||||||
|
destination[prefix + "weight." + k] = v if keep_vars else v.detach()
|
||||||
|
return
|
||||||
|
|
||||||
|
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):
|
||||||
|
quant_state_keys = {k[len(prefix + "weight."):] for k in state_dict.keys() if k.startswith(prefix + "weight.")}
|
||||||
|
|
||||||
|
if any('bitsandbytes' in k for k in quant_state_keys):
|
||||||
|
quant_state_dict = {k: state_dict[prefix + "weight." + k] for k in quant_state_keys}
|
||||||
|
|
||||||
|
self.weight = ForgeParams4bit().from_prequantized(
|
||||||
|
data=state_dict[prefix + 'weight'],
|
||||||
|
quantized_stats=quant_state_dict,
|
||||||
|
requires_grad=False,
|
||||||
|
device=self.dummy.device,
|
||||||
|
module=self
|
||||||
|
)
|
||||||
|
self.quant_state = self.weight.quant_state
|
||||||
|
|
||||||
|
if prefix + 'bias' in state_dict:
|
||||||
|
self.bias = torch.nn.Parameter(state_dict[prefix + 'bias'].to(self.dummy))
|
||||||
|
|
||||||
|
del self.dummy
|
||||||
|
elif hasattr(self, 'dummy'):
|
||||||
|
if prefix + 'weight' in state_dict:
|
||||||
|
self.weight = ForgeParams4bit(
|
||||||
|
state_dict[prefix + 'weight'].to(self.dummy),
|
||||||
|
requires_grad=False,
|
||||||
|
compress_statistics=True,
|
||||||
|
quant_type=self.quant_type,
|
||||||
|
quant_storage=torch.uint8,
|
||||||
|
module=self,
|
||||||
|
)
|
||||||
|
self.quant_state = self.weight.quant_state
|
||||||
|
|
||||||
|
if prefix + 'bias' in state_dict:
|
||||||
|
self.bias = torch.nn.Parameter(state_dict[prefix + 'bias'].to(self.dummy))
|
||||||
|
|
||||||
|
del self.dummy
|
||||||
|
else:
|
||||||
|
super()._load_from_state_dict(state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs)
|
||||||
|
|
||||||
|
current_device = None
|
||||||
|
current_dtype = None
|
||||||
|
current_manual_cast_enabled = False
|
||||||
|
current_bnb_dtype = None
|
||||||
|
|
||||||
|
class OPS(comfy.ops.manual_cast):
|
||||||
|
class Linear(ForgeLoader4Bit):
|
||||||
|
def __init__(self, *args, device=None, dtype=None, **kwargs):
|
||||||
|
super().__init__(device=device, dtype=dtype, quant_type=current_bnb_dtype)
|
||||||
|
self.parameters_manual_cast = current_manual_cast_enabled
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
self.weight.quant_state = self.quant_state
|
||||||
|
|
||||||
|
if self.bias is not None and self.bias.dtype != x.dtype:
|
||||||
|
# Maybe this can also be set to all non-bnb ops since the cost is very low.
|
||||||
|
# And it only invokes one time, and most linear does not have bias
|
||||||
|
self.bias.data = self.bias.data.to(x.dtype)
|
||||||
|
|
||||||
|
if not self.parameters_manual_cast:
|
||||||
|
return functional_linear_4bits(x, self.weight, self.bias)
|
||||||
|
elif not self.weight.bnb_quantized:
|
||||||
|
assert x.device.type == 'cuda', 'BNB Must Use CUDA as Computation Device!'
|
||||||
|
layer_original_device = self.weight.device
|
||||||
|
self.weight = self.weight._quantize(x.device)
|
||||||
|
bias = self.bias.to(x.device) if self.bias is not None else None
|
||||||
|
out = functional_linear_4bits(x, self.weight, bias)
|
||||||
|
self.weight = self.weight.to(layer_original_device)
|
||||||
|
return out
|
||||||
|
else:
|
||||||
|
weight, bias, signal = weights_manual_cast(self, x, skip_weight_dtype=True, skip_bias_dtype=True)
|
||||||
|
with main_stream_worker(weight, bias, signal):
|
||||||
|
return functional_linear_4bits(x, weight, bias)
|
||||||
@@ -5,13 +5,21 @@ import os
|
|||||||
import types
|
import types
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from accelerate import init_empty_weights, load_checkpoint_and_dispatch
|
try:
|
||||||
|
from accelerate import init_empty_weights, load_checkpoint_and_dispatch
|
||||||
|
except:
|
||||||
|
init_empty_weights, load_checkpoint_and_dispatch = None, None
|
||||||
|
|
||||||
import comfy
|
import comfy
|
||||||
|
|
||||||
from .model import BrushNetModel, PowerPaintModel
|
try:
|
||||||
from .model_patch import add_model_patch_option, patch_model_function_wrapper
|
from .model import BrushNetModel, PowerPaintModel
|
||||||
from .powerpaint_utils import TokenizerWrapper, add_tokens
|
from .model_patch import add_model_patch_option, patch_model_function_wrapper
|
||||||
|
from .powerpaint_utils import TokenizerWrapper, add_tokens
|
||||||
|
except:
|
||||||
|
BrushNetModel, PowerPaintModel = None, None
|
||||||
|
add_model_patch_option, patch_model_function_wrapper = None, None
|
||||||
|
TokenizerWrapper, add_tokens = None, None
|
||||||
|
|
||||||
cwd_path = os.path.dirname(os.path.realpath(__file__))
|
cwd_path = os.path.dirname(os.path.realpath(__file__))
|
||||||
brushnet_config_file = os.path.join(cwd_path, 'config', 'brushnet.json')
|
brushnet_config_file = os.path.join(cwd_path, 'config', 'brushnet.json')
|
||||||
@@ -272,11 +280,11 @@ class BrushNet:
|
|||||||
|
|
||||||
# unload vae
|
# unload vae
|
||||||
del vae
|
del vae
|
||||||
for loaded_model in comfy.model_management.current_loaded_models:
|
# for loaded_model in comfy.model_management.current_loaded_models:
|
||||||
if type(loaded_model.model.model) in ModelsToUnload:
|
# if type(loaded_model.model.model) in ModelsToUnload:
|
||||||
comfy.model_management.current_loaded_models.remove(loaded_model)
|
# comfy.model_management.current_loaded_models.remove(loaded_model)
|
||||||
loaded_model.model_unload()
|
# loaded_model.model_unload()
|
||||||
del loaded_model
|
# del loaded_model
|
||||||
|
|
||||||
# prepare embeddings
|
# prepare embeddings
|
||||||
prompt_embeds = positive[0][0].to(dtype=torch_dtype).to(brushnet['brushnet'].device)
|
prompt_embeds = positive[0][0].to(dtype=torch_dtype).to(brushnet['brushnet'].device)
|
||||||
@@ -449,11 +457,11 @@ class BrushNet:
|
|||||||
# unload vae and CLIPs
|
# unload vae and CLIPs
|
||||||
del vae
|
del vae
|
||||||
del clip
|
del clip
|
||||||
for loaded_model in comfy.model_management.current_loaded_models:
|
# for loaded_model in comfy.model_management.current_loaded_models:
|
||||||
if type(loaded_model.model.model) in ModelsToUnload:
|
# if type(loaded_model.model.model) in ModelsToUnload:
|
||||||
comfy.model_management.current_loaded_models.remove(loaded_model)
|
# comfy.model_management.current_loaded_models.remove(loaded_model)
|
||||||
loaded_model.model_unload()
|
# loaded_model.model_unload()
|
||||||
del loaded_model
|
# del loaded_model
|
||||||
|
|
||||||
# apply patch to model
|
# apply patch to model
|
||||||
|
|
||||||
@@ -663,8 +671,16 @@ def add_brushnet_patch(model, brushnet, torch_dtype, conditioning_latents,
|
|||||||
|
|
||||||
is_SDXL = isinstance(model.model.model_config, comfy.supported_models.SDXL)
|
is_SDXL = isinstance(model.model.model_config, comfy.supported_models.SDXL)
|
||||||
|
|
||||||
|
if model.model.model_config.custom_operations is None:
|
||||||
|
fp8 = model.model.model_config.optimizations.get("fp8", model.model.model_config.scaled_fp8 is not None)
|
||||||
|
operations = comfy.ops.pick_operations(model.model.model_config.unet_config.get("dtype", None), model.model.manual_cast_dtype,
|
||||||
|
fp8_optimizations=fp8, scaled_fp8=model.model.model_config.scaled_fp8)
|
||||||
|
else:
|
||||||
|
# such as gguf
|
||||||
|
operations = model.model.model_config.custom_operations
|
||||||
|
|
||||||
if is_SDXL:
|
if is_SDXL:
|
||||||
input_blocks = [[0, comfy.ops.disable_weight_init.Conv2d],
|
input_blocks = [[0, operations.Conv2d],
|
||||||
[1, comfy.ldm.modules.diffusionmodules.openaimodel.ResBlock],
|
[1, comfy.ldm.modules.diffusionmodules.openaimodel.ResBlock],
|
||||||
[2, comfy.ldm.modules.diffusionmodules.openaimodel.ResBlock],
|
[2, comfy.ldm.modules.diffusionmodules.openaimodel.ResBlock],
|
||||||
[3, comfy.ldm.modules.diffusionmodules.openaimodel.Downsample],
|
[3, comfy.ldm.modules.diffusionmodules.openaimodel.Downsample],
|
||||||
@@ -686,7 +702,7 @@ def add_brushnet_patch(model, brushnet, torch_dtype, conditioning_latents,
|
|||||||
[7, comfy.ldm.modules.diffusionmodules.openaimodel.ResBlock],
|
[7, comfy.ldm.modules.diffusionmodules.openaimodel.ResBlock],
|
||||||
[8, comfy.ldm.modules.diffusionmodules.openaimodel.ResBlock]]
|
[8, comfy.ldm.modules.diffusionmodules.openaimodel.ResBlock]]
|
||||||
else:
|
else:
|
||||||
input_blocks = [[0, comfy.ops.disable_weight_init.Conv2d],
|
input_blocks = [[0, operations.Conv2d],
|
||||||
[1, comfy.ldm.modules.attention.SpatialTransformer],
|
[1, comfy.ldm.modules.attention.SpatialTransformer],
|
||||||
[2, comfy.ldm.modules.attention.SpatialTransformer],
|
[2, comfy.ldm.modules.attention.SpatialTransformer],
|
||||||
[3, comfy.ldm.modules.diffusionmodules.openaimodel.Downsample],
|
[3, comfy.ldm.modules.diffusionmodules.openaimodel.Downsample],
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user