Compare commits
394 Commits
developmen
...
wavelet_lo
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ce4c5291a0 | ||
|
|
c101f07834 | ||
|
|
e4526ad4a4 | ||
|
|
4595965e06 | ||
|
|
41edc18750 | ||
|
|
6021a3dbc0 | ||
|
|
71d7a52146 | ||
|
|
45be82d5d6 | ||
|
|
f10937e6da | ||
|
|
ccb66c748f | ||
|
|
2aca2883e7 | ||
|
|
1ad58c5816 | ||
|
|
6dea41b9fc | ||
|
|
9a902c067f | ||
|
|
0bbc69c135 | ||
|
|
6c5eb0cf87 | ||
|
|
e3373671b9 | ||
|
|
aceb3a0f25 | ||
|
|
c8049a483d | ||
|
|
f5aa4232fa | ||
|
|
3a6b24f4c8 | ||
|
|
bbfd6ef0fe | ||
|
|
b829983b16 | ||
|
|
fa187b1208 | ||
|
|
5eb627dd9d | ||
|
|
604e76d34d | ||
|
|
6cde96ae5f | ||
|
|
1be613ed06 | ||
|
|
c52421aab7 | ||
|
|
3812957bc9 | ||
|
|
391329dbdc | ||
|
|
3b45892b4f | ||
|
|
cf4216e6b8 | ||
|
|
31e057d9a3 | ||
|
|
d507b44a7b | ||
|
|
242c04a0b8 | ||
|
|
386e68a422 | ||
|
|
850b8da6e5 | ||
|
|
51ad19b568 | ||
|
|
e6739f7eb2 | ||
|
|
7e37918fbc | ||
|
|
4d88f8f218 | ||
|
|
25341c4613 | ||
|
|
391cf80fea | ||
|
|
4e3bda7c70 | ||
|
|
763128ea42 | ||
|
|
4fe33f51c1 | ||
|
|
aa44828c0c | ||
|
|
6f6fb90812 | ||
|
|
c57434ad7b | ||
|
|
8bb47d1bfe | ||
|
|
e7dbb20f68 | ||
|
|
c5e0c2bbe2 | ||
|
|
1f3f45a48d | ||
|
|
3c8c84f156 | ||
|
|
b001d77efb | ||
|
|
7ae31c9ae9 | ||
|
|
b16819f8e7 | ||
|
|
f5e40dfa62 | ||
|
|
acc79956aa | ||
|
|
60539c0b0f | ||
|
|
dd700f70b3 | ||
|
|
d360e76661 | ||
|
|
6ec23ed226 | ||
|
|
f6e16e582a | ||
|
|
259ded9602 | ||
|
|
440ba5fb3d | ||
|
|
093f14ac19 | ||
|
|
f0fbd8bb53 | ||
|
|
0a981bea2b | ||
|
|
1d0e3a4498 | ||
|
|
3c7daf49f3 | ||
|
|
56d8d6bd81 | ||
|
|
3e49337a58 | ||
|
|
60f848a877 | ||
|
|
b366e46f1c | ||
|
|
a280f78c69 | ||
|
|
6e19e7449e | ||
|
|
a6d46ad9ae | ||
|
|
f3725578dd | ||
|
|
ed99c3c0c8 | ||
|
|
ed84c19205 | ||
|
|
a7a9c11d9e | ||
|
|
f60698d0ee | ||
|
|
5f094fb17a | ||
|
|
a5227cba7b | ||
|
|
77a5e01301 | ||
|
|
4ef5a668c0 | ||
|
|
f081d14527 | ||
|
|
710c6de1c9 | ||
|
|
2b6e66e0cb | ||
|
|
ab641e014f | ||
|
|
ad87f72384 | ||
|
|
d0214c0df9 | ||
|
|
adcf884c0f | ||
|
|
f778d979b5 | ||
|
|
db3ccbba33 | ||
|
|
0d2be18a9b | ||
|
|
bbc340e545 | ||
|
|
33fdfd6091 | ||
|
|
9f6030620f | ||
|
|
b5252b5028 | ||
|
|
b0d8fc220d | ||
|
|
cef7d9e594 | ||
|
|
b13fcc1039 | ||
|
|
b32d7e552b | ||
|
|
4af6c5cf30 | ||
|
|
1f7784510d | ||
|
|
87e557cf1e | ||
|
|
bd8d7dc081 | ||
|
|
2be6926398 | ||
|
|
87ac031859 | ||
|
|
7679105d52 | ||
|
|
2622de1e01 | ||
|
|
8450aca10e | ||
|
|
0b8a32def7 | ||
|
|
787bb37e76 | ||
|
|
10aa7e9d5e | ||
|
|
ed1deb71c4 | ||
|
|
4de6a825fa | ||
|
|
9a7266275d | ||
|
|
d138f07365 | ||
|
|
c6d8eedb94 | ||
|
|
af5e760be1 | ||
|
|
ff3d54bb5b | ||
|
|
0e75724b4d | ||
|
|
376bb1bf6f | ||
|
|
216ab164ce | ||
|
|
e6180d1e1d | ||
|
|
15a57bc89f | ||
|
|
e5355bf8d5 | ||
|
|
34a1c6947a | ||
|
|
2141c6e06c | ||
|
|
1188cf1e8a | ||
|
|
5e663746b8 | ||
|
|
441474e81f | ||
|
|
a6a690f796 | ||
|
|
6191f19e55 | ||
|
|
bbfba0c188 | ||
|
|
e1549ad54d | ||
|
|
04abe57c76 | ||
|
|
89dd041b97 | ||
|
|
29122b1a54 | ||
|
|
6a8e3d8610 | ||
|
|
4c8a9e1b88 | ||
|
|
fadb2f3a76 | ||
|
|
4723f23c0d | ||
|
|
8ef07a9c36 | ||
|
|
92ce93140e | ||
|
|
f213996aa5 | ||
|
|
cbe31eaf0a | ||
|
|
67c2e44edb | ||
|
|
96d418bb95 | ||
|
|
894374b2e9 | ||
|
|
6509ba4484 | ||
|
|
025ee3dd3d | ||
|
|
58f9d01c2b | ||
|
|
e72b59a8e9 | ||
|
|
4aa19b5c1d | ||
|
|
4747716867 | ||
|
|
22cd40d7b9 | ||
|
|
3400882a80 | ||
|
|
9f94c7b61e | ||
|
|
bedb8197a2 | ||
|
|
e3ebd73610 | ||
|
|
dd931757cd | ||
|
|
0640cdf569 | ||
|
|
0b048d0dde | ||
|
|
473d455f44 | ||
|
|
ce759ebd8c | ||
|
|
628a7923a3 | ||
|
|
3922981996 | ||
|
|
ab22674980 | ||
|
|
9452929300 | ||
|
|
a800c9d19e | ||
|
|
28e6f00790 | ||
|
|
67e0aca750 | ||
|
|
f05224970f | ||
|
|
b4f64de4c2 | ||
|
|
2e5f6668dc | ||
|
|
e4c82803e1 | ||
|
|
69aa92bce5 | ||
|
|
a508caad1d | ||
|
|
58537fc92b | ||
|
|
86b5938cf3 | ||
|
|
6b4034122f | ||
|
|
10817696fb | ||
|
|
037ce11740 | ||
|
|
04424fe2d6 | ||
|
|
40a8ff5731 | ||
|
|
2776221497 | ||
|
|
f85ad452c6 | ||
|
|
dd889086f4 | ||
|
|
bc693488eb | ||
|
|
d97c55cd96 | ||
|
|
79b4e04b80 | ||
|
|
951e223481 | ||
|
|
fc34a69bec | ||
|
|
279ee65177 | ||
|
|
3a1f464132 | ||
|
|
5c8fcc8a4e | ||
|
|
121a760c19 | ||
|
|
e5fadddd45 | ||
|
|
d44d4eb61a | ||
|
|
7d9ab22405 | ||
|
|
7ed8c51f20 | ||
|
|
6df33156f0 | ||
|
|
40f5c59da0 | ||
|
|
3e71a99df0 | ||
|
|
562405923f | ||
|
|
f84bd6d7a6 | ||
|
|
60232def91 | ||
|
|
3843e0d148 | ||
|
|
e127c079da | ||
|
|
34db804c76 | ||
|
|
4d35a29c97 | ||
|
|
b322d05fa3 | ||
|
|
8577849eeb | ||
|
|
338c77d677 | ||
|
|
e07a98a50c | ||
|
|
6a754b2710 | ||
|
|
a939cf3730 | ||
|
|
169dbd22ba | ||
|
|
6e7d721382 | ||
|
|
dc6f36cd82 | ||
|
|
5603f9e004 | ||
|
|
c45887192a | ||
|
|
13a965a26c | ||
|
|
77ee7090e8 | ||
|
|
078396ceac | ||
|
|
f944eeaa4d | ||
|
|
81899310f8 | ||
|
|
f9179540d2 | ||
|
|
452e0e286d | ||
|
|
165510ace2 | ||
|
|
0355662e8e | ||
|
|
b99d36dfdb | ||
|
|
9001e5c933 | ||
|
|
7fed4ea761 | ||
|
|
e07bf11727 | ||
|
|
c728cc9a0b | ||
|
|
00bd3d54a3 | ||
|
|
f7cf2f866f | ||
|
|
465bc1e2f8 | ||
|
|
0beca0d4a7 | ||
|
|
418f5f7e8c | ||
|
|
9ee1ef2a0a | ||
|
|
599fafe01f | ||
|
|
af108bb964 | ||
|
|
89d61a3b8e | ||
|
|
a6aa4b2c7d | ||
|
|
f8f0657b68 | ||
|
|
7f0ecdb377 | ||
|
|
fbed8568fb | ||
|
|
6d31c6db73 | ||
|
|
6490a326e5 | ||
|
|
8d48ad4e85 | ||
|
|
ec1ea7aa0e | ||
|
|
fa02e774b0 | ||
|
|
2308ef2868 | ||
|
|
b3e03295ad | ||
|
|
e69a520616 | ||
|
|
acafe9984f | ||
|
|
653fe60f16 | ||
|
|
c2424087d6 | ||
|
|
272c8608c2 | ||
|
|
99f24cfb0c | ||
|
|
187663ab55 | ||
|
|
edb7e827ee | ||
|
|
0ea27011d5 | ||
|
|
f321de7bdb | ||
|
|
88acc28d7f | ||
|
|
de2da96a81 | ||
|
|
9beea1c268 | ||
|
|
369aa143bc | ||
|
|
87ba867fdc | ||
|
|
03613c523f | ||
|
|
47744373f2 | ||
|
|
443c996e7f | ||
|
|
8f0f467c20 | ||
|
|
e81e19fd0f | ||
|
|
0bc4d555c7 | ||
|
|
80aa2dbb80 | ||
|
|
8d799031cf | ||
|
|
6e92922c14 | ||
|
|
c51235c486 | ||
|
|
c2c4e8cf34 | ||
|
|
4c249cf607 | ||
|
|
c2d5f712a3 | ||
|
|
22d2f6e28f | ||
|
|
a2301cf28c | ||
|
|
11e426fdf1 | ||
|
|
58dffd43a8 | ||
|
|
e4558dff4b | ||
|
|
c062b7716c | ||
|
|
c008405480 | ||
|
|
93e5df1d59 | ||
|
|
045e4a6e15 | ||
|
|
76f225a467 | ||
|
|
cab8a1c7b8 | ||
|
|
acb06d6ff3 | ||
|
|
bb57623a35 | ||
|
|
3072d20f17 | ||
|
|
f6b21f47bb | ||
|
|
603ceca3ca | ||
|
|
657fd09f25 | ||
|
|
8407c4deea | ||
|
|
64f2b085b7 | ||
|
|
7165f2d25a | ||
|
|
5d47244c57 | ||
|
|
ada722c9e4 | ||
|
|
696f73c30d | ||
|
|
e3410413b9 | ||
|
|
37cebd9458 | ||
|
|
bd10d2d668 | ||
|
|
cb5d28cba9 | ||
|
|
3f3636b788 | ||
|
|
833c833f28 | ||
|
|
68b7e159bc | ||
|
|
5a45c709cd | ||
|
|
10e1ecf1e8 | ||
|
|
b96913d73c | ||
|
|
5da3613e0b | ||
|
|
5a70b7f38d | ||
|
|
377b81ee3e | ||
|
|
2d0a1be59d | ||
|
|
7284aab7c0 | ||
|
|
427847ac4c | ||
|
|
9c1cc9641e | ||
|
|
89f4bcad2e | ||
|
|
016687bda1 | ||
|
|
72de68d8aa | ||
|
|
d87b49882c | ||
|
|
f415bac7b5 | ||
|
|
f1cb87fe9e | ||
|
|
8f9cd823d1 | ||
|
|
b01e8d889a | ||
|
|
1325613583 | ||
|
|
337945de9a | ||
|
|
561914d8e6 | ||
|
|
b0a0f28191 | ||
|
|
f965a1299f | ||
|
|
1bd94f0f01 | ||
|
|
9ffa8c3711 | ||
|
|
b68c3ef734 | ||
|
|
49c41e6a5f | ||
|
|
2478554c95 | ||
|
|
93b52932c1 | ||
|
|
4ec4025cbb | ||
|
|
e074058faa | ||
|
|
a8481c1670 | ||
|
|
e18e0cb5f8 | ||
|
|
177c7130ec | ||
|
|
1ae1017748 | ||
|
|
92b9c71d44 | ||
|
|
f17ad8d794 | ||
|
|
86c70a2a1f | ||
|
|
655533d4c7 | ||
|
|
eebd3c8212 | ||
|
|
5276975fb0 | ||
|
|
e190fbaeb8 | ||
|
|
290393f7ae | ||
|
|
b2a54c8f36 | ||
|
|
b767d29b3c | ||
|
|
645b27f97a | ||
|
|
65c08b09c3 | ||
|
|
afc231efc1 | ||
|
|
bafacf3b65 | ||
|
|
0892dec4a5 | ||
|
|
eeee4a1620 | ||
|
|
d11ed7f66c | ||
|
|
27ad79053e | ||
|
|
05ae95ca89 | ||
|
|
0f8daa5612 | ||
|
|
7703e3a15e | ||
|
|
0f597f453e | ||
|
|
dfb64b5957 | ||
|
|
82098e5d6e | ||
|
|
b653906715 | ||
|
|
13d32423f6 | ||
|
|
39870411d8 | ||
|
|
e5177833b2 | ||
|
|
eaa0fb6253 | ||
|
|
eaec2f5a52 | ||
|
|
92cb5ae096 | ||
|
|
bd2bce9b92 | ||
|
|
537af79b0d | ||
|
|
7624241032 | ||
|
|
0d5943af91 | ||
|
|
be815f9c47 | ||
|
|
3443d6aafa | ||
|
|
bef10a639c | ||
|
|
3eb3535683 | ||
|
|
3d387103cd |
2
.github/FUNDING.yml
vendored
Normal file
2
.github/FUNDING.yml
vendored
Normal file
@@ -0,0 +1,2 @@
|
||||
github: [ostris]
|
||||
patreon: ostris
|
||||
19
.github/ISSUE_TEMPLATE/bug_report.md
vendored
Normal file
19
.github/ISSUE_TEMPLATE/bug_report.md
vendored
Normal file
@@ -0,0 +1,19 @@
|
||||
---
|
||||
name: Bug Report
|
||||
about: For bugs only. Not for feature requests or questions.
|
||||
title: ''
|
||||
labels: ''
|
||||
assignees: ''
|
||||
---
|
||||
|
||||
## This is for bugs only
|
||||
|
||||
Did you already ask [in the discord](https://discord.gg/VXmU2f5WEU)?
|
||||
|
||||
Yes/No
|
||||
|
||||
You verified that this is a bug and not a feature request or question by asking [in the discord](https://discord.gg/VXmU2f5WEU)?
|
||||
|
||||
Yes/No
|
||||
|
||||
## Describe the bug
|
||||
5
.github/ISSUE_TEMPLATE/config.yml
vendored
Normal file
5
.github/ISSUE_TEMPLATE/config.yml
vendored
Normal file
@@ -0,0 +1,5 @@
|
||||
blank_issues_enabled: false
|
||||
contact_links:
|
||||
- name: Ask in the Discord BEFORE opening an issue
|
||||
url: https://discord.gg/VXmU2f5WEU
|
||||
about: Please ask in the discord before opening a github issue.
|
||||
9
.gitignore
vendored
9
.gitignore
vendored
@@ -161,6 +161,7 @@ cython_debug/
|
||||
|
||||
/env.sh
|
||||
/models
|
||||
/datasets
|
||||
/custom/*
|
||||
!/custom/.gitkeep
|
||||
/.tmp
|
||||
@@ -172,4 +173,10 @@ cython_debug/
|
||||
/output/*
|
||||
!/output/.gitkeep
|
||||
/extensions/*
|
||||
!/extensions/example
|
||||
!/extensions/example
|
||||
/temp
|
||||
/wandb
|
||||
.vscode/settings.json
|
||||
.DS_Store
|
||||
._.DS_Store
|
||||
aitk_db.db
|
||||
4
.gitmodules
vendored
4
.gitmodules
vendored
@@ -1,12 +1,16 @@
|
||||
[submodule "repositories/sd-scripts"]
|
||||
path = repositories/sd-scripts
|
||||
url = https://github.com/kohya-ss/sd-scripts.git
|
||||
commit = b78c0e2a69e52ce6c79abc6c8c82d1a9cabcf05c
|
||||
[submodule "repositories/leco"]
|
||||
path = repositories/leco
|
||||
url = https://github.com/p1atdev/LECO
|
||||
commit = 9294adf40218e917df4516737afb13f069a6789d
|
||||
[submodule "repositories/batch_annotator"]
|
||||
path = repositories/batch_annotator
|
||||
url = https://github.com/ostris/batch-annotator
|
||||
commit = 420e142f6ad3cc14b3ea0500affc2c6c7e7544bf
|
||||
[submodule "repositories/ipadapter"]
|
||||
path = repositories/ipadapter
|
||||
url = https://github.com/tencent-ailab/IP-Adapter.git
|
||||
commit = 5a18b1f3660acaf8bee8250692d6fb3548a19b14
|
||||
|
||||
28
.vscode/launch.json
vendored
Normal file
28
.vscode/launch.json
vendored
Normal file
@@ -0,0 +1,28 @@
|
||||
{
|
||||
"version": "0.2.0",
|
||||
"configurations": [
|
||||
{
|
||||
"name": "Run current config",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${workspaceFolder}/run.py",
|
||||
"args": [
|
||||
"${file}"
|
||||
],
|
||||
"env": {
|
||||
"CUDA_LAUNCH_BLOCKING": "1",
|
||||
"DEBUG_TOOLKIT": "1"
|
||||
},
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false
|
||||
},
|
||||
{
|
||||
"name": "Python: Debug Current File",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${file}",
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false
|
||||
},
|
||||
]
|
||||
}
|
||||
10
FAQ.md
Normal file
10
FAQ.md
Normal file
@@ -0,0 +1,10 @@
|
||||
# FAQ
|
||||
|
||||
WIP. Will continue to add things as they are needed.
|
||||
|
||||
## FLUX.1 Training
|
||||
|
||||
#### How much VRAM is required to train a lora on FLUX.1?
|
||||
|
||||
24GB minimum is required.
|
||||
|
||||
21
LICENSE
Normal file
21
LICENSE
Normal file
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2024 Ostris, LLC
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
40
assets/glif.svg
Normal file
40
assets/glif.svg
Normal file
@@ -0,0 +1,40 @@
|
||||
<svg width="148" height="66" viewBox="0 0 148 66" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<rect opacity="0.3" width="148" height="66" rx="33" fill="#030F2F"/>
|
||||
<g filter="url(#filter0_d_10631_12135)">
|
||||
<path d="M48.8305 21.013H43.5433V53.7839H48.8305V21.013Z" fill="white"/>
|
||||
<path d="M57.0987 32.7729H51.8115V53.7835H57.0987V32.7729Z" fill="white"/>
|
||||
<path d="M73.5495 36.4837L69.6034 32.5067L65.6573 36.4837L69.6034 40.4607L73.5495 36.4837Z" fill="white"/>
|
||||
<path d="M58.4255 24.7602L54.4794 20.7832L50.5333 24.7602L54.4794 28.7372L58.4255 24.7602Z" fill="white"/>
|
||||
<path d="M40.5557 21.0118H35.2685V24.185C33.9942 23.5456 32.5588 23.1832 31.0387 23.1832C25.7911 23.1806 21.5217 27.4834 21.5217 32.7721C21.5217 35.082 22.336 37.2054 23.6921 38.8626H21.5217V44.1912H26.8089V41.3618C28.0832 42.0012 29.5186 42.3635 31.0387 42.3635C36.2863 42.3635 40.5557 38.0607 40.5557 32.7721C40.5557 30.2996 39.6225 28.0429 38.0918 26.3404H40.5557V21.0118ZM31.0387 37.0349C28.707 37.0349 26.8089 35.122 26.8089 32.7721C26.8089 30.4221 28.707 28.5092 31.0387 28.5092C33.3704 28.5092 35.2685 30.4221 35.2685 32.7721C35.2685 35.122 33.3704 37.0349 31.0387 37.0349Z" fill="white"/>
|
||||
<path d="M31.0381 44.1912H26.8083V49.5198H31.0381C33.3697 49.5198 35.2678 51.4327 35.2678 53.7826H40.555C40.555 48.494 36.2856 44.1912 31.0381 44.1912Z" fill="white"/>
|
||||
<path d="M69.6041 26.3416C71.1295 26.3416 72.6073 27.1702 73.3871 28.6968L73.596 28.4864L77.2098 24.8443C75.4253 22.4544 72.6126 21.013 69.6041 21.013C64.4438 21.013 60.2352 25.172 60.0951 30.338L60.0872 53.7839H65.3744V30.6018C65.3744 28.2519 67.2725 26.3389 69.6041 26.3389V26.3416Z" fill="white"/>
|
||||
<path d="M73.5495 36.4837L69.6034 32.5067L65.6573 36.4837L69.6034 40.4607L73.5495 36.4837Z" fill="white"/>
|
||||
<path d="M120.022 53.8259H117.218V32.6354H120.022V35.2219C121.02 33.321 122.702 32.2615 125.102 32.2615C129.371 32.2615 131.397 35.9698 131.397 40.5819C131.397 45.1939 129.371 48.9023 125.102 48.9023C122.702 48.9023 121.02 47.8427 120.022 45.9418V53.8259ZM120.022 39.2419V41.9219C120.022 44.6642 121.581 46.6586 124.385 46.6586C126.722 46.6586 128.436 45.1939 128.436 43.0437V38.12C128.436 35.9698 126.722 34.5052 124.385 34.5052C121.581 34.5052 120.022 36.4996 120.022 39.2419Z" fill="white"/>
|
||||
<path d="M103.267 53.8259H100.463V32.6354H103.267V35.2219C104.265 33.321 105.947 32.2615 108.347 32.2615C112.616 32.2615 114.642 35.9698 114.642 40.5819C114.642 45.1939 112.616 48.9023 108.347 48.9023C105.947 48.9023 104.265 47.8427 103.267 45.9418V53.8259ZM103.267 39.2419V41.9219C103.267 44.6642 104.826 46.6586 107.63 46.6586C109.967 46.6586 111.681 45.1939 111.681 43.0437V38.12C111.681 35.9698 109.967 34.5052 107.63 34.5052C104.826 34.5052 103.267 36.4996 103.267 39.2419Z" fill="white"/>
|
||||
<path d="M87.7844 48.9023C86.2262 48.9023 84.8862 48.4037 83.9825 47.4688C83.1723 46.6274 82.6737 45.4121 82.6737 44.1656C82.6737 41.6726 84.263 39.6782 87.4104 39.3977L92.5834 38.9303V37.6526C92.5834 35.1907 91.2434 34.5052 89.2802 34.5052C87.3169 34.5052 86.0081 35.3466 86.0081 37.3721H83.2035C83.2035 34.3805 85.9146 32.2615 89.3113 32.2615C92.7392 32.2615 95.3257 33.9442 95.3257 37.6838V46.1288H97.694V48.5283H92.9573L92.895 45.599H92.7704C91.8978 47.8427 90.1527 48.9023 87.7844 48.9023ZM88.5011 46.6586C91.0253 46.6586 92.5834 44.6642 92.5834 41.7972V41.0805L86.943 41.6102C85.79 41.7037 85.6341 42.2335 85.6341 43.1372V44.5395C85.6341 46.0041 86.7248 46.6586 88.5011 46.6586Z" fill="white"/>
|
||||
<path d="M80.5279 45.1002V48.5281H77.1V45.1002H80.5279Z" fill="white"/>
|
||||
</g>
|
||||
<path d="M102.683 12.3401L100.874 9.09521H101.922L102.976 11.0436C103.015 11.1169 103.05 11.1852 103.079 11.2487C103.108 11.3122 103.138 11.3757 103.167 11.4392C103.191 11.3952 103.211 11.3537 103.225 11.3147C103.24 11.2756 103.257 11.2365 103.277 11.1975C103.301 11.1535 103.328 11.1022 103.357 11.0436L104.405 9.09521H105.423L103.621 12.3401V14.4497H102.683V12.3401Z" fill="white"/>
|
||||
<path d="M97.1749 9.09521V14.4497H96.2373V9.09521H97.1749ZM98.3615 12.1717H96.8892V11.3806H98.3102C98.5691 11.3806 98.7668 11.3171 98.9036 11.1901C99.0403 11.0583 99.1087 10.8727 99.1087 10.6334C99.1087 10.4039 99.0378 10.2281 98.8962 10.106C98.7546 9.98397 98.5495 9.92292 98.2809 9.92292H96.8599V9.09521H98.3615C98.8938 9.09521 99.3113 9.22462 99.6141 9.48343C99.9168 9.74224 100.068 10.0963 100.068 10.5455C100.068 10.8678 99.9901 11.1389 99.8338 11.3586C99.6776 11.5735 99.4456 11.7297 99.1379 11.8274V11.7248C99.47 11.803 99.7215 11.9495 99.8924 12.1643C100.063 12.3792 100.149 12.6575 100.149 12.9994C100.149 13.3021 100.08 13.5634 99.9437 13.7831C99.807 13.998 99.6067 14.164 99.343 14.2812C99.0842 14.3935 98.7717 14.4497 98.4055 14.4497H96.8599V13.622H98.3615C98.6301 13.622 98.8352 13.5585 98.9768 13.4315C99.1184 13.3046 99.1892 13.1215 99.1892 12.8822C99.1892 12.6575 99.116 12.4842 98.9695 12.3621C98.8279 12.2351 98.6252 12.1717 98.3615 12.1717Z" fill="white"/>
|
||||
<path d="M89.5954 14.4497H87.6689V9.09521H89.5441C90.0715 9.09521 90.5354 9.20997 90.9358 9.43948C91.3363 9.66411 91.6488 9.97908 91.8734 10.3844C92.1029 10.7848 92.2177 11.2512 92.2177 11.7834C92.2177 12.3059 92.1054 12.7699 91.8807 13.1752C91.661 13.5756 91.3534 13.8881 90.9578 14.1128C90.5671 14.3374 90.113 14.4497 89.5954 14.4497ZM88.6065 9.52738V14.0249L88.1597 13.5854H89.5075C89.864 13.5854 90.1716 13.5121 90.4304 13.3656C90.6892 13.2191 90.887 13.0116 91.0237 12.743C91.1605 12.4744 91.2288 12.1546 91.2288 11.7834C91.2288 11.4025 91.158 11.0778 91.0164 10.8092C90.8748 10.5358 90.6721 10.3258 90.4084 10.1793C90.1447 10.0328 89.8273 9.95955 89.4562 9.95955H88.1597L88.6065 9.52738Z" fill="white"/>
|
||||
<path d="M86.0735 14.4497H82.748V9.09521H86.0735V9.95955H83.356L83.6856 9.65923V11.3366H85.8245V12.1643H83.6856V13.8857L83.356 13.5854H86.0735V14.4497Z" fill="white"/>
|
||||
<path d="M78.1926 14.4497H77.255V9.09521H79.2986C79.9042 9.09521 80.3754 9.24171 80.7123 9.53471C81.0542 9.8277 81.2251 10.2379 81.2251 10.7653C81.2251 11.1218 81.1421 11.427 80.976 11.6809C80.8149 11.9299 80.5756 12.1204 80.2582 12.2522L81.2764 14.4497H80.2509L79.3426 12.45H78.1926V14.4497ZM78.1926 9.93025V11.6223H79.2986C79.5965 11.6223 79.8285 11.5466 79.9945 11.3952C80.1605 11.2438 80.2436 11.0339 80.2436 10.7653C80.2436 10.4967 80.1605 10.2916 79.9945 10.15C79.8285 10.0035 79.5965 9.93025 79.2986 9.93025H78.1926Z" fill="white"/>
|
||||
<path d="M75.789 11.7688C75.789 12.3108 75.6792 12.7918 75.4594 13.2118C75.2397 13.6269 74.9345 13.9516 74.5438 14.186C74.1531 14.4204 73.7014 14.5376 73.1887 14.5376C72.6808 14.5376 72.2316 14.4204 71.8409 14.186C71.4503 13.9516 71.1451 13.6269 70.9253 13.2118C70.7105 12.7967 70.603 12.3182 70.603 11.7761C70.603 11.2292 70.7129 10.7482 70.9326 10.3331C71.1524 9.91317 71.4552 9.58599 71.8409 9.35159C72.2316 9.1172 72.6833 9 73.196 9C73.7088 9 74.158 9.1172 74.5438 9.35159C74.9345 9.58599 75.2397 9.91073 75.4594 10.3258C75.6792 10.7409 75.789 11.2219 75.789 11.7688ZM74.8075 11.7688C74.8075 11.3879 74.7416 11.0583 74.6097 10.7799C74.4779 10.5016 74.2923 10.2867 74.053 10.1354C73.8138 9.97909 73.5281 9.90096 73.196 9.90096C72.8689 9.90096 72.5832 9.97909 72.339 10.1354C72.0997 10.2867 71.9142 10.5016 71.7823 10.7799C71.6505 11.0583 71.5846 11.3879 71.5846 11.7688C71.5846 12.1497 71.6505 12.4818 71.7823 12.765C71.9142 13.0433 72.0997 13.2582 72.339 13.4096C72.5832 13.561 72.8689 13.6366 73.196 13.6366C73.5281 13.6366 73.8138 13.561 74.053 13.4096C74.2923 13.2533 74.4779 13.036 74.6097 12.7577C74.7416 12.4744 74.8075 12.1448 74.8075 11.7688Z" fill="white"/>
|
||||
<path d="M65.6821 10.5895C65.6821 10.277 65.7627 10.0011 65.9239 9.76179C66.085 9.52251 66.3072 9.33694 66.5904 9.2051C66.8785 9.06837 67.2106 9 67.5866 9C67.948 9 68.2605 9.06348 68.5242 9.19045C68.7928 9.31741 69.0003 9.49809 69.1468 9.73249C69.2982 9.96688 69.3788 10.2452 69.3885 10.5675H68.4509C68.4412 10.338 68.3582 10.1598 68.2019 10.0328C68.0456 9.90096 67.8356 9.83504 67.572 9.83504C67.2838 9.83504 67.0519 9.90096 66.8761 10.0328C66.7052 10.1598 66.6197 10.3356 66.6197 10.5602C66.6197 10.7506 66.671 10.902 66.7735 11.0143C66.881 11.1218 67.047 11.2023 67.2716 11.2561L68.114 11.4465C68.573 11.5442 68.9148 11.7126 69.1395 11.9519C69.3641 12.1863 69.4764 12.5037 69.4764 12.9042C69.4764 13.2313 69.3958 13.5194 69.2347 13.7685C69.0736 14.0175 68.844 14.2104 68.5462 14.3472C68.2532 14.479 67.9089 14.5449 67.5134 14.5449C67.1373 14.5449 66.8077 14.4814 66.5245 14.3545C66.2413 14.2226 66.0191 14.0395 65.8579 13.8051C65.7017 13.5707 65.6187 13.2948 65.6089 12.9774H66.5465C66.5514 13.202 66.6393 13.3803 66.8102 13.5121C66.986 13.6391 67.2228 13.7026 67.5207 13.7026C67.8332 13.7026 68.0798 13.6391 68.2605 13.5121C68.4461 13.3803 68.5388 13.2069 68.5388 12.9921C68.5388 12.8065 68.49 12.66 68.3923 12.5526C68.2947 12.4402 68.136 12.3621 67.9162 12.3182L67.0665 12.1277C66.6124 12.0301 66.2681 11.8543 66.0337 11.6003C65.7993 11.3415 65.6821 11.0046 65.6821 10.5895Z" fill="white"/>
|
||||
<path d="M60.8331 14.4497H59.9102V9.09521H60.8404L63.6239 13.307H63.3528V9.09521H64.2758V14.4497H63.3528L60.5621 10.2452H60.8331V14.4497Z" fill="white"/>
|
||||
<path d="M58.4443 11.7688C58.4443 12.3108 58.3344 12.7918 58.1147 13.2118C57.8949 13.6269 57.5897 13.9516 57.1991 14.186C56.8084 14.4204 56.3567 14.5376 55.844 14.5376C55.3361 14.5376 54.8869 14.4204 54.4962 14.186C54.1055 13.9516 53.8003 13.6269 53.5806 13.2118C53.3657 12.7967 53.2583 12.3182 53.2583 11.7761C53.2583 11.2292 53.3682 10.7482 53.5879 10.3331C53.8077 9.91317 54.1104 9.58599 54.4962 9.35159C54.8869 9.1172 55.3386 9 55.8513 9C56.364 9 56.8133 9.1172 57.1991 9.35159C57.5897 9.58599 57.8949 9.91073 58.1147 10.3258C58.3344 10.7409 58.4443 11.2219 58.4443 11.7688ZM57.4628 11.7688C57.4628 11.3879 57.3969 11.0583 57.265 10.7799C57.1332 10.5016 56.9476 10.2867 56.7083 10.1354C56.469 9.97909 56.1834 9.90096 55.8513 9.90096C55.5241 9.90096 55.2385 9.97909 54.9943 10.1354C54.755 10.2867 54.5695 10.5016 54.4376 10.7799C54.3058 11.0583 54.2398 11.3879 54.2398 11.7688C54.2398 12.1497 54.3058 12.4818 54.4376 12.765C54.5695 13.0433 54.755 13.2582 54.9943 13.4096C55.2385 13.561 55.5241 13.6366 55.8513 13.6366C56.1834 13.6366 56.469 13.561 56.7083 13.4096C56.9476 13.2533 57.1332 13.036 57.265 12.7577C57.3969 12.4744 57.4628 12.1448 57.4628 11.7688Z" fill="white"/>
|
||||
<path d="M49.2393 9.09521V14.4497H48.3018V9.09521H49.2393ZM50.4186 12.6038H49.0123V11.7688H50.2209C50.5432 11.7688 50.7873 11.6882 50.9534 11.5271C51.1243 11.361 51.2097 11.1315 51.2097 10.8385C51.2097 10.5455 51.1243 10.3209 50.9534 10.1646C50.7873 10.0084 50.5481 9.93025 50.2355 9.93025H48.9244V9.09521H50.4186C50.78 9.09521 51.0925 9.16846 51.3562 9.31496C51.6199 9.46146 51.825 9.66655 51.9715 9.93025C52.118 10.1891 52.1913 10.4943 52.1913 10.8459C52.1913 11.1877 52.118 11.4929 51.9715 11.7615C51.825 12.0252 51.6199 12.2327 51.3562 12.3841C51.0925 12.5306 50.78 12.6038 50.4186 12.6038Z" fill="white"/>
|
||||
<path d="M43.0732 10.5895C43.0732 10.277 43.1538 10.0011 43.315 9.76179C43.4761 9.52251 43.6983 9.33694 43.9815 9.2051C44.2696 9.06837 44.6017 9 44.9777 9C45.3391 9 45.6516 9.06348 45.9153 9.19045C46.1839 9.31741 46.3914 9.49809 46.5379 9.73249C46.6893 9.96688 46.7699 10.2452 46.7796 10.5675H45.842C45.8323 10.338 45.7493 10.1598 45.593 10.0328C45.4367 9.90096 45.2268 9.83504 44.9631 9.83504C44.675 9.83504 44.443 9.90096 44.2672 10.0328C44.0963 10.1598 44.0108 10.3356 44.0108 10.5602C44.0108 10.7506 44.0621 10.902 44.1647 11.0143C44.2721 11.1218 44.4381 11.2023 44.6627 11.2561L45.5051 11.4465C45.9641 11.5442 46.306 11.7126 46.5306 11.9519C46.7552 12.1863 46.8675 12.5037 46.8675 12.9042C46.8675 13.2313 46.787 13.5194 46.6258 13.7685C46.4647 14.0175 46.2352 14.2104 45.9373 14.3472C45.6443 14.479 45.3 14.5449 44.9045 14.5449C44.5285 14.5449 44.1988 14.4814 43.9156 14.3545C43.6324 14.2226 43.4102 14.0395 43.249 13.8051C43.0928 13.5707 43.0098 13.2948 43 12.9774H43.9376C43.9425 13.202 44.0304 13.3803 44.2013 13.5121C44.3771 13.6391 44.6139 13.7026 44.9118 13.7026C45.2243 13.7026 45.4709 13.6391 45.6516 13.5121C45.8372 13.3803 45.9299 13.2069 45.9299 12.9921C45.9299 12.8065 45.8811 12.66 45.7835 12.5526C45.6858 12.4402 45.5271 12.3621 45.3073 12.3182L44.4576 12.1277C44.0035 12.0301 43.6592 11.8543 43.4248 11.6003C43.1904 11.3415 43.0732 11.0046 43.0732 10.5895Z" fill="white"/>
|
||||
<defs>
|
||||
<filter id="filter0_d_10631_12135" x="20.0869" y="19.3484" width="112.745" height="35.9123" filterUnits="userSpaceOnUse" color-interpolation-filters="sRGB">
|
||||
<feFlood flood-opacity="0" result="BackgroundImageFix"/>
|
||||
<feColorMatrix in="SourceAlpha" type="matrix" values="0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 127 0" result="hardAlpha"/>
|
||||
<feOffset/>
|
||||
<feGaussianBlur stdDeviation="0.717391"/>
|
||||
<feComposite in2="hardAlpha" operator="out"/>
|
||||
<feColorMatrix type="matrix" values="0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0.25 0"/>
|
||||
<feBlend mode="normal" in2="BackgroundImageFix" result="effect1_dropShadow_10631_12135"/>
|
||||
<feBlend mode="normal" in="SourceGraphic" in2="effect1_dropShadow_10631_12135" result="shape"/>
|
||||
</filter>
|
||||
</defs>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 13 KiB |
BIN
assets/lora_ease_ui.png
Normal file
BIN
assets/lora_ease_ui.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 340 KiB |
29
build_and_push_docker
Normal file
29
build_and_push_docker
Normal file
@@ -0,0 +1,29 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
# Extract version from version.py
|
||||
if [ -f "version.py" ]; then
|
||||
VERSION=$(python3 -c "from version import VERSION; print(VERSION)")
|
||||
echo "Building version: $VERSION"
|
||||
else
|
||||
echo "Error: version.py not found. Please create a version.py file with VERSION defined."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Docker builds from the repo, not this dir. Make sure changes are pushed to the repo."
|
||||
echo "Building version: $VERSION and latest"
|
||||
# wait 2 seconds
|
||||
sleep 2
|
||||
|
||||
# Build the image with cache busting
|
||||
docker build --build-arg CACHEBUST=$(date +%s) -t aitoolkit:$VERSION -f docker/Dockerfile .
|
||||
|
||||
# Tag with version and latest
|
||||
docker tag aitoolkit:$VERSION ostris/aitoolkit:$VERSION
|
||||
docker tag aitoolkit:$VERSION ostris/aitoolkit:latest
|
||||
|
||||
# Push both tags
|
||||
echo "Pushing images to Docker Hub..."
|
||||
docker push ostris/aitoolkit:$VERSION
|
||||
docker push ostris/aitoolkit:latest
|
||||
|
||||
echo "Successfully built and pushed ostris/aitoolkit:$VERSION and ostris/aitoolkit:latest"
|
||||
96
config/examples/modal/modal_train_lora_flux_24gb.yaml
Normal file
96
config/examples/modal/modal_train_lora_flux_24gb.yaml
Normal file
@@ -0,0 +1,96 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flux_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "/root/ai-toolkit/modal_output" # must match MOUNT_DIR from run_modal.py
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
# your dataset must be placed in /ai-toolkit and /root is for modal to find the dir:
|
||||
- folder_path: "/root/ai-toolkit/your-dataset"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flux
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
# if you get an error, or get stuck while downloading,
|
||||
# check https://github.com/ostris/ai-toolkit/issues/84, download the model locally and
|
||||
# place it like "/root/ai-toolkit/FLUX.1-dev"
|
||||
name_or_path: "black-forest-labs/FLUX.1-dev"
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flux
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 20
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
@@ -0,0 +1,98 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flux_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "/root/ai-toolkit/modal_output" # must match MOUNT_DIR from run_modal.py
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
# your dataset must be placed in /ai-toolkit and /root is for modal to find the dir:
|
||||
- folder_path: "/root/ai-toolkit/your-dataset"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flux
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
# if you get an error, or get stuck while downloading,
|
||||
# check https://github.com/ostris/ai-toolkit/issues/84, download the models locally and
|
||||
# place them like "/root/ai-toolkit/FLUX.1-schnell" and "/root/ai-toolkit/FLUX.1-schnell-training-adapter"
|
||||
name_or_path: "black-forest-labs/FLUX.1-schnell"
|
||||
assistant_lora_path: "ostris/FLUX.1-schnell-training-adapter" # Required for flux schnell training
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
# low_vram is painfully slow to fuse in the adapter avoid it unless absolutely necessary
|
||||
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flux
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 1 # schnell does not do guidance
|
||||
sample_steps: 4 # 1 - 4 works well
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
107
config/examples/train_full_fine_tune_flex.yaml
Normal file
107
config/examples/train_full_fine_tune_flex.yaml
Normal file
@@ -0,0 +1,107 @@
|
||||
---
|
||||
# This configuration requires 48GB of VRAM or more to operate
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flex_finetune_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
save:
|
||||
dtype: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # how many intermittent saves to keep
|
||||
save_format: 'diffusers' # 'diffusers'
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
# cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flex enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
# IMPORTANT! For Flex, you must bypass the guidance embedder during training
|
||||
bypass_guidance_embedding: true
|
||||
|
||||
# can be 'sigmoid', 'linear', or 'lognorm_blend'
|
||||
timestep_type: 'sigmoid'
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flex
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adafactor"
|
||||
lr: 3e-5
|
||||
|
||||
# Paramiter swapping can reduce vram requirements. Set factor from 1.0 to 0.0.
|
||||
# 0.1 is 10% of paramiters active at easc step. Only works with adafactor
|
||||
|
||||
# do_paramiter_swapping: true
|
||||
# paramiter_swapping_factor: 0.9
|
||||
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flex, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "ostris/Flex.1-alpha"
|
||||
is_flux: true # flex is flux architecture
|
||||
# full finetuning quantized models is a crapshoot and results in subpar outputs
|
||||
# quantize: true
|
||||
# you can quantize just the T5 text encoder here to save vram
|
||||
quantize_te: true
|
||||
# only train the transformer blocks
|
||||
only_if_contains:
|
||||
- "transformer.transformer_blocks."
|
||||
- "transformer.single_transformer_blocks."
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flex
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
99
config/examples/train_full_fine_tune_lumina.yaml
Normal file
99
config/examples/train_full_fine_tune_lumina.yaml
Normal file
@@ -0,0 +1,99 @@
|
||||
---
|
||||
# This configuration requires 24GB of VRAM or more to operate
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_lumina_finetune_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
save:
|
||||
dtype: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # how many intermittent saves to keep
|
||||
save_format: 'diffusers' # 'diffusers'
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
# cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # lumina2 enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
|
||||
# can be 'sigmoid', 'linear', or 'lumina2_shift'
|
||||
timestep_type: 'lumina2_shift'
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with lumina2
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adafactor"
|
||||
lr: 3e-5
|
||||
|
||||
# Paramiter swapping can reduce vram requirements. Set factor from 1.0 to 0.0.
|
||||
# 0.1 is 10% of paramiters active at easc step. Only works with adafactor
|
||||
|
||||
# do_paramiter_swapping: true
|
||||
# paramiter_swapping_factor: 0.9
|
||||
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
|
||||
# ema_config:
|
||||
# use_ema: true
|
||||
# ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for lumina2, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Alpha-VLLM/Lumina-Image-2.0"
|
||||
is_lumina2: true # lumina2 architecture
|
||||
# you can quantize just the Gemma2 text encoder here to save vram
|
||||
quantize_te: true
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a cat that is half black and half orange tabby, split down the middle. The cat has on a blue tophat. They are holding a martini glass with a pink ball of yarn in it with green knitting needles sticking out, in one paw. In the other paw, they are holding a DVD case for a movie titled, \"This is a test\" that has a golden robot on it. In the background is a busy night club with a giant mushroom man dancing with a bear."
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4.0
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
101
config/examples/train_lora_flex_24gb.yaml
Normal file
101
config/examples/train_lora_flex_24gb.yaml
Normal file
@@ -0,0 +1,101 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flex_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flex enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
# IMPORTANT! For Flex, you must bypass the guidance embedder during training
|
||||
bypass_guidance_embedding: true
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flex
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flex, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "ostris/Flex.1-alpha"
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
quantize_kwargs:
|
||||
exclude:
|
||||
- "*time_text_embed*" # exclude the time text embedder from quantization
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flex
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
96
config/examples/train_lora_flux_24gb.yaml
Normal file
96
config/examples/train_lora_flux_24gb.yaml
Normal file
@@ -0,0 +1,96 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flux_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flux
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "black-forest-labs/FLUX.1-dev"
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flux
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 20
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
98
config/examples/train_lora_flux_schnell_24gb.yaml
Normal file
98
config/examples/train_lora_flux_schnell_24gb.yaml
Normal file
@@ -0,0 +1,98 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flux_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flux
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new bell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "black-forest-labs/FLUX.1-schnell"
|
||||
assistant_lora_path: "ostris/FLUX.1-schnell-training-adapter" # Required for flux schnell training
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
# low_vram is painfully slow to fuse in the adapter avoid it unless absolutely necessary
|
||||
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flux
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 1 # schnell does not do guidance
|
||||
sample_steps: 4 # 1 - 4 works well
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
96
config/examples/train_lora_lumina.yaml
Normal file
96
config/examples/train_lora_lumina.yaml
Normal file
@@ -0,0 +1,96 @@
|
||||
---
|
||||
# This configuration requires 20GB of VRAM or more to operate
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_lumina_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # how many intermittent saves to keep
|
||||
save_format: 'diffusers' # 'diffusers'
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
# cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # lumina2 enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
|
||||
# can be 'sigmoid', 'linear', or 'lumina2_shift'
|
||||
timestep_type: 'lumina2_shift'
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with lumina2
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for lumina2, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Alpha-VLLM/Lumina-Image-2.0"
|
||||
is_lumina2: true # lumina2 architecture
|
||||
# you can quantize just the Gemma2 text encoder here to save vram
|
||||
quantize_te: true
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a cat that is half black and half orange tabby, split down the middle. The cat has on a blue tophat. They are holding a martini glass with a pink ball of yarn in it with green knitting needles sticking out, in one paw. In the other paw, they are holding a DVD case for a movie titled, \"This is a test\" that has a golden robot on it. In the background is a busy night club with a giant mushroom man dancing with a bear."
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4.0
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
97
config/examples/train_lora_sd35_large_24gb.yaml
Normal file
97
config/examples/train_lora_sd35_large_24gb.yaml
Normal file
@@ -0,0 +1,97 @@
|
||||
---
|
||||
# NOTE!! THIS IS CURRENTLY EXPERIMENTAL AND UNDER DEVELOPMENT. SOME THINGS WILL CHANGE
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_sd3l_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 1024 ]
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # May not fully work with SD3 yet
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch"
|
||||
timestep_type: "linear" # linear or sigmoid
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for sd3, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "stabilityai/stable-diffusion-3.5-large"
|
||||
is_v3: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
101
config/examples/train_lora_wan21_14b_24gb.yaml
Normal file
101
config/examples/train_lora_wan21_14b_24gb.yaml
Normal file
@@ -0,0 +1,101 @@
|
||||
# IMPORTANT: The Wan2.1 14B model is huge. This config should work on 24GB GPUs. It cannot
|
||||
# support keeping the text encoder on GPU while training with 24GB, so it is only good
|
||||
# for training on a single prompt, for example a person with a trigger word.
|
||||
# to train on captions, you need more vran for now.
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_wan21_14b_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# this is probably needed for 24GB cards when offloading TE to CPU
|
||||
trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 32
|
||||
linear_alpha: 32
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
# AI-Toolkit does not currently support video datasets, we will train on 1 frame at a time
|
||||
# it works well for characters, but not as well for "actions"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 632 ] # will be around 480p
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with wan
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
timestep_type: 'sigmoid'
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
optimizer_params:
|
||||
weight_decay: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
dtype: bf16
|
||||
# required for 24GB cards
|
||||
# this will encode your trigger word and use those embeddings for every image in the dataset
|
||||
unload_text_encoder: true
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
arch: 'wan21'
|
||||
# these settings will save as much vram as possible
|
||||
quantize: true
|
||||
quantize_te: true
|
||||
low_vram: true
|
||||
sample:
|
||||
sampler: "flowmatch"
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 832
|
||||
height: 480
|
||||
num_frames: 40
|
||||
fps: 15
|
||||
# samples take a long time. so use them sparingly
|
||||
# samples will be animated webp files, if you don't see them animated, open in a browser.
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 5
|
||||
sample_steps: 30
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
90
config/examples/train_lora_wan21_1b_24gb.yaml
Normal file
90
config/examples/train_lora_wan21_1b_24gb.yaml
Normal file
@@ -0,0 +1,90 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_wan21_1b_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 32
|
||||
linear_alpha: 32
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
# AI-Toolkit does not currently support video datasets, we will train on 1 frame at a time
|
||||
# it works well for characters, but not as well for "actions"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 632 ] # will be around 480p
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with wan
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
timestep_type: 'sigmoid'
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
optimizer_params:
|
||||
weight_decay: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
arch: 'wan21'
|
||||
quantize_te: true # saves vram
|
||||
sample:
|
||||
sampler: "flowmatch"
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 832
|
||||
height: 480
|
||||
num_frames: 40
|
||||
fps: 15
|
||||
# samples take a long time. so use them sparingly
|
||||
# samples will be animated webp files, if you don't see them animated, open in a browser.
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 5
|
||||
sample_steps: 30
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
25
docker-compose.yml
Normal file
25
docker-compose.yml
Normal file
@@ -0,0 +1,25 @@
|
||||
version: "3.8"
|
||||
|
||||
services:
|
||||
ai-toolkit:
|
||||
image: ostris/aitoolkit:latest
|
||||
restart: unless-stopped
|
||||
ports:
|
||||
- "8675:8675"
|
||||
volumes:
|
||||
- ~/.cache/huggingface/hub:/root/.cache/huggingface/hub
|
||||
- ./aitk_db.db:/app/ai-toolkit/aitk_db.db
|
||||
- ./datasets:/app/ai-toolkit/datasets
|
||||
- ./output:/app/ai-toolkit/output
|
||||
- ./config:/app/ai-toolkit/config
|
||||
environment:
|
||||
- AI_TOOLKIT_AUTH=${AI_TOOLKIT_AUTH:-password}
|
||||
- NODE_ENV=production
|
||||
- TZ=UTC
|
||||
deploy:
|
||||
resources:
|
||||
reservations:
|
||||
devices:
|
||||
- driver: nvidia
|
||||
count: all
|
||||
capabilities: [gpu]
|
||||
67
docker/Dockerfile
Normal file
67
docker/Dockerfile
Normal file
@@ -0,0 +1,67 @@
|
||||
FROM nvidia/cuda:12.6.3-base-ubuntu22.04
|
||||
|
||||
LABEL authors="jaret"
|
||||
|
||||
# Set noninteractive to avoid timezone prompts
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# Install dependencies
|
||||
RUN apt-get update && apt-get install --no-install-recommends -y \
|
||||
git \
|
||||
curl \
|
||||
build-essential \
|
||||
cmake \
|
||||
wget \
|
||||
python3.10 \
|
||||
python3-pip \
|
||||
python3-dev \
|
||||
python3-setuptools \
|
||||
python3-wheel \
|
||||
python3-venv \
|
||||
ffmpeg \
|
||||
tmux \
|
||||
htop \
|
||||
nvtop \
|
||||
python3-opencv \
|
||||
&& apt-get clean \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Install nodejs
|
||||
WORKDIR /tmp
|
||||
RUN curl -sL https://deb.nodesource.com/setup_23.x -o nodesource_setup.sh && \
|
||||
bash nodesource_setup.sh && \
|
||||
apt-get update && \
|
||||
apt-get install -y nodejs && \
|
||||
apt-get clean && \
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Set aliases for python and pip
|
||||
RUN ln -s /usr/bin/python3 /usr/bin/python
|
||||
|
||||
# install pytorch before cache bust to avoid redownloading pytorch
|
||||
RUN pip install --no-cache-dir torch==2.6.0 torchvision==0.21.0 --index-url https://download.pytorch.org/whl/cu126
|
||||
|
||||
# Fix cache busting by moving CACHEBUST to right before git clone
|
||||
ARG CACHEBUST=1234
|
||||
RUN echo "Cache bust: ${CACHEBUST}" && \
|
||||
git clone https://github.com/ostris/ai-toolkit.git && \
|
||||
cd ai-toolkit && \
|
||||
git submodule update --init --recursive
|
||||
|
||||
WORKDIR /app/ai-toolkit
|
||||
|
||||
# Install Python dependencies
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# Build UI
|
||||
WORKDIR /app/ai-toolkit/ui
|
||||
RUN npm install && \
|
||||
npm run build && \
|
||||
npm run update_db
|
||||
|
||||
# Expose port (assuming the application runs on port 3000)
|
||||
EXPOSE 8675
|
||||
|
||||
CMD ["npm", "run", "start"]
|
||||
256
extensions_built_in/advanced_generator/Img2ImgGenerator.py
Normal file
256
extensions_built_in/advanced_generator/Img2ImgGenerator.py
Normal file
@@ -0,0 +1,256 @@
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from typing import List
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from diffusers import T2IAdapter
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from torch.utils.data import DataLoader
|
||||
from diffusers import StableDiffusionXLImg2ImgPipeline, PixArtSigmaPipeline
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, preprocess_dataset_raw_config, DatasetConfig
|
||||
from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO
|
||||
from toolkit.sampler import get_sampler
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from toolkit.data_loader import get_dataloader_from_datasets
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from controlnet_aux.midas import MidasDetector
|
||||
from diffusers.utils import load_image
|
||||
from torchvision.transforms import ToTensor
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class GenerateConfig:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.prompts: List[str]
|
||||
self.sampler = kwargs.get('sampler', 'ddpm')
|
||||
self.neg = kwargs.get('neg', '')
|
||||
self.seed = kwargs.get('seed', -1)
|
||||
self.walk_seed = kwargs.get('walk_seed', False)
|
||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||
self.sample_steps = kwargs.get('sample_steps', 20)
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
self.ext = kwargs.get('ext', 'png')
|
||||
self.denoise_strength = kwargs.get('denoise_strength', 0.5)
|
||||
self.trigger_word = kwargs.get('trigger_word', None)
|
||||
|
||||
|
||||
class Img2ImgGenerator(BaseExtensionProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
self.output_folder = self.get_conf('output_folder', required=True)
|
||||
self.copy_inputs_to = self.get_conf('copy_inputs_to', None)
|
||||
self.device = self.get_conf('device', 'cuda')
|
||||
self.model_config = ModelConfig(**self.get_conf('model', required=True))
|
||||
self.generate_config = GenerateConfig(**self.get_conf('generate', required=True))
|
||||
self.is_latents_cached = True
|
||||
raw_datasets = self.get_conf('datasets', None)
|
||||
if raw_datasets is not None and len(raw_datasets) > 0:
|
||||
raw_datasets = preprocess_dataset_raw_config(raw_datasets)
|
||||
self.datasets = None
|
||||
self.datasets_reg = None
|
||||
self.dtype = self.get_conf('dtype', 'float16')
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
self.params = []
|
||||
if raw_datasets is not None and len(raw_datasets) > 0:
|
||||
for raw_dataset in raw_datasets:
|
||||
dataset = DatasetConfig(**raw_dataset)
|
||||
is_caching = dataset.cache_latents or dataset.cache_latents_to_disk
|
||||
if not is_caching:
|
||||
self.is_latents_cached = False
|
||||
if dataset.is_reg:
|
||||
if self.datasets_reg is None:
|
||||
self.datasets_reg = []
|
||||
self.datasets_reg.append(dataset)
|
||||
else:
|
||||
if self.datasets is None:
|
||||
self.datasets = []
|
||||
self.datasets.append(dataset)
|
||||
|
||||
self.progress_bar = None
|
||||
self.sd = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=self.model_config,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
print(f"Using device {self.device}")
|
||||
self.data_loader: DataLoader = None
|
||||
self.adapter: T2IAdapter = None
|
||||
|
||||
def to_pil(self, img):
|
||||
# image comes in -1 to 1. convert to a PIL RGB image
|
||||
img = (img + 1) / 2
|
||||
img = img.clamp(0, 1)
|
||||
img = img[0].permute(1, 2, 0).cpu().numpy()
|
||||
img = (img * 255).astype(np.uint8)
|
||||
image = Image.fromarray(img)
|
||||
return image
|
||||
|
||||
def run(self):
|
||||
with torch.no_grad():
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
device = torch.device(self.device)
|
||||
|
||||
if self.model_config.is_xl:
|
||||
pipe = StableDiffusionXLImg2ImgPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder[0],
|
||||
text_encoder_2=self.sd.text_encoder[1],
|
||||
tokenizer=self.sd.tokenizer[0],
|
||||
tokenizer_2=self.sd.tokenizer[1],
|
||||
scheduler=get_sampler(self.generate_config.sampler),
|
||||
).to(device, dtype=self.torch_dtype)
|
||||
elif self.model_config.is_pixart:
|
||||
pipe = self.sd.pipeline.to(device, dtype=self.torch_dtype)
|
||||
else:
|
||||
raise NotImplementedError("Only XL models are supported")
|
||||
pipe.set_progress_bar_config(disable=True)
|
||||
|
||||
# pipe.unet = torch.compile(pipe.unet, mode="reduce-overhead", fullgraph=True)
|
||||
# midas_depth = torch.compile(midas_depth, mode="reduce-overhead", fullgraph=True)
|
||||
|
||||
self.data_loader = get_dataloader_from_datasets(self.datasets, 1, self.sd)
|
||||
|
||||
num_batches = len(self.data_loader)
|
||||
pbar = tqdm(total=num_batches, desc="Generating images")
|
||||
seed = self.generate_config.seed
|
||||
# load images from datasets, use tqdm
|
||||
for i, batch in enumerate(self.data_loader):
|
||||
batch: DataLoaderBatchDTO = batch
|
||||
|
||||
gen_seed = seed if seed > 0 else random.randint(0, 2 ** 32 - 1)
|
||||
generator = torch.manual_seed(gen_seed)
|
||||
|
||||
file_item: FileItemDTO = batch.file_items[0]
|
||||
img_path = file_item.path
|
||||
img_filename = os.path.basename(img_path)
|
||||
img_filename_no_ext = os.path.splitext(img_filename)[0]
|
||||
img_filename = img_filename_no_ext + '.' + self.generate_config.ext
|
||||
output_path = os.path.join(self.output_folder, img_filename)
|
||||
output_caption_path = os.path.join(self.output_folder, img_filename_no_ext + '.txt')
|
||||
|
||||
if self.copy_inputs_to is not None:
|
||||
output_inputs_path = os.path.join(self.copy_inputs_to, img_filename)
|
||||
output_inputs_caption_path = os.path.join(self.copy_inputs_to, img_filename_no_ext + '.txt')
|
||||
else:
|
||||
output_inputs_path = None
|
||||
output_inputs_caption_path = None
|
||||
|
||||
caption = batch.get_caption_list()[0]
|
||||
if self.generate_config.trigger_word is not None:
|
||||
caption = caption.replace('[trigger]', self.generate_config.trigger_word)
|
||||
|
||||
img: torch.Tensor = batch.tensor.clone()
|
||||
image = self.to_pil(img)
|
||||
|
||||
# image.save(output_depth_path)
|
||||
if self.model_config.is_pixart:
|
||||
pipe: PixArtSigmaPipeline = pipe
|
||||
|
||||
# Encode the full image once
|
||||
encoded_image = pipe.vae.encode(
|
||||
pipe.image_processor.preprocess(image).to(device=pipe.device, dtype=pipe.dtype))
|
||||
if hasattr(encoded_image, "latent_dist"):
|
||||
latents = encoded_image.latent_dist.sample(generator)
|
||||
elif hasattr(encoded_image, "latents"):
|
||||
latents = encoded_image.latents
|
||||
else:
|
||||
raise AttributeError("Could not access latents of provided encoder_output")
|
||||
latents = pipe.vae.config.scaling_factor * latents
|
||||
|
||||
# latents = self.sd.encode_images(img)
|
||||
|
||||
# self.sd.noise_scheduler.set_timesteps(self.generate_config.sample_steps)
|
||||
# start_step = math.floor(self.generate_config.sample_steps * self.generate_config.denoise_strength)
|
||||
# timestep = self.sd.noise_scheduler.timesteps[start_step].unsqueeze(0)
|
||||
# timestep = timestep.to(device, dtype=torch.int32)
|
||||
# latent = latent.to(device, dtype=self.torch_dtype)
|
||||
# noise = torch.randn_like(latent, device=device, dtype=self.torch_dtype)
|
||||
# latent = self.sd.add_noise(latent, noise, timestep)
|
||||
# timesteps_to_use = self.sd.noise_scheduler.timesteps[start_step + 1:]
|
||||
batch_size = 1
|
||||
num_images_per_prompt = 1
|
||||
|
||||
shape = (batch_size, pipe.transformer.config.in_channels, image.height // pipe.vae_scale_factor,
|
||||
image.width // pipe.vae_scale_factor)
|
||||
noise = randn_tensor(shape, generator=generator, device=pipe.device, dtype=pipe.dtype)
|
||||
|
||||
# noise = torch.randn_like(latents, device=device, dtype=self.torch_dtype)
|
||||
num_inference_steps = self.generate_config.sample_steps
|
||||
strength = self.generate_config.denoise_strength
|
||||
# Get timesteps
|
||||
init_timestep = min(int(num_inference_steps * strength), num_inference_steps)
|
||||
t_start = max(num_inference_steps - init_timestep, 0)
|
||||
pipe.scheduler.set_timesteps(num_inference_steps, device="cpu")
|
||||
timesteps = pipe.scheduler.timesteps[t_start:]
|
||||
timestep = timesteps[:1].repeat(batch_size * num_images_per_prompt)
|
||||
latents = pipe.scheduler.add_noise(latents, noise, timestep)
|
||||
|
||||
gen_images = pipe.__call__(
|
||||
prompt=caption,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
latents=latents,
|
||||
timesteps=timesteps,
|
||||
width=image.width,
|
||||
height=image.height,
|
||||
num_inference_steps=num_inference_steps,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
# strength=self.generate_config.denoise_strength,
|
||||
use_resolution_binning=False,
|
||||
output_type="np"
|
||||
).images[0]
|
||||
gen_images = (gen_images * 255).clip(0, 255).astype(np.uint8)
|
||||
gen_images = Image.fromarray(gen_images)
|
||||
else:
|
||||
pipe: StableDiffusionXLImg2ImgPipeline = pipe
|
||||
|
||||
gen_images = pipe.__call__(
|
||||
prompt=caption,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
image=image,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
strength=self.generate_config.denoise_strength,
|
||||
).images[0]
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
gen_images.save(output_path)
|
||||
|
||||
# save caption
|
||||
with open(output_caption_path, 'w') as f:
|
||||
f.write(caption)
|
||||
|
||||
if output_inputs_path is not None:
|
||||
os.makedirs(os.path.dirname(output_inputs_path), exist_ok=True)
|
||||
image.save(output_inputs_path)
|
||||
with open(output_inputs_caption_path, 'w') as f:
|
||||
f.write(caption)
|
||||
|
||||
pbar.update(1)
|
||||
batch.cleanup()
|
||||
|
||||
pbar.close()
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
@@ -7,7 +7,7 @@ import numpy as np
|
||||
from PIL import Image
|
||||
from diffusers import T2IAdapter
|
||||
from torch.utils.data import DataLoader
|
||||
from diffusers import StableDiffusionXLAdapterPipeline
|
||||
from diffusers import StableDiffusionXLAdapterPipeline, StableDiffusionAdapterPipeline
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, preprocess_dataset_raw_config, DatasetConfig
|
||||
@@ -100,25 +100,43 @@ class ReferenceGenerator(BaseExtensionProcess):
|
||||
|
||||
if self.generate_config.t2i_adapter_path is not None:
|
||||
self.adapter = T2IAdapter.from_pretrained(
|
||||
"TencentARC/t2i-adapter-depth-midas-sdxl-1.0", torch_dtype=self.torch_dtype, varient="fp16"
|
||||
self.generate_config.t2i_adapter_path,
|
||||
torch_dtype=self.torch_dtype,
|
||||
varient="fp16"
|
||||
).to(device)
|
||||
|
||||
midas_depth = MidasDetector.from_pretrained(
|
||||
"valhalla/t2iadapter-aux-models", filename="dpt_large_384.pt", model_type="dpt_large"
|
||||
).to(device)
|
||||
|
||||
pipe = StableDiffusionXLAdapterPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder[0],
|
||||
text_encoder_2=self.sd.text_encoder[1],
|
||||
tokenizer=self.sd.tokenizer[0],
|
||||
tokenizer_2=self.sd.tokenizer[1],
|
||||
scheduler=get_sampler(self.generate_config.sampler),
|
||||
adapter=self.adapter,
|
||||
).to(device)
|
||||
if self.model_config.is_xl:
|
||||
pipe = StableDiffusionXLAdapterPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder[0],
|
||||
text_encoder_2=self.sd.text_encoder[1],
|
||||
tokenizer=self.sd.tokenizer[0],
|
||||
tokenizer_2=self.sd.tokenizer[1],
|
||||
scheduler=get_sampler(self.generate_config.sampler),
|
||||
adapter=self.adapter,
|
||||
).to(device, dtype=self.torch_dtype)
|
||||
else:
|
||||
pipe = StableDiffusionAdapterPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder,
|
||||
tokenizer=self.sd.tokenizer,
|
||||
scheduler=get_sampler(self.generate_config.sampler),
|
||||
safety_checker=None,
|
||||
feature_extractor=None,
|
||||
requires_safety_checker=False,
|
||||
adapter=self.adapter,
|
||||
).to(device, dtype=self.torch_dtype)
|
||||
pipe.set_progress_bar_config(disable=True)
|
||||
|
||||
pipe.unet = torch.compile(pipe.unet, mode="reduce-overhead", fullgraph=True)
|
||||
# midas_depth = torch.compile(midas_depth, mode="reduce-overhead", fullgraph=True)
|
||||
|
||||
self.data_loader = get_dataloader_from_datasets(self.datasets, 1, self.sd)
|
||||
|
||||
num_batches = len(self.data_loader)
|
||||
@@ -176,6 +194,7 @@ class ReferenceGenerator(BaseExtensionProcess):
|
||||
adapter_conditioning_scale=self.generate_config.adapter_conditioning_scale,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
).images[0]
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
gen_images.save(output_path)
|
||||
|
||||
# save caption
|
||||
|
||||
@@ -36,7 +36,24 @@ class PureLoraGenerator(Extension):
|
||||
return PureLoraGenerator
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class Img2ImgGeneratorExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "batch_img2img"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Img2ImgGeneratorExtension"
|
||||
|
||||
# This is where your process class is loaded
|
||||
# keep your imports in here so they don't slow down the rest of the program
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .Img2ImgGenerator import Img2ImgGenerator
|
||||
return Img2ImgGenerator
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
AdvancedReferenceGeneratorExtension, PureLoraGenerator
|
||||
AdvancedReferenceGeneratorExtension, PureLoraGenerator, Img2ImgGeneratorExtension
|
||||
]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
234
extensions_built_in/sd_trainer/UITrainer.py
Normal file
234
extensions_built_in/sd_trainer/UITrainer.py
Normal file
@@ -0,0 +1,234 @@
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
import sqlite3
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
from extensions_built_in.sd_trainer.SDTrainer import SDTrainer
|
||||
from typing import Literal, Optional
|
||||
|
||||
|
||||
AITK_Status = Literal["running", "stopped", "error", "completed"]
|
||||
|
||||
|
||||
class UITrainer(SDTrainer):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
|
||||
super(UITrainer, self).__init__(process_id, job, config, **kwargs)
|
||||
self.sqlite_db_path = self.config.get("sqlite_db_path", "./aitk_db.db")
|
||||
if not os.path.exists(self.sqlite_db_path):
|
||||
raise Exception(
|
||||
f"SQLite database not found at {self.sqlite_db_path}")
|
||||
print(f"Using SQLite database at {self.sqlite_db_path}")
|
||||
self.job_id = os.environ.get("AITK_JOB_ID", None)
|
||||
self.job_id = self.job_id.strip() if self.job_id is not None else None
|
||||
print(f"Job ID: \"{self.job_id}\"")
|
||||
if self.job_id is None:
|
||||
raise Exception("AITK_JOB_ID not set")
|
||||
self.is_stopping = False
|
||||
# Create a thread pool for database operations
|
||||
self.thread_pool = concurrent.futures.ThreadPoolExecutor(max_workers=1)
|
||||
# Track all async tasks
|
||||
self._async_tasks = []
|
||||
# Initialize the status
|
||||
self._run_async_operation(self._update_status("running", "Starting"))
|
||||
|
||||
def _run_async_operation(self, coro):
|
||||
"""Helper method to run an async coroutine and track the task."""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
except RuntimeError:
|
||||
# No event loop exists, create a new one
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
|
||||
# Create a task and track it
|
||||
if loop.is_running():
|
||||
task = asyncio.run_coroutine_threadsafe(coro, loop)
|
||||
self._async_tasks.append(asyncio.wrap_future(task))
|
||||
else:
|
||||
task = loop.create_task(coro)
|
||||
self._async_tasks.append(task)
|
||||
loop.run_until_complete(task)
|
||||
|
||||
async def _execute_db_operation(self, operation_func):
|
||||
"""Execute a database operation in a separate thread to avoid blocking."""
|
||||
loop = asyncio.get_event_loop()
|
||||
return await loop.run_in_executor(self.thread_pool, operation_func)
|
||||
|
||||
def _db_connect(self):
|
||||
"""Create a new connection for each operation to avoid locking."""
|
||||
conn = sqlite3.connect(self.sqlite_db_path, timeout=10.0)
|
||||
conn.isolation_level = None # Enable autocommit mode
|
||||
return conn
|
||||
|
||||
def should_stop(self):
|
||||
def _check_stop():
|
||||
with self._db_connect() as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute(
|
||||
"SELECT stop FROM Job WHERE id = ?", (self.job_id,))
|
||||
stop = cursor.fetchone()
|
||||
return False if stop is None else stop[0] == 1
|
||||
|
||||
return _check_stop()
|
||||
|
||||
def maybe_stop(self):
|
||||
if self.should_stop():
|
||||
self._run_async_operation(
|
||||
self._update_status("stopped", "Job stopped"))
|
||||
self.is_stopping = True
|
||||
raise Exception("Job stopped")
|
||||
|
||||
async def _update_key(self, key, value):
|
||||
if not self.accelerator.is_main_process:
|
||||
return
|
||||
|
||||
def _do_update():
|
||||
with self._db_connect() as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute("BEGIN IMMEDIATE")
|
||||
try:
|
||||
# Convert the value to string if it's not already
|
||||
if isinstance(value, str):
|
||||
value_to_insert = value
|
||||
else:
|
||||
value_to_insert = str(value)
|
||||
|
||||
# Use parameterized query for both the column name and value
|
||||
update_query = f"UPDATE Job SET {key} = ? WHERE id = ?"
|
||||
cursor.execute(
|
||||
update_query, (value_to_insert, self.job_id))
|
||||
finally:
|
||||
cursor.execute("COMMIT")
|
||||
|
||||
await self._execute_db_operation(_do_update)
|
||||
|
||||
def update_step(self):
|
||||
"""Non-blocking update of the step count."""
|
||||
if self.accelerator.is_main_process:
|
||||
self._run_async_operation(self._update_key("step", self.step_num))
|
||||
|
||||
def update_db_key(self, key, value):
|
||||
"""Non-blocking update a key in the database."""
|
||||
if self.accelerator.is_main_process:
|
||||
self._run_async_operation(self._update_key(key, value))
|
||||
|
||||
async def _update_status(self, status: AITK_Status, info: Optional[str] = None):
|
||||
if not self.accelerator.is_main_process:
|
||||
return
|
||||
|
||||
def _do_update():
|
||||
with self._db_connect() as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute("BEGIN IMMEDIATE")
|
||||
try:
|
||||
if info is not None:
|
||||
cursor.execute(
|
||||
"UPDATE Job SET status = ?, info = ? WHERE id = ?",
|
||||
(status, info, self.job_id)
|
||||
)
|
||||
else:
|
||||
cursor.execute(
|
||||
"UPDATE Job SET status = ? WHERE id = ?",
|
||||
(status, self.job_id)
|
||||
)
|
||||
finally:
|
||||
cursor.execute("COMMIT")
|
||||
|
||||
await self._execute_db_operation(_do_update)
|
||||
|
||||
def update_status(self, status: AITK_Status, info: Optional[str] = None):
|
||||
"""Non-blocking update of status."""
|
||||
if self.accelerator.is_main_process:
|
||||
self._run_async_operation(self._update_status(status, info))
|
||||
|
||||
async def wait_for_all_async(self):
|
||||
"""Wait for all tracked async operations to complete."""
|
||||
if not self._async_tasks:
|
||||
return
|
||||
|
||||
try:
|
||||
await asyncio.gather(*self._async_tasks)
|
||||
except Exception as e:
|
||||
pass
|
||||
finally:
|
||||
# Clear the task list after completion
|
||||
self._async_tasks.clear()
|
||||
|
||||
def on_error(self, e: Exception):
|
||||
super(UITrainer, self).on_error(e)
|
||||
if self.accelerator.is_main_process and not self.is_stopping:
|
||||
self.update_status("error", str(e))
|
||||
self.update_db_key("step", self.last_save_step)
|
||||
asyncio.run(self.wait_for_all_async())
|
||||
self.thread_pool.shutdown(wait=True)
|
||||
|
||||
def handle_timing_print_hook(self, timing_dict):
|
||||
if "train_loop" not in timing_dict:
|
||||
print("train_loop not found in timing_dict", timing_dict)
|
||||
return
|
||||
seconds_per_iter = timing_dict["train_loop"]
|
||||
# determine iter/sec or sec/iter
|
||||
if seconds_per_iter < 1:
|
||||
iters_per_sec = 1 / seconds_per_iter
|
||||
self.update_db_key("speed_string", f"{iters_per_sec:.2f} iter/sec")
|
||||
else:
|
||||
self.update_db_key(
|
||||
"speed_string", f"{seconds_per_iter:.2f} sec/iter")
|
||||
|
||||
def done_hook(self):
|
||||
super(UITrainer, self).done_hook()
|
||||
self.update_status("completed", "Training completed")
|
||||
# Wait for all async operations to finish before shutting down
|
||||
asyncio.run(self.wait_for_all_async())
|
||||
self.thread_pool.shutdown(wait=True)
|
||||
|
||||
def end_step_hook(self):
|
||||
super(UITrainer, self).end_step_hook()
|
||||
self.update_step()
|
||||
self.maybe_stop()
|
||||
|
||||
def hook_before_model_load(self):
|
||||
super().hook_before_model_load()
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Loading model")
|
||||
|
||||
def before_dataset_load(self):
|
||||
super().before_dataset_load()
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Loading dataset")
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
super().hook_before_train_loop()
|
||||
self.maybe_stop()
|
||||
self.update_step()
|
||||
self.update_status("running", "Training")
|
||||
self.timer.add_after_print_hook(self.handle_timing_print_hook)
|
||||
|
||||
def status_update_hook_func(self, string):
|
||||
self.update_status("running", string)
|
||||
|
||||
def hook_after_sd_init_before_load(self):
|
||||
super().hook_after_sd_init_before_load()
|
||||
self.maybe_stop()
|
||||
self.sd.add_status_update_hook(self.status_update_hook_func)
|
||||
|
||||
def sample_step_hook(self, img_num, total_imgs):
|
||||
super().sample_step_hook(img_num, total_imgs)
|
||||
self.maybe_stop()
|
||||
self.update_status(
|
||||
"running", f"Generating images - {img_num + 1}/{total_imgs}")
|
||||
|
||||
def sample(self, step=None, is_first=False):
|
||||
self.maybe_stop()
|
||||
total_imgs = len(self.sample_config.prompts)
|
||||
self.update_status("running", f"Generating images - 0/{total_imgs}")
|
||||
super().sample(step, is_first)
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Training")
|
||||
|
||||
def save(self, step=None):
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Saving model")
|
||||
super().save(step)
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Training")
|
||||
@@ -18,6 +18,22 @@ class SDTrainerExtension(Extension):
|
||||
from .SDTrainer import SDTrainer
|
||||
return SDTrainer
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class UITrainerExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "ui_trainer"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "UI Trainer"
|
||||
|
||||
# This is where your process class is loaded
|
||||
# keep your imports in here so they don't slow down the rest of the program
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .UITrainer import UITrainer
|
||||
return UITrainer
|
||||
|
||||
|
||||
# for backwards compatability
|
||||
class TextualInversionTrainer(SDTrainerExtension):
|
||||
@@ -26,5 +42,5 @@ class TextualInversionTrainer(SDTrainerExtension):
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
SDTrainerExtension, TextualInversionTrainer
|
||||
SDTrainerExtension, TextualInversionTrainer, UITrainerExtension
|
||||
]
|
||||
|
||||
414
flux_train_ui.py
Normal file
414
flux_train_ui.py
Normal file
@@ -0,0 +1,414 @@
|
||||
import os
|
||||
from huggingface_hub import whoami
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
|
||||
import sys
|
||||
|
||||
# Add the current working directory to the Python path
|
||||
sys.path.insert(0, os.getcwd())
|
||||
|
||||
import gradio as gr
|
||||
from PIL import Image
|
||||
import torch
|
||||
import uuid
|
||||
import os
|
||||
import shutil
|
||||
import json
|
||||
import yaml
|
||||
from slugify import slugify
|
||||
from transformers import AutoProcessor, AutoModelForCausalLM
|
||||
|
||||
sys.path.insert(0, "ai-toolkit")
|
||||
from toolkit.job import get_job
|
||||
|
||||
MAX_IMAGES = 150
|
||||
|
||||
def load_captioning(uploaded_files, concept_sentence):
|
||||
uploaded_images = [file for file in uploaded_files if not file.endswith('.txt')]
|
||||
txt_files = [file for file in uploaded_files if file.endswith('.txt')]
|
||||
txt_files_dict = {os.path.splitext(os.path.basename(txt_file))[0]: txt_file for txt_file in txt_files}
|
||||
updates = []
|
||||
if len(uploaded_images) <= 1:
|
||||
raise gr.Error(
|
||||
"Please upload at least 2 images to train your model (the ideal number with default settings is between 4-30)"
|
||||
)
|
||||
elif len(uploaded_images) > MAX_IMAGES:
|
||||
raise gr.Error(f"For now, only {MAX_IMAGES} or less images are allowed for training")
|
||||
# Update for the captioning_area
|
||||
# for _ in range(3):
|
||||
updates.append(gr.update(visible=True))
|
||||
# Update visibility and image for each captioning row and image
|
||||
for i in range(1, MAX_IMAGES + 1):
|
||||
# Determine if the current row and image should be visible
|
||||
visible = i <= len(uploaded_images)
|
||||
|
||||
# Update visibility of the captioning row
|
||||
updates.append(gr.update(visible=visible))
|
||||
|
||||
# Update for image component - display image if available, otherwise hide
|
||||
image_value = uploaded_images[i - 1] if visible else None
|
||||
updates.append(gr.update(value=image_value, visible=visible))
|
||||
|
||||
corresponding_caption = False
|
||||
if(image_value):
|
||||
base_name = os.path.splitext(os.path.basename(image_value))[0]
|
||||
print(base_name)
|
||||
print(image_value)
|
||||
if base_name in txt_files_dict:
|
||||
print("entrou")
|
||||
with open(txt_files_dict[base_name], 'r') as file:
|
||||
corresponding_caption = file.read()
|
||||
|
||||
# Update value of captioning area
|
||||
text_value = corresponding_caption if visible and corresponding_caption else "[trigger]" if visible and concept_sentence else None
|
||||
updates.append(gr.update(value=text_value, visible=visible))
|
||||
|
||||
# Update for the sample caption area
|
||||
updates.append(gr.update(visible=True))
|
||||
# Update prompt samples
|
||||
updates.append(gr.update(placeholder=f'A portrait of person in a bustling cafe {concept_sentence}', value=f'A person in a bustling cafe {concept_sentence}'))
|
||||
updates.append(gr.update(placeholder=f"A mountainous landscape in the style of {concept_sentence}"))
|
||||
updates.append(gr.update(placeholder=f"A {concept_sentence} in a mall"))
|
||||
updates.append(gr.update(visible=True))
|
||||
return updates
|
||||
|
||||
def hide_captioning():
|
||||
return gr.update(visible=False), gr.update(visible=False), gr.update(visible=False)
|
||||
|
||||
def create_dataset(*inputs):
|
||||
print("Creating dataset")
|
||||
images = inputs[0]
|
||||
destination_folder = str(f"datasets/{uuid.uuid4()}")
|
||||
if not os.path.exists(destination_folder):
|
||||
os.makedirs(destination_folder)
|
||||
|
||||
jsonl_file_path = os.path.join(destination_folder, "metadata.jsonl")
|
||||
with open(jsonl_file_path, "a") as jsonl_file:
|
||||
for index, image in enumerate(images):
|
||||
new_image_path = shutil.copy(image, destination_folder)
|
||||
|
||||
original_caption = inputs[index + 1]
|
||||
file_name = os.path.basename(new_image_path)
|
||||
|
||||
data = {"file_name": file_name, "prompt": original_caption}
|
||||
|
||||
jsonl_file.write(json.dumps(data) + "\n")
|
||||
|
||||
return destination_folder
|
||||
|
||||
|
||||
def run_captioning(images, concept_sentence, *captions):
|
||||
#Load internally to not consume resources for training
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
torch_dtype = torch.float16
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
"multimodalart/Florence-2-large-no-flash-attn", torch_dtype=torch_dtype, trust_remote_code=True
|
||||
).to(device)
|
||||
processor = AutoProcessor.from_pretrained("multimodalart/Florence-2-large-no-flash-attn", trust_remote_code=True)
|
||||
|
||||
captions = list(captions)
|
||||
for i, image_path in enumerate(images):
|
||||
print(captions[i])
|
||||
if isinstance(image_path, str): # If image is a file path
|
||||
image = Image.open(image_path).convert("RGB")
|
||||
|
||||
prompt = "<DETAILED_CAPTION>"
|
||||
inputs = processor(text=prompt, images=image, return_tensors="pt").to(device, torch_dtype)
|
||||
|
||||
generated_ids = model.generate(
|
||||
input_ids=inputs["input_ids"], pixel_values=inputs["pixel_values"], max_new_tokens=1024, num_beams=3
|
||||
)
|
||||
|
||||
generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
|
||||
parsed_answer = processor.post_process_generation(
|
||||
generated_text, task=prompt, image_size=(image.width, image.height)
|
||||
)
|
||||
caption_text = parsed_answer["<DETAILED_CAPTION>"].replace("The image shows ", "")
|
||||
if concept_sentence:
|
||||
caption_text = f"{caption_text} [trigger]"
|
||||
captions[i] = caption_text
|
||||
|
||||
yield captions
|
||||
model.to("cpu")
|
||||
del model
|
||||
del processor
|
||||
|
||||
def recursive_update(d, u):
|
||||
for k, v in u.items():
|
||||
if isinstance(v, dict) and v:
|
||||
d[k] = recursive_update(d.get(k, {}), v)
|
||||
else:
|
||||
d[k] = v
|
||||
return d
|
||||
|
||||
def start_training(
|
||||
lora_name,
|
||||
concept_sentence,
|
||||
steps,
|
||||
lr,
|
||||
rank,
|
||||
model_to_train,
|
||||
low_vram,
|
||||
dataset_folder,
|
||||
sample_1,
|
||||
sample_2,
|
||||
sample_3,
|
||||
use_more_advanced_options,
|
||||
more_advanced_options,
|
||||
):
|
||||
push_to_hub = True
|
||||
if not lora_name:
|
||||
raise gr.Error("You forgot to insert your LoRA name! This name has to be unique.")
|
||||
try:
|
||||
if whoami()["auth"]["accessToken"]["role"] == "write" or "repo.write" in whoami()["auth"]["accessToken"]["fineGrained"]["scoped"][0]["permissions"]:
|
||||
gr.Info(f"Starting training locally {whoami()['name']}. Your LoRA will be available locally and in Hugging Face after it finishes.")
|
||||
else:
|
||||
push_to_hub = False
|
||||
gr.Warning("Started training locally. Your LoRa will only be available locally because you didn't login with a `write` token to Hugging Face")
|
||||
except:
|
||||
push_to_hub = False
|
||||
gr.Warning("Started training locally. Your LoRa will only be available locally because you didn't login with a `write` token to Hugging Face")
|
||||
|
||||
print("Started training")
|
||||
slugged_lora_name = slugify(lora_name)
|
||||
|
||||
# Load the default config
|
||||
with open("config/examples/train_lora_flux_24gb.yaml", "r") as f:
|
||||
config = yaml.safe_load(f)
|
||||
|
||||
# Update the config with user inputs
|
||||
config["config"]["name"] = slugged_lora_name
|
||||
config["config"]["process"][0]["model"]["low_vram"] = low_vram
|
||||
config["config"]["process"][0]["train"]["skip_first_sample"] = True
|
||||
config["config"]["process"][0]["train"]["steps"] = int(steps)
|
||||
config["config"]["process"][0]["train"]["lr"] = float(lr)
|
||||
config["config"]["process"][0]["network"]["linear"] = int(rank)
|
||||
config["config"]["process"][0]["network"]["linear_alpha"] = int(rank)
|
||||
config["config"]["process"][0]["datasets"][0]["folder_path"] = dataset_folder
|
||||
config["config"]["process"][0]["save"]["push_to_hub"] = push_to_hub
|
||||
if(push_to_hub):
|
||||
try:
|
||||
username = whoami()["name"]
|
||||
except:
|
||||
raise gr.Error("Error trying to retrieve your username. Are you sure you are logged in with Hugging Face?")
|
||||
config["config"]["process"][0]["save"]["hf_repo_id"] = f"{username}/{slugged_lora_name}"
|
||||
config["config"]["process"][0]["save"]["hf_private"] = True
|
||||
if concept_sentence:
|
||||
config["config"]["process"][0]["trigger_word"] = concept_sentence
|
||||
|
||||
if sample_1 or sample_2 or sample_3:
|
||||
config["config"]["process"][0]["train"]["disable_sampling"] = False
|
||||
config["config"]["process"][0]["sample"]["sample_every"] = steps
|
||||
config["config"]["process"][0]["sample"]["sample_steps"] = 28
|
||||
config["config"]["process"][0]["sample"]["prompts"] = []
|
||||
if sample_1:
|
||||
config["config"]["process"][0]["sample"]["prompts"].append(sample_1)
|
||||
if sample_2:
|
||||
config["config"]["process"][0]["sample"]["prompts"].append(sample_2)
|
||||
if sample_3:
|
||||
config["config"]["process"][0]["sample"]["prompts"].append(sample_3)
|
||||
else:
|
||||
config["config"]["process"][0]["train"]["disable_sampling"] = True
|
||||
if(model_to_train == "schnell"):
|
||||
config["config"]["process"][0]["model"]["name_or_path"] = "black-forest-labs/FLUX.1-schnell"
|
||||
config["config"]["process"][0]["model"]["assistant_lora_path"] = "ostris/FLUX.1-schnell-training-adapter"
|
||||
config["config"]["process"][0]["sample"]["sample_steps"] = 4
|
||||
if(use_more_advanced_options):
|
||||
more_advanced_options_dict = yaml.safe_load(more_advanced_options)
|
||||
config["config"]["process"][0] = recursive_update(config["config"]["process"][0], more_advanced_options_dict)
|
||||
print(config)
|
||||
|
||||
# Save the updated config
|
||||
# generate a random name for the config
|
||||
random_config_name = str(uuid.uuid4())
|
||||
os.makedirs("tmp", exist_ok=True)
|
||||
config_path = f"tmp/{random_config_name}-{slugged_lora_name}.yaml"
|
||||
with open(config_path, "w") as f:
|
||||
yaml.dump(config, f)
|
||||
|
||||
# run the job locally
|
||||
job = get_job(config_path)
|
||||
job.run()
|
||||
job.cleanup()
|
||||
|
||||
return f"Training completed successfully. Model saved as {slugged_lora_name}"
|
||||
|
||||
config_yaml = '''
|
||||
device: cuda:0
|
||||
model:
|
||||
is_flux: true
|
||||
quantize: true
|
||||
network:
|
||||
linear: 16 #it will overcome the 'rank' parameter
|
||||
linear_alpha: 16 #you can have an alpha different than the ranking if you'd like
|
||||
type: lora
|
||||
sample:
|
||||
guidance_scale: 3.5
|
||||
height: 1024
|
||||
neg: '' #doesn't work for FLUX
|
||||
sample_every: 1000
|
||||
sample_steps: 28
|
||||
sampler: flowmatch
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
width: 1024
|
||||
save:
|
||||
dtype: float16
|
||||
hf_private: true
|
||||
max_step_saves_to_keep: 4
|
||||
push_to_hub: true
|
||||
save_every: 10000
|
||||
train:
|
||||
batch_size: 1
|
||||
dtype: bf16
|
||||
ema_config:
|
||||
ema_decay: 0.99
|
||||
use_ema: true
|
||||
gradient_accumulation_steps: 1
|
||||
gradient_checkpointing: true
|
||||
noise_scheduler: flowmatch
|
||||
optimizer: adamw8bit #options: prodigy, dadaptation, adamw, adamw8bit, lion, lion8bit
|
||||
train_text_encoder: false #probably doesn't work for flux
|
||||
train_unet: true
|
||||
'''
|
||||
|
||||
theme = gr.themes.Monochrome(
|
||||
text_size=gr.themes.Size(lg="18px", md="15px", sm="13px", xl="22px", xs="12px", xxl="24px", xxs="9px"),
|
||||
font=[gr.themes.GoogleFont("Source Sans Pro"), "ui-sans-serif", "system-ui", "sans-serif"],
|
||||
)
|
||||
css = """
|
||||
h1{font-size: 2em}
|
||||
h3{margin-top: 0}
|
||||
#component-1{text-align:center}
|
||||
.main_ui_logged_out{opacity: 0.3; pointer-events: none}
|
||||
.tabitem{border: 0px}
|
||||
.group_padding{padding: .55em}
|
||||
"""
|
||||
with gr.Blocks(theme=theme, css=css) as demo:
|
||||
gr.Markdown(
|
||||
"""# LoRA Ease for FLUX 🧞♂️
|
||||
### Train a high quality FLUX LoRA in a breeze ༄ using [Ostris' AI Toolkit](https://github.com/ostris/ai-toolkit)"""
|
||||
)
|
||||
with gr.Column() as main_ui:
|
||||
with gr.Row():
|
||||
lora_name = gr.Textbox(
|
||||
label="The name of your LoRA",
|
||||
info="This has to be a unique name",
|
||||
placeholder="e.g.: Persian Miniature Painting style, Cat Toy",
|
||||
)
|
||||
concept_sentence = gr.Textbox(
|
||||
label="Trigger word/sentence",
|
||||
info="Trigger word or sentence to be used",
|
||||
placeholder="uncommon word like p3rs0n or trtcrd, or sentence like 'in the style of CNSTLL'",
|
||||
interactive=True,
|
||||
)
|
||||
with gr.Group(visible=True) as image_upload:
|
||||
with gr.Row():
|
||||
images = gr.File(
|
||||
file_types=["image", ".txt"],
|
||||
label="Upload your images",
|
||||
file_count="multiple",
|
||||
interactive=True,
|
||||
visible=True,
|
||||
scale=1,
|
||||
)
|
||||
with gr.Column(scale=3, visible=False) as captioning_area:
|
||||
with gr.Column():
|
||||
gr.Markdown(
|
||||
"""# Custom captioning
|
||||
<p style="margin-top:0">You can optionally add a custom caption for each image (or use an AI model for this). [trigger] will represent your concept sentence/trigger word.</p>
|
||||
""", elem_classes="group_padding")
|
||||
do_captioning = gr.Button("Add AI captions with Florence-2")
|
||||
output_components = [captioning_area]
|
||||
caption_list = []
|
||||
for i in range(1, MAX_IMAGES + 1):
|
||||
locals()[f"captioning_row_{i}"] = gr.Row(visible=False)
|
||||
with locals()[f"captioning_row_{i}"]:
|
||||
locals()[f"image_{i}"] = gr.Image(
|
||||
type="filepath",
|
||||
width=111,
|
||||
height=111,
|
||||
min_width=111,
|
||||
interactive=False,
|
||||
scale=2,
|
||||
show_label=False,
|
||||
show_share_button=False,
|
||||
show_download_button=False,
|
||||
)
|
||||
locals()[f"caption_{i}"] = gr.Textbox(
|
||||
label=f"Caption {i}", scale=15, interactive=True
|
||||
)
|
||||
|
||||
output_components.append(locals()[f"captioning_row_{i}"])
|
||||
output_components.append(locals()[f"image_{i}"])
|
||||
output_components.append(locals()[f"caption_{i}"])
|
||||
caption_list.append(locals()[f"caption_{i}"])
|
||||
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
steps = gr.Number(label="Steps", value=1000, minimum=1, maximum=10000, step=1)
|
||||
lr = gr.Number(label="Learning Rate", value=4e-4, minimum=1e-6, maximum=1e-3, step=1e-6)
|
||||
rank = gr.Number(label="LoRA Rank", value=16, minimum=4, maximum=128, step=4)
|
||||
model_to_train = gr.Radio(["dev", "schnell"], value="dev", label="Model to train")
|
||||
low_vram = gr.Checkbox(label="Low VRAM", value=True)
|
||||
with gr.Accordion("Even more advanced options", open=False):
|
||||
use_more_advanced_options = gr.Checkbox(label="Use more advanced options", value=False)
|
||||
more_advanced_options = gr.Code(config_yaml, language="yaml")
|
||||
|
||||
with gr.Accordion("Sample prompts (optional)", visible=False) as sample:
|
||||
gr.Markdown(
|
||||
"Include sample prompts to test out your trained model. Don't forget to include your trigger word/sentence (optional)"
|
||||
)
|
||||
sample_1 = gr.Textbox(label="Test prompt 1")
|
||||
sample_2 = gr.Textbox(label="Test prompt 2")
|
||||
sample_3 = gr.Textbox(label="Test prompt 3")
|
||||
|
||||
output_components.append(sample)
|
||||
output_components.append(sample_1)
|
||||
output_components.append(sample_2)
|
||||
output_components.append(sample_3)
|
||||
start = gr.Button("Start training", visible=False)
|
||||
output_components.append(start)
|
||||
progress_area = gr.Markdown("")
|
||||
|
||||
dataset_folder = gr.State()
|
||||
|
||||
images.upload(
|
||||
load_captioning,
|
||||
inputs=[images, concept_sentence],
|
||||
outputs=output_components
|
||||
)
|
||||
|
||||
images.delete(
|
||||
load_captioning,
|
||||
inputs=[images, concept_sentence],
|
||||
outputs=output_components
|
||||
)
|
||||
|
||||
images.clear(
|
||||
hide_captioning,
|
||||
outputs=[captioning_area, sample, start]
|
||||
)
|
||||
|
||||
start.click(fn=create_dataset, inputs=[images] + caption_list, outputs=dataset_folder).then(
|
||||
fn=start_training,
|
||||
inputs=[
|
||||
lora_name,
|
||||
concept_sentence,
|
||||
steps,
|
||||
lr,
|
||||
rank,
|
||||
model_to_train,
|
||||
low_vram,
|
||||
dataset_folder,
|
||||
sample_1,
|
||||
sample_2,
|
||||
sample_3,
|
||||
use_more_advanced_options,
|
||||
more_advanced_options
|
||||
],
|
||||
outputs=progress_area,
|
||||
)
|
||||
|
||||
do_captioning.click(fn=run_captioning, inputs=[images, concept_sentence] + caption_list, outputs=caption_list)
|
||||
|
||||
if __name__ == "__main__":
|
||||
demo.launch(share=True, show_error=True)
|
||||
@@ -24,6 +24,9 @@ class BaseProcess(object):
|
||||
self.performance_log_every = self.get_conf('performance_log_every', 0)
|
||||
|
||||
print(json.dumps(self.config, indent=4))
|
||||
|
||||
def on_error(self, e: Exception):
|
||||
pass
|
||||
|
||||
def get_conf(self, key, default=None, required=False, as_type=None):
|
||||
# split key by '.' and recursively get the value
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,7 +1,7 @@
|
||||
import gc
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef, List
|
||||
from typing import ForwardRef, List, Optional, Union
|
||||
|
||||
import torch
|
||||
from safetensors.torch import save_file, load_file
|
||||
@@ -22,6 +22,7 @@ class GenerateConfig:
|
||||
self.sampler = kwargs.get('sampler', 'ddpm')
|
||||
self.width = kwargs.get('width', 512)
|
||||
self.height = kwargs.get('height', 512)
|
||||
self.size_list: Union[List[int], None] = kwargs.get('size_list', None)
|
||||
self.neg = kwargs.get('neg', '')
|
||||
self.seed = kwargs.get('seed', -1)
|
||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||
@@ -30,18 +31,34 @@ class GenerateConfig:
|
||||
self.neg_2 = kwargs.get('neg_2', None)
|
||||
self.prompts = kwargs.get('prompts', None)
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
self.compile = kwargs.get('compile', False)
|
||||
self.ext = kwargs.get('ext', 'png')
|
||||
self.prompt_file = kwargs.get('prompt_file', False)
|
||||
self.num_repeats = kwargs.get('num_repeats', 1)
|
||||
self.prompts_in_file = self.prompts
|
||||
if self.prompts is None:
|
||||
raise ValueError("Prompts must be set")
|
||||
if isinstance(self.prompts, str):
|
||||
if os.path.exists(self.prompts):
|
||||
with open(self.prompts, 'r', encoding='utf-8') as f:
|
||||
self.prompts = f.read().splitlines()
|
||||
self.prompts = [p.strip() for p in self.prompts if len(p.strip()) > 0]
|
||||
self.prompts_in_file = f.read().splitlines()
|
||||
self.prompts_in_file = [p.strip() for p in self.prompts_in_file if len(p.strip()) > 0]
|
||||
else:
|
||||
raise ValueError("Prompts file does not exist, put in list if you want to use a list of prompts")
|
||||
|
||||
self.random_prompts = kwargs.get('random_prompts', False)
|
||||
self.max_random_per_prompt = kwargs.get('max_random_per_prompt', 1)
|
||||
self.max_images = kwargs.get('max_images', 10000)
|
||||
|
||||
if self.random_prompts:
|
||||
self.prompts = []
|
||||
for i in range(self.max_images):
|
||||
num_prompts = random.randint(1, self.max_random_per_prompt)
|
||||
prompt_list = [random.choice(self.prompts_in_file) for _ in range(num_prompts)]
|
||||
self.prompts.append(", ".join(prompt_list))
|
||||
else:
|
||||
self.prompts = self.prompts_in_file
|
||||
|
||||
if kwargs.get('shuffle', False):
|
||||
# shuffle the prompts
|
||||
random.shuffle(self.prompts)
|
||||
@@ -64,6 +81,7 @@ class GenerateProcess(BaseProcess):
|
||||
self.model_config = ModelConfig(**self.get_conf('model', required=True))
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.generate_config = GenerateConfig(**self.get_conf('generate', required=True))
|
||||
self.torch_dtype = get_torch_dtype(self.get_conf('dtype', 'float16'))
|
||||
|
||||
self.progress_bar = None
|
||||
self.sd = StableDiffusion(
|
||||
@@ -71,37 +89,58 @@ class GenerateProcess(BaseProcess):
|
||||
model_config=self.model_config,
|
||||
dtype=self.model_config.dtype,
|
||||
)
|
||||
|
||||
print(f"Using device {self.device}")
|
||||
|
||||
def clean_prompt(self, prompt: str):
|
||||
# remove any non alpha numeric characters or ,'" from prompt
|
||||
return ''.join(e for e in prompt if e.isalnum() or e in ", '\"")
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
with torch.no_grad():
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
self.sd.pipeline.to(self.device, self.torch_dtype)
|
||||
|
||||
print(f"Generating {len(self.generate_config.prompts)} images")
|
||||
# build prompt image configs
|
||||
prompt_image_configs = []
|
||||
for prompt in self.generate_config.prompts:
|
||||
prompt_image_configs.append(GenerateImageConfig(
|
||||
prompt=prompt,
|
||||
prompt_2=self.generate_config.prompt_2,
|
||||
width=self.generate_config.width,
|
||||
height=self.generate_config.height,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
negative_prompt_2=self.generate_config.neg_2,
|
||||
seed=self.generate_config.seed,
|
||||
guidance_rescale=self.generate_config.guidance_rescale,
|
||||
output_ext=self.generate_config.ext,
|
||||
output_folder=self.output_folder,
|
||||
add_prompt_file=self.generate_config.prompt_file
|
||||
))
|
||||
# generate images
|
||||
self.sd.generate_images(prompt_image_configs, sampler=self.generate_config.sampler)
|
||||
print("Compiling model...")
|
||||
# self.sd.unet = torch.compile(self.sd.unet, mode="reduce-overhead", fullgraph=True)
|
||||
if self.generate_config.compile:
|
||||
self.sd.unet = torch.compile(self.sd.unet, mode="reduce-overhead")
|
||||
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
print(f"Generating {len(self.generate_config.prompts)} images")
|
||||
# build prompt image configs
|
||||
prompt_image_configs = []
|
||||
for _ in range(self.generate_config.num_repeats):
|
||||
for prompt in self.generate_config.prompts:
|
||||
width = self.generate_config.width
|
||||
height = self.generate_config.height
|
||||
# prompt = self.clean_prompt(prompt)
|
||||
|
||||
if self.generate_config.size_list is not None:
|
||||
# randomly select a size
|
||||
width, height = random.choice(self.generate_config.size_list)
|
||||
|
||||
prompt_image_configs.append(GenerateImageConfig(
|
||||
prompt=prompt,
|
||||
prompt_2=self.generate_config.prompt_2,
|
||||
width=width,
|
||||
height=height,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
negative_prompt_2=self.generate_config.neg_2,
|
||||
seed=self.generate_config.seed,
|
||||
guidance_rescale=self.generate_config.guidance_rescale,
|
||||
output_ext=self.generate_config.ext,
|
||||
output_folder=self.output_folder,
|
||||
add_prompt_file=self.generate_config.prompt_file
|
||||
))
|
||||
# generate images
|
||||
self.sd.generate_images(prompt_image_configs, sampler=self.generate_config.sampler)
|
||||
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -275,6 +275,8 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
return adapter_tensors
|
||||
|
||||
def hook_train_loop(self, batch: Union['DataLoaderBatchDTO', None]):
|
||||
if isinstance(batch, list):
|
||||
batch = batch[0]
|
||||
# set to eval mode
|
||||
self.sd.set_device_state(self.eval_slider_device_state)
|
||||
with torch.no_grad():
|
||||
@@ -364,14 +366,36 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
denoised_latents = torch.cat([noisy_latents] * self.prompt_chunk_size, dim=0)
|
||||
current_timestep = timesteps
|
||||
else:
|
||||
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
self.train_config.max_denoising_steps, device=self.device_torch
|
||||
)
|
||||
if self.train_config.noise_scheduler == 'flowmatch':
|
||||
linear_timesteps = any([
|
||||
self.train_config.linear_timesteps,
|
||||
self.train_config.linear_timesteps2,
|
||||
self.train_config.timestep_type == 'linear',
|
||||
])
|
||||
|
||||
timestep_type = 'linear' if linear_timesteps else None
|
||||
if timestep_type is None:
|
||||
timestep_type = self.train_config.timestep_type
|
||||
|
||||
# make fake latents
|
||||
l = torch.randn(
|
||||
true_batch_size, 16, height, width
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
self.sd.noise_scheduler.set_train_timesteps(
|
||||
self.train_config.max_denoising_steps,
|
||||
device=self.device_torch,
|
||||
timestep_type=timestep_type,
|
||||
latents=l
|
||||
)
|
||||
else:
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
self.train_config.max_denoising_steps, device=self.device_torch
|
||||
)
|
||||
|
||||
# ger a random number of steps
|
||||
timesteps_to = torch.randint(
|
||||
1, self.train_config.max_denoising_steps, (1,)
|
||||
1, self.train_config.max_denoising_steps - 1, (1,)
|
||||
).item()
|
||||
|
||||
# get noise
|
||||
@@ -389,7 +413,8 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
assert not self.network.is_active
|
||||
self.sd.unet.eval()
|
||||
# pass the multiplier list to the network
|
||||
self.network.multiplier = prompt_pair.multiplier_list
|
||||
# double up since we are doing cfg
|
||||
self.network.multiplier = prompt_pair.multiplier_list + prompt_pair.multiplier_list
|
||||
denoised_latents = self.sd.diffuse_some_steps(
|
||||
latents, # pass simple noise latents
|
||||
train_tools.concat_prompt_embeddings(
|
||||
@@ -507,7 +532,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
for anchor_chunk, denoised_latent_chunk, anchor_target_noise_chunk in zip(
|
||||
anchor_chunks, denoised_latent_chunks, anchor_target_noise_chunks
|
||||
):
|
||||
self.network.multiplier = anchor_chunk.multiplier_list
|
||||
self.network.multiplier = anchor_chunk.multiplier_list + anchor_chunk.multiplier_list
|
||||
|
||||
anchor_pred_noise = get_noise_pred(
|
||||
anchor_chunk.neg_prompt, anchor_chunk.prompt, 1, current_timestep, denoised_latent_chunk
|
||||
@@ -582,7 +607,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
mask_multiplier_chunks,
|
||||
unmasked_target_chunks
|
||||
):
|
||||
self.network.multiplier = prompt_pair_chunk.multiplier_list
|
||||
self.network.multiplier = prompt_pair_chunk.multiplier_list + prompt_pair_chunk.multiplier_list
|
||||
target_latents = get_noise_pred(
|
||||
prompt_pair_chunk.positive_target,
|
||||
prompt_pair_chunk.target_class,
|
||||
@@ -611,6 +636,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
offset_neutral = neutral_latents_chunk
|
||||
# offsets are already adjusted on a per-batch basis
|
||||
offset_neutral += offset
|
||||
offset_neutral = offset_neutral.detach().requires_grad_(False)
|
||||
|
||||
# 16.15 GB RAM for 512x512 -> 4.20GB RAM for 512x512 with new grad_checkpointing
|
||||
loss = torch.nn.functional.mse_loss(target_latents.float(), offset_neutral.float(), reduction="none")
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import copy
|
||||
import glob
|
||||
import os
|
||||
import shutil
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
|
||||
@@ -13,6 +14,7 @@ from torch import nn
|
||||
from torchvision.transforms import transforms
|
||||
|
||||
from jobs.process import BaseTrainProcess
|
||||
from toolkit.image_utils import show_tensors
|
||||
from toolkit.kohya_model_util import load_vae, convert_diffusers_back_to_ldm
|
||||
from toolkit.data_loader import ImageDataset
|
||||
from toolkit.losses import ComparativeTotalVariation, get_gradient_penalty, PatternLoss
|
||||
@@ -25,6 +27,8 @@ from tqdm import tqdm
|
||||
import time
|
||||
import numpy as np
|
||||
from .models.vgg19_critic import Critic
|
||||
from torchvision.transforms import Resize
|
||||
import lpips
|
||||
|
||||
IMAGE_TRANSFORMS = transforms.Compose(
|
||||
[
|
||||
@@ -62,6 +66,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.kld_weight = self.get_conf('kld_weight', 0, as_type=float)
|
||||
self.mse_weight = self.get_conf('mse_weight', 1e0, as_type=float)
|
||||
self.tv_weight = self.get_conf('tv_weight', 1e0, as_type=float)
|
||||
self.lpips_weight = self.get_conf('lpips_weight', 1e0, as_type=float)
|
||||
self.critic_weight = self.get_conf('critic_weight', 1, as_type=float)
|
||||
self.pattern_weight = self.get_conf('pattern_weight', 1, as_type=float)
|
||||
self.optimizer_params = self.get_conf('optimizer_params', {})
|
||||
@@ -71,6 +76,9 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.vgg_19 = None
|
||||
self.style_weight_scalers = []
|
||||
self.content_weight_scalers = []
|
||||
self.lpips_loss:lpips.LPIPS = None
|
||||
|
||||
self.vae_scale_factor = 8
|
||||
|
||||
self.step_num = 0
|
||||
self.epoch_num = 0
|
||||
@@ -137,6 +145,15 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
num_workers=6
|
||||
)
|
||||
|
||||
def remove_oldest_checkpoint(self):
|
||||
max_to_keep = 4
|
||||
folders = glob.glob(os.path.join(self.save_root, f"{self.job.name}*_diffusers"))
|
||||
if len(folders) > max_to_keep:
|
||||
folders.sort(key=os.path.getmtime)
|
||||
for folder in folders[:-max_to_keep]:
|
||||
print(f"Removing {folder}")
|
||||
shutil.rmtree(folder)
|
||||
|
||||
def setup_vgg19(self):
|
||||
if self.vgg_19 is None:
|
||||
self.vgg_19, self.style_losses, self.content_losses, self.vgg19_pool_4 = get_style_model_and_losses(
|
||||
@@ -211,7 +228,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
|
||||
def get_pattern_loss(self, pred, target):
|
||||
if self._pattern_loss is None:
|
||||
self._pattern_loss = PatternLoss(pattern_size=8, dtype=self.torch_dtype).to(self.device,
|
||||
self._pattern_loss = PatternLoss(pattern_size=16, dtype=self.torch_dtype).to(self.device,
|
||||
dtype=self.torch_dtype)
|
||||
loss = torch.mean(self._pattern_loss(pred, target))
|
||||
return loss
|
||||
@@ -226,25 +243,21 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
|
||||
self.update_training_metadata()
|
||||
filename = f'{self.job.name}{step_num}.safetensors'
|
||||
# prepare meta
|
||||
save_meta = get_meta_for_safetensors(self.meta, self.job.name)
|
||||
filename = f'{self.job.name}{step_num}_diffusers'
|
||||
|
||||
state_dict = convert_diffusers_back_to_ldm(self.vae)
|
||||
|
||||
for key in list(state_dict.keys()):
|
||||
v = state_dict[key]
|
||||
v = v.detach().clone().to("cpu").to(torch.float32)
|
||||
state_dict[key] = v
|
||||
|
||||
# having issues with meta
|
||||
save_file(state_dict, os.path.join(self.save_root, filename), save_meta)
|
||||
self.vae = self.vae.to("cpu", dtype=torch.float16)
|
||||
self.vae.save_pretrained(
|
||||
save_directory=os.path.join(self.save_root, filename)
|
||||
)
|
||||
self.vae = self.vae.to(self.device, dtype=self.torch_dtype)
|
||||
|
||||
self.print(f"Saved to {os.path.join(self.save_root, filename)}")
|
||||
|
||||
if self.use_critic:
|
||||
self.critic.save(step)
|
||||
|
||||
self.remove_oldest_checkpoint()
|
||||
|
||||
def sample(self, step=None):
|
||||
sample_folder = os.path.join(self.save_root, 'samples')
|
||||
if not os.path.exists(sample_folder):
|
||||
@@ -280,6 +293,13 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
output_img.paste(input_img, (0, 0))
|
||||
output_img.paste(decoded, (self.resolution, 0))
|
||||
|
||||
scale_up = 2
|
||||
if output_img.height <= 300:
|
||||
scale_up = 4
|
||||
|
||||
# scale up using nearest neighbor
|
||||
output_img = output_img.resize((output_img.width * scale_up, output_img.height * scale_up), Image.NEAREST)
|
||||
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zero-pad 9 digits
|
||||
@@ -294,7 +314,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
path_to_load = self.vae_path
|
||||
# see if we have a checkpoint in out output to resume from
|
||||
self.print(f"Looking for latest checkpoint in {self.save_root}")
|
||||
files = glob.glob(os.path.join(self.save_root, f"{self.job.name}*.safetensors"))
|
||||
files = glob.glob(os.path.join(self.save_root, f"{self.job.name}*_diffusers"))
|
||||
if files and len(files) > 0:
|
||||
latest_file = max(files, key=os.path.getmtime)
|
||||
print(f" - Latest checkpoint is: {latest_file}")
|
||||
@@ -306,13 +326,14 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.print(f"Loading VAE")
|
||||
self.print(f" - Loading VAE: {path_to_load}")
|
||||
if self.vae is None:
|
||||
self.vae = load_vae(path_to_load, dtype=self.torch_dtype)
|
||||
self.vae = AutoencoderKL.from_pretrained(path_to_load)
|
||||
|
||||
# set decoder to train
|
||||
self.vae.to(self.device, dtype=self.torch_dtype)
|
||||
self.vae.requires_grad_(False)
|
||||
self.vae.eval()
|
||||
self.vae.decoder.train()
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config['block_out_channels']) - 1)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
@@ -374,6 +395,10 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
if self.use_critic:
|
||||
self.critic.setup()
|
||||
|
||||
if self.lpips_weight > 0 and self.lpips_loss is None:
|
||||
# self.lpips_loss = lpips.LPIPS(net='vgg')
|
||||
self.lpips_loss = lpips.LPIPS(net='vgg').to(self.device, dtype=self.torch_dtype)
|
||||
|
||||
optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
|
||||
optimizer_params=self.optimizer_params)
|
||||
|
||||
@@ -397,6 +422,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.sample()
|
||||
blank_losses = OrderedDict({
|
||||
"total": [],
|
||||
"lpips": [],
|
||||
"style": [],
|
||||
"content": [],
|
||||
"mse": [],
|
||||
@@ -415,17 +441,29 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
for batch in self.data_loader:
|
||||
if self.step_num >= self.max_steps:
|
||||
break
|
||||
with torch.no_grad():
|
||||
|
||||
batch = batch.to(self.device, dtype=self.torch_dtype)
|
||||
batch = batch.to(self.device, dtype=self.torch_dtype)
|
||||
|
||||
# forward pass
|
||||
dgd = self.vae.encode(batch).latent_dist
|
||||
mu, logvar = dgd.mean, dgd.logvar
|
||||
latents = dgd.sample()
|
||||
latents.requires_grad_(True)
|
||||
# resize so it matches size of vae evenly
|
||||
if batch.shape[2] % self.vae_scale_factor != 0 or batch.shape[3] % self.vae_scale_factor != 0:
|
||||
batch = Resize((batch.shape[2] // self.vae_scale_factor * self.vae_scale_factor,
|
||||
batch.shape[3] // self.vae_scale_factor * self.vae_scale_factor))(batch)
|
||||
|
||||
# forward pass
|
||||
dgd = self.vae.encode(batch).latent_dist
|
||||
mu, logvar = dgd.mean, dgd.logvar
|
||||
latents = dgd.sample()
|
||||
latents.detach().requires_grad_(True)
|
||||
|
||||
pred = self.vae.decode(latents).sample
|
||||
|
||||
with torch.no_grad():
|
||||
show_tensors(
|
||||
pred.clamp(-1, 1).clone(),
|
||||
"combined tensor"
|
||||
)
|
||||
|
||||
# Run through VGG19
|
||||
if self.style_weight > 0 or self.content_weight > 0 or self.use_critic:
|
||||
stacked = torch.cat([pred, batch], dim=0)
|
||||
@@ -441,14 +479,31 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
content_loss = self.get_content_loss() * self.content_weight
|
||||
kld_loss = self.get_kld_loss(mu, logvar) * self.kld_weight
|
||||
mse_loss = self.get_mse_loss(pred, batch) * self.mse_weight
|
||||
if self.lpips_weight > 0:
|
||||
lpips_loss = self.lpips_loss(
|
||||
pred.clamp(-1, 1),
|
||||
batch.clamp(-1, 1)
|
||||
).mean() * self.lpips_weight
|
||||
else:
|
||||
lpips_loss = torch.tensor(0.0, device=self.device, dtype=self.torch_dtype)
|
||||
tv_loss = self.get_tv_loss(pred, batch) * self.tv_weight
|
||||
pattern_loss = self.get_pattern_loss(pred, batch) * self.pattern_weight
|
||||
if self.use_critic:
|
||||
critic_gen_loss = self.critic.get_critic_loss(self.vgg19_pool_4.tensor) * self.critic_weight
|
||||
|
||||
# do not let abs critic gen loss be higher than abs lpips * 0.1 if using it
|
||||
if self.lpips_weight > 0:
|
||||
max_target = lpips_loss.abs() * 0.1
|
||||
with torch.no_grad():
|
||||
crit_g_scaler = 1.0
|
||||
if critic_gen_loss.abs() > max_target:
|
||||
crit_g_scaler = max_target / critic_gen_loss.abs()
|
||||
|
||||
critic_gen_loss *= crit_g_scaler
|
||||
else:
|
||||
critic_gen_loss = torch.tensor(0.0, device=self.device, dtype=self.torch_dtype)
|
||||
|
||||
loss = style_loss + content_loss + kld_loss + mse_loss + tv_loss + critic_gen_loss + pattern_loss
|
||||
loss = style_loss + content_loss + kld_loss + mse_loss + tv_loss + critic_gen_loss + pattern_loss + lpips_loss
|
||||
|
||||
# Backward pass and optimization
|
||||
optimizer.zero_grad()
|
||||
@@ -460,6 +515,8 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
loss_value = loss.item()
|
||||
# get exponent like 3.54e-4
|
||||
loss_string = f"loss: {loss_value:.2e}"
|
||||
if self.lpips_weight > 0:
|
||||
loss_string += f" lpips: {lpips_loss.item():.2e}"
|
||||
if self.content_weight > 0:
|
||||
loss_string += f" cnt: {content_loss.item():.2e}"
|
||||
if self.style_weight > 0:
|
||||
@@ -477,7 +534,8 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
if self.use_critic:
|
||||
loss_string += f" crD: {critic_d_loss:.2e}"
|
||||
|
||||
if self.optimizer_type.startswith('dadaptation'):
|
||||
if self.optimizer_type.startswith('dadaptation') or \
|
||||
self.optimizer_type.lower().startswith('prodigy'):
|
||||
learning_rate = (
|
||||
optimizer.param_groups[0]["d"] *
|
||||
optimizer.param_groups[0]["lr"]
|
||||
@@ -495,6 +553,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.progress_bar.update(1)
|
||||
|
||||
epoch_losses["total"].append(loss_value)
|
||||
epoch_losses["lpips"].append(lpips_loss.item())
|
||||
epoch_losses["style"].append(style_loss.item())
|
||||
epoch_losses["content"].append(content_loss.item())
|
||||
epoch_losses["mse"].append(mse_loss.item())
|
||||
@@ -505,6 +564,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
epoch_losses["crD"].append(critic_d_loss)
|
||||
|
||||
log_losses["total"].append(loss_value)
|
||||
log_losses["lpips"].append(lpips_loss.item())
|
||||
log_losses["style"].append(style_loss.item())
|
||||
log_losses["content"].append(content_loss.item())
|
||||
log_losses["mse"].append(mse_loss.item())
|
||||
|
||||
291
notebooks/FLUX_1_dev_LoRA_Training.ipynb
Normal file
291
notebooks/FLUX_1_dev_LoRA_Training.ipynb
Normal file
@@ -0,0 +1,291 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"collapsed": false,
|
||||
"id": "zl-S0m3pkQC5"
|
||||
},
|
||||
"source": [
|
||||
"# AI Toolkit by Ostris\n",
|
||||
"## FLUX.1-dev Training\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!nvidia-smi"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "BvAG0GKAh59G"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!git clone https://github.com/ostris/ai-toolkit\n",
|
||||
"!mkdir -p /content/dataset"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "UFUW4ZMmnp1V"
|
||||
},
|
||||
"source": [
|
||||
"Put your image dataset in the `/content/dataset` folder"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "XGZqVER_aQJW"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!cd ai-toolkit && git submodule update --init --recursive && pip install -r requirements.txt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "OV0HnOI6o8V6"
|
||||
},
|
||||
"source": [
|
||||
"## Model License\n",
|
||||
"Training currently only works with FLUX.1-dev. Which means anything you train will inherit the non-commercial license. It is also a gated model, so you need to accept the license on HF before using it. Otherwise, this will fail. Here are the required steps to setup a license.\n",
|
||||
"\n",
|
||||
"Sign into HF and accept the model access here [black-forest-labs/FLUX.1-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev)\n",
|
||||
"\n",
|
||||
"[Get a READ key from huggingface](https://huggingface.co/settings/tokens/new?) and place it in the next cell after running it."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "3yZZdhFRoj2m"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import getpass\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"# Prompt for the token\n",
|
||||
"hf_token = getpass.getpass('Enter your HF access token and press enter: ')\n",
|
||||
"\n",
|
||||
"# Set the environment variable\n",
|
||||
"os.environ['HF_TOKEN'] = hf_token\n",
|
||||
"\n",
|
||||
"print(\"HF_TOKEN environment variable has been set.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "9gO2EzQ1kQC8"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"import sys\n",
|
||||
"sys.path.append('/content/ai-toolkit')\n",
|
||||
"from toolkit.job import run_job\n",
|
||||
"from collections import OrderedDict\n",
|
||||
"from PIL import Image\n",
|
||||
"import os\n",
|
||||
"os.environ[\"HF_HUB_ENABLE_HF_TRANSFER\"] = \"1\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "N8UUFzVRigbC"
|
||||
},
|
||||
"source": [
|
||||
"## Setup\n",
|
||||
"\n",
|
||||
"This is your config. It is documented pretty well. Normally you would do this as a yaml file, but for colab, this will work. This will run as is without modification, but feel free to edit as you want."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "_t28QURYjRQO"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from collections import OrderedDict\n",
|
||||
"\n",
|
||||
"job_to_run = OrderedDict([\n",
|
||||
" ('job', 'extension'),\n",
|
||||
" ('config', OrderedDict([\n",
|
||||
" # this name will be the folder and filename name\n",
|
||||
" ('name', 'my_first_flux_lora_v1'),\n",
|
||||
" ('process', [\n",
|
||||
" OrderedDict([\n",
|
||||
" ('type', 'sd_trainer'),\n",
|
||||
" # root folder to save training sessions/samples/weights\n",
|
||||
" ('training_folder', '/content/output'),\n",
|
||||
" # uncomment to see performance stats in the terminal every N steps\n",
|
||||
" #('performance_log_every', 1000),\n",
|
||||
" ('device', 'cuda:0'),\n",
|
||||
" # if a trigger word is specified, it will be added to captions of training data if it does not already exist\n",
|
||||
" # alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word\n",
|
||||
" # ('trigger_word', 'image'),\n",
|
||||
" ('network', OrderedDict([\n",
|
||||
" ('type', 'lora'),\n",
|
||||
" ('linear', 16),\n",
|
||||
" ('linear_alpha', 16)\n",
|
||||
" ])),\n",
|
||||
" ('save', OrderedDict([\n",
|
||||
" ('dtype', 'float16'), # precision to save\n",
|
||||
" ('save_every', 250), # save every this many steps\n",
|
||||
" ('max_step_saves_to_keep', 4) # how many intermittent saves to keep\n",
|
||||
" ])),\n",
|
||||
" ('datasets', [\n",
|
||||
" # datasets are a folder of images. captions need to be txt files with the same name as the image\n",
|
||||
" # for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently\n",
|
||||
" # images will automatically be resized and bucketed into the resolution specified\n",
|
||||
" OrderedDict([\n",
|
||||
" ('folder_path', '/content/dataset'),\n",
|
||||
" ('caption_ext', 'txt'),\n",
|
||||
" ('caption_dropout_rate', 0.05), # will drop out the caption 5% of time\n",
|
||||
" ('shuffle_tokens', False), # shuffle caption order, split by commas\n",
|
||||
" ('cache_latents_to_disk', True), # leave this true unless you know what you're doing\n",
|
||||
" ('resolution', [512, 768, 1024]) # flux enjoys multiple resolutions\n",
|
||||
" ])\n",
|
||||
" ]),\n",
|
||||
" ('train', OrderedDict([\n",
|
||||
" ('batch_size', 1),\n",
|
||||
" ('steps', 2000), # total number of steps to train 500 - 4000 is a good range\n",
|
||||
" ('gradient_accumulation_steps', 1),\n",
|
||||
" ('train_unet', True),\n",
|
||||
" ('train_text_encoder', False), # probably won't work with flux\n",
|
||||
" ('content_or_style', 'balanced'), # content, style, balanced\n",
|
||||
" ('gradient_checkpointing', True), # need the on unless you have a ton of vram\n",
|
||||
" ('noise_scheduler', 'flowmatch'), # for training only\n",
|
||||
" ('optimizer', 'adamw8bit'),\n",
|
||||
" ('lr', 1e-4),\n",
|
||||
"\n",
|
||||
" # uncomment this to skip the pre training sample\n",
|
||||
" # ('skip_first_sample', True),\n",
|
||||
"\n",
|
||||
" # uncomment to completely disable sampling\n",
|
||||
" # ('disable_sampling', True),\n",
|
||||
"\n",
|
||||
" # uncomment to use new vell curved weighting. Experimental but may produce better results\n",
|
||||
" # ('linear_timesteps', True),\n",
|
||||
"\n",
|
||||
" # ema will smooth out learning, but could slow it down. Recommended to leave on.\n",
|
||||
" ('ema_config', OrderedDict([\n",
|
||||
" ('use_ema', True),\n",
|
||||
" ('ema_decay', 0.99)\n",
|
||||
" ])),\n",
|
||||
"\n",
|
||||
" # will probably need this if gpu supports it for flux, other dtypes may not work correctly\n",
|
||||
" ('dtype', 'bf16')\n",
|
||||
" ])),\n",
|
||||
" ('model', OrderedDict([\n",
|
||||
" # huggingface model name or path\n",
|
||||
" ('name_or_path', 'black-forest-labs/FLUX.1-dev'),\n",
|
||||
" ('is_flux', True),\n",
|
||||
" ('quantize', True), # run 8bit mixed precision\n",
|
||||
" #('low_vram', True), # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.\n",
|
||||
" ])),\n",
|
||||
" ('sample', OrderedDict([\n",
|
||||
" ('sampler', 'flowmatch'), # must match train.noise_scheduler\n",
|
||||
" ('sample_every', 250), # sample every this many steps\n",
|
||||
" ('width', 1024),\n",
|
||||
" ('height', 1024),\n",
|
||||
" ('prompts', [\n",
|
||||
" # you can add [trigger] to the prompts here and it will be replaced with the trigger word\n",
|
||||
" #'[trigger] holding a sign that says \\'I LOVE PROMPTS!\\'',\n",
|
||||
" 'woman with red hair, playing chess at the park, bomb going off in the background',\n",
|
||||
" 'a woman holding a coffee cup, in a beanie, sitting at a cafe',\n",
|
||||
" 'a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini',\n",
|
||||
" 'a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background',\n",
|
||||
" 'a bear building a log cabin in the snow covered mountains',\n",
|
||||
" 'woman playing the guitar, on stage, singing a song, laser lights, punk rocker',\n",
|
||||
" 'hipster man with a beard, building a chair, in a wood shop',\n",
|
||||
" 'photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop',\n",
|
||||
" 'a man holding a sign that says, \\'this is a sign\\'',\n",
|
||||
" 'a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle'\n",
|
||||
" ]),\n",
|
||||
" ('neg', ''), # not used on flux\n",
|
||||
" ('seed', 42),\n",
|
||||
" ('walk_seed', True),\n",
|
||||
" ('guidance_scale', 4),\n",
|
||||
" ('sample_steps', 20)\n",
|
||||
" ]))\n",
|
||||
" ])\n",
|
||||
" ])\n",
|
||||
" ])),\n",
|
||||
" # you can add any additional meta info here. [name] is replaced with config name at top\n",
|
||||
" ('meta', OrderedDict([\n",
|
||||
" ('name', '[name]'),\n",
|
||||
" ('version', '1.0')\n",
|
||||
" ]))\n",
|
||||
"])\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "h6F1FlM2Wb3l"
|
||||
},
|
||||
"source": [
|
||||
"## Run it\n",
|
||||
"\n",
|
||||
"Below does all the magic. Check your folders to the left. Items will be in output/LoRA/your_name_v1 In the samples folder, there are preiodic sampled. This doesnt work great with colab. They will be in /content/output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "HkajwI8gteOh"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"run_job(job_to_run)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "Hblgb5uwW5SD"
|
||||
},
|
||||
"source": [
|
||||
"## Done\n",
|
||||
"\n",
|
||||
"Check your ourput dir and get your slider\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"accelerator": "GPU",
|
||||
"colab": {
|
||||
"gpuType": "A100",
|
||||
"machine_shape": "hm",
|
||||
"provenance": []
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0
|
||||
}
|
||||
296
notebooks/FLUX_1_schnell_LoRA_Training.ipynb
Normal file
296
notebooks/FLUX_1_schnell_LoRA_Training.ipynb
Normal file
@@ -0,0 +1,296 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"collapsed": false,
|
||||
"id": "zl-S0m3pkQC5"
|
||||
},
|
||||
"source": [
|
||||
"# AI Toolkit by Ostris\n",
|
||||
"## FLUX.1-schnell Training\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "3cokMT-WC6rG"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!nvidia-smi"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"collapsed": true,
|
||||
"id": "BvAG0GKAh59G"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!git clone https://github.com/ostris/ai-toolkit\n",
|
||||
"!mkdir -p /content/dataset"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "UFUW4ZMmnp1V"
|
||||
},
|
||||
"source": [
|
||||
"Put your image dataset in the `/content/dataset` folder"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"collapsed": true,
|
||||
"id": "XGZqVER_aQJW"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!cd ai-toolkit && git submodule update --init --recursive && pip install -r requirements.txt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "OV0HnOI6o8V6"
|
||||
},
|
||||
"source": [
|
||||
"## Model License\n",
|
||||
"Training currently only works with FLUX.1-dev. Which means anything you train will inherit the non-commercial license. It is also a gated model, so you need to accept the license on HF before using it. Otherwise, this will fail. Here are the required steps to setup a license.\n",
|
||||
"\n",
|
||||
"Sign into HF and accept the model access here [black-forest-labs/FLUX.1-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev)\n",
|
||||
"\n",
|
||||
"[Get a READ key from huggingface](https://huggingface.co/settings/tokens/new?) and place it in the next cell after running it."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "3yZZdhFRoj2m"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import getpass\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"# Prompt for the token\n",
|
||||
"hf_token = getpass.getpass('Enter your HF access token and press enter: ')\n",
|
||||
"\n",
|
||||
"# Set the environment variable\n",
|
||||
"os.environ['HF_TOKEN'] = hf_token\n",
|
||||
"\n",
|
||||
"print(\"HF_TOKEN environment variable has been set.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"metadata": {
|
||||
"id": "9gO2EzQ1kQC8"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"import sys\n",
|
||||
"sys.path.append('/content/ai-toolkit')\n",
|
||||
"from toolkit.job import run_job\n",
|
||||
"from collections import OrderedDict\n",
|
||||
"from PIL import Image\n",
|
||||
"import os\n",
|
||||
"os.environ[\"HF_HUB_ENABLE_HF_TRANSFER\"] = \"1\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "N8UUFzVRigbC"
|
||||
},
|
||||
"source": [
|
||||
"## Setup\n",
|
||||
"\n",
|
||||
"This is your config. It is documented pretty well. Normally you would do this as a yaml file, but for colab, this will work. This will run as is without modification, but feel free to edit as you want."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"metadata": {
|
||||
"id": "_t28QURYjRQO"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from collections import OrderedDict\n",
|
||||
"\n",
|
||||
"job_to_run = OrderedDict([\n",
|
||||
" ('job', 'extension'),\n",
|
||||
" ('config', OrderedDict([\n",
|
||||
" # this name will be the folder and filename name\n",
|
||||
" ('name', 'my_first_flux_lora_v1'),\n",
|
||||
" ('process', [\n",
|
||||
" OrderedDict([\n",
|
||||
" ('type', 'sd_trainer'),\n",
|
||||
" # root folder to save training sessions/samples/weights\n",
|
||||
" ('training_folder', '/content/output'),\n",
|
||||
" # uncomment to see performance stats in the terminal every N steps\n",
|
||||
" #('performance_log_every', 1000),\n",
|
||||
" ('device', 'cuda:0'),\n",
|
||||
" # if a trigger word is specified, it will be added to captions of training data if it does not already exist\n",
|
||||
" # alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word\n",
|
||||
" # ('trigger_word', 'image'),\n",
|
||||
" ('network', OrderedDict([\n",
|
||||
" ('type', 'lora'),\n",
|
||||
" ('linear', 16),\n",
|
||||
" ('linear_alpha', 16)\n",
|
||||
" ])),\n",
|
||||
" ('save', OrderedDict([\n",
|
||||
" ('dtype', 'float16'), # precision to save\n",
|
||||
" ('save_every', 250), # save every this many steps\n",
|
||||
" ('max_step_saves_to_keep', 4) # how many intermittent saves to keep\n",
|
||||
" ])),\n",
|
||||
" ('datasets', [\n",
|
||||
" # datasets are a folder of images. captions need to be txt files with the same name as the image\n",
|
||||
" # for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently\n",
|
||||
" # images will automatically be resized and bucketed into the resolution specified\n",
|
||||
" OrderedDict([\n",
|
||||
" ('folder_path', '/content/dataset'),\n",
|
||||
" ('caption_ext', 'txt'),\n",
|
||||
" ('caption_dropout_rate', 0.05), # will drop out the caption 5% of time\n",
|
||||
" ('shuffle_tokens', False), # shuffle caption order, split by commas\n",
|
||||
" ('cache_latents_to_disk', True), # leave this true unless you know what you're doing\n",
|
||||
" ('resolution', [512, 768, 1024]) # flux enjoys multiple resolutions\n",
|
||||
" ])\n",
|
||||
" ]),\n",
|
||||
" ('train', OrderedDict([\n",
|
||||
" ('batch_size', 1),\n",
|
||||
" ('steps', 2000), # total number of steps to train 500 - 4000 is a good range\n",
|
||||
" ('gradient_accumulation_steps', 1),\n",
|
||||
" ('train_unet', True),\n",
|
||||
" ('train_text_encoder', False), # probably won't work with flux\n",
|
||||
" ('gradient_checkpointing', True), # need the on unless you have a ton of vram\n",
|
||||
" ('noise_scheduler', 'flowmatch'), # for training only\n",
|
||||
" ('optimizer', 'adamw8bit'),\n",
|
||||
" ('lr', 1e-4),\n",
|
||||
"\n",
|
||||
" # uncomment this to skip the pre training sample\n",
|
||||
" # ('skip_first_sample', True),\n",
|
||||
"\n",
|
||||
" # uncomment to completely disable sampling\n",
|
||||
" # ('disable_sampling', True),\n",
|
||||
"\n",
|
||||
" # uncomment to use new vell curved weighting. Experimental but may produce better results\n",
|
||||
" # ('linear_timesteps', True),\n",
|
||||
"\n",
|
||||
" # ema will smooth out learning, but could slow it down. Recommended to leave on.\n",
|
||||
" ('ema_config', OrderedDict([\n",
|
||||
" ('use_ema', True),\n",
|
||||
" ('ema_decay', 0.99)\n",
|
||||
" ])),\n",
|
||||
"\n",
|
||||
" # will probably need this if gpu supports it for flux, other dtypes may not work correctly\n",
|
||||
" ('dtype', 'bf16')\n",
|
||||
" ])),\n",
|
||||
" ('model', OrderedDict([\n",
|
||||
" # huggingface model name or path\n",
|
||||
" ('name_or_path', 'black-forest-labs/FLUX.1-schnell'),\n",
|
||||
" ('assistant_lora_path', 'ostris/FLUX.1-schnell-training-adapter'), # Required for flux schnell training\n",
|
||||
" ('is_flux', True),\n",
|
||||
" ('quantize', True), # run 8bit mixed precision\n",
|
||||
" # low_vram is painfully slow to fuse in the adapter avoid it unless absolutely necessary\n",
|
||||
" #('low_vram', True), # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.\n",
|
||||
" ])),\n",
|
||||
" ('sample', OrderedDict([\n",
|
||||
" ('sampler', 'flowmatch'), # must match train.noise_scheduler\n",
|
||||
" ('sample_every', 250), # sample every this many steps\n",
|
||||
" ('width', 1024),\n",
|
||||
" ('height', 1024),\n",
|
||||
" ('prompts', [\n",
|
||||
" # you can add [trigger] to the prompts here and it will be replaced with the trigger word\n",
|
||||
" #'[trigger] holding a sign that says \\'I LOVE PROMPTS!\\'',\n",
|
||||
" 'woman with red hair, playing chess at the park, bomb going off in the background',\n",
|
||||
" 'a woman holding a coffee cup, in a beanie, sitting at a cafe',\n",
|
||||
" 'a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini',\n",
|
||||
" 'a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background',\n",
|
||||
" 'a bear building a log cabin in the snow covered mountains',\n",
|
||||
" 'woman playing the guitar, on stage, singing a song, laser lights, punk rocker',\n",
|
||||
" 'hipster man with a beard, building a chair, in a wood shop',\n",
|
||||
" 'photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop',\n",
|
||||
" 'a man holding a sign that says, \\'this is a sign\\'',\n",
|
||||
" 'a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle'\n",
|
||||
" ]),\n",
|
||||
" ('neg', ''), # not used on flux\n",
|
||||
" ('seed', 42),\n",
|
||||
" ('walk_seed', True),\n",
|
||||
" ('guidance_scale', 1), # schnell does not do guidance\n",
|
||||
" ('sample_steps', 4) # 1 - 4 works well\n",
|
||||
" ]))\n",
|
||||
" ])\n",
|
||||
" ])\n",
|
||||
" ])),\n",
|
||||
" # you can add any additional meta info here. [name] is replaced with config name at top\n",
|
||||
" ('meta', OrderedDict([\n",
|
||||
" ('name', '[name]'),\n",
|
||||
" ('version', '1.0')\n",
|
||||
" ]))\n",
|
||||
"])\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "h6F1FlM2Wb3l"
|
||||
},
|
||||
"source": [
|
||||
"## Run it\n",
|
||||
"\n",
|
||||
"Below does all the magic. Check your folders to the left. Items will be in output/LoRA/your_name_v1 In the samples folder, there are preiodic sampled. This doesnt work great with colab. They will be in /content/output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "HkajwI8gteOh"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"run_job(job_to_run)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "Hblgb5uwW5SD"
|
||||
},
|
||||
"source": [
|
||||
"## Done\n",
|
||||
"\n",
|
||||
"Check your ourput dir and get your slider\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"accelerator": "GPU",
|
||||
"colab": {
|
||||
"gpuType": "A100",
|
||||
"machine_shape": "hm",
|
||||
"provenance": []
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0
|
||||
}
|
||||
Submodule repositories/ipadapter updated: d8ab37c421...5a18b1f366
@@ -1,8 +1,9 @@
|
||||
torch
|
||||
torchvision
|
||||
torch==2.6.0
|
||||
torchvision==0.21.0
|
||||
torchao==0.9.0
|
||||
safetensors
|
||||
diffusers==0.21.3
|
||||
git+https://github.com/huggingface/transformers.git
|
||||
git+https://github.com/huggingface/diffusers@363d1ab7e24c5ed6c190abb00df66d9edb74383b
|
||||
transformers==4.49.0
|
||||
lycoris-lora==1.8.3
|
||||
flatten_json
|
||||
pyyaml
|
||||
@@ -13,7 +14,8 @@ invisible-watermark
|
||||
einops
|
||||
accelerate
|
||||
toml
|
||||
albumentations
|
||||
albumentations==1.4.15
|
||||
albucore==0.0.16
|
||||
pydantic
|
||||
omegaconf
|
||||
k-diffusion
|
||||
@@ -21,4 +23,16 @@ open_clip_torch
|
||||
timm
|
||||
prodigyopt
|
||||
controlnet_aux==0.0.7
|
||||
python-dotenv
|
||||
python-dotenv
|
||||
bitsandbytes
|
||||
hf_transfer
|
||||
lpips
|
||||
pytorch_fid
|
||||
optimum-quanto==0.2.4
|
||||
sentencepiece
|
||||
huggingface_hub
|
||||
peft
|
||||
gradio
|
||||
python-slugify
|
||||
opencv-python
|
||||
pytorch-wavelets==1.3.0
|
||||
46
run.py
46
run.py
@@ -1,4 +1,5 @@
|
||||
import os
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
|
||||
import sys
|
||||
from typing import Union, OrderedDict
|
||||
from dotenv import load_dotenv
|
||||
@@ -19,20 +20,26 @@ if os.environ.get("DEBUG_TOOLKIT", "0") == "1":
|
||||
torch.autograd.set_detect_anomaly(True)
|
||||
import argparse
|
||||
from toolkit.job import get_job
|
||||
from toolkit.accelerator import get_accelerator
|
||||
from toolkit.print import print_acc, setup_log_to_file
|
||||
|
||||
accelerator = get_accelerator()
|
||||
|
||||
|
||||
def print_end_message(jobs_completed, jobs_failed):
|
||||
if not accelerator.is_main_process:
|
||||
return
|
||||
failure_string = f"{jobs_failed} failure{'' if jobs_failed == 1 else 's'}" if jobs_failed > 0 else ""
|
||||
completed_string = f"{jobs_completed} completed job{'' if jobs_completed == 1 else 's'}"
|
||||
|
||||
print("")
|
||||
print("========================================")
|
||||
print("Result:")
|
||||
print_acc("")
|
||||
print_acc("========================================")
|
||||
print_acc("Result:")
|
||||
if len(completed_string) > 0:
|
||||
print(f" - {completed_string}")
|
||||
print_acc(f" - {completed_string}")
|
||||
if len(failure_string) > 0:
|
||||
print(f" - {failure_string}")
|
||||
print("========================================")
|
||||
print_acc(f" - {failure_string}")
|
||||
print_acc("========================================")
|
||||
|
||||
|
||||
def main():
|
||||
@@ -60,7 +67,17 @@ def main():
|
||||
default=None,
|
||||
help='Name to replace [name] tag in config file, useful for shared config file'
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
'-l', '--log',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Log file to write output to'
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.log is not None:
|
||||
setup_log_to_file(args.log)
|
||||
|
||||
config_file_list = args.config_file_list
|
||||
if len(config_file_list) == 0:
|
||||
@@ -69,7 +86,8 @@ def main():
|
||||
jobs_completed = 0
|
||||
jobs_failed = 0
|
||||
|
||||
print(f"Running {len(config_file_list)} job{'' if len(config_file_list) == 1 else 's'}")
|
||||
if accelerator.is_main_process:
|
||||
print_acc(f"Running {len(config_file_list)} job{'' if len(config_file_list) == 1 else 's'}")
|
||||
|
||||
for config_file in config_file_list:
|
||||
try:
|
||||
@@ -78,8 +96,20 @@ def main():
|
||||
job.cleanup()
|
||||
jobs_completed += 1
|
||||
except Exception as e:
|
||||
print(f"Error running job: {e}")
|
||||
print_acc(f"Error running job: {e}")
|
||||
jobs_failed += 1
|
||||
try:
|
||||
job.process[0].on_error(e)
|
||||
except Exception as e2:
|
||||
print_acc(f"Error running on_error: {e2}")
|
||||
if not args.recover:
|
||||
print_end_message(jobs_completed, jobs_failed)
|
||||
raise e
|
||||
except KeyboardInterrupt as e:
|
||||
try:
|
||||
job.process[0].on_error(e)
|
||||
except Exception as e2:
|
||||
print_acc(f"Error running on_error: {e2}")
|
||||
if not args.recover:
|
||||
print_end_message(jobs_completed, jobs_failed)
|
||||
raise e
|
||||
|
||||
175
run_modal.py
Normal file
175
run_modal.py
Normal file
@@ -0,0 +1,175 @@
|
||||
'''
|
||||
|
||||
ostris/ai-toolkit on https://modal.com
|
||||
Run training with the following command:
|
||||
modal run run_modal.py --config-file-list-str=/root/ai-toolkit/config/whatever_you_want.yml
|
||||
|
||||
'''
|
||||
|
||||
import os
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
|
||||
import sys
|
||||
import modal
|
||||
from dotenv import load_dotenv
|
||||
# Load the .env file if it exists
|
||||
load_dotenv()
|
||||
|
||||
sys.path.insert(0, "/root/ai-toolkit")
|
||||
# must come before ANY torch or fastai imports
|
||||
# import toolkit.cuda_malloc
|
||||
|
||||
# turn off diffusers telemetry until I can figure out how to make it opt-in
|
||||
os.environ['DISABLE_TELEMETRY'] = 'YES'
|
||||
|
||||
# define the volume for storing model outputs, using "creating volumes lazily": https://modal.com/docs/guide/volumes
|
||||
# you will find your model, samples and optimizer stored in: https://modal.com/storage/your-username/main/flux-lora-models
|
||||
model_volume = modal.Volume.from_name("flux-lora-models", create_if_missing=True)
|
||||
|
||||
# modal_output, due to "cannot mount volume on non-empty path" requirement
|
||||
MOUNT_DIR = "/root/ai-toolkit/modal_output" # modal_output, due to "cannot mount volume on non-empty path" requirement
|
||||
|
||||
# define modal app
|
||||
image = (
|
||||
modal.Image.debian_slim(python_version="3.11")
|
||||
# install required system and pip packages, more about this modal approach: https://modal.com/docs/examples/dreambooth_app
|
||||
.apt_install("libgl1", "libglib2.0-0")
|
||||
.pip_install(
|
||||
"python-dotenv",
|
||||
"torch",
|
||||
"diffusers[torch]",
|
||||
"transformers",
|
||||
"ftfy",
|
||||
"torchvision",
|
||||
"oyaml",
|
||||
"opencv-python",
|
||||
"albumentations",
|
||||
"safetensors",
|
||||
"lycoris-lora==1.8.3",
|
||||
"flatten_json",
|
||||
"pyyaml",
|
||||
"tensorboard",
|
||||
"kornia",
|
||||
"invisible-watermark",
|
||||
"einops",
|
||||
"accelerate",
|
||||
"toml",
|
||||
"pydantic",
|
||||
"omegaconf",
|
||||
"k-diffusion",
|
||||
"open_clip_torch",
|
||||
"timm",
|
||||
"prodigyopt",
|
||||
"controlnet_aux==0.0.7",
|
||||
"bitsandbytes",
|
||||
"hf_transfer",
|
||||
"lpips",
|
||||
"pytorch_fid",
|
||||
"optimum-quanto",
|
||||
"sentencepiece",
|
||||
"huggingface_hub",
|
||||
"peft"
|
||||
)
|
||||
)
|
||||
|
||||
# mount for the entire ai-toolkit directory
|
||||
# example: "/Users/username/ai-toolkit" is the local directory, "/root/ai-toolkit" is the remote directory
|
||||
code_mount = modal.Mount.from_local_dir("/Users/username/ai-toolkit", remote_path="/root/ai-toolkit")
|
||||
|
||||
# create the Modal app with the necessary mounts and volumes
|
||||
app = modal.App(name="flux-lora-training", image=image, mounts=[code_mount], volumes={MOUNT_DIR: model_volume})
|
||||
|
||||
# Check if we have DEBUG_TOOLKIT in env
|
||||
if os.environ.get("DEBUG_TOOLKIT", "0") == "1":
|
||||
# Set torch to trace mode
|
||||
import torch
|
||||
torch.autograd.set_detect_anomaly(True)
|
||||
|
||||
import argparse
|
||||
from toolkit.job import get_job
|
||||
|
||||
def print_end_message(jobs_completed, jobs_failed):
|
||||
failure_string = f"{jobs_failed} failure{'' if jobs_failed == 1 else 's'}" if jobs_failed > 0 else ""
|
||||
completed_string = f"{jobs_completed} completed job{'' if jobs_completed == 1 else 's'}"
|
||||
|
||||
print("")
|
||||
print("========================================")
|
||||
print("Result:")
|
||||
if len(completed_string) > 0:
|
||||
print(f" - {completed_string}")
|
||||
if len(failure_string) > 0:
|
||||
print(f" - {failure_string}")
|
||||
print("========================================")
|
||||
|
||||
|
||||
@app.function(
|
||||
# request a GPU with at least 24GB VRAM
|
||||
# more about modal GPU's: https://modal.com/docs/guide/gpu
|
||||
gpu="A100", # gpu="H100"
|
||||
# more about modal timeouts: https://modal.com/docs/guide/timeouts
|
||||
timeout=7200 # 2 hours, increase or decrease if needed
|
||||
)
|
||||
def main(config_file_list_str: str, recover: bool = False, name: str = None):
|
||||
# convert the config file list from a string to a list
|
||||
config_file_list = config_file_list_str.split(",")
|
||||
|
||||
jobs_completed = 0
|
||||
jobs_failed = 0
|
||||
|
||||
print(f"Running {len(config_file_list)} job{'' if len(config_file_list) == 1 else 's'}")
|
||||
|
||||
for config_file in config_file_list:
|
||||
try:
|
||||
job = get_job(config_file, name)
|
||||
|
||||
job.config['process'][0]['training_folder'] = MOUNT_DIR
|
||||
os.makedirs(MOUNT_DIR, exist_ok=True)
|
||||
print(f"Training outputs will be saved to: {MOUNT_DIR}")
|
||||
|
||||
# run the job
|
||||
job.run()
|
||||
|
||||
# commit the volume after training
|
||||
model_volume.commit()
|
||||
|
||||
job.cleanup()
|
||||
jobs_completed += 1
|
||||
except Exception as e:
|
||||
print(f"Error running job: {e}")
|
||||
jobs_failed += 1
|
||||
if not recover:
|
||||
print_end_message(jobs_completed, jobs_failed)
|
||||
raise e
|
||||
|
||||
print_end_message(jobs_completed, jobs_failed)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# require at least one config file
|
||||
parser.add_argument(
|
||||
'config_file_list',
|
||||
nargs='+',
|
||||
type=str,
|
||||
help='Name of config file (eg: person_v1 for config/person_v1.json/yaml), or full path if it is not in config folder, you can pass multiple config files and run them all sequentially'
|
||||
)
|
||||
|
||||
# flag to continue if a job fails
|
||||
parser.add_argument(
|
||||
'-r', '--recover',
|
||||
action='store_true',
|
||||
help='Continue running additional jobs even if a job fails'
|
||||
)
|
||||
|
||||
# optional name replacement for config file
|
||||
parser.add_argument(
|
||||
'-n', '--name',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Name to replace [name] tag in config file, useful for shared config file'
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
# convert list of config files to a comma-separated string for Modal compatibility
|
||||
config_file_list_str = ",".join(args.config_file_list)
|
||||
|
||||
main.call(config_file_list_str=config_file_list_str, recover=args.recover, name=args.name)
|
||||
426
scripts/convert_diffusers_to_comfy.py
Normal file
426
scripts/convert_diffusers_to_comfy.py
Normal file
@@ -0,0 +1,426 @@
|
||||
#######################################################
|
||||
# Convert Diffusers Flux/Flex to all in one ComfyUI safetensors file
|
||||
# The VAE, T5 and clip will all be in the safetensors file
|
||||
# T5 will always be 8bit with the all in one file
|
||||
# You can save the transformer weights as bf16 or 8-bit with the --do_8_bit flag
|
||||
#
|
||||
# Download a reference model from Huggingface
|
||||
# https://huggingface.co/Comfy-Org/flux1-dev/blob/main/flux1-dev-fp8.safetensors
|
||||
#
|
||||
# Call like this for 8-bit transformer weights:
|
||||
# python convert_flux_diffusers_to_orig.py /path/to/diffusers/checkpoint /path/to/flux1-dev-fp8.safetensors /output/path/my_finetune.safetensors --do_8_bit
|
||||
#
|
||||
# Call like this for bf16 transformer weights:
|
||||
# python convert_flux_diffusers_to_orig.py /path/to/diffusers/checkpoint /path/to/flux1-dev-fp8.safetensors /output/path/my_finetune.safetensors
|
||||
#
|
||||
#######################################################
|
||||
|
||||
|
||||
import argparse
|
||||
from datetime import date
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import safetensors
|
||||
import safetensors.torch
|
||||
import torch
|
||||
import tqdm
|
||||
from collections import OrderedDict
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument("diffusers_path", type=str,
|
||||
help="Path to the original Flux diffusers folder.")
|
||||
parser.add_argument("quantized_state_dict_path", type=str,
|
||||
help="Path to the ComfyUI all in one template file.")
|
||||
parser.add_argument("flux_path", type=str,
|
||||
help="Output path for the Flux safetensors file.")
|
||||
parser.add_argument("--do_8_bit", action="store_true",
|
||||
help="Use 8-bit weights instead of bf16.")
|
||||
args = parser.parse_args()
|
||||
|
||||
flux_path = Path(args.flux_path)
|
||||
diffusers_path = Path(args.diffusers_path, "transformer")
|
||||
quantized_state_dict_path = Path(args.quantized_state_dict_path)
|
||||
|
||||
do_8_bit = args.do_8_bit
|
||||
|
||||
if not os.path.exists(flux_path.parent):
|
||||
os.makedirs(flux_path.parent)
|
||||
|
||||
if not diffusers_path.exists():
|
||||
print(f"Error: Missing transformer folder: {diffusers_path}")
|
||||
exit()
|
||||
|
||||
original_json_path = Path.joinpath(
|
||||
diffusers_path, "diffusion_pytorch_model.safetensors.index.json")
|
||||
if not original_json_path.exists():
|
||||
print(f"Error: Missing transformer index json: {original_json_path}")
|
||||
exit()
|
||||
|
||||
if not os.path.exists(quantized_state_dict_path):
|
||||
print(
|
||||
f"Error: Missing quantized state dict file: {args.quantized_state_dict_path}")
|
||||
exit()
|
||||
|
||||
with open(original_json_path, "r", encoding="utf-8") as f:
|
||||
original_json = json.load(f)
|
||||
|
||||
diffusers_map = {
|
||||
"time_in.in_layer.weight": [
|
||||
"time_text_embed.timestep_embedder.linear_1.weight",
|
||||
],
|
||||
"time_in.in_layer.bias": [
|
||||
"time_text_embed.timestep_embedder.linear_1.bias",
|
||||
],
|
||||
"time_in.out_layer.weight": [
|
||||
"time_text_embed.timestep_embedder.linear_2.weight",
|
||||
],
|
||||
"time_in.out_layer.bias": [
|
||||
"time_text_embed.timestep_embedder.linear_2.bias",
|
||||
],
|
||||
"vector_in.in_layer.weight": [
|
||||
"time_text_embed.text_embedder.linear_1.weight",
|
||||
],
|
||||
"vector_in.in_layer.bias": [
|
||||
"time_text_embed.text_embedder.linear_1.bias",
|
||||
],
|
||||
"vector_in.out_layer.weight": [
|
||||
"time_text_embed.text_embedder.linear_2.weight",
|
||||
],
|
||||
"vector_in.out_layer.bias": [
|
||||
"time_text_embed.text_embedder.linear_2.bias",
|
||||
],
|
||||
"guidance_in.in_layer.weight": [
|
||||
"time_text_embed.guidance_embedder.linear_1.weight",
|
||||
],
|
||||
"guidance_in.in_layer.bias": [
|
||||
"time_text_embed.guidance_embedder.linear_1.bias",
|
||||
],
|
||||
"guidance_in.out_layer.weight": [
|
||||
"time_text_embed.guidance_embedder.linear_2.weight",
|
||||
],
|
||||
"guidance_in.out_layer.bias": [
|
||||
"time_text_embed.guidance_embedder.linear_2.bias",
|
||||
],
|
||||
"txt_in.weight": [
|
||||
"context_embedder.weight",
|
||||
],
|
||||
"txt_in.bias": [
|
||||
"context_embedder.bias",
|
||||
],
|
||||
"img_in.weight": [
|
||||
"x_embedder.weight",
|
||||
],
|
||||
"img_in.bias": [
|
||||
"x_embedder.bias",
|
||||
],
|
||||
"double_blocks.().img_mod.lin.weight": [
|
||||
"norm1.linear.weight",
|
||||
],
|
||||
"double_blocks.().img_mod.lin.bias": [
|
||||
"norm1.linear.bias",
|
||||
],
|
||||
"double_blocks.().txt_mod.lin.weight": [
|
||||
"norm1_context.linear.weight",
|
||||
],
|
||||
"double_blocks.().txt_mod.lin.bias": [
|
||||
"norm1_context.linear.bias",
|
||||
],
|
||||
"double_blocks.().img_attn.qkv.weight": [
|
||||
"attn.to_q.weight",
|
||||
"attn.to_k.weight",
|
||||
"attn.to_v.weight",
|
||||
],
|
||||
"double_blocks.().img_attn.qkv.bias": [
|
||||
"attn.to_q.bias",
|
||||
"attn.to_k.bias",
|
||||
"attn.to_v.bias",
|
||||
],
|
||||
"double_blocks.().txt_attn.qkv.weight": [
|
||||
"attn.add_q_proj.weight",
|
||||
"attn.add_k_proj.weight",
|
||||
"attn.add_v_proj.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.qkv.bias": [
|
||||
"attn.add_q_proj.bias",
|
||||
"attn.add_k_proj.bias",
|
||||
"attn.add_v_proj.bias",
|
||||
],
|
||||
"double_blocks.().img_attn.norm.query_norm.scale": [
|
||||
"attn.norm_q.weight",
|
||||
],
|
||||
"double_blocks.().img_attn.norm.key_norm.scale": [
|
||||
"attn.norm_k.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.norm.query_norm.scale": [
|
||||
"attn.norm_added_q.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.norm.key_norm.scale": [
|
||||
"attn.norm_added_k.weight",
|
||||
],
|
||||
"double_blocks.().img_mlp.0.weight": [
|
||||
"ff.net.0.proj.weight",
|
||||
],
|
||||
"double_blocks.().img_mlp.0.bias": [
|
||||
"ff.net.0.proj.bias",
|
||||
],
|
||||
"double_blocks.().img_mlp.2.weight": [
|
||||
"ff.net.2.weight",
|
||||
],
|
||||
"double_blocks.().img_mlp.2.bias": [
|
||||
"ff.net.2.bias",
|
||||
],
|
||||
"double_blocks.().txt_mlp.0.weight": [
|
||||
"ff_context.net.0.proj.weight",
|
||||
],
|
||||
"double_blocks.().txt_mlp.0.bias": [
|
||||
"ff_context.net.0.proj.bias",
|
||||
],
|
||||
"double_blocks.().txt_mlp.2.weight": [
|
||||
"ff_context.net.2.weight",
|
||||
],
|
||||
"double_blocks.().txt_mlp.2.bias": [
|
||||
"ff_context.net.2.bias",
|
||||
],
|
||||
"double_blocks.().img_attn.proj.weight": [
|
||||
"attn.to_out.0.weight",
|
||||
],
|
||||
"double_blocks.().img_attn.proj.bias": [
|
||||
"attn.to_out.0.bias",
|
||||
],
|
||||
"double_blocks.().txt_attn.proj.weight": [
|
||||
"attn.to_add_out.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.proj.bias": [
|
||||
"attn.to_add_out.bias",
|
||||
],
|
||||
"single_blocks.().modulation.lin.weight": [
|
||||
"norm.linear.weight",
|
||||
],
|
||||
"single_blocks.().modulation.lin.bias": [
|
||||
"norm.linear.bias",
|
||||
],
|
||||
"single_blocks.().linear1.weight": [
|
||||
"attn.to_q.weight",
|
||||
"attn.to_k.weight",
|
||||
"attn.to_v.weight",
|
||||
"proj_mlp.weight",
|
||||
],
|
||||
"single_blocks.().linear1.bias": [
|
||||
"attn.to_q.bias",
|
||||
"attn.to_k.bias",
|
||||
"attn.to_v.bias",
|
||||
"proj_mlp.bias",
|
||||
],
|
||||
"single_blocks.().linear2.weight": [
|
||||
"proj_out.weight",
|
||||
],
|
||||
"single_blocks.().norm.query_norm.scale": [
|
||||
"attn.norm_q.weight",
|
||||
],
|
||||
"single_blocks.().norm.key_norm.scale": [
|
||||
"attn.norm_k.weight",
|
||||
],
|
||||
"single_blocks.().linear2.weight": [
|
||||
"proj_out.weight",
|
||||
],
|
||||
"single_blocks.().linear2.bias": [
|
||||
"proj_out.bias",
|
||||
],
|
||||
"final_layer.linear.weight": [
|
||||
"proj_out.weight",
|
||||
],
|
||||
"final_layer.linear.bias": [
|
||||
"proj_out.bias",
|
||||
],
|
||||
"final_layer.adaLN_modulation.1.weight": [
|
||||
"norm_out.linear.weight",
|
||||
],
|
||||
"final_layer.adaLN_modulation.1.bias": [
|
||||
"norm_out.linear.bias",
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def is_in_diffusers_map(k):
|
||||
for values in diffusers_map.values():
|
||||
for value in values:
|
||||
if k.endswith(value):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
diffusers = {k: Path.joinpath(diffusers_path, v)
|
||||
for k, v in original_json["weight_map"].items() if is_in_diffusers_map(k)}
|
||||
|
||||
original_safetensors = set(diffusers.values())
|
||||
|
||||
# determine the number of transformer blocks
|
||||
transformer_blocks = 0
|
||||
single_transformer_blocks = 0
|
||||
for key in diffusers.keys():
|
||||
print(key)
|
||||
if key.startswith("transformer_blocks."):
|
||||
print(key)
|
||||
block = int(key.split(".")[1])
|
||||
if block >= transformer_blocks:
|
||||
transformer_blocks = block + 1
|
||||
elif key.startswith("single_transformer_blocks."):
|
||||
block = int(key.split(".")[1])
|
||||
if block >= single_transformer_blocks:
|
||||
single_transformer_blocks = block + 1
|
||||
|
||||
print(f"Transformer blocks: {transformer_blocks}")
|
||||
print(f"Single transformer blocks: {single_transformer_blocks}")
|
||||
|
||||
for file in original_safetensors:
|
||||
if not file.exists():
|
||||
print(f"Error: Missing transformer safetensors file: {file}")
|
||||
exit()
|
||||
|
||||
original_safetensors = {f: safetensors.safe_open(
|
||||
f, framework="pt", device="cpu") for f in original_safetensors}
|
||||
|
||||
|
||||
def swap_scale_shift(weight):
|
||||
shift, scale = weight.chunk(2, dim=0)
|
||||
new_weight = torch.cat([scale, shift], dim=0)
|
||||
return new_weight
|
||||
|
||||
|
||||
flux_values = {}
|
||||
|
||||
for b in range(transformer_blocks):
|
||||
for key, weights in diffusers_map.items():
|
||||
if key.startswith("double_blocks."):
|
||||
block_prefix = f"transformer_blocks.{b}."
|
||||
found = True
|
||||
for weight in weights:
|
||||
if not (f"{block_prefix}{weight}" in diffusers):
|
||||
found = False
|
||||
if found:
|
||||
flux_values[key.replace("()", f"{b}")] = [
|
||||
f"{block_prefix}{weight}" for weight in weights]
|
||||
for b in range(single_transformer_blocks):
|
||||
for key, weights in diffusers_map.items():
|
||||
if key.startswith("single_blocks."):
|
||||
block_prefix = f"single_transformer_blocks.{b}."
|
||||
found = True
|
||||
for weight in weights:
|
||||
if not (f"{block_prefix}{weight}" in diffusers):
|
||||
found = False
|
||||
if found:
|
||||
flux_values[key.replace("()", f"{b}")] = [
|
||||
f"{block_prefix}{weight}" for weight in weights]
|
||||
|
||||
for key, weights in diffusers_map.items():
|
||||
if not (key.startswith("double_blocks.") or key.startswith("single_blocks.")):
|
||||
found = True
|
||||
for weight in weights:
|
||||
if not (f"{weight}" in diffusers):
|
||||
found = False
|
||||
if found:
|
||||
flux_values[key] = [f"{weight}" for weight in weights]
|
||||
|
||||
flux = {}
|
||||
|
||||
for key, values in tqdm.tqdm(flux_values.items()):
|
||||
if len(values) == 1:
|
||||
flux[key] = original_safetensors[diffusers[values[0]]
|
||||
].get_tensor(values[0]).to("cpu")
|
||||
else:
|
||||
flux[key] = torch.cat(
|
||||
[
|
||||
original_safetensors[diffusers[value]
|
||||
].get_tensor(value).to("cpu")
|
||||
for value in values
|
||||
]
|
||||
)
|
||||
|
||||
if "norm_out.linear.weight" in diffusers:
|
||||
flux["final_layer.adaLN_modulation.1.weight"] = swap_scale_shift(
|
||||
original_safetensors[diffusers["norm_out.linear.weight"]].get_tensor(
|
||||
"norm_out.linear.weight").to("cpu")
|
||||
)
|
||||
if "norm_out.linear.bias" in diffusers:
|
||||
flux["final_layer.adaLN_modulation.1.bias"] = swap_scale_shift(
|
||||
original_safetensors[diffusers["norm_out.linear.bias"]].get_tensor(
|
||||
"norm_out.linear.bias").to("cpu")
|
||||
)
|
||||
|
||||
|
||||
def stochastic_round_to(tensor, dtype=torch.float8_e4m3fn):
|
||||
# Define the float8 range
|
||||
min_val = torch.finfo(dtype).min
|
||||
max_val = torch.finfo(dtype).max
|
||||
|
||||
# Clip values to float8 range
|
||||
tensor = torch.clamp(tensor, min_val, max_val)
|
||||
|
||||
# Convert to float32 for calculations
|
||||
tensor = tensor.float()
|
||||
|
||||
# Get the nearest representable float8 values
|
||||
lower = torch.floor(tensor * 256) / 256
|
||||
upper = torch.ceil(tensor * 256) / 256
|
||||
|
||||
# Calculate the probability of rounding up
|
||||
prob = (tensor - lower) / (upper - lower)
|
||||
|
||||
# Generate random values for stochastic rounding
|
||||
rand = torch.rand_like(tensor)
|
||||
|
||||
# Perform stochastic rounding
|
||||
rounded = torch.where(rand < prob, upper, lower)
|
||||
|
||||
# Convert back to float8
|
||||
return rounded.to(dtype)
|
||||
|
||||
|
||||
# set all the keys to bf16
|
||||
for key in flux.keys():
|
||||
if do_8_bit:
|
||||
flux[key] = stochastic_round_to(
|
||||
flux[key], torch.float8_e4m3fn).to('cpu')
|
||||
else:
|
||||
flux[key] = flux[key].clone().to('cpu', torch.bfloat16)
|
||||
|
||||
# load the quantized state dict
|
||||
quantized_state_dict = safetensors.torch.load_file(quantized_state_dict_path)
|
||||
|
||||
transformer_pre = "model.diffusion_model."
|
||||
did_print = False
|
||||
# remove old parts
|
||||
for key in list(quantized_state_dict.keys()):
|
||||
if key.startswith(transformer_pre):
|
||||
if not did_print:
|
||||
# print("dtype: ", quantized_state_dict[key].dtype)
|
||||
did_print = True
|
||||
del quantized_state_dict[key]
|
||||
|
||||
# add the new parts
|
||||
for key, value in flux.items():
|
||||
quantized_state_dict[transformer_pre + key] = value
|
||||
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = 'pt'
|
||||
# date format like 2024-08-01 YYYY-MM-DD
|
||||
meta['modelspec.date'] = date.today().strftime("%Y-%m-%d")
|
||||
meta['modelspec.title'] = "Flex.1-alpha"
|
||||
meta['modelspec.author'] = "Ostris, LLC"
|
||||
meta['modelspec.license'] = "Apache-2.0"
|
||||
meta['modelspec.implementation'] = "https://github.com/black-forest-labs/flux"
|
||||
meta['modelspec.architecture'] = "Flex.1-alpha"
|
||||
meta['modelspec.description'] = "Flex.1-alpha"
|
||||
|
||||
|
||||
os.makedirs(os.path.dirname(flux_path), exist_ok=True)
|
||||
|
||||
print(f"Saving to {flux_path}")
|
||||
|
||||
safetensors.torch.save_file(quantized_state_dict, flux_path, metadata=meta)
|
||||
|
||||
print("Done.")
|
||||
91
scripts/convert_lora_to_peft_format.py
Normal file
91
scripts/convert_lora_to_peft_format.py
Normal file
@@ -0,0 +1,91 @@
|
||||
# currently only works with flux as support is not quite there yet
|
||||
|
||||
import argparse
|
||||
import os.path
|
||||
from collections import OrderedDict
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
'input_path',
|
||||
type=str,
|
||||
help='Path to original sdxl model'
|
||||
)
|
||||
parser.add_argument(
|
||||
'output_path',
|
||||
type=str,
|
||||
help='output path'
|
||||
)
|
||||
args = parser.parse_args()
|
||||
args.input_path = os.path.abspath(args.input_path)
|
||||
args.output_path = os.path.abspath(args.output_path)
|
||||
|
||||
from safetensors.torch import load_file, save_file
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = 'pt'
|
||||
|
||||
state_dict = load_file(args.input_path)
|
||||
|
||||
# peft doesnt have an alpha so we need to scale the weights
|
||||
alpha_keys = [
|
||||
'lora_transformer_single_transformer_blocks_0_attn_to_q.alpha' # flux
|
||||
]
|
||||
|
||||
# keys where the rank is in the first dimension
|
||||
rank_idx0_keys = [
|
||||
'lora_transformer_single_transformer_blocks_0_attn_to_q.lora_down.weight'
|
||||
# 'transformer.single_transformer_blocks.0.attn.to_q.lora_A.weight'
|
||||
]
|
||||
|
||||
alpha = None
|
||||
rank = None
|
||||
|
||||
for key in rank_idx0_keys:
|
||||
if key in state_dict:
|
||||
rank = int(state_dict[key].shape[0])
|
||||
break
|
||||
|
||||
if rank is None:
|
||||
raise ValueError(f'Could not find rank in state dict')
|
||||
|
||||
for key in alpha_keys:
|
||||
if key in state_dict:
|
||||
alpha = int(state_dict[key])
|
||||
break
|
||||
|
||||
if alpha is None:
|
||||
# set to rank if not found
|
||||
alpha = rank
|
||||
|
||||
|
||||
up_multiplier = alpha / rank
|
||||
|
||||
new_state_dict = {}
|
||||
|
||||
for key, value in state_dict.items():
|
||||
if key.endswith('.alpha'):
|
||||
continue
|
||||
|
||||
orig_dtype = value.dtype
|
||||
|
||||
new_val = value.float() * up_multiplier
|
||||
|
||||
new_key = key
|
||||
new_key = new_key.replace('lora_transformer_', 'transformer.')
|
||||
for i in range(100):
|
||||
new_key = new_key.replace(f'transformer_blocks_{i}_', f'transformer_blocks.{i}.')
|
||||
new_key = new_key.replace('lora_down', 'lora_A')
|
||||
new_key = new_key.replace('lora_up', 'lora_B')
|
||||
new_key = new_key.replace('_lora', '.lora')
|
||||
new_key = new_key.replace('attn_', 'attn.')
|
||||
new_key = new_key.replace('ff_', 'ff.')
|
||||
new_key = new_key.replace('context_net_', 'context.net.')
|
||||
new_key = new_key.replace('0_proj', '0.proj')
|
||||
new_key = new_key.replace('norm_linear', 'norm.linear')
|
||||
new_key = new_key.replace('norm_out_linear', 'norm_out.linear')
|
||||
new_key = new_key.replace('to_out_', 'to_out.')
|
||||
|
||||
new_state_dict[new_key] = new_val.to(orig_dtype)
|
||||
|
||||
save_file(new_state_dict, args.output_path, meta)
|
||||
print(f'Saved to {args.output_path}')
|
||||
245
scripts/extract_lora_from_flex.py
Normal file
245
scripts/extract_lora_from_flex.py
Normal file
@@ -0,0 +1,245 @@
|
||||
import os
|
||||
from tqdm import tqdm
|
||||
import argparse
|
||||
from collections import OrderedDict
|
||||
|
||||
parser = argparse.ArgumentParser(description="Extract LoRA from Flex")
|
||||
parser.add_argument("--base", type=str, default="ostris/Flex.1-alpha", help="Base model path")
|
||||
parser.add_argument("--tuned", type=str, required=True, help="Tuned model path")
|
||||
parser.add_argument("--output", type=str, required=True, help="Output path for lora")
|
||||
parser.add_argument("--rank", type=int, default=32, help="LoRA rank for extraction")
|
||||
parser.add_argument("--gpu", type=int, default=0, help="GPU to process extraction")
|
||||
parser.add_argument("--full", action="store_true", help="Do a full transformer extraction, not just transformer blocks")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if True:
|
||||
# set cuda environment variable
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from lycoris.utils import extract_linear, extract_conv, make_sparse
|
||||
from diffusers import FluxTransformer2DModel
|
||||
|
||||
base = args.base
|
||||
tuned = args.tuned
|
||||
output_path = args.output
|
||||
dim = args.rank
|
||||
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
|
||||
state_dict_base = {}
|
||||
state_dict_tuned = {}
|
||||
|
||||
output_dict = {}
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_diff(
|
||||
base_unet,
|
||||
db_unet,
|
||||
mode="fixed",
|
||||
linear_mode_param=0,
|
||||
conv_mode_param=0,
|
||||
extract_device="cpu",
|
||||
use_bias=False,
|
||||
sparsity=0.98,
|
||||
# small_conv=True,
|
||||
small_conv=False,
|
||||
):
|
||||
UNET_TARGET_REPLACE_MODULE = [
|
||||
"Linear",
|
||||
"Conv2d",
|
||||
"LayerNorm",
|
||||
"GroupNorm",
|
||||
"GroupNorm32",
|
||||
"LoRACompatibleLinear",
|
||||
"LoRACompatibleConv"
|
||||
]
|
||||
LORA_PREFIX_UNET = "transformer"
|
||||
|
||||
def make_state_dict(
|
||||
prefix,
|
||||
root_module: torch.nn.Module,
|
||||
target_module: torch.nn.Module,
|
||||
target_replace_modules,
|
||||
):
|
||||
loras = {}
|
||||
temp = {}
|
||||
|
||||
for name, module in root_module.named_modules():
|
||||
if module.__class__.__name__ in target_replace_modules:
|
||||
temp[name] = module
|
||||
|
||||
for name, module in tqdm(
|
||||
list((n, m) for n, m in target_module.named_modules() if n in temp)
|
||||
):
|
||||
weights = temp[name]
|
||||
lora_name = prefix + "." + name
|
||||
# lora_name = lora_name.replace(".", "_")
|
||||
layer = module.__class__.__name__
|
||||
if 'transformer_blocks' not in lora_name and not args.full:
|
||||
continue
|
||||
|
||||
if layer in {
|
||||
"Linear",
|
||||
"Conv2d",
|
||||
"LayerNorm",
|
||||
"GroupNorm",
|
||||
"GroupNorm32",
|
||||
"Embedding",
|
||||
"LoRACompatibleLinear",
|
||||
"LoRACompatibleConv"
|
||||
}:
|
||||
root_weight = module.weight
|
||||
try:
|
||||
if torch.allclose(root_weight, weights.weight):
|
||||
continue
|
||||
except:
|
||||
continue
|
||||
else:
|
||||
continue
|
||||
module = module.to(extract_device, torch.float32)
|
||||
weights = weights.to(extract_device, torch.float32)
|
||||
|
||||
if mode == "full":
|
||||
decompose_mode = "full"
|
||||
elif layer == "Linear":
|
||||
weight, decompose_mode = extract_linear(
|
||||
(root_weight - weights.weight),
|
||||
mode,
|
||||
linear_mode_param,
|
||||
device=extract_device,
|
||||
)
|
||||
if decompose_mode == "low rank":
|
||||
extract_a, extract_b, diff = weight
|
||||
elif layer == "Conv2d":
|
||||
is_linear = root_weight.shape[2] == 1 and root_weight.shape[3] == 1
|
||||
weight, decompose_mode = extract_conv(
|
||||
(root_weight - weights.weight),
|
||||
mode,
|
||||
linear_mode_param if is_linear else conv_mode_param,
|
||||
device=extract_device,
|
||||
)
|
||||
if decompose_mode == "low rank":
|
||||
extract_a, extract_b, diff = weight
|
||||
if small_conv and not is_linear and decompose_mode == "low rank":
|
||||
dim = extract_a.size(0)
|
||||
(extract_c, extract_a, _), _ = extract_conv(
|
||||
extract_a.transpose(0, 1),
|
||||
"fixed",
|
||||
dim,
|
||||
extract_device,
|
||||
True,
|
||||
)
|
||||
extract_a = extract_a.transpose(0, 1)
|
||||
extract_c = extract_c.transpose(0, 1)
|
||||
loras[f"{lora_name}.lora_mid.weight"] = (
|
||||
extract_c.detach().cpu().contiguous().half()
|
||||
)
|
||||
diff = (
|
||||
(
|
||||
root_weight
|
||||
- torch.einsum(
|
||||
"i j k l, j r, p i -> p r k l",
|
||||
extract_c,
|
||||
extract_a.flatten(1, -1),
|
||||
extract_b.flatten(1, -1),
|
||||
)
|
||||
)
|
||||
.detach()
|
||||
.cpu()
|
||||
.contiguous()
|
||||
)
|
||||
del extract_c
|
||||
else:
|
||||
module = module.to("cpu")
|
||||
weights = weights.to("cpu")
|
||||
continue
|
||||
|
||||
if decompose_mode == "low rank":
|
||||
loras[f"{lora_name}.lora_A.weight"] = (
|
||||
extract_a.detach().cpu().contiguous().half()
|
||||
)
|
||||
loras[f"{lora_name}.lora_B.weight"] = (
|
||||
extract_b.detach().cpu().contiguous().half()
|
||||
)
|
||||
# loras[f"{lora_name}.alpha"] = torch.Tensor([extract_a.shape[0]]).half()
|
||||
if use_bias:
|
||||
diff = diff.detach().cpu().reshape(extract_b.size(0), -1)
|
||||
sparse_diff = make_sparse(diff, sparsity).to_sparse().coalesce()
|
||||
|
||||
indices = sparse_diff.indices().to(torch.int16)
|
||||
values = sparse_diff.values().half()
|
||||
loras[f"{lora_name}.bias_indices"] = indices
|
||||
loras[f"{lora_name}.bias_values"] = values
|
||||
loras[f"{lora_name}.bias_size"] = torch.tensor(diff.shape).to(
|
||||
torch.int16
|
||||
)
|
||||
del extract_a, extract_b, diff
|
||||
elif decompose_mode == "full":
|
||||
if "Norm" in layer:
|
||||
w_key = "w_norm"
|
||||
b_key = "b_norm"
|
||||
else:
|
||||
w_key = "diff"
|
||||
b_key = "diff_b"
|
||||
weight_diff = module.weight - weights.weight
|
||||
loras[f"{lora_name}.{w_key}"] = (
|
||||
weight_diff.detach().cpu().contiguous().half()
|
||||
)
|
||||
if getattr(weights, "bias", None) is not None:
|
||||
bias_diff = module.bias - weights.bias
|
||||
loras[f"{lora_name}.{b_key}"] = (
|
||||
bias_diff.detach().cpu().contiguous().half()
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
module = module.to("cpu", torch.bfloat16)
|
||||
weights = weights.to("cpu", torch.bfloat16)
|
||||
return loras
|
||||
|
||||
all_loras = {}
|
||||
|
||||
all_loras |= make_state_dict(
|
||||
LORA_PREFIX_UNET,
|
||||
base_unet,
|
||||
db_unet,
|
||||
UNET_TARGET_REPLACE_MODULE,
|
||||
)
|
||||
del base_unet, db_unet
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
all_lora_name = set()
|
||||
for k in all_loras:
|
||||
lora_name, weight = k.rsplit(".", 1)
|
||||
all_lora_name.add(lora_name)
|
||||
print(len(all_lora_name))
|
||||
return all_loras
|
||||
|
||||
|
||||
# find all the .safetensors files and load them
|
||||
print("Loading Base")
|
||||
base_model = FluxTransformer2DModel.from_pretrained(base, subfolder="transformer", torch_dtype=torch.bfloat16)
|
||||
|
||||
print("Loading Tuned")
|
||||
tuned_model = FluxTransformer2DModel.from_pretrained(tuned, subfolder="transformer", torch_dtype=torch.bfloat16)
|
||||
|
||||
output_dict = extract_diff(
|
||||
base_model,
|
||||
tuned_model,
|
||||
mode="fixed",
|
||||
linear_mode_param=dim,
|
||||
conv_mode_param=dim,
|
||||
extract_device="cuda",
|
||||
use_bias=False,
|
||||
sparsity=0.98,
|
||||
small_conv=False,
|
||||
)
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = 'pt'
|
||||
|
||||
save_file(output_dict, output_path, metadata=meta)
|
||||
|
||||
print("Done")
|
||||
20
scripts/generate_sampler_step_scales.py
Normal file
20
scripts/generate_sampler_step_scales.py
Normal file
@@ -0,0 +1,20 @@
|
||||
import argparse
|
||||
import torch
|
||||
import os
|
||||
from diffusers import StableDiffusionPipeline
|
||||
import sys
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
# add project root to path
|
||||
sys.path.append(PROJECT_ROOT)
|
||||
|
||||
SAMPLER_SCALES_ROOT = os.path.join(PROJECT_ROOT, 'toolkit', 'samplers_scales')
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser(description='Process some images.')
|
||||
add_arg = parser.add_argument
|
||||
add_arg('--model', type=str, required=True, help='Path to model')
|
||||
add_arg('--sampler', type=str, required=True, help='Name of sampler')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
import argparse
|
||||
from collections import OrderedDict
|
||||
import sys
|
||||
import os
|
||||
ROOT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
sys.path.append(ROOT_DIR)
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
42
scripts/patch_te_adapter.py
Normal file
42
scripts/patch_te_adapter.py
Normal file
@@ -0,0 +1,42 @@
|
||||
import torch
|
||||
from safetensors.torch import save_file, load_file
|
||||
from collections import OrderedDict
|
||||
meta = OrderedDict()
|
||||
meta["format"] ="pt"
|
||||
|
||||
attn_dict = load_file("/mnt/Train/out/ip_adapter/sd15_bigG/sd15_bigG_000266000.safetensors")
|
||||
state_dict = load_file("/home/jaret/Dev/models/hf/OstrisDiffusionV1/unet/diffusion_pytorch_model.safetensors")
|
||||
|
||||
attn_list = []
|
||||
for key, value in state_dict.items():
|
||||
if "attn1" in key:
|
||||
attn_list.append(key)
|
||||
|
||||
attn_names = ['down_blocks.0.attentions.0.transformer_blocks.0.attn2.processor', 'down_blocks.0.attentions.1.transformer_blocks.0.attn2.processor', 'down_blocks.1.attentions.0.transformer_blocks.0.attn2.processor', 'down_blocks.1.attentions.1.transformer_blocks.0.attn2.processor', 'down_blocks.2.attentions.0.transformer_blocks.0.attn2.processor', 'down_blocks.2.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.1.attentions.0.transformer_blocks.0.attn2.processor', 'up_blocks.1.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.1.attentions.2.transformer_blocks.0.attn2.processor', 'up_blocks.2.attentions.0.transformer_blocks.0.attn2.processor', 'up_blocks.2.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.2.attentions.2.transformer_blocks.0.attn2.processor', 'up_blocks.3.attentions.0.transformer_blocks.0.attn2.processor', 'up_blocks.3.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.3.attentions.2.transformer_blocks.0.attn2.processor', 'mid_block.attentions.0.transformer_blocks.0.attn2.processor']
|
||||
|
||||
adapter_names = []
|
||||
for i in range(100):
|
||||
if f'te_adapter.adapter_modules.{i}.to_k_adapter.weight' in attn_dict:
|
||||
adapter_names.append(f"te_adapter.adapter_modules.{i}.adapter")
|
||||
|
||||
|
||||
for i in range(len(adapter_names)):
|
||||
adapter_name = adapter_names[i]
|
||||
attn_name = attn_names[i]
|
||||
adapter_k_name = adapter_name[:-8] + '.to_k_adapter.weight'
|
||||
adapter_v_name = adapter_name[:-8] + '.to_v_adapter.weight'
|
||||
state_k_name = attn_name.replace(".processor", ".to_k.weight")
|
||||
state_v_name = attn_name.replace(".processor", ".to_v.weight")
|
||||
if adapter_k_name in attn_dict:
|
||||
state_dict[state_k_name] = attn_dict[adapter_k_name]
|
||||
state_dict[state_v_name] = attn_dict[adapter_v_name]
|
||||
else:
|
||||
print("adapter_k_name", adapter_k_name)
|
||||
print("state_k_name", state_k_name)
|
||||
|
||||
for key, value in state_dict.items():
|
||||
state_dict[key] = value.cpu().to(torch.float16)
|
||||
|
||||
save_file(state_dict, "/home/jaret/Dev/models/hf/OstrisDiffusionV1/unet/diffusion_pytorch_model.safetensors", metadata=meta)
|
||||
|
||||
print("Done")
|
||||
65
scripts/repair_dataset_folder.py
Normal file
65
scripts/repair_dataset_folder.py
Normal file
@@ -0,0 +1,65 @@
|
||||
import argparse
|
||||
from PIL import Image
|
||||
from PIL.ImageOps import exif_transpose
|
||||
from tqdm import tqdm
|
||||
import os
|
||||
|
||||
parser = argparse.ArgumentParser(description='Process some images.')
|
||||
parser.add_argument("input_folder", type=str, help="Path to folder containing images")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
img_types = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
|
||||
# find all images in the input folder
|
||||
images = []
|
||||
for root, _, files in os.walk(args.input_folder):
|
||||
for file in files:
|
||||
if file.lower().endswith(tuple(img_types)):
|
||||
images.append(os.path.join(root, file))
|
||||
print(f"Found {len(images)} images")
|
||||
|
||||
num_skipped = 0
|
||||
num_repaired = 0
|
||||
num_deleted = 0
|
||||
|
||||
pbar = tqdm(total=len(images), desc=f"Repaired {num_repaired} images", unit="image")
|
||||
for img_path in images:
|
||||
filename = os.path.basename(img_path)
|
||||
filename_no_ext, file_extension = os.path.splitext(filename)
|
||||
# if it is jpg, ignore
|
||||
if file_extension.lower() == '.jpg':
|
||||
num_skipped += 1
|
||||
pbar.update(1)
|
||||
|
||||
continue
|
||||
|
||||
try:
|
||||
img = Image.open(img_path)
|
||||
except Exception as e:
|
||||
print(f"Error opening {img_path}: {e}")
|
||||
# delete it
|
||||
os.remove(img_path)
|
||||
num_deleted += 1
|
||||
pbar.update(1)
|
||||
pbar.set_description(f"Repaired {num_repaired} images, Skipped {num_skipped}, Deleted {num_deleted}")
|
||||
continue
|
||||
|
||||
|
||||
try:
|
||||
img = exif_transpose(img)
|
||||
except Exception as e:
|
||||
print(f"Error rotating {img_path}: {e}")
|
||||
|
||||
new_path = os.path.join(os.path.dirname(img_path), filename_no_ext + '.jpg')
|
||||
|
||||
img = img.convert("RGB")
|
||||
img.save(new_path, quality=95)
|
||||
# remove the old file
|
||||
os.remove(img_path)
|
||||
num_repaired += 1
|
||||
pbar.update(1)
|
||||
# update pbar
|
||||
pbar.set_description(f"Repaired {num_repaired} images, Skipped {num_skipped}, Deleted {num_deleted}")
|
||||
|
||||
print("Done")
|
||||
309
scripts/update_sponsors.py
Normal file
309
scripts/update_sponsors.py
Normal file
@@ -0,0 +1,309 @@
|
||||
import os
|
||||
import requests
|
||||
import json
|
||||
from datetime import datetime
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# Load environment variables from .env file
|
||||
env_path = os.path.join(os.path.dirname(os.path.dirname(__file__)), ".env")
|
||||
load_dotenv(dotenv_path=env_path)
|
||||
|
||||
# API credentials
|
||||
PATREON_TOKEN = os.getenv("PATREON_ACCESS_TOKEN")
|
||||
GITHUB_TOKEN = os.getenv("GITHUB_TOKEN")
|
||||
GITHUB_USERNAME = os.getenv("GITHUB_USERNAME")
|
||||
GITHUB_ORG = os.getenv("GITHUB_ORG") # Organization name (optional)
|
||||
|
||||
# Output file
|
||||
README_PATH = "SUPPORTERS.md"
|
||||
|
||||
def fetch_patreon_supporters():
|
||||
"""Fetch current Patreon supporters"""
|
||||
print("Fetching Patreon supporters...")
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {PATREON_TOKEN}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
url = "https://www.patreon.com/api/oauth2/v2/campaigns"
|
||||
|
||||
try:
|
||||
# First get the campaign ID
|
||||
campaign_response = requests.get(url, headers=headers)
|
||||
campaign_response.raise_for_status()
|
||||
campaign_data = campaign_response.json()
|
||||
|
||||
if not campaign_data.get('data'):
|
||||
print("No campaigns found for this Patreon account")
|
||||
return []
|
||||
|
||||
campaign_id = campaign_data['data'][0]['id']
|
||||
|
||||
# Now get the supporters for this campaign
|
||||
members_url = f"https://www.patreon.com/api/oauth2/v2/campaigns/{campaign_id}/members"
|
||||
params = {
|
||||
"include": "user",
|
||||
"fields[member]": "full_name,is_follower,patron_status", # Removed profile_url
|
||||
"fields[user]": "image_url"
|
||||
}
|
||||
|
||||
supporters = []
|
||||
while members_url:
|
||||
members_response = requests.get(members_url, headers=headers, params=params)
|
||||
members_response.raise_for_status()
|
||||
members_data = members_response.json()
|
||||
|
||||
# Process the response to extract active patrons
|
||||
for member in members_data.get('data', []):
|
||||
attributes = member.get('attributes', {})
|
||||
|
||||
# Only include active patrons
|
||||
if attributes.get('patron_status') == 'active_patron':
|
||||
name = attributes.get('full_name', 'Anonymous Supporter')
|
||||
|
||||
# Get user data which contains the profile image
|
||||
user_id = member.get('relationships', {}).get('user', {}).get('data', {}).get('id')
|
||||
profile_image = None
|
||||
profile_url = None # Removed profile_url since it's not supported
|
||||
|
||||
if user_id:
|
||||
for included in members_data.get('included', []):
|
||||
if included.get('id') == user_id and included.get('type') == 'user':
|
||||
profile_image = included.get('attributes', {}).get('image_url')
|
||||
break
|
||||
|
||||
supporters.append({
|
||||
'name': name,
|
||||
'profile_image': profile_image,
|
||||
'profile_url': profile_url, # This will be None
|
||||
'platform': 'Patreon',
|
||||
'amount': 0 # Placeholder, as Patreon API doesn't provide this in the current response
|
||||
})
|
||||
|
||||
# Handle pagination
|
||||
members_url = members_data.get('links', {}).get('next')
|
||||
|
||||
print(f"Found {len(supporters)} active Patreon supporters")
|
||||
return supporters
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error fetching Patreon data: {e}")
|
||||
print(f"Response content: {e.response.content if hasattr(e, 'response') else 'No response content'}")
|
||||
return []
|
||||
|
||||
def fetch_github_sponsors():
|
||||
"""Fetch current GitHub sponsors for a user or organization"""
|
||||
print("Fetching GitHub sponsors...")
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {GITHUB_TOKEN}",
|
||||
"Accept": "application/vnd.github.v3+json"
|
||||
}
|
||||
|
||||
# Determine if we're fetching for a user or an organization
|
||||
entity_type = "organization" if GITHUB_ORG else "user"
|
||||
entity_name = GITHUB_ORG if GITHUB_ORG else GITHUB_USERNAME
|
||||
|
||||
if not entity_name:
|
||||
print("Error: Neither GITHUB_USERNAME nor GITHUB_ORG is set")
|
||||
return []
|
||||
|
||||
# Different GraphQL query structure based on entity type
|
||||
if entity_type == "user":
|
||||
query = """
|
||||
query {
|
||||
user(login: "%s") {
|
||||
sponsorshipsAsMaintainer(first: 100) {
|
||||
nodes {
|
||||
sponsorEntity {
|
||||
... on User {
|
||||
login
|
||||
name
|
||||
avatarUrl
|
||||
url
|
||||
}
|
||||
... on Organization {
|
||||
login
|
||||
name
|
||||
avatarUrl
|
||||
url
|
||||
}
|
||||
}
|
||||
tier {
|
||||
monthlyPriceInDollars
|
||||
}
|
||||
isOneTimePayment
|
||||
isActive
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
""" % entity_name
|
||||
else: # organization
|
||||
query = """
|
||||
query {
|
||||
organization(login: "%s") {
|
||||
sponsorshipsAsMaintainer(first: 100) {
|
||||
nodes {
|
||||
sponsorEntity {
|
||||
... on User {
|
||||
login
|
||||
name
|
||||
avatarUrl
|
||||
url
|
||||
}
|
||||
... on Organization {
|
||||
login
|
||||
name
|
||||
avatarUrl
|
||||
url
|
||||
}
|
||||
}
|
||||
tier {
|
||||
monthlyPriceInDollars
|
||||
}
|
||||
isOneTimePayment
|
||||
isActive
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
""" % entity_name
|
||||
|
||||
try:
|
||||
response = requests.post(
|
||||
"https://api.github.com/graphql",
|
||||
headers=headers,
|
||||
json={"query": query}
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
# Process the response - the path to the data differs based on entity type
|
||||
if entity_type == "user":
|
||||
sponsors_data = data.get('data', {}).get('user', {}).get('sponsorshipsAsMaintainer', {}).get('nodes', [])
|
||||
else:
|
||||
sponsors_data = data.get('data', {}).get('organization', {}).get('sponsorshipsAsMaintainer', {}).get('nodes', [])
|
||||
|
||||
sponsors = []
|
||||
for sponsor in sponsors_data:
|
||||
# Only include active sponsors
|
||||
if sponsor.get('isActive'):
|
||||
entity = sponsor.get('sponsorEntity', {})
|
||||
name = entity.get('name') or entity.get('login', 'Anonymous Sponsor')
|
||||
profile_image = entity.get('avatarUrl')
|
||||
profile_url = entity.get('url')
|
||||
amount = sponsor.get('tier', {}).get('monthlyPriceInDollars', 0)
|
||||
|
||||
sponsors.append({
|
||||
'name': name,
|
||||
'profile_image': profile_image,
|
||||
'profile_url': profile_url,
|
||||
'platform': 'GitHub Sponsors',
|
||||
'amount': amount
|
||||
})
|
||||
|
||||
print(f"Found {len(sponsors)} active GitHub sponsors for {entity_type} '{entity_name}'")
|
||||
return sponsors
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error fetching GitHub sponsors data: {e}")
|
||||
return []
|
||||
|
||||
def generate_readme(supporters):
|
||||
"""Generate a README.md file with supporter information"""
|
||||
print(f"Generating {README_PATH}...")
|
||||
|
||||
# Sort supporters by amount (descending) and then by name
|
||||
supporters.sort(key=lambda x: (-x['amount'], x['name'].lower()))
|
||||
|
||||
# Determine the proper footer links based on what's configured
|
||||
github_entity = GITHUB_ORG if GITHUB_ORG else GITHUB_USERNAME
|
||||
github_entity_type = "orgs" if GITHUB_ORG else "sponsors"
|
||||
github_sponsor_url = f"https://github.com/{github_entity_type}/{github_entity}"
|
||||
|
||||
with open(README_PATH, "w", encoding="utf-8") as f:
|
||||
f.write("## Support My Work\n\n")
|
||||
f.write("If you enjoy my work, or use it for commercial purposes, please consider sponsoring me so I can continue to maintain it. Every bit helps! \n\n")
|
||||
# Create appropriate call-to-action based on what's configured
|
||||
cta_parts = []
|
||||
if github_entity:
|
||||
cta_parts.append(f"[Become a sponsor on GitHub]({github_sponsor_url})")
|
||||
if PATREON_TOKEN:
|
||||
cta_parts.append("[support me on Patreon](https://www.patreon.com/ostris)")
|
||||
|
||||
if cta_parts:
|
||||
if GITHUB_ORG:
|
||||
f.write(f"{' or '.join(cta_parts)}.\n\n")
|
||||
f.write("Thank you to all my current supporters!\n\n")
|
||||
|
||||
f.write(f"_Last updated: {datetime.now().strftime('%Y-%m-%d')}_\n\n")
|
||||
|
||||
# Write GitHub Sponsors section
|
||||
github_sponsors = [s for s in supporters if s['platform'] == 'GitHub Sponsors']
|
||||
if github_sponsors:
|
||||
f.write("### GitHub Sponsors\n\n")
|
||||
for sponsor in github_sponsors:
|
||||
if sponsor['profile_image']:
|
||||
f.write(f"<a href=\"{sponsor['profile_url']}\" title=\"{sponsor['name']}\"><img src=\"{sponsor['profile_image']}\" width=\"50\" height=\"50\" alt=\"{sponsor['name']}\" style=\"border-radius:50%\"></a> ")
|
||||
else:
|
||||
f.write(f"[{sponsor['name']}]({sponsor['profile_url']}) ")
|
||||
f.write("\n\n")
|
||||
|
||||
# Write Patreon section
|
||||
patreon_supporters = [s for s in supporters if s['platform'] == 'Patreon']
|
||||
if patreon_supporters:
|
||||
f.write("### Patreon Supporters\n\n")
|
||||
for supporter in patreon_supporters:
|
||||
if supporter['profile_image']:
|
||||
f.write(f"<a href=\"{supporter['profile_url']}\" title=\"{supporter['name']}\"><img src=\"{supporter['profile_image']}\" width=\"50\" height=\"50\" alt=\"{supporter['name']}\" style=\"border-radius:50%\"></a> ")
|
||||
else:
|
||||
f.write(f"[{supporter['name']}]({supporter['profile_url']}) ")
|
||||
f.write("\n\n")
|
||||
|
||||
f.write("\n---\n\n")
|
||||
|
||||
|
||||
print(f"Successfully generated {README_PATH} with {len(supporters)} supporters!")
|
||||
|
||||
def main():
|
||||
"""Main function"""
|
||||
print("Starting supporter data collection...")
|
||||
|
||||
# Check if required environment variables are set
|
||||
missing_vars = []
|
||||
if not GITHUB_TOKEN:
|
||||
missing_vars.append("GITHUB_TOKEN")
|
||||
|
||||
# Either username or org is required for GitHub
|
||||
if not GITHUB_USERNAME and not GITHUB_ORG:
|
||||
missing_vars.append("GITHUB_USERNAME or GITHUB_ORG")
|
||||
|
||||
# Patreon token is optional but warn if missing
|
||||
patreon_enabled = bool(PATREON_TOKEN)
|
||||
|
||||
if missing_vars:
|
||||
print(f"Error: Missing required environment variables: {', '.join(missing_vars)}")
|
||||
print("Please add them to your .env file")
|
||||
return
|
||||
|
||||
if not patreon_enabled:
|
||||
print("Warning: PATREON_ACCESS_TOKEN not set. Will only fetch GitHub sponsors.")
|
||||
|
||||
# Fetch data from both platforms
|
||||
patreon_supporters = fetch_patreon_supporters() if PATREON_TOKEN else []
|
||||
github_sponsors = fetch_github_sponsors()
|
||||
|
||||
# Combine supporters from both platforms
|
||||
all_supporters = patreon_supporters + github_sponsors
|
||||
|
||||
if not all_supporters:
|
||||
print("No supporters found on either platform")
|
||||
return
|
||||
|
||||
# Generate README
|
||||
generate_readme(all_supporters)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -54,6 +54,7 @@ parser.add_argument('--name', type=str, default='stable_diffusion', help='name f
|
||||
parser.add_argument('--sdxl', action='store_true', help='is sdxl model')
|
||||
parser.add_argument('--refiner', action='store_true', help='is refiner model')
|
||||
parser.add_argument('--ssd', action='store_true', help='is ssd model')
|
||||
parser.add_argument('--vega', action='store_true', help='is vega model')
|
||||
parser.add_argument('--sd2', action='store_true', help='is sd 2 model')
|
||||
|
||||
args = parser.parse_args()
|
||||
@@ -66,15 +67,15 @@ print(f'Loading diffusers model')
|
||||
|
||||
ignore_ldm_begins_with = []
|
||||
|
||||
diffusers_file_path = file_path
|
||||
diffusers_file_path = file_path if len(args.file_1) == 1 else args.file_1[1]
|
||||
if args.ssd:
|
||||
diffusers_file_path = "segmind/SSD-1B"
|
||||
if args.vega:
|
||||
diffusers_file_path = "segmind/Segmind-Vega"
|
||||
|
||||
# if args.refiner:
|
||||
# diffusers_file_path = "stabilityai/stable-diffusion-xl-refiner-1.0"
|
||||
|
||||
diffusers_file_path = file_path if len(args.file_1) == 1 else args.file_1[1]
|
||||
|
||||
if not args.refiner:
|
||||
|
||||
diffusers_model_config = ModelConfig(
|
||||
@@ -82,6 +83,7 @@ if not args.refiner:
|
||||
is_xl=args.sdxl,
|
||||
is_v2=args.sd2,
|
||||
is_ssd=args.ssd,
|
||||
is_vega=args.vega,
|
||||
dtype=dtype,
|
||||
)
|
||||
diffusers_sd = StableDiffusion(
|
||||
@@ -157,7 +159,7 @@ te_suffix = ''
|
||||
proj_pattern_weight = None
|
||||
proj_pattern_bias = None
|
||||
text_proj_layer = None
|
||||
if args.sdxl or args.ssd:
|
||||
if args.sdxl or args.ssd or args.vega:
|
||||
te_suffix = '1'
|
||||
ldm_res_block_prefix = "conditioner.embedders.1.model.transformer.resblocks"
|
||||
proj_pattern_weight = r"conditioner\.embedders\.1\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_weight"
|
||||
@@ -176,10 +178,13 @@ if args.sd2:
|
||||
proj_pattern_bias = r"cond_stage_model\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_bias"
|
||||
text_proj_layer = "cond_stage_model.model.text_projection"
|
||||
|
||||
if args.sdxl or args.sd2 or args.ssd or args.refiner:
|
||||
if args.sdxl or args.sd2 or args.ssd or args.refiner or args.vega:
|
||||
if "conditioner.embedders.1.model.text_projection" in ldm_dict_keys:
|
||||
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
|
||||
d_model = int(ldm_state_dict["conditioner.embedders.1.model.text_projection"].shape[0])
|
||||
elif "conditioner.embedders.1.model.text_projection.weight" in ldm_dict_keys:
|
||||
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
|
||||
d_model = int(ldm_state_dict["conditioner.embedders.1.model.text_projection.weight"].shape[0])
|
||||
elif "conditioner.embedders.0.model.text_projection" in ldm_dict_keys:
|
||||
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
|
||||
d_model = int(ldm_state_dict["conditioner.embedders.0.model.text_projection"].shape[0])
|
||||
@@ -191,6 +196,8 @@ if args.sdxl or args.sd2 or args.ssd or args.refiner:
|
||||
try:
|
||||
match = re.match(proj_pattern_weight, ldm_key)
|
||||
if match:
|
||||
if ldm_key == "conditioner.embedders.1.model.transformer.resblocks.0.attn.in_proj_weight":
|
||||
print("here")
|
||||
number = int(match.group(1))
|
||||
new_val = torch.cat([
|
||||
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight"],
|
||||
@@ -217,6 +224,8 @@ if args.sdxl or args.sd2 or args.ssd or args.refiner:
|
||||
],
|
||||
}
|
||||
|
||||
matched_ldm_keys.append(ldm_key)
|
||||
|
||||
# text_model_dict[new_key + ".q_proj.weight"] = checkpoint[key][:d_model, :]
|
||||
# text_model_dict[new_key + ".k_proj.weight"] = checkpoint[key][d_model: d_model * 2, :]
|
||||
# text_model_dict[new_key + ".v_proj.weight"] = checkpoint[key][d_model * 2:, :]
|
||||
@@ -266,6 +275,8 @@ if args.sdxl or args.sd2 or args.ssd or args.refiner:
|
||||
],
|
||||
}
|
||||
|
||||
matched_ldm_keys.append(ldm_key)
|
||||
|
||||
# add diffusers operators
|
||||
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.bias"] = {
|
||||
"slice": [
|
||||
@@ -298,6 +309,9 @@ for ldm_key in ldm_dict_keys:
|
||||
ldm_shape_tuple = ldm_state_dict[ldm_key].shape
|
||||
ldm_reduced_shape_tuple = get_reduced_shape(ldm_shape_tuple)
|
||||
for diffusers_key in diffusers_dict_keys:
|
||||
if ldm_key == "conditioner.embedders.1.model.transformer.resblocks.0.attn.in_proj_weight" and diffusers_key == "te1_text_model.encoder.layers.0.self_attn.q_proj.weight":
|
||||
print("here")
|
||||
|
||||
diffusers_shape_tuple = diffusers_state_dict[diffusers_key].shape
|
||||
diffusers_reduced_shape_tuple = get_reduced_shape(diffusers_shape_tuple)
|
||||
|
||||
@@ -356,6 +370,8 @@ if args.sdxl:
|
||||
name += '_sdxl'
|
||||
elif args.ssd:
|
||||
name += '_ssd'
|
||||
elif args.vega:
|
||||
name += '_vega'
|
||||
elif args.refiner:
|
||||
name += '_refiner'
|
||||
elif args.sd2:
|
||||
|
||||
180
testing/merge_in_text_encoder_adapter.py
Normal file
180
testing/merge_in_text_encoder_adapter.py
Normal file
@@ -0,0 +1,180 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
from transformers import T5EncoderModel, T5Tokenizer
|
||||
from diffusers import StableDiffusionPipeline, UNet2DConditionModel, PixArtSigmaPipeline, Transformer2DModel, PixArtTransformer2DModel
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
import json
|
||||
|
||||
# model_path = "/home/jaret/Dev/models/hf/kl-f16-d42_sd15_v01_000527000"
|
||||
# te_path = "google/flan-t5-xl"
|
||||
# te_aug_path = "/mnt/Train/out/ip_adapter/t5xx_sd15_v1/t5xx_sd15_v1_000032000.safetensors"
|
||||
# output_path = "/home/jaret/Dev/models/hf/kl-f16-d42_sd15_t5xl_raw"
|
||||
model_path = "/home/jaret/Dev/models/hf/objective-reality-16ch"
|
||||
te_path = "google/flan-t5-xl"
|
||||
te_aug_path = "/mnt/Train2/out/ip_adapter/t5xl-sd15-16ch_v1/t5xl-sd15-16ch_v1_000115000.safetensors"
|
||||
output_path = "/home/jaret/Dev/models/hf/t5xl-sd15-16ch_sd15_v1"
|
||||
|
||||
|
||||
print("Loading te adapter")
|
||||
te_aug_sd = load_file(te_aug_path)
|
||||
|
||||
print("Loading model")
|
||||
is_diffusers = (not os.path.exists(model_path)) or os.path.isdir(model_path)
|
||||
|
||||
# if "pixart" in model_path.lower():
|
||||
is_pixart = "pixart" in model_path.lower()
|
||||
|
||||
pipeline_class = StableDiffusionPipeline
|
||||
|
||||
# transformer = PixArtTransformer2DModel.from_pretrained('PixArt-alpha/PixArt-Sigma-XL-2-512-MS', subfolder='transformer', torch_dtype=torch.float16)
|
||||
|
||||
if is_pixart:
|
||||
pipeline_class = PixArtSigmaPipeline
|
||||
|
||||
if is_diffusers:
|
||||
sd = pipeline_class.from_pretrained(model_path, torch_dtype=torch.float16)
|
||||
else:
|
||||
sd = pipeline_class.from_single_file(model_path, torch_dtype=torch.float16)
|
||||
|
||||
print("Loading Text Encoder")
|
||||
# Load the text encoder
|
||||
te = T5EncoderModel.from_pretrained(te_path, torch_dtype=torch.float16)
|
||||
|
||||
# patch it
|
||||
sd.text_encoder = te
|
||||
sd.tokenizer = T5Tokenizer.from_pretrained(te_path)
|
||||
|
||||
if is_pixart:
|
||||
unet = sd.transformer
|
||||
unet_sd = sd.transformer.state_dict()
|
||||
else:
|
||||
unet = sd.unet
|
||||
unet_sd = sd.unet.state_dict()
|
||||
|
||||
|
||||
if is_pixart:
|
||||
weight_idx = 0
|
||||
else:
|
||||
weight_idx = 1
|
||||
|
||||
new_cross_attn_dim = None
|
||||
|
||||
# count the num of params in state dict
|
||||
start_params = sum([v.numel() for v in unet_sd.values()])
|
||||
|
||||
print("Building")
|
||||
attn_processor_keys = []
|
||||
if is_pixart:
|
||||
transformer: Transformer2DModel = unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn1")
|
||||
# cross attention
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn2")
|
||||
else:
|
||||
attn_processor_keys = list(unet.attn_processors.keys())
|
||||
|
||||
for name in attn_processor_keys:
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") or name.endswith("attn.1") or name.endswith(
|
||||
"attn1") else \
|
||||
unet.config['cross_attention_dim']
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config['block_out_channels'][-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config['block_out_channels']))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config['block_out_channels'][block_id]
|
||||
elif name.startswith("transformer"):
|
||||
hidden_size = unet.config['cross_attention_dim']
|
||||
else:
|
||||
# they didnt have this, but would lead to undefined below
|
||||
raise ValueError(f"unknown attn processor name: {name}")
|
||||
if cross_attention_dim is None:
|
||||
pass
|
||||
else:
|
||||
layer_name = name.split(".processor")[0]
|
||||
to_k_adapter = unet_sd[layer_name + ".to_k.weight"]
|
||||
to_v_adapter = unet_sd[layer_name + ".to_v.weight"]
|
||||
|
||||
te_aug_name = None
|
||||
while True:
|
||||
if is_pixart:
|
||||
te_aug_name = f"te_adapter.adapter_modules.{weight_idx}.to_k_adapter"
|
||||
else:
|
||||
te_aug_name = f"te_adapter.adapter_modules.{weight_idx}.to_k_adapter"
|
||||
if f"{te_aug_name}.weight" in te_aug_sd:
|
||||
# increment so we dont redo it next time
|
||||
weight_idx += 1
|
||||
break
|
||||
else:
|
||||
weight_idx += 1
|
||||
|
||||
if weight_idx > 1000:
|
||||
raise ValueError("Could not find the next weight")
|
||||
|
||||
orig_weight_shape_k = list(unet_sd[layer_name + ".to_k.weight"].shape)
|
||||
new_weight_shape_k = list(te_aug_sd[te_aug_name + ".weight"].shape)
|
||||
orig_weight_shape_v = list(unet_sd[layer_name + ".to_v.weight"].shape)
|
||||
new_weight_shape_v = list(te_aug_sd[te_aug_name.replace('to_k', 'to_v') + ".weight"].shape)
|
||||
|
||||
unet_sd[layer_name + ".to_k.weight"] = te_aug_sd[te_aug_name + ".weight"]
|
||||
unet_sd[layer_name + ".to_v.weight"] = te_aug_sd[te_aug_name.replace('to_k', 'to_v') + ".weight"]
|
||||
|
||||
if new_cross_attn_dim is None:
|
||||
new_cross_attn_dim = unet_sd[layer_name + ".to_k.weight"].shape[1]
|
||||
|
||||
|
||||
|
||||
if is_pixart:
|
||||
# copy the caption_projection weight
|
||||
del unet_sd['caption_projection.linear_1.bias']
|
||||
del unet_sd['caption_projection.linear_1.weight']
|
||||
del unet_sd['caption_projection.linear_2.bias']
|
||||
del unet_sd['caption_projection.linear_2.weight']
|
||||
|
||||
print("Saving unmodified model")
|
||||
sd = sd.to("cpu", torch.float16)
|
||||
sd.save_pretrained(
|
||||
output_path,
|
||||
safe_serialization=True,
|
||||
)
|
||||
|
||||
# overwrite the unet
|
||||
if is_pixart:
|
||||
unet_folder = os.path.join(output_path, "transformer")
|
||||
else:
|
||||
unet_folder = os.path.join(output_path, "unet")
|
||||
|
||||
# move state_dict to cpu
|
||||
unet_sd = {k: v.clone().cpu().to(torch.float16) for k, v in unet_sd.items()}
|
||||
|
||||
meta = OrderedDict()
|
||||
meta["format"] = "pt"
|
||||
|
||||
print("Patching")
|
||||
|
||||
save_file(unet_sd, os.path.join(unet_folder, "diffusion_pytorch_model.safetensors"), meta)
|
||||
|
||||
# load the json file
|
||||
with open(os.path.join(unet_folder, "config.json"), 'r') as f:
|
||||
config = json.load(f)
|
||||
|
||||
config['cross_attention_dim'] = new_cross_attn_dim
|
||||
|
||||
if is_pixart:
|
||||
config['caption_channels'] = None
|
||||
|
||||
# save it
|
||||
with open(os.path.join(unet_folder, "config.json"), 'w') as f:
|
||||
json.dump(config, f, indent=2)
|
||||
|
||||
print("Done")
|
||||
|
||||
new_params = sum([v.numel() for v in unet_sd.values()])
|
||||
|
||||
# print new and old params with , formatted
|
||||
print(f"Old params: {start_params:,}")
|
||||
print(f"New params: {new_params:,}")
|
||||
62
testing/shrink_pixart.py
Normal file
62
testing/shrink_pixart.py
Normal file
@@ -0,0 +1,62 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
model_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model_orig.safetensors"
|
||||
output_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model.safetensors"
|
||||
|
||||
state_dict = load_file(model_path)
|
||||
|
||||
meta = OrderedDict()
|
||||
meta["format"] = "pt"
|
||||
|
||||
new_state_dict = {}
|
||||
|
||||
# Move non-blocks over
|
||||
for key, value in state_dict.items():
|
||||
if not key.startswith("transformer_blocks."):
|
||||
new_state_dict[key] = value
|
||||
|
||||
block_names = ['transformer_blocks.{idx}.attn1.to_k.bias', 'transformer_blocks.{idx}.attn1.to_k.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_out.0.bias', 'transformer_blocks.{idx}.attn1.to_out.0.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_q.bias', 'transformer_blocks.{idx}.attn1.to_q.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_v.bias', 'transformer_blocks.{idx}.attn1.to_v.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_k.bias', 'transformer_blocks.{idx}.attn2.to_k.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_out.0.bias', 'transformer_blocks.{idx}.attn2.to_out.0.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_q.bias', 'transformer_blocks.{idx}.attn2.to_q.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_v.bias', 'transformer_blocks.{idx}.attn2.to_v.weight',
|
||||
'transformer_blocks.{idx}.ff.net.0.proj.bias', 'transformer_blocks.{idx}.ff.net.0.proj.weight',
|
||||
'transformer_blocks.{idx}.ff.net.2.bias', 'transformer_blocks.{idx}.ff.net.2.weight',
|
||||
'transformer_blocks.{idx}.scale_shift_table']
|
||||
|
||||
# New block idx 0, 1, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 27
|
||||
|
||||
current_idx = 0
|
||||
for i in range(28):
|
||||
if i not in [0, 1, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 27]:
|
||||
# todo merge in with previous block
|
||||
for name in block_names:
|
||||
try:
|
||||
new_state_dict_key = name.format(idx=current_idx - 1)
|
||||
old_state_dict_key = name.format(idx=i)
|
||||
new_state_dict[new_state_dict_key] = (new_state_dict[new_state_dict_key] * 0.5) + (state_dict[old_state_dict_key] * 0.5)
|
||||
except KeyError:
|
||||
raise KeyError(f"KeyError: {name.format(idx=current_idx)}")
|
||||
else:
|
||||
for name in block_names:
|
||||
new_state_dict[name.format(idx=current_idx)] = state_dict[name.format(idx=i)]
|
||||
current_idx += 1
|
||||
|
||||
|
||||
# make sure they are all fp16 and on cpu
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.to(torch.float16).cpu()
|
||||
|
||||
# save the new state dict
|
||||
save_file(new_state_dict, output_path, metadata=meta)
|
||||
|
||||
new_param_count = sum([v.numel() for v in new_state_dict.values()])
|
||||
old_param_count = sum([v.numel() for v in state_dict.values()])
|
||||
|
||||
print(f"Old param count: {old_param_count:,}")
|
||||
print(f"New param count: {new_param_count:,}")
|
||||
81
testing/shrink_pixart2.py
Normal file
81
testing/shrink_pixart2.py
Normal file
@@ -0,0 +1,81 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
model_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model_orig.safetensors"
|
||||
output_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model.safetensors"
|
||||
|
||||
state_dict = load_file(model_path)
|
||||
|
||||
meta = OrderedDict()
|
||||
meta["format"] = "pt"
|
||||
|
||||
new_state_dict = {}
|
||||
|
||||
# Move non-blocks over
|
||||
for key, value in state_dict.items():
|
||||
if not key.startswith("transformer_blocks."):
|
||||
new_state_dict[key] = value
|
||||
|
||||
block_names = ['transformer_blocks.{idx}.attn1.to_k.bias', 'transformer_blocks.{idx}.attn1.to_k.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_out.0.bias', 'transformer_blocks.{idx}.attn1.to_out.0.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_q.bias', 'transformer_blocks.{idx}.attn1.to_q.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_v.bias', 'transformer_blocks.{idx}.attn1.to_v.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_k.bias', 'transformer_blocks.{idx}.attn2.to_k.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_out.0.bias', 'transformer_blocks.{idx}.attn2.to_out.0.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_q.bias', 'transformer_blocks.{idx}.attn2.to_q.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_v.bias', 'transformer_blocks.{idx}.attn2.to_v.weight',
|
||||
'transformer_blocks.{idx}.ff.net.0.proj.bias', 'transformer_blocks.{idx}.ff.net.0.proj.weight',
|
||||
'transformer_blocks.{idx}.ff.net.2.bias', 'transformer_blocks.{idx}.ff.net.2.weight',
|
||||
'transformer_blocks.{idx}.scale_shift_table']
|
||||
|
||||
# Blocks to keep
|
||||
# keep_blocks = [0, 1, 2, 6, 10, 14, 18, 22, 26, 27]
|
||||
keep_blocks = [0, 1, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 27]
|
||||
|
||||
|
||||
def weighted_merge(kept_block, removed_block, weight):
|
||||
return kept_block * (1 - weight) + removed_block * weight
|
||||
|
||||
|
||||
# First, copy all kept blocks to new_state_dict
|
||||
for i, old_idx in enumerate(keep_blocks):
|
||||
for name in block_names:
|
||||
old_key = name.format(idx=old_idx)
|
||||
new_key = name.format(idx=i)
|
||||
new_state_dict[new_key] = state_dict[old_key].clone()
|
||||
|
||||
# Then, merge information from removed blocks
|
||||
for i in range(28):
|
||||
if i not in keep_blocks:
|
||||
# Find the nearest kept blocks
|
||||
prev_kept = max([b for b in keep_blocks if b < i])
|
||||
next_kept = min([b for b in keep_blocks if b > i])
|
||||
|
||||
# Calculate the weight based on position
|
||||
weight = (i - prev_kept) / (next_kept - prev_kept)
|
||||
|
||||
for name in block_names:
|
||||
removed_key = name.format(idx=i)
|
||||
prev_new_key = name.format(idx=keep_blocks.index(prev_kept))
|
||||
next_new_key = name.format(idx=keep_blocks.index(next_kept))
|
||||
|
||||
# Weighted merge for previous kept block
|
||||
new_state_dict[prev_new_key] = weighted_merge(new_state_dict[prev_new_key], state_dict[removed_key], weight)
|
||||
|
||||
# Weighted merge for next kept block
|
||||
new_state_dict[next_new_key] = weighted_merge(new_state_dict[next_new_key], state_dict[removed_key],
|
||||
1 - weight)
|
||||
|
||||
# Convert to fp16 and move to CPU
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.to(torch.float16).cpu()
|
||||
|
||||
# Save the new state dict
|
||||
save_file(new_state_dict, output_path, metadata=meta)
|
||||
|
||||
new_param_count = sum([v.numel() for v in new_state_dict.values()])
|
||||
old_param_count = sum([v.numel() for v in state_dict.values()])
|
||||
|
||||
print(f"Old param count: {old_param_count:,}")
|
||||
print(f"New param count: {new_param_count:,}")
|
||||
84
testing/shrink_pixart_sm.py
Normal file
84
testing/shrink_pixart_sm.py
Normal file
@@ -0,0 +1,84 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = "pt"
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
def reduce_weight(weight, target_size):
|
||||
weight = weight.to(device, torch.float32)
|
||||
original_shape = weight.shape
|
||||
flattened = weight.view(-1, original_shape[-1])
|
||||
|
||||
if flattened.shape[1] <= target_size:
|
||||
return weight
|
||||
|
||||
U, S, V = torch.svd(flattened)
|
||||
reduced = torch.mm(U[:, :target_size], torch.diag(S[:target_size]))
|
||||
|
||||
if reduced.shape[1] < target_size:
|
||||
padding = torch.zeros(reduced.shape[0], target_size - reduced.shape[1], device=device)
|
||||
reduced = torch.cat((reduced, padding), dim=1)
|
||||
|
||||
return reduced.view(original_shape[:-1] + (target_size,))
|
||||
|
||||
|
||||
def reduce_bias(bias, target_size):
|
||||
bias = bias.to(device, torch.float32)
|
||||
original_size = bias.shape[0]
|
||||
|
||||
if original_size <= target_size:
|
||||
return torch.nn.functional.pad(bias, (0, target_size - original_size))
|
||||
else:
|
||||
return bias.view(-1, original_size // target_size).mean(dim=1)[:target_size]
|
||||
|
||||
|
||||
# Load your original state dict
|
||||
state_dict = load_file(
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.orig.safetensors")
|
||||
|
||||
# Create a new state dict for the reduced model
|
||||
new_state_dict = {}
|
||||
|
||||
source_hidden_size = 1152
|
||||
target_hidden_size = 1024
|
||||
|
||||
for key, value in state_dict.items():
|
||||
value = value.to(device, torch.float32)
|
||||
if 'weight' in key or 'scale_shift_table' in key:
|
||||
if value.shape[0] == source_hidden_size:
|
||||
value = value[:target_hidden_size]
|
||||
elif value.shape[0] == source_hidden_size * 4:
|
||||
value = value[:target_hidden_size * 4]
|
||||
elif value.shape[0] == source_hidden_size * 6:
|
||||
value = value[:target_hidden_size * 6]
|
||||
|
||||
if len(value.shape) > 1 and value.shape[
|
||||
1] == source_hidden_size and 'attn2.to_k.weight' not in key and 'attn2.to_v.weight' not in key:
|
||||
value = value[:, :target_hidden_size]
|
||||
elif len(value.shape) > 1 and value.shape[1] == source_hidden_size * 4:
|
||||
value = value[:, :target_hidden_size * 4]
|
||||
|
||||
elif 'bias' in key:
|
||||
if value.shape[0] == source_hidden_size:
|
||||
value = value[:target_hidden_size]
|
||||
elif value.shape[0] == source_hidden_size * 4:
|
||||
value = value[:target_hidden_size * 4]
|
||||
elif value.shape[0] == source_hidden_size * 6:
|
||||
value = value[:target_hidden_size * 6]
|
||||
|
||||
new_state_dict[key] = value
|
||||
|
||||
# Move all to CPU and convert to float16
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.cpu().to(torch.float16)
|
||||
|
||||
# Save the new state dict
|
||||
save_file(new_state_dict,
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.safetensors",
|
||||
metadata=meta)
|
||||
|
||||
print("Done!")
|
||||
110
testing/shrink_pixart_sm2.py
Normal file
110
testing/shrink_pixart_sm2.py
Normal file
@@ -0,0 +1,110 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = "pt"
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
def reduce_weight(weight, target_size):
|
||||
weight = weight.to(device, torch.float32)
|
||||
original_shape = weight.shape
|
||||
|
||||
if len(original_shape) == 1:
|
||||
# For 1D tensors, simply truncate
|
||||
return weight[:target_size]
|
||||
|
||||
if original_shape[0] <= target_size:
|
||||
return weight
|
||||
|
||||
# Reshape the tensor to 2D
|
||||
flattened = weight.reshape(original_shape[0], -1)
|
||||
|
||||
# Perform SVD
|
||||
U, S, V = torch.svd(flattened)
|
||||
|
||||
# Reduce the dimensions
|
||||
reduced = torch.mm(U[:target_size, :], torch.diag(S)).mm(V.t())
|
||||
|
||||
# Reshape back to the original shape with reduced first dimension
|
||||
new_shape = (target_size,) + original_shape[1:]
|
||||
return reduced.reshape(new_shape)
|
||||
|
||||
|
||||
def reduce_bias(bias, target_size):
|
||||
bias = bias.to(device, torch.float32)
|
||||
return bias[:target_size]
|
||||
|
||||
|
||||
# Load your original state dict
|
||||
state_dict = load_file(
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.orig.safetensors")
|
||||
|
||||
# Create a new state dict for the reduced model
|
||||
new_state_dict = {}
|
||||
|
||||
for key, value in state_dict.items():
|
||||
value = value.to(device, torch.float32)
|
||||
|
||||
if 'weight' in key or 'scale_shift_table' in key:
|
||||
if value.shape[0] == 1152:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (512, orig_shape[1], orig_shape[2], orig_shape[3]) # reshape to (1152, -1)
|
||||
# reshape to (1152, -1)
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 512)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
# value = reduce_weight(value.t(), 576).t().contiguous()
|
||||
value = reduce_weight(value, 512)
|
||||
pass
|
||||
elif value.shape[0] == 4608:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (2048, orig_shape[1], orig_shape[2], orig_shape[3])
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 2048)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
value = reduce_weight(value, 2048)
|
||||
elif value.shape[0] == 6912:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (3072, orig_shape[1], orig_shape[2], orig_shape[3])
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 3072)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
value = reduce_weight(value, 3072)
|
||||
|
||||
if len(value.shape) > 1 and value.shape[
|
||||
1] == 1152 and 'attn2.to_k.weight' not in key and 'attn2.to_v.weight' not in key:
|
||||
value = reduce_weight(value.t(), 512).t().contiguous() # Transpose before and after reduction
|
||||
pass
|
||||
elif len(value.shape) > 1 and value.shape[1] == 4608:
|
||||
value = reduce_weight(value.t(), 2048).t().contiguous() # Transpose before and after reduction
|
||||
pass
|
||||
|
||||
elif 'bias' in key:
|
||||
if value.shape[0] == 1152:
|
||||
value = reduce_bias(value, 512)
|
||||
elif value.shape[0] == 4608:
|
||||
value = reduce_bias(value, 2048)
|
||||
elif value.shape[0] == 6912:
|
||||
value = reduce_bias(value, 3072)
|
||||
|
||||
new_state_dict[key] = value
|
||||
|
||||
# Move all to CPU and convert to float16
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.cpu().to(torch.float16)
|
||||
|
||||
# Save the new state dict
|
||||
save_file(new_state_dict,
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.safetensors",
|
||||
metadata=meta)
|
||||
|
||||
print("Done!")
|
||||
100
testing/shrink_pixart_sm3.py
Normal file
100
testing/shrink_pixart_sm3.py
Normal file
@@ -0,0 +1,100 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = "pt"
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
def reduce_weight(weight, target_size):
|
||||
weight = weight.to(device, torch.float32)
|
||||
# resize so target_size is the first dimension
|
||||
tmp_weight = weight.view(1, 1, weight.shape[0], weight.shape[1])
|
||||
|
||||
# use interpolate to resize the tensor
|
||||
new_weight = torch.nn.functional.interpolate(tmp_weight, size=(target_size, weight.shape[1]), mode='bicubic', align_corners=True)
|
||||
|
||||
# reshape back to original shape
|
||||
return new_weight.view(target_size, weight.shape[1])
|
||||
|
||||
|
||||
def reduce_bias(bias, target_size):
|
||||
bias = bias.view(1, 1, bias.shape[0], 1)
|
||||
|
||||
new_bias = torch.nn.functional.interpolate(bias, size=(target_size, 1), mode='bicubic', align_corners=True)
|
||||
|
||||
return new_bias.view(target_size)
|
||||
|
||||
|
||||
# Load your original state dict
|
||||
state_dict = load_file(
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.orig.safetensors")
|
||||
|
||||
# Create a new state dict for the reduced model
|
||||
new_state_dict = {}
|
||||
|
||||
for key, value in state_dict.items():
|
||||
value = value.to(device, torch.float32)
|
||||
|
||||
if 'weight' in key or 'scale_shift_table' in key:
|
||||
if value.shape[0] == 1152:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (512, orig_shape[1], orig_shape[2], orig_shape[3]) # reshape to (1152, -1)
|
||||
# reshape to (1152, -1)
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 512)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
# value = reduce_weight(value.t(), 576).t().contiguous()
|
||||
value = reduce_weight(value, 512)
|
||||
pass
|
||||
elif value.shape[0] == 4608:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (2048, orig_shape[1], orig_shape[2], orig_shape[3])
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 2048)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
value = reduce_weight(value, 2048)
|
||||
elif value.shape[0] == 6912:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (3072, orig_shape[1], orig_shape[2], orig_shape[3])
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 3072)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
value = reduce_weight(value, 3072)
|
||||
|
||||
if len(value.shape) > 1 and value.shape[
|
||||
1] == 1152 and 'attn2.to_k.weight' not in key and 'attn2.to_v.weight' not in key:
|
||||
value = reduce_weight(value.t(), 512).t().contiguous() # Transpose before and after reduction
|
||||
pass
|
||||
elif len(value.shape) > 1 and value.shape[1] == 4608:
|
||||
value = reduce_weight(value.t(), 2048).t().contiguous() # Transpose before and after reduction
|
||||
pass
|
||||
|
||||
elif 'bias' in key:
|
||||
if value.shape[0] == 1152:
|
||||
value = reduce_bias(value, 512)
|
||||
elif value.shape[0] == 4608:
|
||||
value = reduce_bias(value, 2048)
|
||||
elif value.shape[0] == 6912:
|
||||
value = reduce_bias(value, 3072)
|
||||
|
||||
new_state_dict[key] = value
|
||||
|
||||
# Move all to CPU and convert to float16
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.cpu().to(torch.float16)
|
||||
|
||||
# Save the new state dict
|
||||
save_file(new_state_dict,
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.safetensors",
|
||||
metadata=meta)
|
||||
|
||||
print("Done!")
|
||||
@@ -7,11 +7,13 @@ from torchvision import transforms
|
||||
import sys
|
||||
import os
|
||||
import cv2
|
||||
import random
|
||||
from transformers import CLIPImageProcessor
|
||||
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from toolkit.paths import SD_SCRIPTS_ROOT
|
||||
|
||||
from toolkit.image_utils import show_img
|
||||
import torchvision.transforms.functional
|
||||
from toolkit.image_utils import save_tensors, show_img, show_tensors
|
||||
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
@@ -21,83 +23,118 @@ from toolkit.data_loader import AiToolkitDataset, get_dataloader_from_datasets,
|
||||
trigger_dataloader_setup_epoch
|
||||
from toolkit.config_modules import DatasetConfig
|
||||
import argparse
|
||||
from tqdm import tqdm
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('dataset_folder', type=str, default='input')
|
||||
parser.add_argument('--epochs', type=int, default=1)
|
||||
parser.add_argument('--num_frames', type=int, default=1)
|
||||
parser.add_argument('--output_path', type=str, default=None)
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.output_path is not None:
|
||||
args.output_path = os.path.abspath(args.output_path)
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
|
||||
dataset_folder = args.dataset_folder
|
||||
resolution = 1024
|
||||
resolution = 512
|
||||
bucket_tolerance = 64
|
||||
batch_size = 1
|
||||
|
||||
clip_processor = CLIPImageProcessor.from_pretrained("openai/clip-vit-base-patch16")
|
||||
|
||||
class FakeAdapter:
|
||||
def __init__(self):
|
||||
self.clip_image_processor = clip_processor
|
||||
|
||||
|
||||
## make fake sd
|
||||
class FakeSD:
|
||||
def __init__(self):
|
||||
self.adapter = FakeAdapter()
|
||||
|
||||
|
||||
|
||||
##
|
||||
|
||||
dataset_config = DatasetConfig(
|
||||
dataset_path=dataset_folder,
|
||||
# clip_image_path=dataset_folder,
|
||||
# square_crop=True,
|
||||
resolution=resolution,
|
||||
caption_ext='json',
|
||||
# caption_ext='json',
|
||||
default_caption='default',
|
||||
# clip_image_path='/mnt/Datasets2/regs/yetibear_xl_v14/random_aspect/',
|
||||
buckets=True,
|
||||
bucket_tolerance=bucket_tolerance,
|
||||
poi='person',
|
||||
augmentations=[
|
||||
{
|
||||
'method': 'RandomBrightnessContrast',
|
||||
'brightness_limit': (-0.3, 0.3),
|
||||
'contrast_limit': (-0.3, 0.3),
|
||||
'brightness_by_max': False,
|
||||
'p': 1.0
|
||||
},
|
||||
{
|
||||
'method': 'HueSaturationValue',
|
||||
'hue_shift_limit': (-0, 0),
|
||||
'sat_shift_limit': (-40, 40),
|
||||
'val_shift_limit': (-40, 40),
|
||||
'p': 1.0
|
||||
},
|
||||
# {
|
||||
# 'method': 'RGBShift',
|
||||
# 'r_shift_limit': (-20, 20),
|
||||
# 'g_shift_limit': (-20, 20),
|
||||
# 'b_shift_limit': (-20, 20),
|
||||
# 'p': 1.0
|
||||
# },
|
||||
]
|
||||
|
||||
|
||||
shrink_video_to_frames=True,
|
||||
num_frames=args.num_frames,
|
||||
# poi='person',
|
||||
# shuffle_augmentations=True,
|
||||
# augmentations=[
|
||||
# {
|
||||
# 'method': 'Posterize',
|
||||
# 'num_bits': [(0, 4), (0, 4), (0, 4)],
|
||||
# 'p': 1.0
|
||||
# },
|
||||
#
|
||||
# ]
|
||||
)
|
||||
|
||||
dataloader: DataLoader = get_dataloader_from_datasets([dataset_config], batch_size=batch_size)
|
||||
dataloader: DataLoader = get_dataloader_from_datasets([dataset_config], batch_size=batch_size, sd=FakeSD())
|
||||
|
||||
|
||||
# run through an epoch ang check sizes
|
||||
dataloader_iterator = iter(dataloader)
|
||||
idx = 0
|
||||
for epoch in range(args.epochs):
|
||||
for batch in dataloader:
|
||||
for batch in tqdm(dataloader):
|
||||
batch: 'DataLoaderBatchDTO'
|
||||
img_batch = batch.tensor
|
||||
frames = 1
|
||||
if len(img_batch.shape) == 5:
|
||||
frames = img_batch.shape[1]
|
||||
batch_size, frames, channels, height, width = img_batch.shape
|
||||
else:
|
||||
batch_size, channels, height, width = img_batch.shape
|
||||
|
||||
chunks = torch.chunk(img_batch, batch_size, dim=0)
|
||||
# put them so they are size by side
|
||||
big_img = torch.cat(chunks, dim=3)
|
||||
big_img = big_img.squeeze(0)
|
||||
# img_batch = color_block_imgs(img_batch, neg1_1=True)
|
||||
|
||||
min_val = big_img.min()
|
||||
max_val = big_img.max()
|
||||
# chunks = torch.chunk(img_batch, batch_size, dim=0)
|
||||
# # put them so they are size by side
|
||||
# big_img = torch.cat(chunks, dim=3)
|
||||
# big_img = big_img.squeeze(0)
|
||||
#
|
||||
# control_chunks = torch.chunk(batch.clip_image_tensor, batch_size, dim=0)
|
||||
# big_control_img = torch.cat(control_chunks, dim=3)
|
||||
# big_control_img = big_control_img.squeeze(0) * 2 - 1
|
||||
#
|
||||
#
|
||||
# # resize control image
|
||||
# big_control_img = torchvision.transforms.Resize((width, height))(big_control_img)
|
||||
#
|
||||
# big_img = torch.cat([big_img, big_control_img], dim=2)
|
||||
#
|
||||
# min_val = big_img.min()
|
||||
# max_val = big_img.max()
|
||||
#
|
||||
# big_img = (big_img / 2 + 0.5).clamp(0, 1)
|
||||
|
||||
big_img = (big_img / 2 + 0.5).clamp(0, 1)
|
||||
big_img = img_batch
|
||||
# big_img = big_img.clamp(-1, 1)
|
||||
if args.output_path is not None:
|
||||
save_tensors(big_img, os.path.join(args.output_path, f'{idx}.png'))
|
||||
else:
|
||||
show_tensors(big_img)
|
||||
|
||||
# convert to image
|
||||
img = transforms.ToPILImage()(big_img)
|
||||
# convert to image
|
||||
# img = transforms.ToPILImage()(big_img)
|
||||
#
|
||||
# show_img(img)
|
||||
|
||||
show_img(img)
|
||||
|
||||
time.sleep(1.0)
|
||||
time.sleep(0.2)
|
||||
idx += 1
|
||||
# if not last epoch
|
||||
if epoch < args.epochs - 1:
|
||||
trigger_dataloader_setup_epoch(dataloader)
|
||||
|
||||
130
testing/test_vae.py
Normal file
130
testing/test_vae.py
Normal file
@@ -0,0 +1,130 @@
|
||||
import argparse
|
||||
import os
|
||||
from PIL import Image
|
||||
import torch
|
||||
from torchvision.transforms import Resize, ToTensor
|
||||
from diffusers import AutoencoderKL
|
||||
from pytorch_fid import fid_score
|
||||
from skimage.metrics import peak_signal_noise_ratio as psnr
|
||||
import lpips
|
||||
from tqdm import tqdm
|
||||
from torchvision import transforms
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
def load_images(folder_path):
|
||||
images = []
|
||||
for filename in os.listdir(folder_path):
|
||||
if filename.lower().endswith(('.png', '.jpg', '.jpeg')):
|
||||
img_path = os.path.join(folder_path, filename)
|
||||
images.append(img_path)
|
||||
return images
|
||||
|
||||
|
||||
def paramiter_count(model):
|
||||
state_dict = model.state_dict()
|
||||
paramiter_count = 0
|
||||
for key in state_dict:
|
||||
paramiter_count += torch.numel(state_dict[key])
|
||||
return int(paramiter_count)
|
||||
|
||||
|
||||
def calculate_metrics(vae, images, max_imgs=-1, save_output=False):
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
vae = vae.to(device)
|
||||
lpips_model = lpips.LPIPS(net='alex').to(device)
|
||||
|
||||
rfid_scores = []
|
||||
psnr_scores = []
|
||||
lpips_scores = []
|
||||
|
||||
# transform = transforms.Compose([
|
||||
# transforms.Resize(256, antialias=True),
|
||||
# transforms.CenterCrop(256)
|
||||
# ])
|
||||
# needs values between -1 and 1
|
||||
to_tensor = ToTensor()
|
||||
|
||||
# remove _reconstructed.png files
|
||||
images = [img for img in images if not img.endswith("_reconstructed.png")]
|
||||
|
||||
if max_imgs > 0 and len(images) > max_imgs:
|
||||
images = images[:max_imgs]
|
||||
|
||||
for img_path in tqdm(images):
|
||||
try:
|
||||
img = Image.open(img_path).convert('RGB')
|
||||
# img_tensor = to_tensor(transform(img)).unsqueeze(0).to(device)
|
||||
img_tensor = to_tensor(img).unsqueeze(0).to(device)
|
||||
img_tensor = 2 * img_tensor - 1
|
||||
# if width or height is not divisible by 8, crop it
|
||||
if img_tensor.shape[2] % 8 != 0 or img_tensor.shape[3] % 8 != 0:
|
||||
img_tensor = img_tensor[:, :, :img_tensor.shape[2] // 8 * 8, :img_tensor.shape[3] // 8 * 8]
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error processing {img_path}: {e}")
|
||||
continue
|
||||
|
||||
|
||||
with torch.no_grad():
|
||||
reconstructed = vae.decode(vae.encode(img_tensor).latent_dist.sample()).sample
|
||||
|
||||
# Calculate rFID
|
||||
# rfid = fid_score.calculate_frechet_distance(vae, img_tensor, reconstructed)
|
||||
# rfid_scores.append(rfid)
|
||||
|
||||
# Calculate PSNR
|
||||
psnr_val = psnr(img_tensor.cpu().numpy(), reconstructed.cpu().numpy())
|
||||
psnr_scores.append(psnr_val)
|
||||
|
||||
# Calculate LPIPS
|
||||
lpips_val = lpips_model(img_tensor, reconstructed).item()
|
||||
lpips_scores.append(lpips_val)
|
||||
|
||||
# avg_rfid = sum(rfid_scores) / len(rfid_scores)
|
||||
avg_rfid = 0
|
||||
avg_psnr = sum(psnr_scores) / len(psnr_scores)
|
||||
avg_lpips = sum(lpips_scores) / len(lpips_scores)
|
||||
|
||||
if save_output:
|
||||
filename_no_ext = os.path.splitext(os.path.basename(img_path))[0]
|
||||
folder = os.path.dirname(img_path)
|
||||
save_path = os.path.join(folder, filename_no_ext + "_reconstructed.png")
|
||||
reconstructed = (reconstructed + 1) / 2
|
||||
reconstructed = reconstructed.clamp(0, 1)
|
||||
reconstructed = transforms.ToPILImage()(reconstructed[0].cpu())
|
||||
reconstructed.save(save_path)
|
||||
|
||||
return avg_rfid, avg_psnr, avg_lpips
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Calculate average rFID, PSNR, and LPIPS for VAE reconstructions")
|
||||
parser.add_argument("--vae_path", type=str, required=True, help="Path to the VAE model")
|
||||
parser.add_argument("--image_folder", type=str, required=True, help="Path to the folder containing images")
|
||||
parser.add_argument("--max_imgs", type=int, default=-1, help="Max num of images. Default is -1 for all images.")
|
||||
# boolean store true
|
||||
parser.add_argument("--save_output", action="store_true", help="Save the output images")
|
||||
args = parser.parse_args()
|
||||
|
||||
if os.path.isfile(args.vae_path):
|
||||
vae = AutoencoderKL.from_single_file(args.vae_path)
|
||||
else:
|
||||
try:
|
||||
vae = AutoencoderKL.from_pretrained(args.vae_path)
|
||||
except:
|
||||
vae = AutoencoderKL.from_pretrained(args.vae_path, subfolder="vae")
|
||||
vae.eval()
|
||||
vae = vae.to(device)
|
||||
print(f"Model has {paramiter_count(vae)} parameters")
|
||||
images = load_images(args.image_folder)
|
||||
|
||||
avg_rfid, avg_psnr, avg_lpips = calculate_metrics(vae, images, args.max_imgs, args.save_output)
|
||||
|
||||
# print(f"Average rFID: {avg_rfid}")
|
||||
print(f"Average PSNR: {avg_psnr}")
|
||||
print(f"Average LPIPS: {avg_lpips}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
17
toolkit/accelerator.py
Normal file
17
toolkit/accelerator.py
Normal file
@@ -0,0 +1,17 @@
|
||||
from accelerate import Accelerator
|
||||
from diffusers.utils.torch_utils import is_compiled_module
|
||||
|
||||
global_accelerator = None
|
||||
|
||||
|
||||
def get_accelerator() -> Accelerator:
|
||||
global global_accelerator
|
||||
if global_accelerator is None:
|
||||
global_accelerator = Accelerator()
|
||||
return global_accelerator
|
||||
|
||||
def unwrap_model(model):
|
||||
accelerator = get_accelerator()
|
||||
model = accelerator.unwrap_model(model)
|
||||
model = model._orig_mod if is_compiled_module(model) else model
|
||||
return model
|
||||
55
toolkit/assistant_lora.py
Normal file
55
toolkit/assistant_lora.py
Normal file
@@ -0,0 +1,55 @@
|
||||
from typing import TYPE_CHECKING
|
||||
from toolkit.config_modules import NetworkConfig
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
from safetensors.torch import load_file
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
def load_assistant_lora_from_path(adapter_path, sd: 'StableDiffusion') -> LoRASpecialNetwork:
|
||||
if not sd.is_flux:
|
||||
raise ValueError("Only Flux models can load assistant adapters currently.")
|
||||
pipe = sd.pipeline
|
||||
print(f"Loading assistant adapter from {adapter_path}")
|
||||
adapter_name = adapter_path.split("/")[-1].split(".")[0]
|
||||
lora_state_dict = load_file(adapter_path)
|
||||
|
||||
linear_dim = int(lora_state_dict['transformer.single_transformer_blocks.0.attn.to_k.lora_A.weight'].shape[0])
|
||||
# linear_alpha = int(lora_state_dict['lora_transformer_single_transformer_blocks_0_attn_to_k.alpha'].item())
|
||||
linear_alpha = linear_dim
|
||||
transformer_only = 'transformer.proj_out.alpha' not in lora_state_dict
|
||||
# get dim and scale
|
||||
network_config = NetworkConfig(
|
||||
linear=linear_dim,
|
||||
linear_alpha=linear_alpha,
|
||||
transformer_only=transformer_only,
|
||||
)
|
||||
|
||||
network = LoRASpecialNetwork(
|
||||
text_encoder=pipe.text_encoder,
|
||||
unet=pipe.transformer,
|
||||
lora_dim=network_config.linear,
|
||||
multiplier=1.0,
|
||||
alpha=network_config.linear_alpha,
|
||||
train_unet=True,
|
||||
train_text_encoder=False,
|
||||
is_flux=True,
|
||||
network_config=network_config,
|
||||
network_type=network_config.type,
|
||||
transformer_only=network_config.transformer_only,
|
||||
is_assistant_adapter=True
|
||||
)
|
||||
network.apply_to(
|
||||
pipe.text_encoder,
|
||||
pipe.transformer,
|
||||
apply_text_encoder=False,
|
||||
apply_unet=True
|
||||
)
|
||||
network.force_to(sd.device_torch, dtype=sd.torch_dtype)
|
||||
network.eval()
|
||||
network._update_torch_multiplier()
|
||||
network.load_weights(lora_state_dict)
|
||||
network.is_active = True
|
||||
|
||||
return network
|
||||
@@ -31,12 +31,18 @@ def get_mean_std(tensor):
|
||||
def adain(content_features, style_features):
|
||||
# Assumes that the content and style features are of shape (batch_size, channels, width, height)
|
||||
|
||||
dims = [2, 3]
|
||||
if len(content_features.shape) == 3:
|
||||
# content_features = content_features.unsqueeze(0)
|
||||
# style_features = style_features.unsqueeze(0)
|
||||
dims = [1]
|
||||
|
||||
# Step 1: Calculate mean and variance of content features
|
||||
content_mean, content_var = torch.mean(content_features, dim=[2, 3], keepdim=True), torch.var(content_features,
|
||||
dim=[2, 3],
|
||||
content_mean, content_var = torch.mean(content_features, dim=dims, keepdim=True), torch.var(content_features,
|
||||
dim=dims,
|
||||
keepdim=True)
|
||||
# Step 2: Calculate mean and variance of style features
|
||||
style_mean, style_var = torch.mean(style_features, dim=[2, 3], keepdim=True), torch.var(style_features, dim=[2, 3],
|
||||
style_mean, style_var = torch.mean(style_features, dim=dims, keepdim=True), torch.var(style_features, dim=dims,
|
||||
keepdim=True)
|
||||
|
||||
# Step 3: Normalize content features
|
||||
|
||||
@@ -51,9 +51,11 @@ resolutions_1024: List[BucketResolution] = [
|
||||
{"width": 512, "height": 1920},
|
||||
{"width": 512, "height": 1984},
|
||||
{"width": 512, "height": 2048},
|
||||
# extra wides
|
||||
{"width": 8192, "height": 128},
|
||||
{"width": 128, "height": 8192},
|
||||
]
|
||||
|
||||
|
||||
def get_bucket_sizes(resolution: int = 512, divisibility: int = 8) -> List[BucketResolution]:
|
||||
# determine scaler form 1024 to resolution
|
||||
scaler = resolution / 1024
|
||||
@@ -124,4 +126,4 @@ def get_bucket_for_image_size(
|
||||
if closest_bucket is None:
|
||||
raise ValueError("No suitable bucket found")
|
||||
|
||||
return closest_bucket
|
||||
return closest_bucket
|
||||
406
toolkit/clip_vision_adapter.py
Normal file
406
toolkit/clip_vision_adapter.py
Normal file
@@ -0,0 +1,406 @@
|
||||
from typing import TYPE_CHECKING, Mapping, Any
|
||||
|
||||
import torch
|
||||
import weakref
|
||||
|
||||
from toolkit.config_modules import AdapterConfig
|
||||
from toolkit.models.clip_fusion import ZipperBlock
|
||||
from toolkit.models.zipper_resampler import ZipperModule
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
from transformers import (
|
||||
CLIPImageProcessor,
|
||||
CLIPVisionModelWithProjection,
|
||||
CLIPVisionModel
|
||||
)
|
||||
|
||||
from toolkit.resampler import Resampler
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class Embedder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
num_input_tokens: int = 1,
|
||||
input_dim: int = 1024,
|
||||
num_output_tokens: int = 8,
|
||||
output_dim: int = 768,
|
||||
mid_dim: int = 1024
|
||||
):
|
||||
super(Embedder, self).__init__()
|
||||
self.num_output_tokens = num_output_tokens
|
||||
self.num_input_tokens = num_input_tokens
|
||||
self.input_dim = input_dim
|
||||
self.output_dim = output_dim
|
||||
|
||||
self.layer_norm = nn.LayerNorm(input_dim)
|
||||
self.fc1 = nn.Linear(input_dim, mid_dim)
|
||||
self.gelu = nn.GELU()
|
||||
# self.fc2 = nn.Linear(mid_dim, mid_dim)
|
||||
self.fc2 = nn.Linear(mid_dim, mid_dim)
|
||||
|
||||
self.fc2.weight.data.zero_()
|
||||
|
||||
self.layer_norm2 = nn.LayerNorm(mid_dim)
|
||||
self.fc3 = nn.Linear(mid_dim, mid_dim)
|
||||
self.gelu2 = nn.GELU()
|
||||
self.fc4 = nn.Linear(mid_dim, output_dim * num_output_tokens)
|
||||
|
||||
# set the weights to 0
|
||||
self.fc3.weight.data.zero_()
|
||||
self.fc4.weight.data.zero_()
|
||||
|
||||
|
||||
# self.static_tokens = nn.Parameter(torch.zeros(num_output_tokens, output_dim))
|
||||
# self.scaler = nn.Parameter(torch.zeros(num_output_tokens, output_dim))
|
||||
|
||||
def forward(self, x):
|
||||
if len(x.shape) == 2:
|
||||
x = x.unsqueeze(1)
|
||||
x = self.layer_norm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.gelu(x)
|
||||
x = self.fc2(x)
|
||||
x = self.layer_norm2(x)
|
||||
x = self.fc3(x)
|
||||
x = self.gelu2(x)
|
||||
x = self.fc4(x)
|
||||
|
||||
x = x.view(-1, self.num_output_tokens, self.output_dim)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class ClipVisionAdapter(torch.nn.Module):
|
||||
def __init__(self, sd: 'StableDiffusion', adapter_config: AdapterConfig):
|
||||
super().__init__()
|
||||
self.config = adapter_config
|
||||
self.trigger = adapter_config.trigger
|
||||
self.trigger_class_name = adapter_config.trigger_class_name
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
# embedding stuff
|
||||
self.text_encoder_list = sd.text_encoder if isinstance(sd.text_encoder, list) else [sd.text_encoder]
|
||||
self.tokenizer_list = sd.tokenizer if isinstance(sd.tokenizer, list) else [sd.tokenizer]
|
||||
placeholder_tokens = [self.trigger]
|
||||
|
||||
# add dummy tokens for multi-vector
|
||||
additional_tokens = []
|
||||
for i in range(1, self.config.num_tokens):
|
||||
additional_tokens.append(f"{self.trigger}_{i}")
|
||||
placeholder_tokens += additional_tokens
|
||||
|
||||
# handle dual tokenizer
|
||||
self.tokenizer_list = self.sd_ref().tokenizer if isinstance(self.sd_ref().tokenizer, list) else [
|
||||
self.sd_ref().tokenizer]
|
||||
self.text_encoder_list = self.sd_ref().text_encoder if isinstance(self.sd_ref().text_encoder, list) else [
|
||||
self.sd_ref().text_encoder]
|
||||
|
||||
self.placeholder_token_ids = []
|
||||
self.embedding_tokens = []
|
||||
|
||||
print(f"Adding {placeholder_tokens} tokens to tokenizer")
|
||||
print(f"Adding {self.config.num_tokens} tokens to tokenizer")
|
||||
|
||||
|
||||
for text_encoder, tokenizer in zip(self.text_encoder_list, self.tokenizer_list):
|
||||
num_added_tokens = tokenizer.add_tokens(placeholder_tokens)
|
||||
if num_added_tokens != self.config.num_tokens:
|
||||
raise ValueError(
|
||||
f"The tokenizer already contains the token {self.trigger}. Please pass a different"
|
||||
f" `placeholder_token` that is not already in the tokenizer. Only added {num_added_tokens}"
|
||||
)
|
||||
|
||||
# Convert the initializer_token, placeholder_token to ids
|
||||
init_token_ids = tokenizer.encode(self.config.trigger_class_name, add_special_tokens=False)
|
||||
# if length of token ids is more than number of orm embedding tokens fill with *
|
||||
if len(init_token_ids) > self.config.num_tokens:
|
||||
init_token_ids = init_token_ids[:self.config.num_tokens]
|
||||
elif len(init_token_ids) < self.config.num_tokens:
|
||||
pad_token_id = tokenizer.encode(["*"], add_special_tokens=False)
|
||||
init_token_ids += pad_token_id * (self.config.num_tokens - len(init_token_ids))
|
||||
|
||||
placeholder_token_ids = tokenizer.encode(placeholder_tokens, add_special_tokens=False)
|
||||
self.placeholder_token_ids.append(placeholder_token_ids)
|
||||
|
||||
# Resize the token embeddings as we are adding new special tokens to the tokenizer
|
||||
text_encoder.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
# Initialise the newly added placeholder token with the embeddings of the initializer token
|
||||
token_embeds = text_encoder.get_input_embeddings().weight.data
|
||||
with torch.no_grad():
|
||||
for initializer_token_id, token_id in zip(init_token_ids, placeholder_token_ids):
|
||||
token_embeds[token_id] = token_embeds[initializer_token_id].clone()
|
||||
|
||||
# replace "[name] with this. on training. This is automatically generated in pipeline on inference
|
||||
self.embedding_tokens.append(" ".join(tokenizer.convert_ids_to_tokens(placeholder_token_ids)))
|
||||
|
||||
# backup text encoder embeddings
|
||||
self.orig_embeds_params = [x.get_input_embeddings().weight.data.clone() for x in self.text_encoder_list]
|
||||
|
||||
try:
|
||||
self.clip_image_processor = CLIPImageProcessor.from_pretrained(self.config.image_encoder_path)
|
||||
except EnvironmentError:
|
||||
self.clip_image_processor = CLIPImageProcessor()
|
||||
self.device = self.sd_ref().unet.device
|
||||
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
||||
self.config.image_encoder_path,
|
||||
ignore_mismatched_sizes=True
|
||||
).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
if self.config.train_image_encoder:
|
||||
self.image_encoder.train()
|
||||
else:
|
||||
self.image_encoder.eval()
|
||||
|
||||
# max_seq_len = CLIP tokens + CLS token
|
||||
image_encoder_state_dict = self.image_encoder.state_dict()
|
||||
in_tokens = 257
|
||||
if "vision_model.embeddings.position_embedding.weight" in image_encoder_state_dict:
|
||||
# clip
|
||||
in_tokens = int(image_encoder_state_dict["vision_model.embeddings.position_embedding.weight"].shape[0])
|
||||
|
||||
if hasattr(self.image_encoder.config, 'hidden_sizes'):
|
||||
embedding_dim = self.image_encoder.config.hidden_sizes[-1]
|
||||
else:
|
||||
embedding_dim = self.image_encoder.config.target_hidden_size
|
||||
|
||||
if self.config.clip_layer == 'image_embeds':
|
||||
in_tokens = 1
|
||||
embedding_dim = self.image_encoder.config.projection_dim
|
||||
|
||||
self.embedder = Embedder(
|
||||
num_output_tokens=self.config.num_tokens,
|
||||
num_input_tokens=in_tokens,
|
||||
input_dim=embedding_dim,
|
||||
output_dim=self.sd_ref().unet.config['cross_attention_dim'],
|
||||
mid_dim=embedding_dim * self.config.num_tokens,
|
||||
).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
|
||||
self.embedder.train()
|
||||
|
||||
def state_dict(self, *args, destination=None, prefix='', keep_vars=False):
|
||||
state_dict = {
|
||||
'embedder': self.embedder.state_dict(*args, destination=destination, prefix=prefix, keep_vars=keep_vars)
|
||||
}
|
||||
if self.config.train_image_encoder:
|
||||
state_dict['image_encoder'] = self.image_encoder.state_dict(
|
||||
*args, destination=destination, prefix=prefix,
|
||||
keep_vars=keep_vars)
|
||||
|
||||
return state_dict
|
||||
|
||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict: bool = True):
|
||||
self.embedder.load_state_dict(state_dict["embedder"], strict=strict)
|
||||
if self.config.train_image_encoder and 'image_encoder' in state_dict:
|
||||
self.image_encoder.load_state_dict(state_dict["image_encoder"], strict=strict)
|
||||
|
||||
def parameters(self, *args, **kwargs):
|
||||
yield from self.embedder.parameters(*args, **kwargs)
|
||||
|
||||
def named_parameters(self, *args, **kwargs):
|
||||
yield from self.embedder.named_parameters(*args, **kwargs)
|
||||
|
||||
def get_clip_image_embeds_from_tensors(
|
||||
self, tensors_0_1: torch.Tensor, drop=False,
|
||||
is_training=False,
|
||||
has_been_preprocessed=False
|
||||
) -> torch.Tensor:
|
||||
with torch.no_grad():
|
||||
if not has_been_preprocessed:
|
||||
# tensors should be 0-1
|
||||
if tensors_0_1.ndim == 3:
|
||||
tensors_0_1 = tensors_0_1.unsqueeze(0)
|
||||
# training tensors are 0 - 1
|
||||
tensors_0_1 = tensors_0_1.to(self.device, dtype=torch.float16)
|
||||
|
||||
# if images are out of this range throw error
|
||||
if tensors_0_1.min() < -0.3 or tensors_0_1.max() > 1.3:
|
||||
raise ValueError("image tensor values must be between 0 and 1. Got min: {}, max: {}".format(
|
||||
tensors_0_1.min(), tensors_0_1.max()
|
||||
))
|
||||
# unconditional
|
||||
if drop:
|
||||
if self.clip_noise_zero:
|
||||
tensors_0_1 = torch.rand_like(tensors_0_1).detach()
|
||||
noise_scale = torch.rand([tensors_0_1.shape[0], 1, 1, 1], device=self.device,
|
||||
dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
tensors_0_1 = tensors_0_1 * noise_scale
|
||||
else:
|
||||
tensors_0_1 = torch.zeros_like(tensors_0_1).detach()
|
||||
# tensors_0_1 = tensors_0_1 * 0
|
||||
clip_image = self.clip_image_processor(
|
||||
images=tensors_0_1,
|
||||
return_tensors="pt",
|
||||
do_resize=True,
|
||||
do_rescale=False,
|
||||
).pixel_values
|
||||
else:
|
||||
if drop:
|
||||
# scale the noise down
|
||||
if self.clip_noise_zero:
|
||||
tensors_0_1 = torch.rand_like(tensors_0_1).detach()
|
||||
noise_scale = torch.rand([tensors_0_1.shape[0], 1, 1, 1], device=self.device,
|
||||
dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
tensors_0_1 = tensors_0_1 * noise_scale
|
||||
else:
|
||||
tensors_0_1 = torch.zeros_like(tensors_0_1).detach()
|
||||
# tensors_0_1 = tensors_0_1 * 0
|
||||
mean = torch.tensor(self.clip_image_processor.image_mean).to(
|
||||
self.device, dtype=get_torch_dtype(self.sd_ref().dtype)
|
||||
).detach()
|
||||
std = torch.tensor(self.clip_image_processor.image_std).to(
|
||||
self.device, dtype=get_torch_dtype(self.sd_ref().dtype)
|
||||
).detach()
|
||||
tensors_0_1 = torch.clip((255. * tensors_0_1), 0, 255).round() / 255.0
|
||||
clip_image = (tensors_0_1 - mean.view([1, 3, 1, 1])) / std.view([1, 3, 1, 1])
|
||||
|
||||
else:
|
||||
clip_image = tensors_0_1
|
||||
clip_image = clip_image.to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype)).detach()
|
||||
with torch.set_grad_enabled(is_training):
|
||||
if is_training:
|
||||
self.image_encoder.train()
|
||||
else:
|
||||
self.image_encoder.eval()
|
||||
clip_output = self.image_encoder(clip_image, output_hidden_states=True)
|
||||
|
||||
if self.config.clip_layer == 'penultimate_hidden_states':
|
||||
# they skip last layer for ip+
|
||||
# https://github.com/tencent-ailab/IP-Adapter/blob/f4b6742db35ea6d81c7b829a55b0a312c7f5a677/tutorial_train_plus.py#L403C26-L403C26
|
||||
clip_image_embeds = clip_output.hidden_states[-2]
|
||||
elif self.config.clip_layer == 'last_hidden_state':
|
||||
clip_image_embeds = clip_output.hidden_states[-1]
|
||||
else:
|
||||
clip_image_embeds = clip_output.image_embeds
|
||||
return clip_image_embeds
|
||||
|
||||
import torch
|
||||
|
||||
def set_vec(self, new_vector, text_encoder_idx=0):
|
||||
# Get the embedding layer
|
||||
embedding_layer = self.text_encoder_list[text_encoder_idx].get_input_embeddings()
|
||||
|
||||
# Indices to replace in the embeddings
|
||||
indices_to_replace = self.placeholder_token_ids[text_encoder_idx]
|
||||
|
||||
# Replace the specified embeddings with new_vector
|
||||
for idx in indices_to_replace:
|
||||
vector_idx = idx - indices_to_replace[0]
|
||||
embedding_layer.weight[idx] = new_vector[vector_idx]
|
||||
|
||||
# adds it to the tokenizer
|
||||
def forward(self, clip_image_embeds: torch.Tensor) -> PromptEmbeds:
|
||||
clip_image_embeds = clip_image_embeds.to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
if clip_image_embeds.ndim == 2:
|
||||
# expand the token dimension
|
||||
clip_image_embeds = clip_image_embeds.unsqueeze(1)
|
||||
image_prompt_embeds = self.embedder(clip_image_embeds)
|
||||
# todo add support for multiple batch sizes
|
||||
if image_prompt_embeds.shape[0] != 1:
|
||||
raise ValueError("Batch size must be 1 for embedder for now")
|
||||
|
||||
# output on sd1.5 is bs, num_tokens, 768
|
||||
if len(self.text_encoder_list) == 1:
|
||||
# add it to the text encoder
|
||||
self.set_vec(image_prompt_embeds[0], text_encoder_idx=0)
|
||||
elif len(self.text_encoder_list) == 2:
|
||||
if self.text_encoder_list[0].config.target_hidden_size + self.text_encoder_list[1].config.target_hidden_size != \
|
||||
image_prompt_embeds.shape[2]:
|
||||
raise ValueError("Something went wrong. The embeddings do not match the text encoder sizes")
|
||||
# sdxl variants
|
||||
# image_prompt_embeds = 2048
|
||||
# te1 = 768
|
||||
# te2 = 1280
|
||||
te1_embeds = image_prompt_embeds[:, :, :self.text_encoder_list[0].config.target_hidden_size]
|
||||
te2_embeds = image_prompt_embeds[:, :, self.text_encoder_list[0].config.target_hidden_size:]
|
||||
self.set_vec(te1_embeds[0], text_encoder_idx=0)
|
||||
self.set_vec(te2_embeds[0], text_encoder_idx=1)
|
||||
else:
|
||||
|
||||
raise ValueError("Unsupported number of text encoders")
|
||||
# just a place to put a breakpoint
|
||||
pass
|
||||
|
||||
def restore_embeddings(self):
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(
|
||||
self.text_encoder_list,
|
||||
self.tokenizer_list,
|
||||
self.orig_embeds_params,
|
||||
self.placeholder_token_ids
|
||||
):
|
||||
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
index_no_updates[
|
||||
min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
|
||||
with torch.no_grad():
|
||||
text_encoder.get_input_embeddings().weight[
|
||||
index_no_updates
|
||||
] = orig_embeds[index_no_updates]
|
||||
# detach it all
|
||||
text_encoder.get_input_embeddings().weight.detach_()
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
self.image_encoder.gradient_checkpointing = True
|
||||
|
||||
def inject_trigger_into_prompt(self, prompt, expand_token=False, to_replace_list=None, add_if_not_present=True):
|
||||
output_prompt = prompt
|
||||
embedding_tokens = self.embedding_tokens[0] # shoudl be the same
|
||||
default_replacements = ["[name]", "[trigger]"]
|
||||
|
||||
replace_with = embedding_tokens if expand_token else self.trigger
|
||||
if to_replace_list is None:
|
||||
to_replace_list = default_replacements
|
||||
else:
|
||||
to_replace_list += default_replacements
|
||||
|
||||
# remove duplicates
|
||||
to_replace_list = list(set(to_replace_list))
|
||||
|
||||
# replace them all
|
||||
for to_replace in to_replace_list:
|
||||
# replace it
|
||||
output_prompt = output_prompt.replace(to_replace, replace_with)
|
||||
|
||||
# see how many times replace_with is in the prompt
|
||||
num_instances = output_prompt.count(replace_with)
|
||||
|
||||
if num_instances == 0 and add_if_not_present:
|
||||
# add it to the beginning of the prompt
|
||||
output_prompt = replace_with + " " + output_prompt
|
||||
|
||||
if num_instances > 1:
|
||||
print(
|
||||
f"Warning: {replace_with} token appears {num_instances} times in prompt {output_prompt}. This may cause issues.")
|
||||
|
||||
return output_prompt
|
||||
|
||||
# reverses injection with class name. useful for normalizations
|
||||
def inject_trigger_class_name_into_prompt(self, prompt):
|
||||
output_prompt = prompt
|
||||
embedding_tokens = self.embedding_tokens[0] # shoudl be the same
|
||||
|
||||
default_replacements = ["[name]", "[trigger]", embedding_tokens, self.trigger]
|
||||
|
||||
replace_with = self.config.trigger_class_name
|
||||
to_replace_list = default_replacements
|
||||
|
||||
# remove duplicates
|
||||
to_replace_list = list(set(to_replace_list))
|
||||
|
||||
# replace them all
|
||||
for to_replace in to_replace_list:
|
||||
# replace it
|
||||
output_prompt = output_prompt.replace(to_replace, replace_with)
|
||||
|
||||
# see how many times replace_with is in the prompt
|
||||
num_instances = output_prompt.count(replace_with)
|
||||
|
||||
if num_instances > 1:
|
||||
print(
|
||||
f"Warning: {replace_with} token appears {num_instances} times in prompt {output_prompt}. This may cause issues.")
|
||||
|
||||
return output_prompt
|
||||
@@ -43,9 +43,7 @@ def preprocess_config(config: OrderedDict, name: str = None):
|
||||
if "name" not in config["config"] and name is None:
|
||||
raise ValueError("config file must have a config.name key")
|
||||
# we need to replace tags. For now just [name]
|
||||
if name is not None:
|
||||
config["config"]["name"] = name
|
||||
else:
|
||||
if name is None:
|
||||
name = config["config"]["name"]
|
||||
config_string = json.dumps(config)
|
||||
config_string = config_string.replace("[name]", name)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import os
|
||||
import time
|
||||
from typing import List, Optional, Literal, Union
|
||||
from typing import List, Optional, Literal, Union, TYPE_CHECKING, Dict
|
||||
import random
|
||||
|
||||
import torch
|
||||
@@ -11,22 +11,31 @@ ImgExt = Literal['jpg', 'png', 'webp']
|
||||
|
||||
SaveFormat = Literal['safetensors', 'diffusers']
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.guidance import GuidanceType
|
||||
from toolkit.logging import EmptyLogger
|
||||
else:
|
||||
EmptyLogger = None
|
||||
|
||||
class SaveConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.save_every: int = kwargs.get('save_every', 1000)
|
||||
self.dtype: str = kwargs.get('save_dtype', 'float16')
|
||||
self.dtype: str = kwargs.get('dtype', 'float16')
|
||||
self.max_step_saves_to_keep: int = kwargs.get('max_step_saves_to_keep', 5)
|
||||
self.save_format: SaveFormat = kwargs.get('save_format', 'safetensors')
|
||||
if self.save_format not in ['safetensors', 'diffusers']:
|
||||
raise ValueError(f"save_format must be safetensors or diffusers, got {self.save_format}")
|
||||
self.push_to_hub: bool = kwargs.get("push_to_hub", False)
|
||||
self.hf_repo_id: Optional[str] = kwargs.get("hf_repo_id", None)
|
||||
self.hf_private: Optional[str] = kwargs.get("hf_private", False)
|
||||
|
||||
|
||||
class LogingConfig:
|
||||
class LoggingConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.log_every: int = kwargs.get('log_every', 100)
|
||||
self.verbose: bool = kwargs.get('verbose', False)
|
||||
self.use_wandb: bool = kwargs.get('use_wandb', False)
|
||||
self.project_name: str = kwargs.get('project_name', 'ai-toolkit')
|
||||
self.run_name: str = kwargs.get('run_name', None)
|
||||
|
||||
|
||||
class SampleConfig:
|
||||
@@ -45,7 +54,14 @@ class SampleConfig:
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
self.ext: ImgExt = kwargs.get('format', 'jpg')
|
||||
self.adapter_conditioning_scale = kwargs.get('adapter_conditioning_scale', 1.0)
|
||||
self.refiner_start_at = kwargs.get('refiner_start_at', 0.5) # step to start using refiner on sample if it exists
|
||||
self.refiner_start_at = kwargs.get('refiner_start_at',
|
||||
0.5) # step to start using refiner on sample if it exists
|
||||
self.extra_values = kwargs.get('extra_values', [])
|
||||
self.num_frames = kwargs.get('num_frames', 1)
|
||||
self.fps: int = kwargs.get('fps', 16)
|
||||
if self.num_frames > 1 and self.ext not in ['webp']:
|
||||
print("Changing sample extention to animated webp")
|
||||
self.ext = 'webp'
|
||||
|
||||
|
||||
class LormModuleSettingsConfig:
|
||||
@@ -90,7 +106,7 @@ class LoRMConfig:
|
||||
})
|
||||
|
||||
|
||||
NetworkType = Literal['lora', 'locon', 'lorm']
|
||||
NetworkType = Literal['lora', 'locon', 'lorm', 'lokr']
|
||||
|
||||
|
||||
class NetworkConfig:
|
||||
@@ -109,6 +125,7 @@ class NetworkConfig:
|
||||
self.linear_alpha: float = kwargs.get('linear_alpha', self.alpha)
|
||||
self.conv_alpha: float = kwargs.get('conv_alpha', self.conv)
|
||||
self.dropout: Union[float, None] = kwargs.get('dropout', None)
|
||||
self.network_kwargs: dict = kwargs.get('network_kwargs', {})
|
||||
|
||||
self.lorm_config: Union[LoRMConfig, None] = None
|
||||
lorm = kwargs.get('lorm', None)
|
||||
@@ -122,24 +139,120 @@ class NetworkConfig:
|
||||
if self.lorm_config.do_conv:
|
||||
self.conv = 4
|
||||
|
||||
self.transformer_only = kwargs.get('transformer_only', True)
|
||||
|
||||
self.lokr_full_rank = kwargs.get('lokr_full_rank', False)
|
||||
if self.lokr_full_rank and self.type.lower() == 'lokr':
|
||||
self.linear = 9999999999
|
||||
self.linear_alpha = 9999999999
|
||||
self.conv = 9999999999
|
||||
self.conv_alpha = 9999999999
|
||||
# -1 automatically finds the largest factor
|
||||
self.lokr_factor = kwargs.get('lokr_factor', -1)
|
||||
|
||||
AdapterTypes = Literal['t2i', 'ip', 'ip+']
|
||||
|
||||
AdapterTypes = Literal['t2i', 'ip', 'ip+', 'clip', 'ilora', 'photo_maker', 'control_net', 'control_lora']
|
||||
|
||||
CLIPLayer = Literal['penultimate_hidden_states', 'image_embeds', 'last_hidden_state']
|
||||
|
||||
|
||||
class AdapterConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.type: AdapterTypes = kwargs.get('type', 't2i') # t2i, ip
|
||||
self.type: AdapterTypes = kwargs.get('type', 't2i') # t2i, ip, clip, control_net
|
||||
self.in_channels: int = kwargs.get('in_channels', 3)
|
||||
self.channels: List[int] = kwargs.get('channels', [320, 640, 1280, 1280])
|
||||
self.num_res_blocks: int = kwargs.get('num_res_blocks', 2)
|
||||
self.downscale_factor: int = kwargs.get('downscale_factor', 8)
|
||||
self.adapter_type: str = kwargs.get('adapter_type', 'full_adapter')
|
||||
self.image_dir: str = kwargs.get('image_dir', None)
|
||||
self.test_img_path: str = kwargs.get('test_img_path', None)
|
||||
self.test_img_path: List[str] = kwargs.get('test_img_path', None)
|
||||
if self.test_img_path is not None:
|
||||
if isinstance(self.test_img_path, str):
|
||||
self.test_img_path = self.test_img_path.split(',')
|
||||
self.test_img_path = [p.strip() for p in self.test_img_path]
|
||||
self.test_img_path = [p for p in self.test_img_path if p != '']
|
||||
|
||||
self.train: str = kwargs.get('train', False)
|
||||
self.image_encoder_path: str = kwargs.get('image_encoder_path', None)
|
||||
self.name_or_path = kwargs.get('name_or_path', None)
|
||||
|
||||
num_tokens = kwargs.get('num_tokens', None)
|
||||
if num_tokens is None and self.type.startswith('ip'):
|
||||
if self.type == 'ip+':
|
||||
num_tokens = 16
|
||||
num_tokens = 16
|
||||
elif self.type == 'ip':
|
||||
num_tokens = 4
|
||||
|
||||
self.num_tokens: int = num_tokens
|
||||
self.train_image_encoder: bool = kwargs.get('train_image_encoder', False)
|
||||
self.train_only_image_encoder: bool = kwargs.get('train_only_image_encoder', False)
|
||||
if self.train_only_image_encoder:
|
||||
self.train_image_encoder = True
|
||||
self.train_only_image_encoder_positional_embedding: bool = kwargs.get(
|
||||
'train_only_image_encoder_positional_embedding', False)
|
||||
self.image_encoder_arch: str = kwargs.get('image_encoder_arch', 'clip') # clip vit vit_hybrid, safe
|
||||
self.safe_reducer_channels: int = kwargs.get('safe_reducer_channels', 512)
|
||||
self.safe_channels: int = kwargs.get('safe_channels', 2048)
|
||||
self.safe_tokens: int = kwargs.get('safe_tokens', 8)
|
||||
self.quad_image: bool = kwargs.get('quad_image', False)
|
||||
|
||||
# clip vision
|
||||
self.trigger = kwargs.get('trigger', 'tri993r')
|
||||
self.trigger_class_name = kwargs.get('trigger_class_name', None)
|
||||
|
||||
self.class_names = kwargs.get('class_names', [])
|
||||
|
||||
self.clip_layer: CLIPLayer = kwargs.get('clip_layer', None)
|
||||
if self.clip_layer is None:
|
||||
if self.type.startswith('ip+'):
|
||||
self.clip_layer = 'penultimate_hidden_states'
|
||||
else:
|
||||
self.clip_layer = 'last_hidden_state'
|
||||
|
||||
# text encoder
|
||||
self.text_encoder_path: str = kwargs.get('text_encoder_path', None)
|
||||
self.text_encoder_arch: str = kwargs.get('text_encoder_arch', 'clip') # clip t5
|
||||
|
||||
self.train_scaler: bool = kwargs.get('train_scaler', False)
|
||||
self.scaler_lr: Optional[float] = kwargs.get('scaler_lr', None)
|
||||
|
||||
# trains with a scaler to easy channel bias but merges it in on save
|
||||
self.merge_scaler: bool = kwargs.get('merge_scaler', False)
|
||||
|
||||
# for ilora
|
||||
self.head_dim: int = kwargs.get('head_dim', 1024)
|
||||
self.num_heads: int = kwargs.get('num_heads', 1)
|
||||
self.ilora_down: bool = kwargs.get('ilora_down', True)
|
||||
self.ilora_mid: bool = kwargs.get('ilora_mid', True)
|
||||
self.ilora_up: bool = kwargs.get('ilora_up', True)
|
||||
|
||||
self.pixtral_max_image_size: int = kwargs.get('pixtral_max_image_size', 512)
|
||||
self.pixtral_random_image_size: int = kwargs.get('pixtral_random_image_size', False)
|
||||
|
||||
self.flux_only_double: bool = kwargs.get('flux_only_double', False)
|
||||
|
||||
# train and use a conv layer to pool the embedding
|
||||
self.conv_pooling: bool = kwargs.get('conv_pooling', False)
|
||||
self.conv_pooling_stacks: int = kwargs.get('conv_pooling_stacks', 1)
|
||||
self.sparse_autoencoder_dim: Optional[int] = kwargs.get('sparse_autoencoder_dim', None)
|
||||
|
||||
# for llm adapter
|
||||
self.num_cloned_blocks: int = kwargs.get('num_cloned_blocks', 0)
|
||||
self.quantize_llm: bool = kwargs.get('quantize_llm', False)
|
||||
|
||||
# for control lora only
|
||||
lora_config: dict = kwargs.get('lora_config', None)
|
||||
if lora_config is not None:
|
||||
self.lora_config: NetworkConfig = NetworkConfig(**lora_config)
|
||||
else:
|
||||
self.lora_config = None
|
||||
self.num_control_images: int = kwargs.get('num_control_images', 1)
|
||||
# decimal for how often the control is dropped out and replaced with noise 1.0 is 100%
|
||||
self.control_image_dropout: float = kwargs.get('control_image_dropout', 0.0)
|
||||
self.has_inpainting_input: bool = kwargs.get('has_inpainting_input', False)
|
||||
self.invert_inpaint_mask_chance: float = kwargs.get('invert_inpaint_mask_chance', 0.0)
|
||||
|
||||
|
||||
class EmbeddingConfig:
|
||||
def __init__(self, **kwargs):
|
||||
@@ -147,6 +260,12 @@ class EmbeddingConfig:
|
||||
self.tokens = kwargs.get('tokens', 4)
|
||||
self.init_words = kwargs.get('init_words', '*')
|
||||
self.save_format = kwargs.get('save_format', 'safetensors')
|
||||
self.trigger_class_name = kwargs.get('trigger_class_name', None) # used for inverted masked prior
|
||||
|
||||
|
||||
class DecoratorConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.num_tokens: str = kwargs.get('num_tokens', 4)
|
||||
|
||||
|
||||
ContentOrStyleType = Literal['balanced', 'style', 'content']
|
||||
@@ -157,6 +276,7 @@ class TrainConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.noise_scheduler = kwargs.get('noise_scheduler', 'ddpm')
|
||||
self.content_or_style: ContentOrStyleType = kwargs.get('content_or_style', 'balanced')
|
||||
self.content_or_style_reg: ContentOrStyleType = kwargs.get('content_or_style', 'balanced')
|
||||
self.steps: int = kwargs.get('steps', 1000)
|
||||
self.lr = kwargs.get('lr', 1e-6)
|
||||
self.unet_lr = kwargs.get('unet_lr', self.lr)
|
||||
@@ -171,12 +291,15 @@ class TrainConfig:
|
||||
self.min_denoising_steps: int = kwargs.get('min_denoising_steps', 0)
|
||||
self.max_denoising_steps: int = kwargs.get('max_denoising_steps', 1000)
|
||||
self.batch_size: int = kwargs.get('batch_size', 1)
|
||||
self.orig_batch_size: int = self.batch_size
|
||||
self.dtype: str = kwargs.get('dtype', 'fp32')
|
||||
self.xformers = kwargs.get('xformers', False)
|
||||
self.sdp = kwargs.get('sdp', False)
|
||||
self.train_unet = kwargs.get('train_unet', True)
|
||||
self.train_text_encoder = kwargs.get('train_text_encoder', True)
|
||||
self.train_text_encoder = kwargs.get('train_text_encoder', False)
|
||||
self.train_refiner = kwargs.get('train_refiner', True)
|
||||
self.train_turbo = kwargs.get('train_turbo', False)
|
||||
self.show_turbo_outputs = kwargs.get('show_turbo_outputs', False)
|
||||
self.min_snr_gamma = kwargs.get('min_snr_gamma', None)
|
||||
self.snr_gamma = kwargs.get('snr_gamma', None)
|
||||
# trains a gamma, offset, and scale to adjust loss to adapt to timestep differentials
|
||||
@@ -184,6 +307,7 @@ class TrainConfig:
|
||||
self.learnable_snr_gos = kwargs.get('learnable_snr_gos', False)
|
||||
self.noise_offset = kwargs.get('noise_offset', 0.0)
|
||||
self.skip_first_sample = kwargs.get('skip_first_sample', False)
|
||||
self.force_first_sample = kwargs.get('force_first_sample', False)
|
||||
self.gradient_checkpointing = kwargs.get('gradient_checkpointing', True)
|
||||
self.weight_jitter = kwargs.get('weight_jitter', 0.0)
|
||||
self.merge_network_on_save = kwargs.get('merge_network_on_save', False)
|
||||
@@ -191,12 +315,20 @@ class TrainConfig:
|
||||
self.start_step = kwargs.get('start_step', None)
|
||||
self.free_u = kwargs.get('free_u', False)
|
||||
self.adapter_assist_name_or_path: Optional[str] = kwargs.get('adapter_assist_name_or_path', None)
|
||||
self.adapter_assist_type: Optional[str] = kwargs.get('adapter_assist_type', 't2i') # t2i, control_net
|
||||
self.noise_multiplier = kwargs.get('noise_multiplier', 1.0)
|
||||
self.target_noise_multiplier = kwargs.get('target_noise_multiplier', 1.0)
|
||||
self.img_multiplier = kwargs.get('img_multiplier', 1.0)
|
||||
self.noisy_latent_multiplier = kwargs.get('noisy_latent_multiplier', 1.0)
|
||||
self.latent_multiplier = kwargs.get('latent_multiplier', 1.0)
|
||||
self.negative_prompt = kwargs.get('negative_prompt', None)
|
||||
self.max_negative_prompts = kwargs.get('max_negative_prompts', 1)
|
||||
# multiplier applied to loos on regularization images
|
||||
self.reg_weight = kwargs.get('reg_weight', 1.0)
|
||||
self.num_train_timesteps = kwargs.get('num_train_timesteps', 1000)
|
||||
self.random_noise_shift = kwargs.get('random_noise_shift', 0.0)
|
||||
# automatically adapte the vae scaling based on the image norm
|
||||
self.adaptive_scaling_factor = kwargs.get('adaptive_scaling_factor', False)
|
||||
|
||||
# dropout that happens before encoding. It functions independently per text encoder
|
||||
self.prompt_dropout_prob = kwargs.get('prompt_dropout_prob', 0.0)
|
||||
@@ -208,8 +340,16 @@ class TrainConfig:
|
||||
|
||||
# set to -1 to accumulate gradients for entire epoch
|
||||
# warning, only do this with a small dataset or you will run out of memory
|
||||
# This is legacy but left in for backwards compatibility
|
||||
self.gradient_accumulation_steps = kwargs.get('gradient_accumulation_steps', 1)
|
||||
|
||||
# this will do proper gradient accumulation where you will not see a step until the end of the accumulation
|
||||
# the method above will show a step every accumulation
|
||||
self.gradient_accumulation = kwargs.get('gradient_accumulation', 1)
|
||||
if self.gradient_accumulation > 1:
|
||||
if self.gradient_accumulation_steps != 1:
|
||||
raise ValueError("gradient_accumulation and gradient_accumulation_steps are mutually exclusive")
|
||||
|
||||
# short long captions will double your batch size. This only works when a dataset is
|
||||
# prepared with a json caption file that has both short and long captions in it. It will
|
||||
# Double up every image and run it through with both short and long captions. The idea
|
||||
@@ -233,24 +373,123 @@ class TrainConfig:
|
||||
# unmasked reign. It is unmasked regularization basically
|
||||
self.inverted_mask_prior = kwargs.get('inverted_mask_prior', False)
|
||||
self.inverted_mask_prior_multiplier = kwargs.get('inverted_mask_prior_multiplier', 0.5)
|
||||
|
||||
# DOP will will run the same image and prompt through the network without the trigger word blank and use it as a target
|
||||
self.diff_output_preservation = kwargs.get('diff_output_preservation', False)
|
||||
self.diff_output_preservation_multiplier = kwargs.get('diff_output_preservation_multiplier', 1.0)
|
||||
# If the trigger word is in the prompt, we will use this class name to replace it eg. "sks woman" -> "woman"
|
||||
self.diff_output_preservation_class = kwargs.get('diff_output_preservation_class', '')
|
||||
|
||||
# legacy
|
||||
if match_adapter_assist and self.match_adapter_chance == 0.0:
|
||||
self.match_adapter_chance = 1.0
|
||||
|
||||
# standardize inputs to the meand std of the model knowledge
|
||||
self.standardize_images = kwargs.get('standardize_images', False)
|
||||
self.standardize_latents = kwargs.get('standardize_latents', False)
|
||||
|
||||
# if self.train_turbo and not self.noise_scheduler.startswith("euler"):
|
||||
# raise ValueError(f"train_turbo is only supported with euler and wuler_a noise schedulers")
|
||||
|
||||
self.dynamic_noise_offset = kwargs.get('dynamic_noise_offset', False)
|
||||
self.do_cfg = kwargs.get('do_cfg', False)
|
||||
self.do_random_cfg = kwargs.get('do_random_cfg', False)
|
||||
self.cfg_scale = kwargs.get('cfg_scale', 1.0)
|
||||
self.max_cfg_scale = kwargs.get('max_cfg_scale', self.cfg_scale)
|
||||
self.cfg_rescale = kwargs.get('cfg_rescale', None)
|
||||
if self.cfg_rescale is None:
|
||||
self.cfg_rescale = self.cfg_scale
|
||||
|
||||
# applies the inverse of the prediction mean and std to the target to correct
|
||||
# for norm drift
|
||||
self.correct_pred_norm = kwargs.get('correct_pred_norm', False)
|
||||
self.correct_pred_norm_multiplier = kwargs.get('correct_pred_norm_multiplier', 1.0)
|
||||
|
||||
self.loss_type = kwargs.get('loss_type', 'mse') # mse, mae, wavelet
|
||||
|
||||
# scale the prediction by this. Increase for more detail, decrease for less
|
||||
self.pred_scaler = kwargs.get('pred_scaler', 1.0)
|
||||
|
||||
# repeats the prompt a few times to saturate the encoder
|
||||
self.prompt_saturation_chance = kwargs.get('prompt_saturation_chance', 0.0)
|
||||
|
||||
# applies negative loss on the prior to encourage network to diverge from it
|
||||
self.do_prior_divergence = kwargs.get('do_prior_divergence', False)
|
||||
|
||||
ema_config: Union[Dict, None] = kwargs.get('ema_config', None)
|
||||
# if it is set explicitly to false, leave it false.
|
||||
if ema_config is not None and ema_config.get('use_ema', None) is not None:
|
||||
ema_config['use_ema'] = True
|
||||
print(f"Using EMA")
|
||||
else:
|
||||
ema_config = {'use_ema': False}
|
||||
|
||||
self.ema_config: EMAConfig = EMAConfig(**ema_config)
|
||||
|
||||
# adds an additional loss to the network to encourage it output a normalized standard deviation
|
||||
self.target_norm_std = kwargs.get('target_norm_std', None)
|
||||
self.target_norm_std_value = kwargs.get('target_norm_std_value', 1.0)
|
||||
self.timestep_type = kwargs.get('timestep_type', 'sigmoid') # sigmoid, linear, lognorm_blend
|
||||
self.linear_timesteps = kwargs.get('linear_timesteps', False)
|
||||
self.linear_timesteps2 = kwargs.get('linear_timesteps2', False)
|
||||
self.disable_sampling = kwargs.get('disable_sampling', False)
|
||||
|
||||
# will cache a blank prompt or the trigger word, and unload the text encoder to cpu
|
||||
# will make training faster and use less vram
|
||||
self.unload_text_encoder = kwargs.get('unload_text_encoder', False)
|
||||
# for swapping which parameters are trained during training
|
||||
self.do_paramiter_swapping = kwargs.get('do_paramiter_swapping', False)
|
||||
# 0.1 is 10% of the parameters active at a time lower is less vram, higher is more
|
||||
self.paramiter_swapping_factor = kwargs.get('paramiter_swapping_factor', 0.1)
|
||||
# bypass the guidance embedding for training. For open flux with guidance embedding
|
||||
self.bypass_guidance_embedding = kwargs.get('bypass_guidance_embedding', False)
|
||||
|
||||
# diffusion feature extractor
|
||||
self.diffusion_feature_extractor_path = kwargs.get('diffusion_feature_extractor_path', None)
|
||||
self.diffusion_feature_extractor_weight = kwargs.get('diffusion_feature_extractor_weight', 1.0)
|
||||
|
||||
# optimal noise pairing
|
||||
self.optimal_noise_pairing_samples = kwargs.get('optimal_noise_pairing_samples', 1)
|
||||
|
||||
# forces same noise for the same image at a given size.
|
||||
self.force_consistent_noise = kwargs.get('force_consistent_noise', False)
|
||||
|
||||
|
||||
ModelArch = Literal['sd1', 'sd2', 'sd3', 'sdxl', 'pixart', 'pixart_sigma', 'auraflow', 'flux', 'flex2', 'lumina2', 'vega', 'ssd', 'wan21']
|
||||
|
||||
|
||||
class ModelConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.name_or_path: str = kwargs.get('name_or_path', None)
|
||||
# name or path is updated on fine tuning. Keep a copy of the original
|
||||
self.name_or_path_original: str = self.name_or_path
|
||||
self.is_v2: bool = kwargs.get('is_v2', False)
|
||||
self.is_xl: bool = kwargs.get('is_xl', False)
|
||||
self.is_pixart: bool = kwargs.get('is_pixart', False)
|
||||
self.is_pixart_sigma: bool = kwargs.get('is_pixart_sigma', False)
|
||||
self.is_auraflow: bool = kwargs.get('is_auraflow', False)
|
||||
self.is_v3: bool = kwargs.get('is_v3', False)
|
||||
self.is_flux: bool = kwargs.get('is_flux', False)
|
||||
self.is_flex2: bool = kwargs.get('is_flex2', False)
|
||||
if self.is_flex2:
|
||||
self.is_flux = True
|
||||
self.is_lumina2: bool = kwargs.get('is_lumina2', False)
|
||||
if self.is_pixart_sigma:
|
||||
self.is_pixart = True
|
||||
self.use_flux_cfg = kwargs.get('use_flux_cfg', False)
|
||||
self.is_ssd: bool = kwargs.get('is_ssd', False)
|
||||
self.is_vega: bool = kwargs.get('is_vega', False)
|
||||
self.is_v_pred: bool = kwargs.get('is_v_pred', False)
|
||||
self.dtype: str = kwargs.get('dtype', 'float16')
|
||||
self.vae_path = kwargs.get('vae_path', None)
|
||||
self.refiner_name_or_path = kwargs.get('refiner_name_or_path', None)
|
||||
self._original_refiner_name_or_path = self.refiner_name_or_path
|
||||
self.refiner_start_at = kwargs.get('refiner_start_at', 0.5)
|
||||
self.lora_path = kwargs.get('lora_path', None)
|
||||
# mainly for decompression loras for distilled models
|
||||
self.assistant_lora_path = kwargs.get('assistant_lora_path', None)
|
||||
self.inference_lora_path = kwargs.get('inference_lora_path', None)
|
||||
self.latent_space_version = kwargs.get('latent_space_version', None)
|
||||
|
||||
# only for SDXL models for now
|
||||
self.use_text_encoder_1: bool = kwargs.get('use_text_encoder_1', True)
|
||||
@@ -265,6 +504,109 @@ class ModelConfig:
|
||||
# sed sdxl as true since it is mostly the same architecture
|
||||
self.is_xl = True
|
||||
|
||||
if self.is_vega:
|
||||
self.is_xl = True
|
||||
|
||||
# for text encoder quant. Only works with pixart currently
|
||||
self.text_encoder_bits = kwargs.get('text_encoder_bits', 16) # 16, 8, 4
|
||||
self.unet_path = kwargs.get("unet_path", None)
|
||||
self.unet_sample_size = kwargs.get("unet_sample_size", None)
|
||||
self.vae_device = kwargs.get("vae_device", None)
|
||||
self.vae_dtype = kwargs.get("vae_dtype", self.dtype)
|
||||
self.te_device = kwargs.get("te_device", None)
|
||||
self.te_dtype = kwargs.get("te_dtype", self.dtype)
|
||||
|
||||
# only for flux for now
|
||||
self.quantize = kwargs.get("quantize", False)
|
||||
self.quantize_te = kwargs.get("quantize_te", self.quantize)
|
||||
self.qtype = kwargs.get("qtype", "qfloat8")
|
||||
self.qtype_te = kwargs.get("qtype_te", "qfloat8")
|
||||
self.low_vram = kwargs.get("low_vram", False)
|
||||
self.attn_masking = kwargs.get("attn_masking", False)
|
||||
if self.attn_masking and not self.is_flux:
|
||||
raise ValueError("attn_masking is only supported with flux models currently")
|
||||
# for targeting a specific layers
|
||||
self.ignore_if_contains: Optional[List[str]] = kwargs.get("ignore_if_contains", None)
|
||||
self.only_if_contains: Optional[List[str]] = kwargs.get("only_if_contains", None)
|
||||
self.quantize_kwargs = kwargs.get("quantize_kwargs", {})
|
||||
|
||||
# splits the model over the available gpus WIP
|
||||
self.split_model_over_gpus = kwargs.get("split_model_over_gpus", False)
|
||||
if self.split_model_over_gpus and not self.is_flux:
|
||||
raise ValueError("split_model_over_gpus is only supported with flux models currently")
|
||||
self.split_model_other_module_param_count_scale = kwargs.get("split_model_other_module_param_count_scale", 0.3)
|
||||
|
||||
self.te_name_or_path = kwargs.get("te_name_or_path", None)
|
||||
|
||||
self.arch: ModelArch = kwargs.get("arch", None)
|
||||
|
||||
# handle migrating to new model arch
|
||||
if self.arch is not None:
|
||||
# reverse the arch to the old style
|
||||
if self.arch == 'sd2':
|
||||
self.is_v2 = True
|
||||
elif self.arch == 'sd3':
|
||||
self.is_v3 = True
|
||||
elif self.arch == 'sdxl':
|
||||
self.is_xl = True
|
||||
elif self.arch == 'pixart':
|
||||
self.is_pixart = True
|
||||
elif self.arch == 'pixart_sigma':
|
||||
self.is_pixart_sigma = True
|
||||
elif self.arch == 'auraflow':
|
||||
self.is_auraflow = True
|
||||
elif self.arch == 'flux':
|
||||
self.is_flux = True
|
||||
elif self.arch == 'flex2':
|
||||
self.is_flex2 = True
|
||||
elif self.arch == 'lumina2':
|
||||
self.is_lumina2 = True
|
||||
elif self.arch == 'vega':
|
||||
self.is_vega = True
|
||||
elif self.arch == 'ssd':
|
||||
self.is_ssd = True
|
||||
else:
|
||||
pass
|
||||
if self.arch is None:
|
||||
if kwargs.get('is_v2', False):
|
||||
self.arch = 'sd2'
|
||||
elif kwargs.get('is_v3', False):
|
||||
self.arch = 'sd3'
|
||||
elif kwargs.get('is_xl', False):
|
||||
self.arch = 'sdxl'
|
||||
elif kwargs.get('is_pixart', False):
|
||||
self.arch = 'pixart'
|
||||
elif kwargs.get('is_pixart_sigma', False):
|
||||
self.arch = 'pixart_sigma'
|
||||
elif kwargs.get('is_auraflow', False):
|
||||
self.arch = 'auraflow'
|
||||
elif kwargs.get('is_flux', False):
|
||||
self.arch = 'flux'
|
||||
elif kwargs.get('is_flex2', False):
|
||||
self.arch = 'flex2'
|
||||
elif kwargs.get('is_lumina2', False):
|
||||
self.arch = 'lumina2'
|
||||
elif kwargs.get('is_vega', False):
|
||||
self.arch = 'vega'
|
||||
elif kwargs.get('is_ssd', False):
|
||||
self.arch = 'ssd'
|
||||
else:
|
||||
self.arch = 'sd1'
|
||||
|
||||
|
||||
|
||||
class EMAConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.use_ema: bool = kwargs.get('use_ema', False)
|
||||
self.ema_decay: float = kwargs.get('ema_decay', 0.999)
|
||||
# feeds back the decay difference into the parameter
|
||||
self.use_feedback: bool = kwargs.get('use_feedback', False)
|
||||
|
||||
# every update, the params are multiplied by this amount
|
||||
# only use for things without a bias like lora
|
||||
# similar to a decay in an optimizer but the opposite
|
||||
self.param_multiplier: float = kwargs.get('param_multiplier', 1.0)
|
||||
|
||||
|
||||
class ReferenceDatasetConfig:
|
||||
def __init__(self, **kwargs):
|
||||
@@ -352,29 +694,46 @@ class DatasetConfig:
|
||||
self.dataset_path: str = kwargs.get('dataset_path', None)
|
||||
|
||||
self.default_caption: str = kwargs.get('default_caption', None)
|
||||
self.random_triggers: List[str] = kwargs.get('random_triggers', [])
|
||||
# trigger word for just this dataset
|
||||
self.trigger_word: str = kwargs.get('trigger_word', None)
|
||||
random_triggers = kwargs.get('random_triggers', [])
|
||||
# if they are a string, load them from a file
|
||||
if isinstance(random_triggers, str) and os.path.exists(random_triggers):
|
||||
with open(random_triggers, 'r') as f:
|
||||
random_triggers = f.read().splitlines()
|
||||
# remove empty lines
|
||||
random_triggers = [line for line in random_triggers if line.strip() != '']
|
||||
self.random_triggers: List[str] = random_triggers
|
||||
self.random_triggers_max: int = kwargs.get('random_triggers_max', 1)
|
||||
self.caption_ext: str = kwargs.get('caption_ext', None)
|
||||
self.random_scale: bool = kwargs.get('random_scale', False)
|
||||
self.random_crop: bool = kwargs.get('random_crop', False)
|
||||
self.resolution: int = kwargs.get('resolution', 512)
|
||||
self.scale: float = kwargs.get('scale', 1.0)
|
||||
self.buckets: bool = kwargs.get('buckets', False)
|
||||
self.buckets: bool = kwargs.get('buckets', True)
|
||||
self.bucket_tolerance: int = kwargs.get('bucket_tolerance', 64)
|
||||
self.is_reg: bool = kwargs.get('is_reg', False)
|
||||
self.network_weight: float = float(kwargs.get('network_weight', 1.0))
|
||||
self.token_dropout_rate: float = float(kwargs.get('token_dropout_rate', 0.0))
|
||||
self.shuffle_tokens: bool = kwargs.get('shuffle_tokens', False)
|
||||
self.caption_dropout_rate: float = float(kwargs.get('caption_dropout_rate', 0.0))
|
||||
self.keep_tokens: int = kwargs.get('keep_tokens', 0) # #of first tokens to always keep unless caption dropped
|
||||
self.flip_x: bool = kwargs.get('flip_x', False)
|
||||
self.flip_y: bool = kwargs.get('flip_y', False)
|
||||
self.augments: List[str] = kwargs.get('augments', [])
|
||||
self.control_path: str = kwargs.get('control_path', None) # depth maps, etc
|
||||
self.control_path: Union[str,List[str]] = kwargs.get('control_path', None) # depth maps, etc
|
||||
# inpaint images should be webp/png images with alpha channel. The alpha 0 (invisible) section will
|
||||
# be the part conditioned to be inpainted. The alpha 1 (visible) section will be the part that is ignored
|
||||
self.inpaint_path: Union[str,List[str]] = kwargs.get('inpaint_path', None)
|
||||
# instead of cropping ot match image, it will serve the full size control image (clip images ie for ip adapters)
|
||||
self.full_size_control_images: bool = kwargs.get('full_size_control_images', False)
|
||||
self.alpha_mask: bool = kwargs.get('alpha_mask', False) # if true, will use alpha channel as mask
|
||||
self.mask_path: str = kwargs.get('mask_path',
|
||||
None) # focus mask (black and white. White has higher loss than black)
|
||||
self.unconditional_path: str = kwargs.get('unconditional_path', None) # path where matching unconditional images are located
|
||||
self.unconditional_path: str = kwargs.get('unconditional_path',
|
||||
None) # path where matching unconditional images are located
|
||||
self.invert_mask: bool = kwargs.get('invert_mask', False) # invert mask
|
||||
self.mask_min_value: float = kwargs.get('mask_min_value', 0.01) # min value for . 0 - 1
|
||||
self.mask_min_value: float = kwargs.get('mask_min_value', 0.0) # min value for . 0 - 1
|
||||
self.poi: Union[str, None] = kwargs.get('poi',
|
||||
None) # if one is set and in json data, will be used as auto crop scale point of interes
|
||||
self.num_repeats: int = kwargs.get('num_repeats', 1) # number of times to repeat dataset
|
||||
@@ -382,6 +741,9 @@ class DatasetConfig:
|
||||
self.cache_latents: bool = kwargs.get('cache_latents', False)
|
||||
# cache latents to disk will store them on disk. If both are true, it will save to disk, but keep in memory
|
||||
self.cache_latents_to_disk: bool = kwargs.get('cache_latents_to_disk', False)
|
||||
self.cache_clip_vision_to_disk: bool = kwargs.get('cache_clip_vision_to_disk', False)
|
||||
|
||||
self.standardize_images: bool = kwargs.get('standardize_images', False)
|
||||
|
||||
# https://albumentations.ai/docs/api_reference/augmentations/transforms
|
||||
# augmentations are returned as a separate image and cannot currently be cached
|
||||
@@ -400,6 +762,39 @@ class DatasetConfig:
|
||||
if legacy_caption_type:
|
||||
self.caption_ext = legacy_caption_type
|
||||
self.caption_type = self.caption_ext
|
||||
self.guidance_type: GuidanceType = kwargs.get('guidance_type', 'targeted')
|
||||
|
||||
# ip adapter / reference dataset
|
||||
self.clip_image_path: str = kwargs.get('clip_image_path', None) # depth maps, etc
|
||||
# get the clip image randomly from the same folder as the image. Useful for folder grouped pairs.
|
||||
self.clip_image_from_same_folder: bool = kwargs.get('clip_image_from_same_folder', False)
|
||||
self.clip_image_augmentations: List[dict] = kwargs.get('clip_image_augmentations', None)
|
||||
self.clip_image_shuffle_augmentations: bool = kwargs.get('clip_image_shuffle_augmentations', False)
|
||||
self.replacements: List[str] = kwargs.get('replacements', [])
|
||||
self.loss_multiplier: float = kwargs.get('loss_multiplier', 1.0)
|
||||
|
||||
self.num_workers: int = kwargs.get('num_workers', 2)
|
||||
self.prefetch_factor: int = kwargs.get('prefetch_factor', 2)
|
||||
self.extra_values: List[float] = kwargs.get('extra_values', [])
|
||||
self.square_crop: bool = kwargs.get('square_crop', False)
|
||||
# apply same augmentations to control images. Usually want this true unless special case
|
||||
self.replay_transforms: bool = kwargs.get('replay_transforms', True)
|
||||
|
||||
# for video
|
||||
# if num_frames is greater than 1, the dataloader will look for video files.
|
||||
# num_frames will be the number of frames in the training batch. If num_frames is 1, it will look for images
|
||||
self.num_frames: int = kwargs.get('num_frames', 1)
|
||||
# if true, will shrink video to our frames. For instance, if we have a video with 100 frames and num_frames is 10,
|
||||
# we would pull frame 0, 10, 20, 30, 40, 50, 60, 70, 80, 90 so they are evenly spaced
|
||||
self.shrink_video_to_frames: bool = kwargs.get('shrink_video_to_frames', True)
|
||||
# fps is only used if shrink_video_to_frames is false. This will attempt to pull the num_frames at the given fps
|
||||
# it will select a random start frame and pull the frames at the given fps
|
||||
# this could have various issues with shorter videos and videos with variable fps
|
||||
# I recommend trimming your videos to the desired length and using shrink_video_to_frames(default)
|
||||
self.fps: int = kwargs.get('fps', 16)
|
||||
|
||||
# debug the frame count and frame selection. You dont need this. It is for debugging.
|
||||
self.debug: bool = kwargs.get('debug', False)
|
||||
|
||||
|
||||
def preprocess_dataset_raw_config(raw_config: List[dict]) -> List[dict]:
|
||||
@@ -448,6 +843,11 @@ class GenerateImageConfig:
|
||||
latents: Union[torch.Tensor | None] = None, # input latent to start with,
|
||||
extra_kwargs: dict = None, # extra data to save with prompt file
|
||||
refiner_start_at: float = 0.5, # start at this percentage of a step. 0.0 to 1.0 . 1.0 is the end
|
||||
extra_values: List[float] = None, # extra values to save with prompt file
|
||||
logger: Optional[EmptyLogger] = None,
|
||||
num_frames: int = 1,
|
||||
fps: int = 15,
|
||||
ctrl_idx: int = 0
|
||||
):
|
||||
self.width: int = width
|
||||
self.height: int = height
|
||||
@@ -475,6 +875,11 @@ class GenerateImageConfig:
|
||||
self.adapter_conditioning_scale: float = adapter_conditioning_scale
|
||||
self.extra_kwargs = extra_kwargs if extra_kwargs is not None else {}
|
||||
self.refiner_start_at = refiner_start_at
|
||||
self.extra_values = extra_values if extra_values is not None else []
|
||||
self.num_frames = num_frames
|
||||
self.fps = fps
|
||||
self.ctrl_idx = ctrl_idx
|
||||
|
||||
|
||||
# prompt string will override any settings above
|
||||
self._process_prompt_string()
|
||||
@@ -484,7 +889,7 @@ class GenerateImageConfig:
|
||||
self.negative_prompt_2 = negative_prompt
|
||||
|
||||
if prompt_2 is None:
|
||||
self.prompt_2 = prompt
|
||||
self.prompt_2 = self.prompt
|
||||
|
||||
# parse prompt paths
|
||||
if self.output_path is None and self.output_folder is None:
|
||||
@@ -504,6 +909,8 @@ class GenerateImageConfig:
|
||||
self.height = max(64, self.height - self.height % 8) # round to divisible by 8
|
||||
self.width = max(64, self.width - self.width % 8) # round to divisible by 8
|
||||
|
||||
self.logger = logger
|
||||
|
||||
def set_gen_time(self, gen_time: int = None):
|
||||
if gen_time is not None:
|
||||
self.gen_time = gen_time
|
||||
@@ -539,11 +946,30 @@ class GenerateImageConfig:
|
||||
# make parent dirs
|
||||
os.makedirs(self.output_folder, exist_ok=True)
|
||||
self.set_gen_time()
|
||||
# TODO save image gen header info for A1111 and us, our seeds probably wont match
|
||||
image.save(self.get_image_path(count, max_count))
|
||||
# do prompt file
|
||||
if self.add_prompt_file:
|
||||
self.save_prompt_file(count, max_count)
|
||||
if isinstance(image, list):
|
||||
# video
|
||||
if self.num_frames == 1:
|
||||
raise ValueError(f"Expected 1 img but got a list {len(image)}")
|
||||
if self.output_ext == 'webp':
|
||||
# save as animated webp
|
||||
duration = 1000 // self.fps # Convert fps to milliseconds per frame
|
||||
image[0].save(
|
||||
self.get_image_path(count, max_count),
|
||||
format='WEBP',
|
||||
append_images=image[1:],
|
||||
save_all=True,
|
||||
duration=duration, # Duration per frame in milliseconds
|
||||
loop=0, # 0 means loop forever
|
||||
quality=80 # Quality setting (0-100)
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported video format {self.output_ext}")
|
||||
else:
|
||||
# TODO save image gen header info for A1111 and us, our seeds probably wont match
|
||||
image.save(self.get_image_path(count, max_count))
|
||||
# do prompt file
|
||||
if self.add_prompt_file:
|
||||
self.save_prompt_file(count, max_count)
|
||||
|
||||
def save_prompt_file(self, count: int = 0, max_count=0):
|
||||
# save prompt file
|
||||
@@ -564,7 +990,10 @@ class GenerateImageConfig:
|
||||
prompt += ' --gr ' + str(self.guidance_rescale)
|
||||
|
||||
# get gen info
|
||||
f.write(self.prompt)
|
||||
try:
|
||||
f.write(self.prompt)
|
||||
except Exception as e:
|
||||
print(f"Error writing prompt file. Prompt contains non-unicode characters. {e}")
|
||||
|
||||
def _process_prompt_string(self):
|
||||
# we will try to support all sd-scripts where we can
|
||||
@@ -633,6 +1062,18 @@ class GenerateImageConfig:
|
||||
self.adapter_conditioning_scale = float(content)
|
||||
elif flag == 'ref':
|
||||
self.refiner_start_at = float(content)
|
||||
elif flag == 'ev':
|
||||
# split by comma
|
||||
self.extra_values = [float(val) for val in content.split(',')]
|
||||
elif flag == 'extra_values':
|
||||
# split by comma
|
||||
self.extra_values = [float(val) for val in content.split(',')]
|
||||
elif flag == 'frames':
|
||||
self.num_frames = int(content)
|
||||
elif flag == 'fps':
|
||||
self.fps = int(content)
|
||||
elif flag == 'ctrl_idx':
|
||||
self.ctrl_idx = int(content)
|
||||
|
||||
def post_process_embeddings(
|
||||
self,
|
||||
@@ -641,3 +1082,23 @@ class GenerateImageConfig:
|
||||
):
|
||||
# this is called after prompt embeds are encoded. We can override them in the future here
|
||||
pass
|
||||
|
||||
def log_image(self, image, count: int = 0, max_count=0):
|
||||
if self.logger is None:
|
||||
return
|
||||
|
||||
self.logger.log_image(image, count, self.prompt)
|
||||
|
||||
|
||||
def validate_configs(
|
||||
train_config: TrainConfig,
|
||||
model_config: ModelConfig,
|
||||
save_config: SaveConfig,
|
||||
):
|
||||
if model_config.is_flux:
|
||||
if save_config.save_format != 'diffusers':
|
||||
# make it diffusers
|
||||
save_config.save_format = 'diffusers'
|
||||
if model_config.use_flux_cfg:
|
||||
# bypass the embedding
|
||||
train_config.bypass_guidance_embedding = True
|
||||
|
||||
1252
toolkit/custom_adapter.py
Normal file
1252
toolkit/custom_adapter.py
Normal file
File diff suppressed because it is too large
Load Diff
@@ -8,6 +8,7 @@ from typing import List, TYPE_CHECKING
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
from PIL.ImageOps import exif_transpose
|
||||
from torchvision import transforms
|
||||
@@ -17,11 +18,61 @@ import albumentations as A
|
||||
|
||||
from toolkit.buckets import get_bucket_for_image_size, BucketResolution
|
||||
from toolkit.config_modules import DatasetConfig, preprocess_dataset_raw_config
|
||||
from toolkit.dataloader_mixins import CaptionMixin, BucketsMixin, LatentCachingMixin, Augments
|
||||
from toolkit.dataloader_mixins import CaptionMixin, BucketsMixin, LatentCachingMixin, Augments, CLIPCachingMixin
|
||||
from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO
|
||||
from toolkit.print import print_acc
|
||||
from toolkit.accelerator import get_accelerator
|
||||
|
||||
import platform
|
||||
|
||||
def is_native_windows():
|
||||
return platform.system() == "Windows" and platform.release() != "2"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
image_extensions = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
video_extensions = ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv']
|
||||
|
||||
|
||||
class RescaleTransform:
|
||||
"""Transform to rescale images to the range [-1, 1]."""
|
||||
|
||||
def __call__(self, image):
|
||||
return image * 2 - 1
|
||||
|
||||
|
||||
class NormalizeSDXLTransform:
|
||||
"""
|
||||
Transforms the range from 0 to 1 to SDXL mean and std per channel based on avgs over thousands of images
|
||||
|
||||
Mean: tensor([ 0.0002, -0.1034, -0.1879])
|
||||
Standard Deviation: tensor([0.5436, 0.5116, 0.5033])
|
||||
"""
|
||||
|
||||
def __call__(self, image):
|
||||
return transforms.Normalize(
|
||||
mean=[0.0002, -0.1034, -0.1879],
|
||||
std=[0.5436, 0.5116, 0.5033],
|
||||
)(image)
|
||||
|
||||
|
||||
class NormalizeSD15Transform:
|
||||
"""
|
||||
Transforms the range from 0 to 1 to SDXL mean and std per channel based on avgs over thousands of images
|
||||
|
||||
Mean: tensor([-0.1600, -0.2450, -0.3227])
|
||||
Standard Deviation: tensor([0.5319, 0.4997, 0.5139])
|
||||
|
||||
"""
|
||||
|
||||
def __call__(self, image):
|
||||
return transforms.Normalize(
|
||||
mean=[-0.1600, -0.2450, -0.3227],
|
||||
std=[0.5319, 0.4997, 0.5139],
|
||||
)(image)
|
||||
|
||||
|
||||
|
||||
class ImageDataset(Dataset, CaptionMixin):
|
||||
@@ -45,7 +96,7 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))]
|
||||
|
||||
# this might take a while
|
||||
print(f" - Preprocessing image dimensions")
|
||||
print_acc(f" - Preprocessing image dimensions")
|
||||
new_file_list = []
|
||||
bad_count = 0
|
||||
for file in tqdm(self.file_list):
|
||||
@@ -57,13 +108,13 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
|
||||
self.file_list = new_file_list
|
||||
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
print(f" - Found {bad_count} images that are too small")
|
||||
print_acc(f" - Found {len(self.file_list)} images")
|
||||
print_acc(f" - Found {bad_count} images that are too small")
|
||||
assert len(self.file_list) > 0, f"no images found in {self.path}"
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
RescaleTransform(),
|
||||
])
|
||||
|
||||
def get_config(self, key, default=None, required=False):
|
||||
@@ -80,7 +131,13 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
|
||||
def __getitem__(self, index):
|
||||
img_path = self.file_list[index]
|
||||
img = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
try:
|
||||
img = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
except Exception as e:
|
||||
print_acc(f"Error opening image: {img_path}")
|
||||
print_acc(e)
|
||||
# make a noise image if we can't open it
|
||||
img = Image.fromarray(np.random.randint(0, 255, (1024, 1024, 3), dtype=np.uint8))
|
||||
|
||||
# Downscale the source image first
|
||||
img = img.resize((int(img.size[0] * self.scale), int(img.size[1] * self.scale)), Image.BICUBIC)
|
||||
@@ -89,7 +146,7 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
if self.random_crop:
|
||||
if self.random_scale and min_img_size > self.resolution:
|
||||
if min_img_size < self.resolution:
|
||||
print(
|
||||
print_acc(
|
||||
f"Unexpected values: min_img_size={min_img_size}, self.resolution={self.resolution}, image file={img_path}")
|
||||
scale_size = self.resolution
|
||||
else:
|
||||
@@ -192,15 +249,15 @@ class PairedImageDataset(Dataset):
|
||||
matched_files = [t for t in (set(tuple(i) for i in matched_files))]
|
||||
|
||||
self.file_list = matched_files
|
||||
print(f" - Found {len(self.file_list)} matching pairs")
|
||||
print_acc(f" - Found {len(self.file_list)} matching pairs")
|
||||
else:
|
||||
self.file_list = [os.path.join(self.path, file) for file in os.listdir(self.path) if
|
||||
file.lower().endswith(supported_exts)]
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
print_acc(f" - Found {len(self.file_list)} images")
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
RescaleTransform(),
|
||||
])
|
||||
|
||||
def get_all_prompts(self):
|
||||
@@ -315,7 +372,7 @@ class PairedImageDataset(Dataset):
|
||||
return img, prompt, (self.neg_weight, self.pos_weight)
|
||||
|
||||
|
||||
class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -323,8 +380,9 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
batch_size=1,
|
||||
sd: 'StableDiffusion' = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dataset_config = dataset_config
|
||||
self.is_video = dataset_config.num_frames > 1
|
||||
super().__init__()
|
||||
folder_path = dataset_config.folder_path
|
||||
self.dataset_path = dataset_config.dataset_path
|
||||
if self.dataset_path is None:
|
||||
@@ -333,6 +391,7 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
self.is_caching_latents = dataset_config.cache_latents or dataset_config.cache_latents_to_disk
|
||||
self.is_caching_latents_to_memory = dataset_config.cache_latents
|
||||
self.is_caching_latents_to_disk = dataset_config.cache_latents_to_disk
|
||||
self.is_caching_clip_vision_to_disk = dataset_config.cache_clip_vision_to_disk
|
||||
self.epoch_num = 0
|
||||
|
||||
self.sd = sd
|
||||
@@ -353,10 +412,11 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
|
||||
# check if dataset_path is a folder or json
|
||||
if os.path.isdir(self.dataset_path):
|
||||
file_list = [
|
||||
os.path.join(self.dataset_path, file) for file in os.listdir(self.dataset_path) if
|
||||
file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))
|
||||
]
|
||||
extensions = image_extensions
|
||||
if self.is_video:
|
||||
# only look for videos
|
||||
extensions = video_extensions
|
||||
file_list = [os.path.join(root, file) for root, _, files in os.walk(self.dataset_path) for file in files if file.lower().endswith(tuple(extensions))]
|
||||
else:
|
||||
# assume json
|
||||
with open(self.dataset_path, 'r') as f:
|
||||
@@ -368,29 +428,88 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
# repeat the list
|
||||
file_list = file_list * self.dataset_config.num_repeats
|
||||
|
||||
if self.dataset_config.standardize_images:
|
||||
if self.sd.is_xl or self.sd.is_vega or self.sd.is_ssd:
|
||||
NormalizeMethod = NormalizeSDXLTransform
|
||||
else:
|
||||
NormalizeMethod = NormalizeSD15Transform
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
RescaleTransform(),
|
||||
NormalizeMethod(),
|
||||
])
|
||||
else:
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
RescaleTransform(),
|
||||
])
|
||||
|
||||
# this might take a while
|
||||
print(f" - Preprocessing image dimensions")
|
||||
print_acc(f"Dataset: {self.dataset_path}")
|
||||
if self.is_video:
|
||||
print_acc(f" - Preprocessing video dimensions")
|
||||
else:
|
||||
print_acc(f" - Preprocessing image dimensions")
|
||||
dataset_folder = self.dataset_path
|
||||
if not os.path.isdir(self.dataset_path):
|
||||
dataset_folder = os.path.dirname(dataset_folder)
|
||||
|
||||
dataset_size_file = os.path.join(dataset_folder, '.aitk_size.json')
|
||||
dataloader_version = "0.1.1"
|
||||
if os.path.exists(dataset_size_file):
|
||||
try:
|
||||
with open(dataset_size_file, 'r') as f:
|
||||
self.size_database = json.load(f)
|
||||
|
||||
if "__version__" not in self.size_database or self.size_database["__version__"] != dataloader_version:
|
||||
print_acc("Upgrading size database to new version")
|
||||
# old version, delete and recreate
|
||||
self.size_database = {}
|
||||
except Exception as e:
|
||||
print_acc(f"Error loading size database: {dataset_size_file}")
|
||||
print_acc(e)
|
||||
self.size_database = {}
|
||||
else:
|
||||
self.size_database = {}
|
||||
|
||||
self.size_database["__version__"] = dataloader_version
|
||||
|
||||
bad_count = 0
|
||||
for file in tqdm(file_list):
|
||||
try:
|
||||
file_item = FileItemDTO(
|
||||
sd=self.sd,
|
||||
path=file,
|
||||
dataset_config=dataset_config
|
||||
dataset_config=dataset_config,
|
||||
dataloader_transforms=self.transform,
|
||||
size_database=self.size_database,
|
||||
dataset_root=dataset_folder,
|
||||
)
|
||||
self.file_list.append(file_item)
|
||||
except Exception as e:
|
||||
print(traceback.format_exc())
|
||||
print(f"Error processing image: {file}")
|
||||
print(e)
|
||||
print_acc(traceback.format_exc())
|
||||
if self.is_video:
|
||||
print_acc(f"Error processing video: {file}")
|
||||
else:
|
||||
print_acc(f"Error processing image: {file}")
|
||||
print_acc(e)
|
||||
bad_count += 1
|
||||
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
# print(f" - Found {bad_count} images that are too small")
|
||||
assert len(self.file_list) > 0, f"no images found in {self.dataset_path}"
|
||||
# save the size database
|
||||
with open(dataset_size_file, 'w') as f:
|
||||
json.dump(self.size_database, f)
|
||||
|
||||
if self.is_video:
|
||||
print_acc(f" - Found {len(self.file_list)} videos")
|
||||
assert len(self.file_list) > 0, f"no videos found in {self.dataset_path}"
|
||||
else:
|
||||
print_acc(f" - Found {len(self.file_list)} images")
|
||||
assert len(self.file_list) > 0, f"no images found in {self.dataset_path}"
|
||||
|
||||
# handle x axis flips
|
||||
if self.dataset_config.flip_x:
|
||||
print(" - adding x axis flips")
|
||||
print_acc(" - adding x axis flips")
|
||||
current_file_list = [x for x in self.file_list]
|
||||
for file_item in current_file_list:
|
||||
# create a copy that is flipped on the x axis
|
||||
@@ -400,7 +519,7 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
|
||||
# handle y axis flips
|
||||
if self.dataset_config.flip_y:
|
||||
print(" - adding y axis flips")
|
||||
print_acc(" - adding y axis flips")
|
||||
current_file_list = [x for x in self.file_list]
|
||||
for file_item in current_file_list:
|
||||
# create a copy that is flipped on the y axis
|
||||
@@ -409,12 +528,10 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
self.file_list.append(new_file_item)
|
||||
|
||||
if self.dataset_config.flip_x or self.dataset_config.flip_y:
|
||||
print(f" - Found {len(self.file_list)} images after adding flips")
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
])
|
||||
if self.is_video:
|
||||
print_acc(f" - Found {len(self.file_list)} videos after adding flips")
|
||||
else:
|
||||
print_acc(f" - Found {len(self.file_list)} images after adding flips")
|
||||
|
||||
self.setup_epoch()
|
||||
|
||||
@@ -427,6 +544,8 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
self.setup_buckets()
|
||||
if self.is_caching_latents:
|
||||
self.cache_latents_all_latents()
|
||||
if self.is_caching_clip_vision_to_disk:
|
||||
self.cache_clip_vision_to_disk()
|
||||
else:
|
||||
if self.dataset_config.poi is not None:
|
||||
# handle cropping to a specific point of interest
|
||||
@@ -440,7 +559,7 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
return len(self.file_list)
|
||||
|
||||
def _get_single_item(self, index) -> 'FileItemDTO':
|
||||
file_item = copy.deepcopy(self.file_list[index])
|
||||
file_item: 'FileItemDTO' = copy.deepcopy(self.file_list[index])
|
||||
file_item.load_and_process_image(self.transform)
|
||||
file_item.load_caption(self.caption_dict)
|
||||
return file_item
|
||||
@@ -509,6 +628,13 @@ def get_dataloader_from_datasets(
|
||||
|
||||
# check if is caching latents
|
||||
|
||||
dataloader_kwargs = {}
|
||||
|
||||
if is_native_windows():
|
||||
dataloader_kwargs['num_workers'] = 0
|
||||
else:
|
||||
dataloader_kwargs['num_workers'] = dataset_config_list[0].num_workers
|
||||
dataloader_kwargs['prefetch_factor'] = dataset_config_list[0].prefetch_factor
|
||||
|
||||
if has_buckets:
|
||||
# make sure they all have buckets
|
||||
@@ -521,15 +647,15 @@ def get_dataloader_from_datasets(
|
||||
drop_last=False,
|
||||
shuffle=True,
|
||||
collate_fn=dto_collation, # Use the custom collate function
|
||||
num_workers=4
|
||||
**dataloader_kwargs
|
||||
)
|
||||
else:
|
||||
data_loader = DataLoader(
|
||||
concatenated_dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=True,
|
||||
num_workers=4,
|
||||
collate_fn=dto_collation
|
||||
collate_fn=dto_collation,
|
||||
**dataloader_kwargs
|
||||
)
|
||||
return data_loader
|
||||
|
||||
@@ -556,3 +682,19 @@ def trigger_dataloader_setup_epoch(dataloader: DataLoader):
|
||||
if hasattr(sub_dataset, 'setup_epoch'):
|
||||
sub_dataset.setup_epoch()
|
||||
sub_dataset.len = None
|
||||
|
||||
def get_dataloader_datasets(dataloader: DataLoader):
|
||||
# hacky but needed because of different types of datasets and dataloaders
|
||||
if isinstance(dataloader.dataset, list):
|
||||
datasets = []
|
||||
for dataset in dataloader.dataset:
|
||||
if hasattr(dataset, 'datasets'):
|
||||
for sub_dataset in dataset.datasets:
|
||||
datasets.append(sub_dataset)
|
||||
else:
|
||||
datasets.append(dataset)
|
||||
return datasets
|
||||
elif hasattr(dataloader.dataset, 'datasets'):
|
||||
return dataloader.dataset.datasets
|
||||
else:
|
||||
return [dataloader.dataset]
|
||||
|
||||
@@ -1,4 +1,8 @@
|
||||
import os
|
||||
import weakref
|
||||
from _weakref import ReferenceType
|
||||
from typing import TYPE_CHECKING, List, Union
|
||||
import cv2
|
||||
import torch
|
||||
import random
|
||||
|
||||
@@ -8,10 +12,12 @@ from PIL.ImageOps import exif_transpose
|
||||
from toolkit import image_utils
|
||||
from toolkit.dataloader_mixins import CaptionProcessingDTOMixin, ImageProcessingDTOMixin, LatentCachingFileItemDTOMixin, \
|
||||
ControlFileItemDTOMixin, ArgBreakMixin, PoiFileItemDTOMixin, MaskFileItemDTOMixin, AugmentationFileItemDTOMixin, \
|
||||
UnconditionalFileItemDTOMixin
|
||||
UnconditionalFileItemDTOMixin, ClipImageFileItemDTOMixin, InpaintControlFileItemDTOMixin
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.config_modules import DatasetConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
printed_messages = []
|
||||
|
||||
@@ -28,6 +34,8 @@ class FileItemDTO(
|
||||
CaptionProcessingDTOMixin,
|
||||
ImageProcessingDTOMixin,
|
||||
ControlFileItemDTOMixin,
|
||||
InpaintControlFileItemDTOMixin,
|
||||
ClipImageFileItemDTOMixin,
|
||||
MaskFileItemDTOMixin,
|
||||
AugmentationFileItemDTOMixin,
|
||||
UnconditionalFileItemDTOMixin,
|
||||
@@ -35,18 +43,47 @@ class FileItemDTO(
|
||||
ArgBreakMixin,
|
||||
):
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.path = kwargs.get('path', None)
|
||||
self.path = kwargs.get('path', '')
|
||||
self.dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
|
||||
# process width and height
|
||||
try:
|
||||
w, h = image_utils.get_image_size(self.path)
|
||||
except image_utils.UnknownImageFormat:
|
||||
print_once(f'Warning: Some images in the dataset cannot be fast read. ' + \
|
||||
f'This process is faster for png, jpeg')
|
||||
self.is_video = self.dataset_config.num_frames > 1
|
||||
size_database = kwargs.get('size_database', {})
|
||||
dataset_root = kwargs.get('dataset_root', None)
|
||||
if dataset_root is not None:
|
||||
# remove dataset root from path
|
||||
file_key = self.path.replace(dataset_root, '')
|
||||
else:
|
||||
file_key = os.path.basename(self.path)
|
||||
if file_key in size_database:
|
||||
w, h = size_database[file_key]
|
||||
elif self.is_video:
|
||||
# Open the video file
|
||||
video = cv2.VideoCapture(self.path)
|
||||
|
||||
# Check if video opened successfully
|
||||
if not video.isOpened():
|
||||
raise Exception(f"Error: Could not open video file {self.path}")
|
||||
|
||||
# Get width and height
|
||||
width = int(video.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(video.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
# Release the video capture object immediately
|
||||
video.release()
|
||||
size_database[file_key] = (width, height)
|
||||
else:
|
||||
# original method is significantly faster, but some images are read sideways. Not sure why. Do slow method for now.
|
||||
# process width and height
|
||||
# try:
|
||||
# w, h = image_utils.get_image_size(self.path)
|
||||
# except image_utils.UnknownImageFormat:
|
||||
# print_once(f'Warning: Some images in the dataset cannot be fast read. ' + \
|
||||
# f'This process is faster for png, jpeg')
|
||||
img = exif_transpose(Image.open(self.path))
|
||||
h, w = img.size
|
||||
w, h = img.size
|
||||
size_database[file_key] = (w, h)
|
||||
self.width: int = w
|
||||
self.height: int = h
|
||||
self.dataloader_transforms = kwargs.get('dataloader_transforms', None)
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
# self.caption_path: str = kwargs.get('caption_path', None)
|
||||
@@ -62,6 +99,7 @@ class FileItemDTO(
|
||||
self.flip_x: bool = kwargs.get('flip_x', False)
|
||||
self.flip_y: bool = kwargs.get('flip_x', False)
|
||||
self.augments: List[str] = self.dataset_config.augments
|
||||
self.loss_multiplier: float = self.dataset_config.loss_multiplier
|
||||
|
||||
self.network_weight: float = self.dataset_config.network_weight
|
||||
self.is_reg = self.dataset_config.is_reg
|
||||
@@ -71,6 +109,8 @@ class FileItemDTO(
|
||||
self.tensor = None
|
||||
self.cleanup_latent()
|
||||
self.cleanup_control()
|
||||
self.cleanup_inpaint()
|
||||
self.cleanup_clip_image()
|
||||
self.cleanup_mask()
|
||||
self.cleanup_unconditional()
|
||||
|
||||
@@ -83,11 +123,15 @@ class DataLoaderBatchDTO:
|
||||
self.tensor: Union[torch.Tensor, None] = None
|
||||
self.latents: Union[torch.Tensor, None] = None
|
||||
self.control_tensor: Union[torch.Tensor, None] = None
|
||||
self.clip_image_tensor: Union[torch.Tensor, None] = None
|
||||
self.mask_tensor: Union[torch.Tensor, None] = None
|
||||
self.unaugmented_tensor: Union[torch.Tensor, None] = None
|
||||
self.unconditional_tensor: Union[torch.Tensor, None] = None
|
||||
self.unconditional_latents: Union[torch.Tensor, None] = None
|
||||
self.clip_image_embeds: Union[List[dict], None] = None
|
||||
self.clip_image_embeds_unconditional: Union[List[dict], None] = None
|
||||
self.sigmas: Union[torch.Tensor, None] = None # can be added elseware and passed along training code
|
||||
self.extra_values: Union[torch.Tensor, None] = torch.tensor([x.extra_values for x in self.file_items]) if len(self.file_items[0].extra_values) > 0 else None
|
||||
if not is_latents_cached:
|
||||
# only return a tensor if latents are not cached
|
||||
self.tensor: torch.Tensor = torch.cat([x.tensor.unsqueeze(0) for x in self.file_items])
|
||||
@@ -112,6 +156,39 @@ class DataLoaderBatchDTO:
|
||||
else:
|
||||
control_tensors.append(x.control_tensor)
|
||||
self.control_tensor = torch.cat([x.unsqueeze(0) for x in control_tensors])
|
||||
|
||||
self.inpaint_tensor: Union[torch.Tensor, None] = None
|
||||
if any([x.inpaint_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
base_inpaint_tensor = None
|
||||
for x in self.file_items:
|
||||
if x.inpaint_tensor is not None:
|
||||
base_inpaint_tensor = x.inpaint_tensor
|
||||
break
|
||||
inpaint_tensors = []
|
||||
for x in self.file_items:
|
||||
if x.inpaint_tensor is None:
|
||||
inpaint_tensors.append(torch.zeros_like(base_inpaint_tensor))
|
||||
else:
|
||||
inpaint_tensors.append(x.inpaint_tensor)
|
||||
self.inpaint_tensor = torch.cat([x.unsqueeze(0) for x in inpaint_tensors])
|
||||
|
||||
self.loss_multiplier_list: List[float] = [x.loss_multiplier for x in self.file_items]
|
||||
|
||||
if any([x.clip_image_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
base_clip_image_tensor = None
|
||||
for x in self.file_items:
|
||||
if x.clip_image_tensor is not None:
|
||||
base_clip_image_tensor = x.clip_image_tensor
|
||||
break
|
||||
clip_image_tensors = []
|
||||
for x in self.file_items:
|
||||
if x.clip_image_tensor is None:
|
||||
clip_image_tensors.append(torch.zeros_like(base_clip_image_tensor))
|
||||
else:
|
||||
clip_image_tensors.append(x.clip_image_tensor)
|
||||
self.clip_image_tensor = torch.cat([x.unsqueeze(0) for x in clip_image_tensors])
|
||||
|
||||
if any([x.mask_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
@@ -159,6 +236,23 @@ class DataLoaderBatchDTO:
|
||||
else:
|
||||
unconditional_tensor.append(x.unconditional_tensor)
|
||||
self.unconditional_tensor = torch.cat([x.unsqueeze(0) for x in unconditional_tensor])
|
||||
|
||||
if any([x.clip_image_embeds is not None for x in self.file_items]):
|
||||
self.clip_image_embeds = []
|
||||
for x in self.file_items:
|
||||
if x.clip_image_embeds is not None:
|
||||
self.clip_image_embeds.append(x.clip_image_embeds)
|
||||
else:
|
||||
raise Exception("clip_image_embeds is None for some file items")
|
||||
|
||||
if any([x.clip_image_embeds_unconditional is not None for x in self.file_items]):
|
||||
self.clip_image_embeds_unconditional = []
|
||||
for x in self.file_items:
|
||||
if x.clip_image_embeds_unconditional is not None:
|
||||
self.clip_image_embeds_unconditional.append(x.clip_image_embeds_unconditional)
|
||||
else:
|
||||
raise Exception("clip_image_embeds_unconditional is None for some file items")
|
||||
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise e
|
||||
@@ -175,11 +269,7 @@ class DataLoaderBatchDTO:
|
||||
to_replace_list=None,
|
||||
add_if_not_present=True
|
||||
):
|
||||
return [x.get_caption(
|
||||
trigger=trigger,
|
||||
to_replace_list=to_replace_list,
|
||||
add_if_not_present=add_if_not_present
|
||||
) for x in self.file_items]
|
||||
return [x.caption for x in self.file_items]
|
||||
|
||||
def get_caption_short_list(
|
||||
self,
|
||||
@@ -187,12 +277,7 @@ class DataLoaderBatchDTO:
|
||||
to_replace_list=None,
|
||||
add_if_not_present=True
|
||||
):
|
||||
return [x.get_caption(
|
||||
trigger=trigger,
|
||||
to_replace_list=to_replace_list,
|
||||
add_if_not_present=add_if_not_present,
|
||||
short_caption=True
|
||||
) for x in self.file_items]
|
||||
return [x.caption_short for x in self.file_items]
|
||||
|
||||
def cleanup(self):
|
||||
del self.latents
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
88
toolkit/dequantize.py
Normal file
88
toolkit/dequantize.py
Normal file
@@ -0,0 +1,88 @@
|
||||
|
||||
|
||||
from functools import partial
|
||||
from optimum.quanto.tensor import QTensor
|
||||
import torch
|
||||
|
||||
|
||||
def hacked_state_dict(self, *args, **kwargs):
|
||||
orig_state_dict = self.orig_state_dict(*args, **kwargs)
|
||||
new_state_dict = {}
|
||||
for key, value in orig_state_dict.items():
|
||||
if key.endswith("._scale"):
|
||||
continue
|
||||
if key.endswith(".input_scale"):
|
||||
continue
|
||||
if key.endswith(".output_scale"):
|
||||
continue
|
||||
if key.endswith("._data"):
|
||||
key = key[:-6]
|
||||
scale = orig_state_dict[key + "._scale"]
|
||||
# scale is the original dtype
|
||||
dtype = scale.dtype
|
||||
scale = scale.float()
|
||||
value = value.float()
|
||||
dequantized = value * scale
|
||||
|
||||
# handle input and output scaling if they exist
|
||||
input_scale = orig_state_dict.get(key + ".input_scale")
|
||||
|
||||
if input_scale is not None:
|
||||
# make sure the tensor is 1.0
|
||||
if input_scale.item() != 1.0:
|
||||
raise ValueError("Input scale is not 1.0, cannot dequantize")
|
||||
|
||||
output_scale = orig_state_dict.get(key + ".output_scale")
|
||||
|
||||
if output_scale is not None:
|
||||
# make sure the tensor is 1.0
|
||||
if output_scale.item() != 1.0:
|
||||
raise ValueError("Output scale is not 1.0, cannot dequantize")
|
||||
|
||||
new_state_dict[key] = dequantized.to('cpu', dtype=dtype)
|
||||
else:
|
||||
new_state_dict[key] = value
|
||||
return new_state_dict
|
||||
|
||||
# hacks the state dict so we can dequantize before saving
|
||||
def patch_dequantization_on_save(model):
|
||||
model.orig_state_dict = model.state_dict
|
||||
model.state_dict = partial(hacked_state_dict, model)
|
||||
|
||||
|
||||
def dequantize_parameter(module: torch.nn.Module, param_name: str) -> bool:
|
||||
"""
|
||||
Convert a quantized parameter back to a regular Parameter with floating point values.
|
||||
|
||||
Args:
|
||||
module: The module containing the parameter to unquantize
|
||||
param_name: Name of the parameter to unquantize (e.g., 'weight', 'bias')
|
||||
|
||||
Returns:
|
||||
bool: True if parameter was unquantized, False if it was already unquantized
|
||||
"""
|
||||
|
||||
# Check if the parameter exists
|
||||
if not hasattr(module, param_name):
|
||||
raise AttributeError(f"Module has no parameter named '{param_name}'")
|
||||
|
||||
param = getattr(module, param_name)
|
||||
|
||||
# If it's not a parameter or not quantized, nothing to do
|
||||
if not isinstance(param, torch.nn.Parameter):
|
||||
raise TypeError(f"'{param_name}' is not a Parameter")
|
||||
if not isinstance(param, QTensor):
|
||||
return False
|
||||
|
||||
# Convert to float tensor while preserving device and requires_grad
|
||||
with torch.no_grad():
|
||||
float_tensor = param.float()
|
||||
new_param = torch.nn.Parameter(
|
||||
float_tensor,
|
||||
requires_grad=param.requires_grad
|
||||
)
|
||||
|
||||
# Replace the parameter
|
||||
setattr(module, param_name, new_param)
|
||||
|
||||
return True
|
||||
346
toolkit/ema.py
Normal file
346
toolkit/ema.py
Normal file
@@ -0,0 +1,346 @@
|
||||
from __future__ import division
|
||||
from __future__ import unicode_literals
|
||||
|
||||
from typing import Iterable, Optional
|
||||
import weakref
|
||||
import copy
|
||||
import contextlib
|
||||
from toolkit.optimizers.optimizer_utils import copy_stochastic
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
# Partially based on:
|
||||
# https://github.com/tensorflow/tensorflow/blob/r1.13/tensorflow/python/training/moving_averages.py
|
||||
class ExponentialMovingAverage:
|
||||
"""
|
||||
Maintains (exponential) moving average of a set of parameters.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter` (typically from
|
||||
`model.parameters()`).
|
||||
Note that EMA is computed on *all* provided parameters,
|
||||
regardless of whether or not they have `requires_grad = True`;
|
||||
this allows a single EMA object to be consistantly used even
|
||||
if which parameters are trainable changes step to step.
|
||||
|
||||
If you want to some parameters in the EMA, do not pass them
|
||||
to the object in the first place. For example:
|
||||
|
||||
ExponentialMovingAverage(
|
||||
parameters=[p for p in model.parameters() if p.requires_grad],
|
||||
decay=0.9
|
||||
)
|
||||
|
||||
will ignore parameters that do not require grad.
|
||||
|
||||
decay: The exponential decay.
|
||||
|
||||
use_num_updates: Whether to use number of updates when computing
|
||||
averages.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
parameters: Iterable[torch.nn.Parameter] = None,
|
||||
decay: float = 0.995,
|
||||
use_num_updates: bool = False,
|
||||
# feeds back the decat to the parameter
|
||||
use_feedback: bool = False,
|
||||
param_multiplier: float = 1.0
|
||||
):
|
||||
if parameters is None:
|
||||
raise ValueError("parameters must be provided")
|
||||
if decay < 0.0 or decay > 1.0:
|
||||
raise ValueError('Decay must be between 0 and 1')
|
||||
self.decay = decay
|
||||
self.num_updates = 0 if use_num_updates else None
|
||||
self.use_feedback = use_feedback
|
||||
self.param_multiplier = param_multiplier
|
||||
parameters = list(parameters)
|
||||
self.shadow_params = [
|
||||
p.clone().detach()
|
||||
for p in parameters
|
||||
]
|
||||
self.collected_params = None
|
||||
self._is_train_mode = True
|
||||
# By maintaining only a weakref to each parameter,
|
||||
# we maintain the old GC behaviour of ExponentialMovingAverage:
|
||||
# if the model goes out of scope but the ExponentialMovingAverage
|
||||
# is kept, no references to the model or its parameters will be
|
||||
# maintained, and the model will be cleaned up.
|
||||
self._params_refs = [weakref.ref(p) for p in parameters]
|
||||
|
||||
def _get_parameters(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]]
|
||||
) -> Iterable[torch.nn.Parameter]:
|
||||
if parameters is None:
|
||||
parameters = [p() for p in self._params_refs]
|
||||
if any(p is None for p in parameters):
|
||||
raise ValueError(
|
||||
"(One of) the parameters with which this "
|
||||
"ExponentialMovingAverage "
|
||||
"was initialized no longer exists (was garbage collected);"
|
||||
" please either provide `parameters` explicitly or keep "
|
||||
"the model to which they belong from being garbage "
|
||||
"collected."
|
||||
)
|
||||
return parameters
|
||||
else:
|
||||
parameters = list(parameters)
|
||||
if len(parameters) != len(self.shadow_params):
|
||||
raise ValueError(
|
||||
"Number of parameters passed as argument is different "
|
||||
"from number of shadow parameters maintained by this "
|
||||
"ExponentialMovingAverage"
|
||||
)
|
||||
return parameters
|
||||
|
||||
def update(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
) -> None:
|
||||
"""
|
||||
Update currently maintained parameters.
|
||||
|
||||
Call this every time the parameters are updated, such as the result of
|
||||
the `optimizer.step()` call.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; usually the same set of
|
||||
parameters used to initialize this object. If `None`, the
|
||||
parameters with which this `ExponentialMovingAverage` was
|
||||
initialized will be used.
|
||||
"""
|
||||
parameters = self._get_parameters(parameters)
|
||||
decay = self.decay
|
||||
if self.num_updates is not None:
|
||||
self.num_updates += 1
|
||||
decay = min(
|
||||
decay,
|
||||
(1 + self.num_updates) / (10 + self.num_updates)
|
||||
)
|
||||
one_minus_decay = 1.0 - decay
|
||||
with torch.no_grad():
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
s_param_float = s_param.float()
|
||||
if s_param.dtype != torch.float32:
|
||||
s_param_float = s_param_float.to(torch.float32)
|
||||
param_float = param
|
||||
if param.dtype != torch.float32:
|
||||
param_float = param_float.to(torch.float32)
|
||||
tmp = (s_param_float - param_float)
|
||||
# tmp will be a new tensor so we can do in-place
|
||||
tmp.mul_(one_minus_decay)
|
||||
s_param_float.sub_(tmp)
|
||||
|
||||
update_param = False
|
||||
if self.use_feedback:
|
||||
param_float.add_(tmp)
|
||||
update_param = True
|
||||
|
||||
if self.param_multiplier != 1.0:
|
||||
param_float.mul_(self.param_multiplier)
|
||||
update_param = True
|
||||
|
||||
if s_param.dtype != torch.float32:
|
||||
copy_stochastic(s_param, s_param_float)
|
||||
|
||||
if update_param and param.dtype != torch.float32:
|
||||
copy_stochastic(param, param_float)
|
||||
|
||||
|
||||
def copy_to(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
) -> None:
|
||||
"""
|
||||
Copy current averaged parameters into given collection of parameters.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored moving averages. If `None`, the
|
||||
parameters with which this `ExponentialMovingAverage` was
|
||||
initialized will be used.
|
||||
"""
|
||||
parameters = self._get_parameters(parameters)
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
param.data.copy_(s_param.data)
|
||||
|
||||
def store(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
) -> None:
|
||||
"""
|
||||
Save the current parameters for restoring later.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
temporarily stored. If `None`, the parameters of with which this
|
||||
`ExponentialMovingAverage` was initialized will be used.
|
||||
"""
|
||||
parameters = self._get_parameters(parameters)
|
||||
self.collected_params = [
|
||||
param.clone()
|
||||
for param in parameters
|
||||
]
|
||||
|
||||
def restore(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
) -> None:
|
||||
"""
|
||||
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. If `None`, the
|
||||
parameters with which this `ExponentialMovingAverage` was
|
||||
initialized will be used.
|
||||
"""
|
||||
if self.collected_params is None:
|
||||
raise RuntimeError(
|
||||
"This ExponentialMovingAverage has no `store()`ed weights "
|
||||
"to `restore()`"
|
||||
)
|
||||
parameters = self._get_parameters(parameters)
|
||||
for c_param, param in zip(self.collected_params, parameters):
|
||||
param.data.copy_(c_param.data)
|
||||
|
||||
@contextlib.contextmanager
|
||||
def average_parameters(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
):
|
||||
r"""
|
||||
Context manager for validation/inference with averaged parameters.
|
||||
|
||||
Equivalent to:
|
||||
|
||||
ema.store()
|
||||
ema.copy_to()
|
||||
try:
|
||||
...
|
||||
finally:
|
||||
ema.restore()
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored parameters. If `None`, the
|
||||
parameters with which this `ExponentialMovingAverage` was
|
||||
initialized will be used.
|
||||
"""
|
||||
parameters = self._get_parameters(parameters)
|
||||
self.store(parameters)
|
||||
self.copy_to(parameters)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self.restore(parameters)
|
||||
|
||||
def to(self, device=None, dtype=None) -> None:
|
||||
r"""Move internal buffers of the ExponentialMovingAverage to `device`.
|
||||
|
||||
Args:
|
||||
device: like `device` argument to `torch.Tensor.to`
|
||||
"""
|
||||
# .to() on the tensors handles None correctly
|
||||
self.shadow_params = [
|
||||
p.to(device=device, dtype=dtype)
|
||||
if p.is_floating_point()
|
||||
else p.to(device=device)
|
||||
for p in self.shadow_params
|
||||
]
|
||||
if self.collected_params is not None:
|
||||
self.collected_params = [
|
||||
p.to(device=device, dtype=dtype)
|
||||
if p.is_floating_point()
|
||||
else p.to(device=device)
|
||||
for p in self.collected_params
|
||||
]
|
||||
return
|
||||
|
||||
def state_dict(self) -> dict:
|
||||
r"""Returns the state of the ExponentialMovingAverage as a dict."""
|
||||
# Following PyTorch conventions, references to tensors are returned:
|
||||
# "returns a reference to the state and not its copy!" -
|
||||
# https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict
|
||||
return {
|
||||
"decay": self.decay,
|
||||
"num_updates": self.num_updates,
|
||||
"shadow_params": self.shadow_params,
|
||||
"collected_params": self.collected_params
|
||||
}
|
||||
|
||||
def load_state_dict(self, state_dict: dict) -> None:
|
||||
r"""Loads the ExponentialMovingAverage state.
|
||||
|
||||
Args:
|
||||
state_dict (dict): EMA state. Should be an object returned
|
||||
from a call to :meth:`state_dict`.
|
||||
"""
|
||||
# deepcopy, to be consistent with module API
|
||||
state_dict = copy.deepcopy(state_dict)
|
||||
self.decay = state_dict["decay"]
|
||||
if self.decay < 0.0 or self.decay > 1.0:
|
||||
raise ValueError('Decay must be between 0 and 1')
|
||||
self.num_updates = state_dict["num_updates"]
|
||||
assert self.num_updates is None or isinstance(self.num_updates, int), \
|
||||
"Invalid num_updates"
|
||||
|
||||
self.shadow_params = state_dict["shadow_params"]
|
||||
assert isinstance(self.shadow_params, list), \
|
||||
"shadow_params must be a list"
|
||||
assert all(
|
||||
isinstance(p, torch.Tensor) for p in self.shadow_params
|
||||
), "shadow_params must all be Tensors"
|
||||
|
||||
self.collected_params = state_dict["collected_params"]
|
||||
if self.collected_params is not None:
|
||||
assert isinstance(self.collected_params, list), \
|
||||
"collected_params must be a list"
|
||||
assert all(
|
||||
isinstance(p, torch.Tensor) for p in self.collected_params
|
||||
), "collected_params must all be Tensors"
|
||||
assert len(self.collected_params) == len(self.shadow_params), \
|
||||
"collected_params and shadow_params had different lengths"
|
||||
|
||||
if len(self.shadow_params) == len(self._params_refs):
|
||||
# Consistant with torch.optim.Optimizer, cast things to consistant
|
||||
# device and dtype with the parameters
|
||||
params = [p() for p in self._params_refs]
|
||||
# If parameters have been garbage collected, just load the state
|
||||
# we were given without change.
|
||||
if not any(p is None for p in params):
|
||||
# ^ parameter references are still good
|
||||
for i, p in enumerate(params):
|
||||
self.shadow_params[i] = self.shadow_params[i].to(
|
||||
device=p.device, dtype=p.dtype
|
||||
)
|
||||
if self.collected_params is not None:
|
||||
self.collected_params[i] = self.collected_params[i].to(
|
||||
device=p.device, dtype=p.dtype
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Tried to `load_state_dict()` with the wrong number of "
|
||||
"parameters in the saved state."
|
||||
)
|
||||
|
||||
def eval(self):
|
||||
if self._is_train_mode:
|
||||
with torch.no_grad():
|
||||
self.store()
|
||||
self.copy_to()
|
||||
self._is_train_mode = False
|
||||
|
||||
def train(self):
|
||||
if not self._is_train_mode:
|
||||
with torch.no_grad():
|
||||
self.restore()
|
||||
self._is_train_mode = True
|
||||
@@ -86,18 +86,19 @@ class Embedding:
|
||||
self.orig_embeds_params = [x.get_input_embeddings().weight.data.clone() for x in self.text_encoder_list]
|
||||
|
||||
def restore_embeddings(self):
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(self.text_encoder_list,
|
||||
self.tokenizer_list,
|
||||
self.orig_embeds_params,
|
||||
self.placeholder_token_ids):
|
||||
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
index_no_updates[
|
||||
min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
|
||||
with torch.no_grad():
|
||||
with torch.no_grad():
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(self.text_encoder_list,
|
||||
self.tokenizer_list,
|
||||
self.orig_embeds_params,
|
||||
self.placeholder_token_ids):
|
||||
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
index_no_updates[ min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
|
||||
text_encoder.get_input_embeddings().weight[
|
||||
index_no_updates
|
||||
] = orig_embeds[index_no_updates]
|
||||
weight = text_encoder.get_input_embeddings().weight
|
||||
pass
|
||||
|
||||
def get_trainable_params(self):
|
||||
params = []
|
||||
|
||||
827
toolkit/guidance.py
Normal file
827
toolkit/guidance.py
Normal file
@@ -0,0 +1,827 @@
|
||||
import torch
|
||||
from typing import Literal, Optional
|
||||
|
||||
from toolkit.basic import value_map
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
from toolkit.prompt_utils import PromptEmbeds, concat_prompt_embeds
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from toolkit.config_modules import TrainConfig
|
||||
|
||||
GuidanceType = Literal["targeted", "polarity", "targeted_polarity", "direct"]
|
||||
|
||||
DIFFERENTIAL_SCALER = 0.2
|
||||
|
||||
|
||||
# DIFFERENTIAL_SCALER = 0.25
|
||||
|
||||
|
||||
def get_differential_mask(
|
||||
conditional_latents: torch.Tensor,
|
||||
unconditional_latents: torch.Tensor,
|
||||
threshold: float = 0.2,
|
||||
gradient: bool = False,
|
||||
):
|
||||
# make a differential mask
|
||||
differential_mask = torch.abs(conditional_latents - unconditional_latents)
|
||||
if len(differential_mask.shape) == 4:
|
||||
max_differential = \
|
||||
differential_mask.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
|
||||
elif len(differential_mask.shape) == 5:
|
||||
max_differential = \
|
||||
differential_mask.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0].max(dim=4, keepdim=True)[0]
|
||||
differential_scaler = 1.0 / max_differential
|
||||
differential_mask = differential_mask * differential_scaler
|
||||
|
||||
if gradient:
|
||||
# wew need to scale it to 0-1
|
||||
# differential_mask = differential_mask - differential_mask.min()
|
||||
# differential_mask = differential_mask / differential_mask.max()
|
||||
# add 0.2 threshold to both sides and clip
|
||||
differential_mask = value_map(
|
||||
differential_mask,
|
||||
differential_mask.min(),
|
||||
differential_mask.max(),
|
||||
0 - threshold,
|
||||
1 + threshold
|
||||
)
|
||||
differential_mask = torch.clamp(differential_mask, 0.0, 1.0)
|
||||
else:
|
||||
|
||||
# make everything less than 0.2 be 0.0 and everything else be 1.0
|
||||
differential_mask = torch.where(
|
||||
differential_mask < threshold,
|
||||
torch.zeros_like(differential_mask),
|
||||
torch.ones_like(differential_mask)
|
||||
)
|
||||
return differential_mask
|
||||
|
||||
|
||||
def get_targeted_polarity_loss(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
**kwargs
|
||||
):
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
with torch.no_grad():
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
# inputs_abs_mean = torch.abs(conditional_latents).mean(dim=[1, 2, 3], keepdim=True)
|
||||
# noise_abs_mean = torch.abs(noise).mean(dim=[1, 2, 3], keepdim=True)
|
||||
differential_scaler = DIFFERENTIAL_SCALER
|
||||
|
||||
unconditional_diff = (unconditional_latents - conditional_latents)
|
||||
unconditional_diff_noise = unconditional_diff * differential_scaler
|
||||
conditional_diff = (conditional_latents - unconditional_latents)
|
||||
conditional_diff_noise = conditional_diff * differential_scaler
|
||||
conditional_diff_noise = conditional_diff_noise.detach().requires_grad_(False)
|
||||
unconditional_diff_noise = unconditional_diff_noise.detach().requires_grad_(False)
|
||||
#
|
||||
baseline_conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
baseline_unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
conditional_noise = noise + unconditional_diff_noise
|
||||
unconditional_noise = noise + conditional_diff_noise
|
||||
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
conditional_noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
unconditional_noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
# double up everything to run it through all at once
|
||||
cat_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
cat_latents = torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0)
|
||||
cat_timesteps = torch.cat([timesteps, timesteps], dim=0)
|
||||
# cat_baseline_noisy_latents = torch.cat(
|
||||
# [baseline_conditional_noisy_latents, baseline_unconditional_noisy_latents],
|
||||
# dim=0
|
||||
# )
|
||||
|
||||
# Disable the LoRA network so we can predict parent network knowledge without it
|
||||
# sd.network.is_active = False
|
||||
# sd.unet.eval()
|
||||
|
||||
# Predict noise to get a baseline of what the parent network wants to do with the latents + noise.
|
||||
# This acts as our control to preserve the unaltered parts of the image.
|
||||
# baseline_prediction = sd.predict_noise(
|
||||
# latents=cat_baseline_noisy_latents.to(device, dtype=dtype).detach(),
|
||||
# conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
|
||||
# timestep=cat_timesteps,
|
||||
# guidance_scale=1.0,
|
||||
# **pred_kwargs # adapter residuals in here
|
||||
# ).detach()
|
||||
|
||||
# conditional_baseline_prediction, unconditional_baseline_prediction = torch.chunk(baseline_prediction, 2, dim=0)
|
||||
|
||||
# negative_network_weights = [weight * -1.0 for weight in network_weight_list]
|
||||
# positive_network_weights = [weight * 1.0 for weight in network_weight_list]
|
||||
# cat_network_weight_list = positive_network_weights + negative_network_weights
|
||||
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
# sd.network.is_active = True
|
||||
|
||||
# sd.network.multiplier = cat_network_weight_list
|
||||
|
||||
# do our prediction with LoRA active on the scaled guidance latents
|
||||
prediction = sd.predict_noise(
|
||||
latents=cat_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=cat_timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
|
||||
# prediction = prediction - baseline_prediction
|
||||
|
||||
pred_pos, pred_neg = torch.chunk(prediction, 2, dim=0)
|
||||
# pred_pos = pred_pos - conditional_baseline_prediction
|
||||
# pred_neg = pred_neg - unconditional_baseline_prediction
|
||||
|
||||
pred_loss = torch.nn.functional.mse_loss(
|
||||
pred_pos.float(),
|
||||
conditional_noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
pred_loss = pred_loss.mean([1, 2, 3])
|
||||
|
||||
pred_neg_loss = torch.nn.functional.mse_loss(
|
||||
pred_neg.float(),
|
||||
unconditional_noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
pred_neg_loss = pred_neg_loss.mean([1, 2, 3])
|
||||
|
||||
loss = pred_loss + pred_neg_loss
|
||||
|
||||
loss = loss.mean()
|
||||
loss.backward()
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
loss.requires_grad_(True)
|
||||
|
||||
return loss
|
||||
|
||||
def get_direct_guidance_loss(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: 'PromptEmbeds',
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
unconditional_embeds: Optional[PromptEmbeds] = None,
|
||||
mask_multiplier=None,
|
||||
prior_pred=None,
|
||||
**kwargs
|
||||
):
|
||||
with torch.no_grad():
|
||||
# Perform targeted guidance (working title)
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
|
||||
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
# target_noise,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
# sd.network.is_active = True
|
||||
|
||||
# sd.network.multiplier = network_weight_list
|
||||
# do our prediction with LoRA active on the scaled guidance latents
|
||||
if unconditional_embeds is not None:
|
||||
unconditional_embeds = unconditional_embeds.to(device, dtype=dtype).detach()
|
||||
unconditional_embeds = concat_prompt_embeds([unconditional_embeds, unconditional_embeds])
|
||||
|
||||
prediction = sd.predict_noise(
|
||||
latents=torch.cat([unconditional_noisy_latents, conditional_noisy_latents]).to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=concat_prompt_embeds([conditional_embeds,conditional_embeds]).to(device, dtype=dtype).detach(),
|
||||
unconditional_embeddings=unconditional_embeds,
|
||||
timestep=torch.cat([timesteps, timesteps]),
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
|
||||
noise_pred_uncond, noise_pred_cond = torch.chunk(prediction, 2, dim=0)
|
||||
|
||||
guidance_scale = 1.1
|
||||
guidance_pred = noise_pred_uncond + guidance_scale * (
|
||||
noise_pred_cond - noise_pred_uncond
|
||||
)
|
||||
|
||||
guidance_loss = torch.nn.functional.mse_loss(
|
||||
guidance_pred.float(),
|
||||
noise.detach().float(),
|
||||
reduction="none"
|
||||
)
|
||||
if mask_multiplier is not None:
|
||||
guidance_loss = guidance_loss * mask_multiplier
|
||||
|
||||
guidance_loss = guidance_loss.mean([1, 2, 3])
|
||||
|
||||
guidance_loss = guidance_loss.mean()
|
||||
|
||||
# loss = guidance_loss + masked_noise_loss
|
||||
loss = guidance_loss
|
||||
|
||||
loss.backward()
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
loss.requires_grad_(True)
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
# targeted
|
||||
def get_targeted_guidance_loss(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: 'PromptEmbeds',
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
**kwargs
|
||||
):
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
# Encode the unconditional image into latents
|
||||
unconditional_noisy_latents = sd.noise_scheduler.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
)
|
||||
conditional_noisy_latents = sd.noise_scheduler.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
)
|
||||
|
||||
# was_network_active = self.network.is_active
|
||||
sd.network.is_active = False
|
||||
sd.unet.eval()
|
||||
|
||||
target_differential = unconditional_latents - conditional_latents
|
||||
# scale our loss by the differential scaler
|
||||
target_differential_abs = target_differential.abs()
|
||||
target_differential_abs_min = \
|
||||
target_differential_abs.min(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
|
||||
target_differential_abs_max = \
|
||||
target_differential_abs.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
|
||||
|
||||
min_guidance = 1.0
|
||||
max_guidance = 2.0
|
||||
|
||||
differential_scaler = value_map(
|
||||
target_differential_abs,
|
||||
target_differential_abs_min,
|
||||
target_differential_abs_max,
|
||||
min_guidance,
|
||||
max_guidance
|
||||
).detach()
|
||||
|
||||
|
||||
# With LoRA network bypassed, predict noise to get a baseline of what the network
|
||||
# wants to do with the latents + noise. Pass our target latents here for the input.
|
||||
target_unconditional = sd.predict_noise(
|
||||
latents=unconditional_noisy_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=conditional_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
).detach()
|
||||
prior_prediction_loss = torch.nn.functional.mse_loss(
|
||||
target_unconditional.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
).detach().clone()
|
||||
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
sd.network.is_active = True
|
||||
sd.network.multiplier = network_weight_list + [x + -1.0 for x in network_weight_list]
|
||||
|
||||
# with LoRA active, predict the noise with the scaled differential latents added. This will allow us
|
||||
# the opportunity to predict the differential + noise that was added to the latents.
|
||||
prediction = sd.predict_noise(
|
||||
latents=torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0).to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=concat_prompt_embeds([conditional_embeds, conditional_embeds]).to(device, dtype=dtype).detach(),
|
||||
timestep=torch.cat([timesteps, timesteps], dim=0),
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
|
||||
prediction_conditional, prediction_unconditional = torch.chunk(prediction, 2, dim=0)
|
||||
|
||||
conditional_loss = torch.nn.functional.mse_loss(
|
||||
prediction_conditional.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
unconditional_loss = torch.nn.functional.mse_loss(
|
||||
prediction_unconditional.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
positive_loss = torch.abs(
|
||||
conditional_loss.float() - prior_prediction_loss.float(),
|
||||
)
|
||||
# scale our loss by the differential scaler
|
||||
positive_loss = positive_loss * differential_scaler
|
||||
|
||||
positive_loss = positive_loss.mean([1, 2, 3])
|
||||
|
||||
polar_loss = torch.abs(
|
||||
conditional_loss.float() - unconditional_loss.float(),
|
||||
).mean([1, 2, 3])
|
||||
|
||||
|
||||
positive_loss = positive_loss.mean() + polar_loss.mean()
|
||||
|
||||
|
||||
positive_loss.backward()
|
||||
# loss = positive_loss.detach() + negative_loss.detach()
|
||||
loss = positive_loss.detach()
|
||||
|
||||
# add a grad so other backward does not fail
|
||||
loss.requires_grad_(True)
|
||||
|
||||
# restore network
|
||||
sd.network.multiplier = network_weight_list
|
||||
|
||||
return loss
|
||||
|
||||
def get_guided_loss_polarity(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
train_config: 'TrainConfig',
|
||||
scaler=None,
|
||||
**kwargs
|
||||
):
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(dtype)
|
||||
noise = noise.to(device, dtype=dtype).detach()
|
||||
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
target_pos = noise
|
||||
target_neg = noise
|
||||
|
||||
if sd.is_flow_matching:
|
||||
linear_timesteps = any([
|
||||
train_config.linear_timesteps,
|
||||
train_config.linear_timesteps2,
|
||||
train_config.timestep_type == 'linear',
|
||||
])
|
||||
|
||||
timestep_type = 'linear' if linear_timesteps else None
|
||||
if timestep_type is None:
|
||||
timestep_type = train_config.timestep_type
|
||||
|
||||
sd.noise_scheduler.set_train_timesteps(
|
||||
1000,
|
||||
device=device,
|
||||
timestep_type=timestep_type,
|
||||
latents=conditional_latents
|
||||
)
|
||||
target_pos = (noise - conditional_latents).detach()
|
||||
target_neg = (noise - unconditional_latents).detach()
|
||||
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
# double up everything to run it through all at once
|
||||
cat_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
cat_latents = torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0)
|
||||
cat_timesteps = torch.cat([timesteps, timesteps], dim=0)
|
||||
|
||||
negative_network_weights = [weight * -1.0 for weight in network_weight_list]
|
||||
positive_network_weights = [weight * 1.0 for weight in network_weight_list]
|
||||
cat_network_weight_list = positive_network_weights + negative_network_weights
|
||||
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
sd.network.is_active = True
|
||||
|
||||
sd.network.multiplier = cat_network_weight_list
|
||||
|
||||
# do our prediction with LoRA active on the scaled guidance latents
|
||||
prediction = sd.predict_noise(
|
||||
latents=cat_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=cat_timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
|
||||
pred_pos, pred_neg = torch.chunk(prediction, 2, dim=0)
|
||||
|
||||
pred_loss = torch.nn.functional.mse_loss(
|
||||
pred_pos.float(),
|
||||
target_pos.float(),
|
||||
reduction="none"
|
||||
)
|
||||
# pred_loss = pred_loss.mean([1, 2, 3])
|
||||
|
||||
pred_neg_loss = torch.nn.functional.mse_loss(
|
||||
pred_neg.float(),
|
||||
target_neg.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
loss = pred_loss + pred_neg_loss
|
||||
|
||||
loss = loss.mean([1, 2, 3])
|
||||
loss = loss.mean()
|
||||
if scaler is not None:
|
||||
scaler.scale(loss).backward()
|
||||
else:
|
||||
loss.backward()
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
loss.requires_grad_(True)
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
|
||||
def get_guided_tnt(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
prior_pred: torch.Tensor = None,
|
||||
**kwargs
|
||||
):
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(dtype)
|
||||
noise = noise.to(device, dtype=dtype).detach()
|
||||
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
# double up everything to run it through all at once
|
||||
cat_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
cat_latents = torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0)
|
||||
cat_timesteps = torch.cat([timesteps, timesteps], dim=0)
|
||||
|
||||
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
if sd.network is not None:
|
||||
cat_network_weight_list = [weight for weight in network_weight_list * 2]
|
||||
sd.network.multiplier = cat_network_weight_list
|
||||
sd.network.is_active = True
|
||||
|
||||
|
||||
prediction = sd.predict_noise(
|
||||
latents=cat_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=cat_timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
this_prediction, that_prediction = torch.chunk(prediction, 2, dim=0)
|
||||
|
||||
this_loss = torch.nn.functional.mse_loss(
|
||||
this_prediction.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
that_loss = torch.nn.functional.mse_loss(
|
||||
that_prediction.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
this_loss = this_loss.mean([1, 2, 3])
|
||||
# negative loss on that
|
||||
that_loss = -that_loss.mean([1, 2, 3])
|
||||
|
||||
with torch.no_grad():
|
||||
# match that loss with this loss so it is not a negative value and same scale
|
||||
that_loss_scaler = torch.abs(this_loss) / torch.abs(that_loss)
|
||||
|
||||
that_loss = that_loss * that_loss_scaler * 0.01
|
||||
|
||||
loss = this_loss + that_loss
|
||||
|
||||
loss = loss.mean()
|
||||
|
||||
loss.backward()
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
loss.requires_grad_(True)
|
||||
|
||||
return loss
|
||||
|
||||
def targeted_flow_guidance(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: 'PromptEmbeds',
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
unconditional_embeds: Optional[PromptEmbeds] = None,
|
||||
mask_multiplier=None,
|
||||
prior_pred=None,
|
||||
scaler=None,
|
||||
train_config=None,
|
||||
**kwargs
|
||||
):
|
||||
if not sd.is_flow_matching:
|
||||
raise ValueError("targeted_flow only works on flow matching models")
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(dtype)
|
||||
noise = noise.to(device, dtype=dtype).detach()
|
||||
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
# get a mask on the differential of the latents
|
||||
# this will be scaled from 0.0-1.0 with 1.0 being the largest differential
|
||||
abs_differential_mask = get_differential_mask(
|
||||
conditional_latents,
|
||||
unconditional_latents,
|
||||
gradient=True
|
||||
)
|
||||
|
||||
# get noisy latents for both conditional and unconditional predictions
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
# disable the lora to get a baseline prediction
|
||||
sd.network.is_active = False
|
||||
sd.unet.eval()
|
||||
|
||||
# get a baseline prediction of the model knowledge without the lora network
|
||||
# we do this with the unconditional noisy latents
|
||||
baseline_prediction = sd.predict_noise(
|
||||
latents=unconditional_noisy_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=conditional_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs
|
||||
).detach()
|
||||
|
||||
# This is our normal flowmatching target
|
||||
# target = noise - latents
|
||||
# we need to target the baseline noise but with our conditional latents
|
||||
# to do this we first have to determine the baseline_prediction noise by reversing the flowmatching target
|
||||
baseline_predicted_noise = baseline_prediction + unconditional_latents
|
||||
|
||||
# baseline_predicted_noise is now the noise prediction our model would make with a the unconditional image.
|
||||
# we use this as our new noise target to preserve the existing knowledge of the image.
|
||||
# we apply a mask to this noise to only allow the differential of the conditional latents to be learned
|
||||
baseline_predicted_noise = (1 - abs_differential_mask) * baseline_predicted_noise
|
||||
masked_noise = abs_differential_mask * noise
|
||||
target_noise = masked_noise + baseline_predicted_noise
|
||||
|
||||
# compute our new target prediction using our current knowledge noise with our conditional latents
|
||||
# this makes it so the only new information is the differential of our conditional and unconditional latents
|
||||
# forcing the network to preserve existing knowledge, but learn only our changes
|
||||
target_pred = (target_noise - conditional_latents).detach()
|
||||
|
||||
# make a prediction with the lora network active
|
||||
sd.unet.train()
|
||||
sd.network.is_active = True
|
||||
sd.network.multiplier = network_weight_list
|
||||
prediction = sd.predict_noise(
|
||||
latents=conditional_noisy_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=conditional_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs
|
||||
)
|
||||
|
||||
# target our baseline + diffirential noise target
|
||||
pred_loss = torch.nn.functional.mse_loss(
|
||||
prediction.float(),
|
||||
target_pred.float()
|
||||
)
|
||||
|
||||
return pred_loss
|
||||
|
||||
|
||||
# this processes all guidance losses based on the batch information
|
||||
def get_guidance_loss(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: 'PromptEmbeds',
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
unconditional_embeds: Optional[PromptEmbeds] = None,
|
||||
mask_multiplier=None,
|
||||
prior_pred=None,
|
||||
scaler=None,
|
||||
train_config=None,
|
||||
**kwargs
|
||||
):
|
||||
# TODO add others and process individual batch items separately
|
||||
guidance_type: GuidanceType = batch.file_items[0].dataset_config.guidance_type
|
||||
|
||||
if guidance_type == "targeted":
|
||||
assert unconditional_embeds is None, "Unconditional embeds are not supported for targeted guidance"
|
||||
return get_targeted_guidance_loss(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
**kwargs
|
||||
)
|
||||
elif guidance_type == "polarity":
|
||||
assert unconditional_embeds is None, "Unconditional embeds are not supported for polarity guidance"
|
||||
return get_guided_loss_polarity(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
scaler=scaler,
|
||||
train_config=train_config,
|
||||
**kwargs
|
||||
)
|
||||
elif guidance_type == "tnt":
|
||||
assert unconditional_embeds is None, "Unconditional embeds are not supported for polarity guidance"
|
||||
return get_guided_tnt(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
prior_pred=prior_pred,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
elif guidance_type == "targeted_polarity":
|
||||
assert unconditional_embeds is None, "Unconditional embeds are not supported for targeted polarity guidance"
|
||||
return get_targeted_polarity_loss(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
**kwargs
|
||||
)
|
||||
elif guidance_type == "direct":
|
||||
return get_direct_guidance_loss(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
unconditional_embeds=unconditional_embeds,
|
||||
mask_multiplier=mask_multiplier,
|
||||
prior_pred=prior_pred,
|
||||
**kwargs
|
||||
)
|
||||
elif guidance_type == "targeted_flow":
|
||||
return targeted_flow_guidance(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
unconditional_embeds=unconditional_embeds,
|
||||
mask_multiplier=mask_multiplier,
|
||||
prior_pred=prior_pred,
|
||||
scaler=scaler,
|
||||
train_config=train_config,
|
||||
**kwargs
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Guidance type {guidance_type} is not implemented")
|
||||
@@ -5,12 +5,14 @@ import json
|
||||
import os
|
||||
import io
|
||||
import struct
|
||||
import threading
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import AutoencoderTiny
|
||||
from PIL import Image as PILImage
|
||||
|
||||
FILE_UNKNOWN = "Sorry, don't know how to get size for this file."
|
||||
|
||||
@@ -425,43 +427,82 @@ def main(argv=None):
|
||||
|
||||
|
||||
is_window_shown = False
|
||||
display_lock = threading.Lock()
|
||||
current_img = None
|
||||
update_event = threading.Event()
|
||||
|
||||
def update_image(img, name):
|
||||
global current_img
|
||||
with display_lock:
|
||||
current_img = (img, name)
|
||||
update_event.set()
|
||||
|
||||
def display_image_in_thread():
|
||||
global is_window_shown
|
||||
|
||||
def display_img():
|
||||
global current_img
|
||||
while True:
|
||||
update_event.wait()
|
||||
with display_lock:
|
||||
if current_img:
|
||||
img, name = current_img
|
||||
cv2.imshow(name, img)
|
||||
current_img = None
|
||||
update_event.clear()
|
||||
if cv2.waitKey(1) & 0xFF == 27: # Esc key to stop
|
||||
cv2.destroyAllWindows()
|
||||
print('\nESC pressed, stopping')
|
||||
break
|
||||
|
||||
if not is_window_shown:
|
||||
is_window_shown = True
|
||||
threading.Thread(target=display_img, daemon=True).start()
|
||||
|
||||
|
||||
def show_img(img, name='AI Toolkit'):
|
||||
global is_window_shown
|
||||
|
||||
img = np.clip(img, 0, 255).astype(np.uint8)
|
||||
cv2.imshow(name, img[:, :, ::-1])
|
||||
k = cv2.waitKey(10) & 0xFF
|
||||
if k == 27: # Esc key to stop
|
||||
print('\nESC pressed, stopping')
|
||||
raise KeyboardInterrupt
|
||||
update_image(img[:, :, ::-1], name)
|
||||
if not is_window_shown:
|
||||
is_window_shown = True
|
||||
|
||||
display_image_in_thread()
|
||||
|
||||
|
||||
def show_tensors(imgs: torch.Tensor, name='AI Toolkit'):
|
||||
# if rank is 4
|
||||
if len(imgs.shape) == 4:
|
||||
img_list = torch.chunk(imgs, imgs.shape[0], dim=0)
|
||||
else:
|
||||
img_list = [imgs]
|
||||
# put images side by side
|
||||
|
||||
img = torch.cat(img_list, dim=3)
|
||||
# img is -1 to 1, convert to 0 to 255
|
||||
img = img / 2 + 0.5
|
||||
img_numpy = img.to(torch.float32).detach().cpu().numpy()
|
||||
img_numpy = np.clip(img_numpy, 0, 1) * 255
|
||||
# convert to numpy Move channel to last
|
||||
img_numpy = img_numpy.transpose(0, 2, 3, 1)
|
||||
# convert to uint8
|
||||
img_numpy = img_numpy.astype(np.uint8)
|
||||
show_img(img_numpy[0], name=name)
|
||||
|
||||
show_img(img_numpy[0], name=name)
|
||||
|
||||
def save_tensors(imgs: torch.Tensor, path='output.png'):
|
||||
if len(imgs.shape) == 5 and imgs.shape[0] == 1:
|
||||
imgs = imgs.squeeze(0)
|
||||
if len(imgs.shape) == 4:
|
||||
img_list = torch.chunk(imgs, imgs.shape[0], dim=0)
|
||||
else:
|
||||
img_list = [imgs]
|
||||
|
||||
img = torch.cat(img_list, dim=3)
|
||||
img = img / 2 + 0.5
|
||||
img_numpy = img.to(torch.float32).detach().cpu().numpy()
|
||||
img_numpy = np.clip(img_numpy, 0, 1) * 255
|
||||
img_numpy = img_numpy.transpose(0, 2, 3, 1)
|
||||
img_numpy = img_numpy.astype(np.uint8)
|
||||
# concat images to one
|
||||
img_numpy = np.concatenate(img_numpy, axis=1)
|
||||
# conver to pil
|
||||
img_pil = PILImage.fromarray(img_numpy)
|
||||
img_pil.save(path)
|
||||
|
||||
def show_latents(latents: torch.Tensor, vae: 'AutoencoderTiny', name='AI Toolkit'):
|
||||
# decode latents
|
||||
if vae.device == 'cpu':
|
||||
vae.to(latents.device)
|
||||
latents = latents / vae.config['scaling_factor']
|
||||
@@ -469,12 +510,24 @@ def show_latents(latents: torch.Tensor, vae: 'AutoencoderTiny', name='AI Toolkit
|
||||
show_tensors(imgs, name=name)
|
||||
|
||||
|
||||
|
||||
def on_exit():
|
||||
if is_window_shown:
|
||||
cv2.destroyAllWindows()
|
||||
|
||||
|
||||
def reduce_contrast(tensor, factor):
|
||||
# Ensure factor is between 0 and 1
|
||||
factor = max(0, min(factor, 1))
|
||||
|
||||
# Calculate the mean of the tensor
|
||||
mean = torch.mean(tensor)
|
||||
|
||||
# Reduce contrast
|
||||
adjusted_tensor = (tensor - mean) * factor + mean
|
||||
|
||||
# Clip values to ensure they stay within -1 to 1 range
|
||||
return torch.clamp(adjusted_tensor, -1.0, 1.0)
|
||||
|
||||
atexit.register(on_exit)
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
3039
toolkit/keymaps/stable_diffusion_vega.json
Normal file
3039
toolkit/keymaps/stable_diffusion_vega.json
Normal file
File diff suppressed because it is too large
Load Diff
BIN
toolkit/keymaps/stable_diffusion_vega_ldm_base.safetensors
Normal file
BIN
toolkit/keymaps/stable_diffusion_vega_ldm_base.safetensors
Normal file
Binary file not shown.
84
toolkit/logging.py
Normal file
84
toolkit/logging.py
Normal file
@@ -0,0 +1,84 @@
|
||||
from typing import OrderedDict, Optional
|
||||
from PIL import Image
|
||||
|
||||
from toolkit.config_modules import LoggingConfig
|
||||
|
||||
# Base logger class
|
||||
# This class does nothing, it's just a placeholder
|
||||
class EmptyLogger:
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
pass
|
||||
|
||||
# start logging the training
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
# collect the log to send
|
||||
def log(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
# send the log
|
||||
def commit(self, step: Optional[int] = None):
|
||||
pass
|
||||
|
||||
# log image
|
||||
def log_image(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
# finish logging
|
||||
def finish(self):
|
||||
pass
|
||||
|
||||
# Wandb logger class
|
||||
# This class logs the data to wandb
|
||||
class WandbLogger(EmptyLogger):
|
||||
def __init__(self, project: str, run_name: str | None, config: OrderedDict) -> None:
|
||||
self.project = project
|
||||
self.run_name = run_name
|
||||
self.config = config
|
||||
|
||||
def start(self):
|
||||
try:
|
||||
import wandb
|
||||
except ImportError:
|
||||
raise ImportError("Failed to import wandb. Please install wandb by running `pip install wandb`")
|
||||
|
||||
# send the whole config to wandb
|
||||
run = wandb.init(project=self.project, name=self.run_name, config=self.config)
|
||||
self.run = run
|
||||
self._log = wandb.log # log function
|
||||
self._image = wandb.Image # image object
|
||||
|
||||
def log(self, *args, **kwargs):
|
||||
# when commit is False, wandb increments the step,
|
||||
# but we don't want that to happen, so we set commit=False
|
||||
self._log(*args, **kwargs, commit=False)
|
||||
|
||||
def commit(self, step: Optional[int] = None):
|
||||
# after overall one step is done, we commit the log
|
||||
# by log empty object with commit=True
|
||||
self._log({}, step=step, commit=True)
|
||||
|
||||
def log_image(
|
||||
self,
|
||||
image: Image,
|
||||
id, # sample index
|
||||
caption: str | None = None, # positive prompt
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
# create a wandb image object and log it
|
||||
image = self._image(image, caption=caption, *args, **kwargs)
|
||||
self._log({f"sample_{id}": image}, commit=False)
|
||||
|
||||
def finish(self):
|
||||
self.run.finish()
|
||||
|
||||
# create logger based on the logging config
|
||||
def create_logger(logging_config: LoggingConfig, all_config: OrderedDict):
|
||||
if logging_config.use_wandb:
|
||||
project_name = logging_config.project_name
|
||||
run_name = logging_config.run_name
|
||||
return WandbLogger(project=project_name, run_name=run_name, config=all_config)
|
||||
else:
|
||||
return EmptyLogger()
|
||||
@@ -1,11 +1,15 @@
|
||||
import copy
|
||||
import json
|
||||
import math
|
||||
import weakref
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from typing import List, Optional, Dict, Type, Union
|
||||
import torch
|
||||
from diffusers import UNet2DConditionModel, PixArtTransformer2DModel, AuraFlowTransformer2DModel
|
||||
from transformers import CLIPTextModel
|
||||
from toolkit.models.lokr import LokrModule
|
||||
|
||||
from .config_modules import NetworkConfig
|
||||
from .lorm import count_parameters
|
||||
@@ -15,21 +19,28 @@ from .paths import SD_SCRIPTS_ROOT
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
from networks.lora import LoRANetwork, get_block_index
|
||||
from toolkit.models.DoRA import DoRAModule
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers|attentions)_(\d+)_")
|
||||
|
||||
|
||||
# diffusers specific stuff
|
||||
LINEAR_MODULES = [
|
||||
'Linear',
|
||||
'LoRACompatibleLinear'
|
||||
'LoRACompatibleLinear',
|
||||
'QLinear',
|
||||
# 'GroupNorm',
|
||||
]
|
||||
CONV_MODULES = [
|
||||
'Conv2d',
|
||||
'LoRACompatibleConv'
|
||||
'LoRACompatibleConv',
|
||||
'QConv2d',
|
||||
]
|
||||
|
||||
class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
|
||||
@@ -51,11 +62,13 @@ class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
|
||||
use_bias: bool = False,
|
||||
**kwargs
|
||||
):
|
||||
self.can_merge_in = True
|
||||
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
||||
ToolkitModuleMixin.__init__(self, network=network)
|
||||
torch.nn.Module.__init__(self)
|
||||
self.lora_name = lora_name
|
||||
self.scalar = torch.tensor(1.0)
|
||||
self.orig_module_ref = weakref.ref(org_module)
|
||||
self.scalar = torch.tensor(1.0, device=org_module.weight.device)
|
||||
# check if parent has bias. if not force use_bias to False
|
||||
if org_module.bias is None:
|
||||
use_bias = False
|
||||
@@ -111,10 +124,14 @@ class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
|
||||
class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
NUM_OF_BLOCKS = 12 # フルモデル相当でのup,downの層の数
|
||||
|
||||
UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel"]
|
||||
UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["ResnetBlock2D", "Downsample2D", "Upsample2D"]
|
||||
# UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel"]
|
||||
# UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel", "ResnetBlock2D"]
|
||||
UNET_TARGET_REPLACE_MODULE = ["UNet2DConditionModel"]
|
||||
# UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["ResnetBlock2D", "Downsample2D", "Upsample2D"]
|
||||
UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["UNet2DConditionModel"]
|
||||
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPMLP"]
|
||||
LORA_PREFIX_UNET = "lora_unet"
|
||||
PEFT_PREFIX_UNET = "unet"
|
||||
LORA_PREFIX_TEXT_ENCODER = "lora_te"
|
||||
|
||||
# SDXL: must starts with LORA_PREFIX_TEXT_ENCODER
|
||||
@@ -147,12 +164,26 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
train_unet: Optional[bool] = True,
|
||||
is_sdxl=False,
|
||||
is_v2=False,
|
||||
is_v3=False,
|
||||
is_pixart: bool = False,
|
||||
is_auraflow: bool = False,
|
||||
is_flux: bool = False,
|
||||
is_lumina2: bool = False,
|
||||
use_bias: bool = False,
|
||||
is_lorm: bool = False,
|
||||
ignore_if_contains = None,
|
||||
only_if_contains = None,
|
||||
parameter_threshold: float = 0.0,
|
||||
attn_only: bool = False,
|
||||
target_lin_modules=LoRANetwork.UNET_TARGET_REPLACE_MODULE,
|
||||
target_conv_modules=LoRANetwork.UNET_TARGET_REPLACE_MODULE_CONV2D_3X3,
|
||||
network_type: str = "lora",
|
||||
full_train_in_out: bool = False,
|
||||
transformer_only: bool = False,
|
||||
peft_format: bool = False,
|
||||
is_assistant_adapter: bool = False,
|
||||
is_transformer: bool = False,
|
||||
base_model: 'StableDiffusion' = None,
|
||||
**kwargs
|
||||
) -> None:
|
||||
"""
|
||||
@@ -177,6 +208,13 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
if ignore_if_contains is None:
|
||||
ignore_if_contains = []
|
||||
self.ignore_if_contains = ignore_if_contains
|
||||
self.transformer_only = transformer_only
|
||||
self.base_model_ref = None
|
||||
if base_model is not None:
|
||||
self.base_model_ref = weakref.ref(base_model)
|
||||
|
||||
self.only_if_contains: Union[List, None] = only_if_contains
|
||||
|
||||
self.lora_dim = lora_dim
|
||||
self.alpha = alpha
|
||||
self.conv_lora_dim = conv_lora_dim
|
||||
@@ -192,6 +230,39 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
self.multiplier = multiplier
|
||||
self.is_sdxl = is_sdxl
|
||||
self.is_v2 = is_v2
|
||||
self.is_v3 = is_v3
|
||||
self.is_pixart = is_pixart
|
||||
self.is_auraflow = is_auraflow
|
||||
self.is_flux = is_flux
|
||||
self.is_lumina2 = is_lumina2
|
||||
self.network_type = network_type
|
||||
self.is_assistant_adapter = is_assistant_adapter
|
||||
if self.network_type.lower() == "dora":
|
||||
self.module_class = DoRAModule
|
||||
module_class = DoRAModule
|
||||
elif self.network_type.lower() == "lokr":
|
||||
self.module_class = LokrModule
|
||||
module_class = LokrModule
|
||||
self.network_config: NetworkConfig = kwargs.get("network_config", None)
|
||||
|
||||
self.peft_format = peft_format
|
||||
self.is_transformer = is_transformer
|
||||
|
||||
|
||||
# always do peft for flux only for now
|
||||
if self.is_flux or self.is_v3 or self.is_lumina2 or is_transformer:
|
||||
# don't do peft format for lokr
|
||||
if self.network_type.lower() != "lokr":
|
||||
self.peft_format = True
|
||||
|
||||
if self.peft_format:
|
||||
# no alpha for peft
|
||||
self.alpha = self.lora_dim
|
||||
alpha = self.alpha
|
||||
self.conv_alpha = self.conv_lora_dim
|
||||
conv_alpha = self.conv_alpha
|
||||
|
||||
self.full_train_in_out = full_train_in_out
|
||||
|
||||
if modules_dim is not None:
|
||||
print(f"create LoRA network from weights")
|
||||
@@ -219,8 +290,16 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
root_module: torch.nn.Module,
|
||||
target_replace_modules: List[torch.nn.Module],
|
||||
) -> List[LoRAModule]:
|
||||
unet_prefix = self.LORA_PREFIX_UNET
|
||||
if self.peft_format:
|
||||
unet_prefix = self.PEFT_PREFIX_UNET
|
||||
if is_pixart or is_v3 or is_auraflow or is_flux or is_lumina2 or self.is_transformer:
|
||||
unet_prefix = f"lora_transformer"
|
||||
if self.peft_format:
|
||||
unet_prefix = "transformer"
|
||||
|
||||
prefix = (
|
||||
self.LORA_PREFIX_UNET
|
||||
unet_prefix
|
||||
if is_unet
|
||||
else (
|
||||
self.LORA_PREFIX_TEXT_ENCODER
|
||||
@@ -230,6 +309,8 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
)
|
||||
loras = []
|
||||
skipped = []
|
||||
attached_modules = []
|
||||
lora_shape_dict = {}
|
||||
for name, module in root_module.named_modules():
|
||||
if module.__class__.__name__ in target_replace_modules:
|
||||
for child_name, child_module in module.named_modules():
|
||||
@@ -237,17 +318,55 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
is_conv2d = child_module.__class__.__name__ in CONV_MODULES
|
||||
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
|
||||
|
||||
|
||||
lora_name = [prefix, name, child_name]
|
||||
# filter out blank
|
||||
lora_name = [x for x in lora_name if x and x != ""]
|
||||
lora_name = ".".join(lora_name)
|
||||
# if it doesnt have a name, it wil have two dots
|
||||
lora_name.replace("..", ".")
|
||||
clean_name = lora_name
|
||||
if self.peft_format:
|
||||
# we replace this on saving
|
||||
lora_name = lora_name.replace(".", "$$")
|
||||
else:
|
||||
lora_name = lora_name.replace(".", "_")
|
||||
|
||||
skip = False
|
||||
if any([word in child_name for word in self.ignore_if_contains]):
|
||||
if any([word in clean_name for word in self.ignore_if_contains]):
|
||||
skip = True
|
||||
|
||||
# see if it is over threshold
|
||||
if count_parameters(child_module) < parameter_threshold:
|
||||
skip = True
|
||||
|
||||
if self.transformer_only and self.is_pixart and is_unet:
|
||||
if "transformer_blocks" not in lora_name:
|
||||
skip = True
|
||||
if self.transformer_only and self.is_flux and is_unet:
|
||||
if "transformer_blocks" not in lora_name:
|
||||
skip = True
|
||||
if self.transformer_only and self.is_lumina2 and is_unet:
|
||||
if "layers$$" not in lora_name and "noise_refiner$$" not in lora_name and "context_refiner$$" not in lora_name:
|
||||
skip = True
|
||||
if self.transformer_only and self.is_v3 and is_unet:
|
||||
if "transformer_blocks" not in lora_name:
|
||||
skip = True
|
||||
|
||||
# handle custom models
|
||||
if self.transformer_only and is_unet and hasattr(root_module, 'transformer_blocks'):
|
||||
if "transformer_blocks" not in lora_name:
|
||||
skip = True
|
||||
|
||||
if self.transformer_only and is_unet and hasattr(root_module, 'blocks'):
|
||||
if "blocks" not in lora_name:
|
||||
skip = True
|
||||
|
||||
if (is_linear or is_conv2d) and not skip:
|
||||
lora_name = prefix + "." + name + "." + child_name
|
||||
lora_name = lora_name.replace(".", "_")
|
||||
|
||||
if self.only_if_contains is not None:
|
||||
if not any([word in clean_name for word in self.only_if_contains]) and not any([word in lora_name for word in self.only_if_contains]):
|
||||
continue
|
||||
|
||||
dim = None
|
||||
alpha = None
|
||||
@@ -281,6 +400,11 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
self.conv_lora_dim is not None or conv_block_dims is not None):
|
||||
skipped.append(lora_name)
|
||||
continue
|
||||
|
||||
module_kwargs = {}
|
||||
|
||||
if self.network_type.lower() == "lokr":
|
||||
module_kwargs["factor"] = self.network_config.lokr_factor
|
||||
|
||||
lora = module_class(
|
||||
lora_name,
|
||||
@@ -294,8 +418,16 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
network=self,
|
||||
parent=module,
|
||||
use_bias=use_bias,
|
||||
**module_kwargs
|
||||
)
|
||||
loras.append(lora)
|
||||
if self.network_type.lower() == "lokr":
|
||||
try:
|
||||
lora_shape_dict[lora_name] = [list(lora.lokr_w1.weight.shape), list(lora.lokr_w2.weight.shape)]
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
lora_shape_dict[lora_name] = [list(lora.lora_down.weight.shape), list(lora.lora_up.weight.shape)]
|
||||
return loras, skipped
|
||||
|
||||
text_encoders = text_encoder if type(text_encoder) == list else [text_encoder]
|
||||
@@ -317,8 +449,12 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
index = None
|
||||
print(f"create LoRA for Text Encoder:")
|
||||
|
||||
text_encoder_loras, skipped = create_modules(False, index, text_encoder,
|
||||
LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE)
|
||||
replace_modules = LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE
|
||||
|
||||
if self.is_pixart:
|
||||
replace_modules = ["T5EncoderModel"]
|
||||
|
||||
text_encoder_loras, skipped = create_modules(False, index, text_encoder, replace_modules)
|
||||
self.text_encoder_loras.extend(text_encoder_loras)
|
||||
skipped_te += skipped
|
||||
print(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
|
||||
@@ -328,6 +464,21 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
if modules_dim is not None or self.conv_lora_dim is not None or conv_block_dims is not None:
|
||||
target_modules += target_conv_modules
|
||||
|
||||
if is_v3:
|
||||
target_modules = ["SD3Transformer2DModel"]
|
||||
|
||||
if is_pixart:
|
||||
target_modules = ["PixArtTransformer2DModel"]
|
||||
|
||||
if is_auraflow:
|
||||
target_modules = ["AuraFlowTransformer2DModel"]
|
||||
|
||||
if is_flux:
|
||||
target_modules = ["FluxTransformer2DModel"]
|
||||
|
||||
if is_lumina2:
|
||||
target_modules = ["Lumina2Transformer2DModel"]
|
||||
|
||||
if train_unet:
|
||||
self.unet_loras, skipped_un = create_modules(True, None, unet, target_modules)
|
||||
else:
|
||||
@@ -353,3 +504,49 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
for lora in self.text_encoder_loras + self.unet_loras:
|
||||
assert lora.lora_name not in names, f"duplicated lora name: {lora.lora_name}"
|
||||
names.add(lora.lora_name)
|
||||
|
||||
if self.full_train_in_out:
|
||||
print("full train in out")
|
||||
# we are going to retrain the main in out layers for VAE change usually
|
||||
if self.is_pixart:
|
||||
transformer: PixArtTransformer2DModel = unet
|
||||
self.transformer_pos_embed = copy.deepcopy(transformer.pos_embed)
|
||||
self.transformer_proj_out = copy.deepcopy(transformer.proj_out)
|
||||
|
||||
transformer.pos_embed = self.transformer_pos_embed
|
||||
transformer.proj_out = self.transformer_proj_out
|
||||
|
||||
elif self.is_auraflow:
|
||||
transformer: AuraFlowTransformer2DModel = unet
|
||||
self.transformer_pos_embed = copy.deepcopy(transformer.pos_embed)
|
||||
self.transformer_proj_out = copy.deepcopy(transformer.proj_out)
|
||||
|
||||
transformer.pos_embed = self.transformer_pos_embed
|
||||
transformer.proj_out = self.transformer_proj_out
|
||||
|
||||
else:
|
||||
unet: UNet2DConditionModel = unet
|
||||
unet_conv_in: torch.nn.Conv2d = unet.conv_in
|
||||
unet_conv_out: torch.nn.Conv2d = unet.conv_out
|
||||
|
||||
# clone these and replace their forwards with ours
|
||||
self.unet_conv_in = copy.deepcopy(unet_conv_in)
|
||||
self.unet_conv_out = copy.deepcopy(unet_conv_out)
|
||||
unet.conv_in = self.unet_conv_in
|
||||
unet.conv_out = self.unet_conv_out
|
||||
|
||||
def prepare_optimizer_params(self, text_encoder_lr, unet_lr, default_lr):
|
||||
# call Lora prepare_optimizer_params
|
||||
all_params = super().prepare_optimizer_params(text_encoder_lr, unet_lr, default_lr)
|
||||
|
||||
if self.full_train_in_out:
|
||||
if self.is_pixart or self.is_auraflow or self.is_flux:
|
||||
all_params.append({"lr": unet_lr, "params": list(self.transformer_pos_embed.parameters())})
|
||||
all_params.append({"lr": unet_lr, "params": list(self.transformer_proj_out.parameters())})
|
||||
else:
|
||||
all_params.append({"lr": unet_lr, "params": list(self.unet_conv_in.parameters())})
|
||||
all_params.append({"lr": unet_lr, "params": list(self.unet_conv_out.parameters())})
|
||||
|
||||
return all_params
|
||||
|
||||
|
||||
|
||||
@@ -354,7 +354,8 @@ def convert_diffusers_unet_to_lorm(
|
||||
elif child_module.__class__.__name__ in LINEAR_MODULES:
|
||||
if count_parameters(child_module) > parameter_threshold:
|
||||
|
||||
dtype = child_module.weight.dtype
|
||||
# dtype = child_module.weight.dtype
|
||||
dtype = torch.float32
|
||||
# extract and convert
|
||||
down_weight, up_weight, lora_dim, diff = extract_linear(
|
||||
weight=child_module.weight.clone().detach().float(),
|
||||
|
||||
@@ -23,6 +23,8 @@ def get_meta_for_safetensors(meta: OrderedDict, name=None, add_software_info=Tru
|
||||
# if not float, int, bool, or str, convert to json string
|
||||
if not isinstance(value, str):
|
||||
save_meta[key] = json.dumps(value)
|
||||
# add the pt format
|
||||
save_meta["format"] = "pt"
|
||||
return save_meta
|
||||
|
||||
|
||||
|
||||
146
toolkit/models/DoRA.py
Normal file
146
toolkit/models/DoRA.py
Normal file
@@ -0,0 +1,146 @@
|
||||
#based off https://github.com/catid/dora/blob/main/dora.py
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from typing import TYPE_CHECKING, Union, List
|
||||
|
||||
from optimum.quanto import QBytesTensor, QTensor
|
||||
|
||||
from toolkit.network_mixins import ToolkitModuleMixin, ExtractableModuleMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
|
||||
# diffusers specific stuff
|
||||
LINEAR_MODULES = [
|
||||
'Linear',
|
||||
'LoRACompatibleLinear'
|
||||
# 'GroupNorm',
|
||||
]
|
||||
CONV_MODULES = [
|
||||
'Conv2d',
|
||||
'LoRACompatibleConv'
|
||||
]
|
||||
|
||||
def transpose(weight, fan_in_fan_out):
|
||||
if not fan_in_fan_out:
|
||||
return weight
|
||||
|
||||
if isinstance(weight, torch.nn.Parameter):
|
||||
return torch.nn.Parameter(weight.T)
|
||||
return weight.T
|
||||
|
||||
class DoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
|
||||
# def __init__(self, d_in, d_out, rank=4, weight=None, bias=None):
|
||||
def __init__(
|
||||
self,
|
||||
lora_name,
|
||||
org_module: torch.nn.Module,
|
||||
multiplier=1.0,
|
||||
lora_dim=4,
|
||||
alpha=1,
|
||||
dropout=None,
|
||||
rank_dropout=None,
|
||||
module_dropout=None,
|
||||
network: 'LoRASpecialNetwork' = None,
|
||||
use_bias: bool = False,
|
||||
**kwargs
|
||||
):
|
||||
self.can_merge_in = False
|
||||
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
||||
ToolkitModuleMixin.__init__(self, network=network)
|
||||
torch.nn.Module.__init__(self)
|
||||
self.lora_name = lora_name
|
||||
self.scalar = torch.tensor(1.0)
|
||||
|
||||
self.lora_dim = lora_dim
|
||||
|
||||
if org_module.__class__.__name__ in CONV_MODULES:
|
||||
raise NotImplementedError("Convolutional layers are not supported yet")
|
||||
|
||||
if type(alpha) == torch.Tensor:
|
||||
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
|
||||
alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
|
||||
self.scale = alpha / self.lora_dim
|
||||
# self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える eng: treat as constant
|
||||
|
||||
self.multiplier: Union[float, List[float]] = multiplier
|
||||
# wrap the original module so it doesn't get weights updated
|
||||
self.org_module = [org_module]
|
||||
self.dropout = dropout
|
||||
self.rank_dropout = rank_dropout
|
||||
self.module_dropout = module_dropout
|
||||
self.is_checkpointing = False
|
||||
|
||||
d_out = org_module.out_features
|
||||
d_in = org_module.in_features
|
||||
|
||||
std_dev = 1 / torch.sqrt(torch.tensor(self.lora_dim).float())
|
||||
# self.lora_up = nn.Parameter(torch.randn(d_out, self.lora_dim) * std_dev) # lora_A
|
||||
# self.lora_down = nn.Parameter(torch.zeros(self.lora_dim, d_in)) # lora_B
|
||||
self.lora_up = nn.Linear(self.lora_dim, d_out, bias=False) # lora_B
|
||||
# self.lora_up.weight.data = torch.randn_like(self.lora_up.weight.data) * std_dev
|
||||
self.lora_up.weight.data = torch.zeros_like(self.lora_up.weight.data)
|
||||
# self.lora_A[adapter_name] = nn.Linear(self.in_features, r, bias=False)
|
||||
# self.lora_B[adapter_name] = nn.Linear(r, self.out_features, bias=False)
|
||||
self.lora_down = nn.Linear(d_in, self.lora_dim, bias=False) # lora_A
|
||||
# self.lora_down.weight.data = torch.zeros_like(self.lora_down.weight.data)
|
||||
self.lora_down.weight.data = torch.randn_like(self.lora_down.weight.data) * std_dev
|
||||
|
||||
# m = Magnitude column-wise across output dimension
|
||||
weight = self.get_orig_weight()
|
||||
weight = weight.to(self.lora_up.weight.device, dtype=self.lora_up.weight.dtype)
|
||||
lora_weight = self.lora_up.weight @ self.lora_down.weight
|
||||
weight_norm = self._get_weight_norm(weight, lora_weight)
|
||||
self.magnitude = nn.Parameter(weight_norm.detach().clone(), requires_grad=True)
|
||||
|
||||
def apply_to(self):
|
||||
self.org_forward = self.org_module[0].forward
|
||||
self.org_module[0].forward = self.forward
|
||||
# del self.org_module
|
||||
|
||||
def get_orig_weight(self):
|
||||
weight = self.org_module[0].weight
|
||||
if isinstance(weight, QTensor) or isinstance(weight, QBytesTensor):
|
||||
return weight.dequantize().data.detach()
|
||||
else:
|
||||
return weight.data.detach()
|
||||
|
||||
def get_orig_bias(self):
|
||||
if hasattr(self.org_module[0], 'bias') and self.org_module[0].bias is not None:
|
||||
return self.org_module[0].bias.data.detach()
|
||||
return None
|
||||
|
||||
# def dora_forward(self, x, *args, **kwargs):
|
||||
# lora = torch.matmul(self.lora_A, self.lora_B)
|
||||
# adapted = self.get_orig_weight() + lora
|
||||
# column_norm = adapted.norm(p=2, dim=0, keepdim=True)
|
||||
# norm_adapted = adapted / column_norm
|
||||
# calc_weights = self.magnitude * norm_adapted
|
||||
# return F.linear(x, calc_weights, self.get_orig_bias())
|
||||
|
||||
def _get_weight_norm(self, weight, scaled_lora_weight) -> torch.Tensor:
|
||||
# calculate L2 norm of weight matrix, column-wise
|
||||
weight = weight + scaled_lora_weight.to(weight.device)
|
||||
weight_norm = torch.linalg.norm(weight, dim=1)
|
||||
return weight_norm
|
||||
|
||||
def apply_dora(self, x, scaled_lora_weight):
|
||||
# ref https://github.com/huggingface/peft/blob/1e6d1d73a0850223b0916052fd8d2382a90eae5a/src/peft/tuners/lora/layer.py#L192
|
||||
# lora weight is already scaled
|
||||
|
||||
# magnitude = self.lora_magnitude_vector[active_adapter]
|
||||
weight = self.get_orig_weight()
|
||||
weight = weight.to(scaled_lora_weight.device, dtype=scaled_lora_weight.dtype)
|
||||
weight_norm = self._get_weight_norm(weight, scaled_lora_weight)
|
||||
# see section 4.3 of DoRA (https://arxiv.org/abs/2402.09353)
|
||||
# "[...] we suggest treating ||V +∆V ||_c in
|
||||
# Eq. (5) as a constant, thereby detaching it from the gradient
|
||||
# graph. This means that while ||V + ∆V ||_c dynamically
|
||||
# reflects the updates of ∆V , it won’t receive any gradient
|
||||
# during backpropagation"
|
||||
weight_norm = weight_norm.detach()
|
||||
dora_weight = transpose(weight + scaled_lora_weight, False)
|
||||
return (self.magnitude / weight_norm - 1).view(1, -1) * F.linear(x.to(dora_weight.dtype), dora_weight)
|
||||
267
toolkit/models/LoRAFormer.py
Normal file
267
toolkit/models/LoRAFormer.py
Normal file
@@ -0,0 +1,267 @@
|
||||
import math
|
||||
import weakref
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import TYPE_CHECKING, List, Dict, Any
|
||||
from toolkit.models.clip_fusion import ZipperBlock
|
||||
from toolkit.models.zipper_resampler import ZipperModule, ZipperResampler
|
||||
import sys
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from ipadapter.ip_adapter.resampler import Resampler
|
||||
from collections import OrderedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRAModule
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
def __init__(self, d_model, nhead, dim_feedforward):
|
||||
super().__init__()
|
||||
self.self_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
|
||||
self.cross_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
|
||||
self.feed_forward = nn.Sequential(
|
||||
nn.Linear(d_model, dim_feedforward),
|
||||
nn.ReLU(),
|
||||
nn.Linear(dim_feedforward, d_model)
|
||||
)
|
||||
self.norm1 = nn.LayerNorm(d_model)
|
||||
self.norm2 = nn.LayerNorm(d_model)
|
||||
self.norm3 = nn.LayerNorm(d_model)
|
||||
|
||||
def forward(self, x, cross_attn_input):
|
||||
# Self-attention
|
||||
attn_output, _ = self.self_attn(x, x, x)
|
||||
x = self.norm1(x + attn_output)
|
||||
|
||||
# Cross-attention
|
||||
cross_attn_output, _ = self.cross_attn(x, cross_attn_input, cross_attn_input)
|
||||
x = self.norm2(x + cross_attn_output)
|
||||
|
||||
# Feed-forward
|
||||
ff_output = self.feed_forward(x)
|
||||
x = self.norm3(x + ff_output)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class InstantLoRAMidModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
index: int,
|
||||
lora_module: 'LoRAModule',
|
||||
instant_lora_module: 'InstantLoRAModule',
|
||||
up_shape: list = None,
|
||||
down_shape: list = None,
|
||||
):
|
||||
super(InstantLoRAMidModule, self).__init__()
|
||||
self.up_shape = up_shape
|
||||
self.down_shape = down_shape
|
||||
self.index = index
|
||||
self.lora_module_ref = weakref.ref(lora_module)
|
||||
self.instant_lora_module_ref = weakref.ref(instant_lora_module)
|
||||
|
||||
self.embed = None
|
||||
|
||||
def down_forward(self, x, *args, **kwargs):
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
down_size = math.prod(self.down_shape)
|
||||
down_weight = self.embed[:, :down_size]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if down_weight.shape[0] * 2 == batch_size:
|
||||
down_weight = torch.cat([down_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(down_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
weight_chunk = weight_chunk.view(self.down_shape)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
padding = 0
|
||||
if weight_chunk.shape[-1] == 3:
|
||||
padding = 1
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
def up_forward(self, x, *args, **kwargs):
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
up_size = math.prod(self.up_shape)
|
||||
up_weight = self.embed[:, -up_size:]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if up_weight.shape[0] * 2 == batch_size:
|
||||
up_weight = torch.cat([up_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(up_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
weight_chunk = weight_chunk.view(self.up_shape)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
padding = 0
|
||||
if weight_chunk.shape[-1] == 3:
|
||||
padding = 1
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
# Initialize the network
|
||||
# num_blocks = 8
|
||||
# d_model = 1024 # Adjust as needed
|
||||
# nhead = 16 # Adjust as needed
|
||||
# dim_feedforward = 4096 # Adjust as needed
|
||||
# latent_dim = 1695744
|
||||
|
||||
class LoRAFormer(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
num_blocks,
|
||||
d_model=1024,
|
||||
nhead=16,
|
||||
dim_feedforward=4096,
|
||||
sd: 'StableDiffusion'=None,
|
||||
):
|
||||
super(LoRAFormer, self).__init__()
|
||||
# self.linear = torch.nn.Linear(2, 1)
|
||||
self.sd_ref = weakref.ref(sd)
|
||||
self.dim = sd.network.lora_dim
|
||||
|
||||
# stores the projection vector. Grabbed by modules
|
||||
self.img_embeds: List[torch.Tensor] = None
|
||||
|
||||
# disable merging in. It is slower on inference
|
||||
self.sd_ref().network.can_merge_in = False
|
||||
|
||||
self.ilora_modules = torch.nn.ModuleList()
|
||||
|
||||
lora_modules = self.sd_ref().network.get_all_modules()
|
||||
|
||||
output_size = 0
|
||||
|
||||
self.embed_lengths = []
|
||||
self.weight_mapping = []
|
||||
|
||||
for idx, lora_module in enumerate(lora_modules):
|
||||
module_dict = lora_module.state_dict()
|
||||
down_shape = list(module_dict['lora_down.weight'].shape)
|
||||
up_shape = list(module_dict['lora_up.weight'].shape)
|
||||
|
||||
self.weight_mapping.append([lora_module.lora_name, [down_shape, up_shape]])
|
||||
|
||||
module_size = math.prod(down_shape) + math.prod(up_shape)
|
||||
output_size += module_size
|
||||
self.embed_lengths.append(module_size)
|
||||
|
||||
|
||||
# add a new mid module that will take the original forward and add a vector to it
|
||||
# this will be used to add the vector to the original forward
|
||||
instant_module = InstantLoRAMidModule(
|
||||
idx,
|
||||
lora_module,
|
||||
self,
|
||||
up_shape=up_shape,
|
||||
down_shape=down_shape
|
||||
)
|
||||
|
||||
self.ilora_modules.append(instant_module)
|
||||
|
||||
# replace the LoRA forwards
|
||||
lora_module.lora_down.forward = instant_module.down_forward
|
||||
lora_module.lora_up.forward = instant_module.up_forward
|
||||
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
self.latent = nn.Parameter(torch.randn(1, output_size))
|
||||
self.latent_proj = nn.Linear(output_size, d_model)
|
||||
self.blocks = nn.ModuleList([
|
||||
TransformerBlock(d_model, nhead, dim_feedforward)
|
||||
for _ in range(num_blocks)
|
||||
])
|
||||
self.final_proj = nn.Linear(d_model, output_size)
|
||||
|
||||
self.migrate_weight_mapping()
|
||||
|
||||
def migrate_weight_mapping(self):
|
||||
return
|
||||
# # changes the names of the modules to common ones
|
||||
# keymap = self.sd_ref().network.get_keymap()
|
||||
# save_keymap = {}
|
||||
# if keymap is not None:
|
||||
# for ldm_key, diffusers_key in keymap.items():
|
||||
# # invert them
|
||||
# save_keymap[diffusers_key] = ldm_key
|
||||
#
|
||||
# new_keymap = {}
|
||||
# for key, value in self.weight_mapping:
|
||||
# if key in save_keymap:
|
||||
# new_keymap[save_keymap[key]] = value
|
||||
# else:
|
||||
# print(f"Key {key} not found in keymap")
|
||||
# new_keymap[key] = value
|
||||
# self.weight_mapping = new_keymap
|
||||
# else:
|
||||
# print("No keymap found. Using default names")
|
||||
# return
|
||||
|
||||
|
||||
def forward(self, img_embeds):
|
||||
# expand token rank if only rank 2
|
||||
if len(img_embeds.shape) == 2:
|
||||
img_embeds = img_embeds.unsqueeze(1)
|
||||
|
||||
# resample the image embeddings
|
||||
img_embeds = self.resampler(img_embeds)
|
||||
img_embeds = self.proj_module(img_embeds)
|
||||
if len(img_embeds.shape) == 3:
|
||||
# merge the heads
|
||||
img_embeds = img_embeds.mean(dim=1)
|
||||
|
||||
self.img_embeds = []
|
||||
# get all the slices
|
||||
start = 0
|
||||
for length in self.embed_lengths:
|
||||
self.img_embeds.append(img_embeds[:, start:start+length])
|
||||
start += length
|
||||
|
||||
|
||||
def get_additional_save_metadata(self) -> Dict[str, Any]:
|
||||
# save the weight mapping
|
||||
return {
|
||||
"weight_mapping": self.weight_mapping,
|
||||
"num_heads": self.num_heads,
|
||||
"vision_hidden_size": self.vision_hidden_size,
|
||||
"head_dim": self.head_dim,
|
||||
"vision_tokens": self.vision_tokens,
|
||||
"output_size": self.output_size,
|
||||
}
|
||||
|
||||
127
toolkit/models/auraflow.py
Normal file
127
toolkit/models/auraflow.py
Normal file
@@ -0,0 +1,127 @@
|
||||
import math
|
||||
from functools import partial
|
||||
|
||||
from torch import nn
|
||||
import torch
|
||||
|
||||
|
||||
class AuraFlowPatchEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
height=224,
|
||||
width=224,
|
||||
patch_size=16,
|
||||
in_channels=3,
|
||||
embed_dim=768,
|
||||
pos_embed_max_size=None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.num_patches = (height // patch_size) * (width // patch_size)
|
||||
self.pos_embed_max_size = pos_embed_max_size
|
||||
|
||||
self.proj = nn.Linear(patch_size * patch_size * in_channels, embed_dim)
|
||||
self.pos_embed = nn.Parameter(torch.randn(1, pos_embed_max_size, embed_dim) * 0.1)
|
||||
|
||||
self.patch_size = patch_size
|
||||
self.height, self.width = height // patch_size, width // patch_size
|
||||
self.base_size = height // patch_size
|
||||
|
||||
def forward(self, latent):
|
||||
batch_size, num_channels, height, width = latent.size()
|
||||
latent = latent.view(
|
||||
batch_size,
|
||||
num_channels,
|
||||
height // self.patch_size,
|
||||
self.patch_size,
|
||||
width // self.patch_size,
|
||||
self.patch_size,
|
||||
)
|
||||
latent = latent.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2)
|
||||
latent = self.proj(latent)
|
||||
try:
|
||||
return latent + self.pos_embed
|
||||
except RuntimeError:
|
||||
raise RuntimeError(
|
||||
f"Positional embeddings are too small for the number of patches. "
|
||||
f"Please increase `pos_embed_max_size` to at least {self.num_patches}."
|
||||
)
|
||||
|
||||
|
||||
# comfy
|
||||
# def apply_pos_embeds(self, x, h, w):
|
||||
# h = (h + 1) // self.patch_size
|
||||
# w = (w + 1) // self.patch_size
|
||||
# max_dim = max(h, w)
|
||||
#
|
||||
# cur_dim = self.h_max
|
||||
# pos_encoding = self.positional_encoding.reshape(1, cur_dim, cur_dim, -1).to(device=x.device, dtype=x.dtype)
|
||||
#
|
||||
# if max_dim > cur_dim:
|
||||
# pos_encoding = F.interpolate(pos_encoding.movedim(-1, 1), (max_dim, max_dim), mode="bilinear").movedim(1,
|
||||
# -1)
|
||||
# cur_dim = max_dim
|
||||
#
|
||||
# from_h = (cur_dim - h) // 2
|
||||
# from_w = (cur_dim - w) // 2
|
||||
# pos_encoding = pos_encoding[:, from_h:from_h + h, from_w:from_w + w]
|
||||
# return x + pos_encoding.reshape(1, -1, self.positional_encoding.shape[-1])
|
||||
|
||||
# def patchify(self, x):
|
||||
# B, C, H, W = x.size()
|
||||
# pad_h = (self.patch_size - H % self.patch_size) % self.patch_size
|
||||
# pad_w = (self.patch_size - W % self.patch_size) % self.patch_size
|
||||
#
|
||||
# x = torch.nn.functional.pad(x, (0, pad_w, 0, pad_h), mode='reflect')
|
||||
# x = x.view(
|
||||
# B,
|
||||
# C,
|
||||
# (H + 1) // self.patch_size,
|
||||
# self.patch_size,
|
||||
# (W + 1) // self.patch_size,
|
||||
# self.patch_size,
|
||||
# )
|
||||
# x = x.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2)
|
||||
# return x
|
||||
|
||||
def patch_auraflow_pos_embed(pos_embed):
|
||||
# we need to hijack the forward and replace with a custom one. Self is the model
|
||||
def new_forward(self, latent):
|
||||
batch_size, num_channels, height, width = latent.size()
|
||||
|
||||
# add padding to the latent to make it match pos_embed
|
||||
latent_size = height * width * num_channels / 16 # todo check where 16 comes from?
|
||||
pos_embed_size = self.pos_embed.shape[1]
|
||||
if latent_size < pos_embed_size:
|
||||
total_padding = int(pos_embed_size - math.floor(latent_size))
|
||||
total_padding = total_padding // 16
|
||||
pad_height = total_padding // 2
|
||||
pad_width = total_padding - pad_height
|
||||
# mirror padding on the right side
|
||||
padding = (0, pad_width, 0, pad_height)
|
||||
latent = torch.nn.functional.pad(latent, padding, mode='reflect')
|
||||
elif latent_size > pos_embed_size:
|
||||
amount_to_remove = latent_size - pos_embed_size
|
||||
latent = latent[:, :, :-amount_to_remove]
|
||||
|
||||
batch_size, num_channels, height, width = latent.size()
|
||||
|
||||
latent = latent.view(
|
||||
batch_size,
|
||||
num_channels,
|
||||
height // self.patch_size,
|
||||
self.patch_size,
|
||||
width // self.patch_size,
|
||||
self.patch_size,
|
||||
)
|
||||
latent = latent.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2)
|
||||
latent = self.proj(latent)
|
||||
try:
|
||||
return latent + self.pos_embed
|
||||
except RuntimeError:
|
||||
raise RuntimeError(
|
||||
f"Positional embeddings are too small for the number of patches. "
|
||||
f"Please increase `pos_embed_max_size` to at least {self.num_patches}."
|
||||
)
|
||||
|
||||
pos_embed.forward = partial(new_forward, pos_embed)
|
||||
1433
toolkit/models/base_model.py
Normal file
1433
toolkit/models/base_model.py
Normal file
File diff suppressed because it is too large
Load Diff
162
toolkit/models/clip_fusion.py
Normal file
162
toolkit/models/clip_fusion.py
Normal file
@@ -0,0 +1,162 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from toolkit.models.zipper_resampler import ContextualAlphaMask
|
||||
|
||||
|
||||
# Conv1d MLP
|
||||
# MLP that can alternately be used as a conv1d on dim 1
|
||||
class MLPC(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_dim,
|
||||
out_dim,
|
||||
hidden_dim,
|
||||
do_conv=False,
|
||||
use_residual=True
|
||||
):
|
||||
super().__init__()
|
||||
self.do_conv = do_conv
|
||||
if use_residual:
|
||||
assert in_dim == out_dim
|
||||
# dont normalize if using conv
|
||||
if not do_conv:
|
||||
self.layernorm = nn.LayerNorm(in_dim)
|
||||
|
||||
if do_conv:
|
||||
self.fc1 = nn.Conv1d(in_dim, hidden_dim, 1)
|
||||
self.fc2 = nn.Conv1d(hidden_dim, out_dim, 1)
|
||||
else:
|
||||
self.fc1 = nn.Linear(in_dim, hidden_dim)
|
||||
self.fc2 = nn.Linear(hidden_dim, out_dim)
|
||||
|
||||
self.use_residual = use_residual
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
if not self.do_conv:
|
||||
x = self.layernorm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
|
||||
class ZipperBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_size,
|
||||
in_tokens,
|
||||
out_size,
|
||||
out_tokens,
|
||||
hidden_size,
|
||||
hidden_tokens,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_size = in_size
|
||||
self.in_tokens = in_tokens
|
||||
self.out_size = out_size
|
||||
self.out_tokens = out_tokens
|
||||
self.hidden_size = hidden_size
|
||||
self.hidden_tokens = hidden_tokens
|
||||
# permute to (batch_size, out_size, in_tokens)
|
||||
|
||||
self.zip_token = MLPC(
|
||||
in_dim=self.in_tokens,
|
||||
out_dim=self.out_tokens,
|
||||
hidden_dim=self.hidden_tokens,
|
||||
do_conv=True, # no need to permute
|
||||
use_residual=False
|
||||
)
|
||||
|
||||
# permute to (batch_size, out_tokens, out_size)
|
||||
|
||||
# in shpae: (batch_size, in_tokens, in_size)
|
||||
self.zip_size = MLPC(
|
||||
in_dim=self.in_size,
|
||||
out_dim=self.out_size,
|
||||
hidden_dim=self.hidden_size,
|
||||
use_residual=False
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.zip_token(x)
|
||||
x = self.zip_size(x)
|
||||
return x
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# CLIPFusionModule
|
||||
# Fuses any size of vision and text embeddings into a single embedding.
|
||||
# remaps tokens and vectors.
|
||||
class CLIPFusionModule(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
text_hidden_size: int = 768,
|
||||
text_tokens: int = 77,
|
||||
vision_hidden_size: int = 1024,
|
||||
vision_tokens: int = 257,
|
||||
num_blocks: int = 1,
|
||||
):
|
||||
super(CLIPFusionModule, self).__init__()
|
||||
|
||||
self.text_hidden_size = text_hidden_size
|
||||
self.text_tokens = text_tokens
|
||||
self.vision_hidden_size = vision_hidden_size
|
||||
self.vision_tokens = vision_tokens
|
||||
|
||||
self.resampler = ZipperBlock(
|
||||
in_size=self.vision_hidden_size,
|
||||
in_tokens=self.vision_tokens,
|
||||
out_size=self.text_hidden_size,
|
||||
out_tokens=self.text_tokens,
|
||||
hidden_size=self.vision_hidden_size * 2,
|
||||
hidden_tokens=self.vision_tokens * 2
|
||||
)
|
||||
|
||||
self.zipper_blocks = torch.nn.ModuleList([
|
||||
ZipperBlock(
|
||||
in_size=self.text_hidden_size * 2,
|
||||
in_tokens=self.text_tokens,
|
||||
out_size=self.text_hidden_size,
|
||||
out_tokens=self.text_tokens,
|
||||
hidden_size=self.text_hidden_size * 2,
|
||||
hidden_tokens=self.text_tokens * 2
|
||||
) for i in range(num_blocks)
|
||||
])
|
||||
|
||||
self.ctx_alpha = ContextualAlphaMask(
|
||||
dim=self.text_hidden_size,
|
||||
)
|
||||
|
||||
self.alpha = nn.Parameter(torch.zeros([text_tokens]) + 0.01)
|
||||
|
||||
def forward(self, text_embeds, vision_embeds):
|
||||
# text_embeds = (batch_size, 77, 768)
|
||||
# vision_embeds = (batch_size, 257, 1024)
|
||||
# output = (batch_size, 77, 768)
|
||||
|
||||
vision_embeds = self.resampler(vision_embeds)
|
||||
x = vision_embeds
|
||||
for i, block in enumerate(self.zipper_blocks):
|
||||
res = x
|
||||
x = torch.cat([text_embeds, x], dim=-1)
|
||||
x = block(x)
|
||||
x = x + res
|
||||
|
||||
# alpha mask
|
||||
ctx_alpha = self.ctx_alpha(text_embeds)
|
||||
# reshape alpha to (1, 77, 1)
|
||||
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
|
||||
|
||||
x = ctx_alpha * x * alpha
|
||||
|
||||
x = x + text_embeds
|
||||
|
||||
return x
|
||||
123
toolkit/models/clip_pre_processor.py
Normal file
123
toolkit/models/clip_pre_processor.py
Normal file
@@ -0,0 +1,123 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class UpsampleBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.conv_in = nn.Sequential(
|
||||
nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1),
|
||||
nn.GELU()
|
||||
)
|
||||
self.conv_up = nn.Sequential(
|
||||
nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2),
|
||||
nn.GELU()
|
||||
)
|
||||
|
||||
self.conv_out = nn.Sequential(
|
||||
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_in(x)
|
||||
x = self.conv_up(x)
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class CLIPImagePreProcessor(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_size=896,
|
||||
clip_input_size=224,
|
||||
downscale_factor: int = 16,
|
||||
):
|
||||
super().__init__()
|
||||
# make sure they are evenly divisible
|
||||
assert input_size % clip_input_size == 0
|
||||
in_channels = 3
|
||||
|
||||
self.input_size = input_size
|
||||
self.clip_input_size = clip_input_size
|
||||
self.downscale_factor = downscale_factor
|
||||
|
||||
subpixel_channels = in_channels * downscale_factor ** 2 # 3 * 16 ** 2 = 768
|
||||
channels = subpixel_channels
|
||||
|
||||
upscale_factor = downscale_factor / int((input_size / clip_input_size)) # 16 / (896 / 224) = 4
|
||||
|
||||
num_upsample_blocks = int(upscale_factor // 2) # 4 // 2 = 2
|
||||
|
||||
# make the residual down up blocks
|
||||
self.upsample_blocks = nn.ModuleList()
|
||||
self.subpixel_blocks = nn.ModuleList()
|
||||
current_channels = channels
|
||||
current_downscale = downscale_factor
|
||||
for _ in range(num_upsample_blocks):
|
||||
# determine the reshuffled channel count for this dimension
|
||||
output_downscale = current_downscale // 2
|
||||
out_channels = in_channels * output_downscale ** 2
|
||||
# out_channels = current_channels // 2
|
||||
self.upsample_blocks.append(UpsampleBlock(current_channels, out_channels))
|
||||
current_channels = out_channels
|
||||
current_downscale = output_downscale
|
||||
self.subpixel_blocks.append(nn.PixelUnshuffle(current_downscale))
|
||||
|
||||
# (bs, 768, 56, 56) -> (bs, 192, 112, 112)
|
||||
# (bs, 192, 112, 112) -> (bs, 48, 224, 224)
|
||||
|
||||
self.conv_out = nn.Conv2d(
|
||||
current_channels,
|
||||
out_channels=3,
|
||||
kernel_size=3,
|
||||
padding=1
|
||||
) # (bs, 48, 224, 224) -> (bs, 3, 224, 224)
|
||||
|
||||
# do a pooling layer to downscale the input to 1/3 of the size
|
||||
# (bs, 3, 896, 896) -> (bs, 3, 224, 224)
|
||||
kernel_size = input_size // clip_input_size
|
||||
self.res_down = nn.AvgPool2d(
|
||||
kernel_size=kernel_size,
|
||||
stride=kernel_size
|
||||
) # (bs, 3, 896, 896) -> (bs, 3, 224, 224)
|
||||
|
||||
# make a blending for output residual with near 0 weight
|
||||
self.res_blend = nn.Parameter(torch.tensor(0.001)) # (bs, 3, 224, 224) -> (bs, 3, 224, 224)
|
||||
|
||||
self.unshuffle = nn.PixelUnshuffle(downscale_factor) # (bs, 3, 896, 896) -> (bs, 768, 56, 56)
|
||||
|
||||
self.conv_in = nn.Sequential(
|
||||
nn.Conv2d(
|
||||
subpixel_channels,
|
||||
channels,
|
||||
kernel_size=3,
|
||||
padding=1
|
||||
),
|
||||
nn.GELU()
|
||||
) # (bs, 768, 56, 56) -> (bs, 768, 56, 56)
|
||||
|
||||
# make 2 deep blocks
|
||||
|
||||
def forward(self, x):
|
||||
inputs = x
|
||||
# resize to input_size x input_size
|
||||
x = nn.functional.interpolate(x, size=(self.input_size, self.input_size), mode='bicubic')
|
||||
|
||||
res = self.res_down(inputs)
|
||||
|
||||
x = self.unshuffle(x)
|
||||
x = self.conv_in(x)
|
||||
for up, subpixel in zip(self.upsample_blocks, self.subpixel_blocks):
|
||||
x = up(x)
|
||||
block_res = subpixel(inputs)
|
||||
x = x + block_res
|
||||
x = self.conv_out(x)
|
||||
# blend residual
|
||||
x = x * self.res_blend + res
|
||||
return x
|
||||
466
toolkit/models/cogview4.py
Normal file
466
toolkit/models/cogview4.py
Normal file
@@ -0,0 +1,466 @@
|
||||
# DONT USE THIS!. IT DOES NOT WORK YET!
|
||||
# Will revisit this when they release more info on how it was trained.
|
||||
|
||||
import weakref
|
||||
from diffusers import CogView4Pipeline
|
||||
import torch
|
||||
import yaml
|
||||
|
||||
from toolkit.basic import flush
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
from toolkit.models.base_model import BaseModel
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
|
||||
import os
|
||||
import copy
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, ModelArch
|
||||
import torch
|
||||
import diffusers
|
||||
from diffusers import AutoencoderKL, CogView4Transformer2DModel, CogView4Pipeline
|
||||
from optimum.quanto import freeze, qfloat8, QTensor, qint4
|
||||
from toolkit.util.quantize import quantize, get_qtype
|
||||
from transformers import GlmModel, AutoTokenizer
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from typing import TYPE_CHECKING
|
||||
from toolkit.accelerator import unwrap_model
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
|
||||
# remove this after a bug is fixed in diffusers code. This is a workaround.
|
||||
|
||||
|
||||
class FakeModel:
|
||||
def __init__(self, model):
|
||||
self.model_ref = weakref.ref(model)
|
||||
pass
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return self.model_ref().device
|
||||
|
||||
|
||||
scheduler_config = {
|
||||
"base_image_seq_len": 256,
|
||||
"base_shift": 0.25,
|
||||
"invert_sigmas": False,
|
||||
"max_image_seq_len": 4096,
|
||||
"max_shift": 0.75,
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 1.0,
|
||||
"shift_terminal": None,
|
||||
"time_shift_type": "linear",
|
||||
"use_beta_sigmas": False,
|
||||
"use_dynamic_shifting": True,
|
||||
"use_exponential_sigmas": False,
|
||||
"use_karras_sigmas": False
|
||||
}
|
||||
|
||||
|
||||
class CogView4(BaseModel):
|
||||
def __init__(
|
||||
self,
|
||||
device,
|
||||
model_config: ModelConfig,
|
||||
dtype='bf16',
|
||||
custom_pipeline=None,
|
||||
noise_scheduler=None,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(device, model_config, dtype,
|
||||
custom_pipeline, noise_scheduler, **kwargs)
|
||||
self.is_flow_matching = True
|
||||
self.is_transformer = True
|
||||
self.target_lora_modules = ['CogView4Transformer2DModel']
|
||||
|
||||
# cache for holding noise
|
||||
self.effective_noise = None
|
||||
|
||||
# static method to get the scheduler
|
||||
@staticmethod
|
||||
def get_train_scheduler():
|
||||
scheduler = CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
|
||||
return scheduler
|
||||
|
||||
def load_model(self):
|
||||
dtype = self.torch_dtype
|
||||
base_model_path = "THUDM/CogView4-6B"
|
||||
model_path = self.model_config.name_or_path
|
||||
|
||||
self.print_and_status_update("Loading CogView4 model")
|
||||
# base_model_path = "black-forest-labs/FLUX.1-schnell"
|
||||
base_model_path = self.model_config.name_or_path_original
|
||||
subfolder = 'transformer'
|
||||
transformer_path = model_path
|
||||
if os.path.exists(transformer_path):
|
||||
subfolder = None
|
||||
transformer_path = os.path.join(transformer_path, 'transformer')
|
||||
# check if the path is a full checkpoint.
|
||||
te_folder_path = os.path.join(model_path, 'text_encoder')
|
||||
# if we have the te, this folder is a full checkpoint, use it as the base
|
||||
if os.path.exists(te_folder_path):
|
||||
base_model_path = model_path
|
||||
|
||||
self.print_and_status_update("Loading GlmModel")
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
base_model_path, subfolder="tokenizer", torch_dtype=dtype)
|
||||
text_encoder = GlmModel.from_pretrained(
|
||||
base_model_path, subfolder="text_encoder", torch_dtype=dtype)
|
||||
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
flush()
|
||||
|
||||
if self.model_config.quantize_te:
|
||||
self.print_and_status_update("Quantizing GlmModel")
|
||||
quantize(text_encoder, weights=get_qtype(self.model_config.qtype))
|
||||
freeze(text_encoder)
|
||||
flush()
|
||||
|
||||
# hack to fix diffusers bug workaround
|
||||
text_encoder.model = FakeModel(text_encoder)
|
||||
|
||||
self.print_and_status_update("Loading transformer")
|
||||
transformer = CogView4Transformer2DModel.from_pretrained(
|
||||
transformer_path,
|
||||
subfolder=subfolder,
|
||||
torch_dtype=dtype,
|
||||
)
|
||||
|
||||
if self.model_config.split_model_over_gpus:
|
||||
raise ValueError(
|
||||
"Splitting model over gpus is not supported for CogViewModels models")
|
||||
|
||||
transformer.to(self.quantize_device, dtype=dtype)
|
||||
flush()
|
||||
|
||||
if self.model_config.assistant_lora_path is not None or self.model_config.inference_lora_path is not None:
|
||||
raise ValueError(
|
||||
"Assistant LoRA is not supported for CogViewModels models currently")
|
||||
|
||||
if self.model_config.lora_path is not None:
|
||||
raise ValueError(
|
||||
"Loading LoRA is not supported for CogViewModels models currently")
|
||||
|
||||
flush()
|
||||
|
||||
if self.model_config.quantize:
|
||||
quantization_args = self.model_config.quantize_kwargs
|
||||
if 'exclude' not in quantization_args:
|
||||
quantization_args['exclude'] = []
|
||||
if 'include' not in quantization_args:
|
||||
quantization_args['include'] = []
|
||||
|
||||
# Be more specific with the include pattern to exactly match transformer blocks
|
||||
quantization_args['include'] += ["transformer_blocks.*"]
|
||||
|
||||
# Exclude all LayerNorm layers within transformer blocks
|
||||
quantization_args['exclude'] += [
|
||||
"transformer_blocks.*.norm1",
|
||||
"transformer_blocks.*.norm2",
|
||||
"transformer_blocks.*.norm2_context",
|
||||
"transformer_blocks.*.attn1.norm_q",
|
||||
"transformer_blocks.*.attn1.norm_k"
|
||||
]
|
||||
|
||||
# patch the state dict method
|
||||
patch_dequantization_on_save(transformer)
|
||||
quantization_type = get_qtype(self.model_config.qtype)
|
||||
self.print_and_status_update("Quantizing transformer")
|
||||
quantize(transformer, weights=quantization_type, **quantization_args)
|
||||
freeze(transformer)
|
||||
transformer.to(self.device_torch)
|
||||
else:
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
|
||||
flush()
|
||||
|
||||
scheduler = CogView4.get_train_scheduler()
|
||||
self.print_and_status_update("Loading VAE")
|
||||
vae = AutoencoderKL.from_pretrained(
|
||||
base_model_path, subfolder="vae", torch_dtype=dtype)
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Making pipe")
|
||||
pipe: CogView4Pipeline = CogView4Pipeline(
|
||||
scheduler=scheduler,
|
||||
text_encoder=None,
|
||||
tokenizer=tokenizer,
|
||||
vae=vae,
|
||||
transformer=None,
|
||||
)
|
||||
pipe.text_encoder = text_encoder
|
||||
pipe.transformer = transformer
|
||||
|
||||
self.print_and_status_update("Preparing Model")
|
||||
|
||||
text_encoder = pipe.text_encoder
|
||||
tokenizer = pipe.tokenizer
|
||||
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
|
||||
flush()
|
||||
text_encoder.to(self.device_torch)
|
||||
text_encoder.requires_grad_(False)
|
||||
text_encoder.eval()
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
flush()
|
||||
self.pipeline = pipe
|
||||
self.model = transformer
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def get_generation_pipeline(self):
|
||||
scheduler = CogView4.get_train_scheduler()
|
||||
pipeline = CogView4Pipeline(
|
||||
vae=self.vae,
|
||||
transformer=self.unet,
|
||||
text_encoder=self.text_encoder,
|
||||
tokenizer=self.tokenizer,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
return pipeline
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: CogView4Pipeline,
|
||||
gen_config: GenerateImageConfig,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
unconditional_embeds: PromptEmbeds,
|
||||
generator: torch.Generator,
|
||||
extra: dict,
|
||||
):
|
||||
img = pipeline(
|
||||
prompt_embeds=conditional_embeds.text_embeds.to(
|
||||
self.device_torch, dtype=self.torch_dtype),
|
||||
negative_prompt_embeds=unconditional_embeds.text_embeds.to(
|
||||
self.device_torch, dtype=self.torch_dtype),
|
||||
height=gen_config.height,
|
||||
width=gen_config.width,
|
||||
num_inference_steps=gen_config.num_inference_steps,
|
||||
guidance_scale=gen_config.guidance_scale,
|
||||
latents=gen_config.latents,
|
||||
generator=generator,
|
||||
**extra
|
||||
).images[0]
|
||||
return img
|
||||
|
||||
def get_noise_prediction(
|
||||
self,
|
||||
latent_model_input: torch.Tensor,
|
||||
timestep: torch.Tensor, # 0 to 1000 scale
|
||||
text_embeddings: PromptEmbeds,
|
||||
**kwargs
|
||||
):
|
||||
# target_size = (height, width)
|
||||
target_size = latent_model_input.shape[-2:]
|
||||
# multiply by 8
|
||||
target_size = (target_size[0] * 8, target_size[1] * 8)
|
||||
crops_coords_top_left = torch.tensor(
|
||||
[(0, 0)], dtype=self.torch_dtype, device=self.device_torch)
|
||||
|
||||
original_size = torch.tensor(
|
||||
[target_size], dtype=self.torch_dtype, device=self.device_torch)
|
||||
target_size = original_size.clone()
|
||||
noise_pred_cond = self.model(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
timestep=timestep,
|
||||
original_size=original_size,
|
||||
target_size=target_size,
|
||||
crop_coords=crops_coords_top_left,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
return noise_pred_cond
|
||||
|
||||
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
|
||||
prompt_embeds, _ = self.pipeline.encode_prompt(
|
||||
prompt,
|
||||
do_classifier_free_guidance=False,
|
||||
device=self.device_torch,
|
||||
dtype=self.torch_dtype,
|
||||
)
|
||||
return PromptEmbeds(prompt_embeds)
|
||||
|
||||
def get_model_has_grad(self):
|
||||
return self.model.proj_out.weight.requires_grad
|
||||
|
||||
def get_te_has_grad(self):
|
||||
return self.text_encoder.layers[0].mlp.down_proj.weight.requires_grad
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
# only save the unet
|
||||
transformer: CogView4Transformer2DModel = unwrap_model(self.model)
|
||||
transformer.save_pretrained(
|
||||
save_directory=os.path.join(output_path, 'transformer'),
|
||||
safe_serialization=True,
|
||||
)
|
||||
|
||||
meta_path = os.path.join(output_path, 'aitk_meta.yaml')
|
||||
with open(meta_path, 'w') as f:
|
||||
yaml.dump(meta, f)
|
||||
|
||||
def get_loss_target(self, *args, **kwargs):
|
||||
noise = kwargs.get('noise')
|
||||
effective_noise = self.effective_noise
|
||||
batch = kwargs.get('batch')
|
||||
if batch is None:
|
||||
raise ValueError("Batch is not provided")
|
||||
if noise is None:
|
||||
raise ValueError("Noise is not provided")
|
||||
# return batch.latents
|
||||
# return (batch.latents - noise).detach()
|
||||
return (noise - batch.latents).detach()
|
||||
# return (batch.latents).detach()
|
||||
# return (effective_noise - batch.latents).detach()
|
||||
|
||||
def _get_low_res_latents(self, latents):
|
||||
# todo prevent needing to do this and grab the tensor another way.
|
||||
with torch.no_grad():
|
||||
# Decode latents to image space
|
||||
images = self.decode_latents(
|
||||
latents, device=latents.device, dtype=latents.dtype)
|
||||
|
||||
# Downsample by a factor of 2 using bilinear interpolation
|
||||
B, C, H, W = images.shape
|
||||
low_res_images = torch.nn.functional.interpolate(
|
||||
images,
|
||||
size=(H // 2, W // 2),
|
||||
mode="bilinear",
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
# Upsample back to original resolution to match expected VAE input dimensions
|
||||
upsampled_low_res_images = torch.nn.functional.interpolate(
|
||||
low_res_images,
|
||||
size=(H, W),
|
||||
mode="bilinear",
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
# Encode the low-resolution images back to latent space
|
||||
low_res_latents = self.encode_images(
|
||||
upsampled_low_res_images, device=latents.device, dtype=latents.dtype)
|
||||
return low_res_latents
|
||||
|
||||
# def add_noise(
|
||||
# self,
|
||||
# original_samples: torch.FloatTensor,
|
||||
# noise: torch.FloatTensor,
|
||||
# timesteps: torch.IntTensor,
|
||||
# **kwargs,
|
||||
# ) -> torch.FloatTensor:
|
||||
# relay_start_point = 500
|
||||
|
||||
# # Store original samples for loss calculation
|
||||
# self.original_samples = original_samples
|
||||
|
||||
# # Prepare chunks for batch processing
|
||||
# original_samples_chunks = torch.chunk(
|
||||
# original_samples, original_samples.shape[0], dim=0)
|
||||
# noise_chunks = torch.chunk(noise, noise.shape[0], dim=0)
|
||||
# timesteps_chunks = torch.chunk(timesteps, timesteps.shape[0], dim=0)
|
||||
|
||||
# # Get the low res latents only if needed
|
||||
# low_res_latents_chunks = None
|
||||
|
||||
# # Handle case where timesteps is a single value for all samples
|
||||
# if len(timesteps_chunks) == 1 and len(timesteps_chunks) != len(original_samples_chunks):
|
||||
# timesteps_chunks = [timesteps_chunks[0]] * len(original_samples_chunks)
|
||||
|
||||
# noisy_latents_chunks = []
|
||||
# effective_noise_chunks = [] # Store the effective noise for each sample
|
||||
|
||||
# for idx in range(original_samples.shape[0]):
|
||||
# t = timesteps_chunks[idx]
|
||||
# t_01 = (t / 1000).to(original_samples_chunks[idx].device)
|
||||
|
||||
# # Flowmatching interpolation between original and noise
|
||||
# if t > relay_start_point:
|
||||
# # Standard flowmatching - direct linear interpolation
|
||||
# noisy_latents = (1 - t_01) * original_samples_chunks[idx] + t_01 * noise_chunks[idx]
|
||||
# effective_noise_chunks.append(noise_chunks[idx]) # Effective noise is just the noise
|
||||
# else:
|
||||
# # Relay flowmatching case - only compute low_res_latents if needed
|
||||
# if low_res_latents_chunks is None:
|
||||
# low_res_latents = self._get_low_res_latents(original_samples)
|
||||
# low_res_latents_chunks = torch.chunk(low_res_latents, low_res_latents.shape[0], dim=0)
|
||||
|
||||
# # Calculate the relay ratio (0 to 1)
|
||||
# t_ratio = t.float() / relay_start_point
|
||||
# t_ratio = torch.clamp(t_ratio, 0.0, 1.0)
|
||||
|
||||
# # First blend between original and low-res based on t_ratio
|
||||
# z0_t = (1 - t_ratio) * original_samples_chunks[idx] + t_ratio * low_res_latents_chunks[idx]
|
||||
|
||||
# added_lor_res_noise = z0_t - original_samples_chunks[idx]
|
||||
|
||||
# # Then apply flowmatching interpolation between this blended state and noise
|
||||
# noisy_latents = (1 - t_01) * z0_t + t_01 * noise_chunks[idx]
|
||||
|
||||
# # For prediction target, we need to store the effective "source"
|
||||
# effective_noise_chunks.append(noise_chunks[idx] + added_lor_res_noise)
|
||||
|
||||
# noisy_latents_chunks.append(noisy_latents)
|
||||
|
||||
# noisy_latents = torch.cat(noisy_latents_chunks, dim=0)
|
||||
# self.effective_noise = torch.cat(effective_noise_chunks, dim=0) # Store for loss calculation
|
||||
|
||||
# return noisy_latents
|
||||
|
||||
# def add_noise(
|
||||
# self,
|
||||
# original_samples: torch.FloatTensor,
|
||||
# noise: torch.FloatTensor,
|
||||
# timesteps: torch.IntTensor,
|
||||
# **kwargs,
|
||||
# ) -> torch.FloatTensor:
|
||||
# relay_start_point = 500
|
||||
|
||||
# # Store original samples for loss calculation
|
||||
# self.original_samples = original_samples
|
||||
|
||||
# # Prepare chunks for batch processing
|
||||
# original_samples_chunks = torch.chunk(
|
||||
# original_samples, original_samples.shape[0], dim=0)
|
||||
# noise_chunks = torch.chunk(noise, noise.shape[0], dim=0)
|
||||
# timesteps_chunks = torch.chunk(timesteps, timesteps.shape[0], dim=0)
|
||||
|
||||
# # Get the low res latents only if needed
|
||||
# low_res_latents = self._get_low_res_latents(original_samples)
|
||||
# low_res_latents_chunks = torch.chunk(low_res_latents, low_res_latents.shape[0], dim=0)
|
||||
|
||||
# # Handle case where timesteps is a single value for all samples
|
||||
# if len(timesteps_chunks) == 1 and len(timesteps_chunks) != len(original_samples_chunks):
|
||||
# timesteps_chunks = [timesteps_chunks[0]] * len(original_samples_chunks)
|
||||
|
||||
# noisy_latents_chunks = []
|
||||
# effective_noise_chunks = [] # Store the effective noise for each sample
|
||||
|
||||
# for idx in range(original_samples.shape[0]):
|
||||
# t = timesteps_chunks[idx]
|
||||
# t_01 = (t / 1000).to(original_samples_chunks[idx].device)
|
||||
|
||||
# lrln = low_res_latents_chunks[idx] - original_samples_chunks[idx]
|
||||
# # lrln = lrln * (1 - t_01)
|
||||
|
||||
# # make the noise an interpolation between noise and low_res_latents with
|
||||
# # being noise at t_01=1 and low_res_latents at t_01=0
|
||||
# new_noise = t_01 * noise_chunks[idx] + (1 - t_01) * lrln
|
||||
# # new_noise = noise_chunks[idx] + lrln
|
||||
# # new_noise = noise_chunks[idx] + lrln
|
||||
|
||||
# # Then apply flowmatching interpolation between this blended state and noise
|
||||
# noisy_latents = (1 - t_01) * original_samples + t_01 * new_noise
|
||||
|
||||
# # For prediction target, we need to store the effective "source"
|
||||
# effective_noise_chunks.append(new_noise)
|
||||
|
||||
# noisy_latents_chunks.append(noisy_latents)
|
||||
|
||||
# noisy_latents = torch.cat(noisy_latents_chunks, dim=0)
|
||||
# self.effective_noise = torch.cat(effective_noise_chunks, dim=0) # Store for loss calculation
|
||||
|
||||
# return noisy_latents
|
||||
272
toolkit/models/control_lora_adapter.py
Normal file
272
toolkit/models/control_lora_adapter.py
Normal file
@@ -0,0 +1,272 @@
|
||||
import inspect
|
||||
import weakref
|
||||
import torch
|
||||
from typing import TYPE_CHECKING
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
from diffusers import FluxTransformer2DModel
|
||||
# weakref
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.config_modules import AdapterConfig, TrainConfig, ModelConfig
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
|
||||
|
||||
# after each step we concat the control image with the latents
|
||||
# latent_model_input = torch.cat([latents, control_image], dim=2)
|
||||
# the x_embedder has a full rank lora to handle the additional channels
|
||||
# this replaces the x_embedder with a full rank lora. on flux this is
|
||||
# x_embedder(diffusers) or img_in(bfl)
|
||||
|
||||
# Flux
|
||||
# img_in.lora_A.weight [128, 128]
|
||||
# img_in.lora_B.bias [3 072]
|
||||
# img_in.lora_B.weight [3 072, 128]
|
||||
|
||||
|
||||
class ImgEmbedder(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'ControlLoraAdapter',
|
||||
orig_layer: torch.nn.Linear,
|
||||
in_channels=64,
|
||||
out_channels=3072
|
||||
):
|
||||
super().__init__()
|
||||
# only do the weight for the new input. We combine with the original linear layer
|
||||
init = torch.randn(out_channels, in_channels, device=orig_layer.weight.device, dtype=orig_layer.weight.dtype) * 0.01
|
||||
self.weight = torch.nn.Parameter(init)
|
||||
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.orig_layer_ref: weakref.ref = weakref.ref(orig_layer)
|
||||
|
||||
@classmethod
|
||||
def from_model(
|
||||
cls,
|
||||
model: FluxTransformer2DModel,
|
||||
adapter: 'ControlLoraAdapter',
|
||||
num_control_images=1,
|
||||
has_inpainting_input=False
|
||||
):
|
||||
if model.__class__.__name__ == 'FluxTransformer2DModel':
|
||||
num_adapter_in_channels = model.x_embedder.in_features * num_control_images
|
||||
|
||||
if has_inpainting_input:
|
||||
# inpainting has the mask before packing latents. it is normally 16 ch + 1ch mask
|
||||
# packed it is 64ch + 4ch mask
|
||||
# so we need to add 4 to the input channels
|
||||
num_adapter_in_channels += 4
|
||||
|
||||
x_embedder: torch.nn.Linear = model.x_embedder
|
||||
img_embedder = cls(
|
||||
adapter,
|
||||
orig_layer=x_embedder,
|
||||
in_channels=num_adapter_in_channels,
|
||||
out_channels=x_embedder.out_features,
|
||||
)
|
||||
|
||||
# hijack the forward method
|
||||
x_embedder._orig_ctrl_lora_forward = x_embedder.forward
|
||||
x_embedder.forward = img_embedder.forward
|
||||
|
||||
# update the config of the transformer
|
||||
model.config.in_channels = model.config.in_channels * (num_control_images + 1)
|
||||
model.config["in_channels"] = model.config.in_channels
|
||||
|
||||
return img_embedder
|
||||
else:
|
||||
raise ValueError("Model not supported")
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
|
||||
def forward(self, x):
|
||||
if not self.is_active:
|
||||
# make sure lora is not active
|
||||
if self.adapter_ref().control_lora is not None:
|
||||
self.adapter_ref().control_lora.is_active = False
|
||||
return self.orig_layer_ref()._orig_ctrl_lora_forward(x)
|
||||
|
||||
# make sure lora is active
|
||||
if self.adapter_ref().control_lora is not None:
|
||||
self.adapter_ref().control_lora.is_active = True
|
||||
|
||||
orig_device = x.device
|
||||
orig_dtype = x.dtype
|
||||
|
||||
x = x.to(self.weight.device, dtype=self.weight.dtype)
|
||||
|
||||
orig_weight = self.orig_layer_ref().weight.data.detach()
|
||||
orig_weight = orig_weight.to(self.weight.device, dtype=self.weight.dtype)
|
||||
linear_weight = torch.cat([orig_weight, self.weight], dim=1)
|
||||
|
||||
bias = None
|
||||
if self.orig_layer_ref().bias is not None:
|
||||
bias = self.orig_layer_ref().bias.data.detach().to(self.weight.device, dtype=self.weight.dtype)
|
||||
|
||||
x = torch.nn.functional.linear(x, linear_weight, bias)
|
||||
|
||||
x = x.to(orig_device, dtype=orig_dtype)
|
||||
return x
|
||||
|
||||
|
||||
|
||||
class ControlLoraAdapter(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'CustomAdapter',
|
||||
sd: 'StableDiffusion',
|
||||
config: 'AdapterConfig',
|
||||
train_config: 'TrainConfig'
|
||||
):
|
||||
super().__init__()
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.sd_ref = weakref.ref(sd)
|
||||
self.model_config: ModelConfig = sd.model_config
|
||||
self.network_config = config.lora_config
|
||||
self.train_config = train_config
|
||||
self.device_torch = sd.device_torch
|
||||
self.control_lora = None
|
||||
|
||||
if self.network_config is not None:
|
||||
|
||||
network_kwargs = {} if self.network_config.network_kwargs is None else self.network_config.network_kwargs
|
||||
if hasattr(sd, 'target_lora_modules'):
|
||||
network_kwargs['target_lin_modules'] = self.sd.target_lora_modules
|
||||
|
||||
if 'ignore_if_contains' not in network_kwargs:
|
||||
network_kwargs['ignore_if_contains'] = []
|
||||
|
||||
# always ignore x_embedder
|
||||
network_kwargs['ignore_if_contains'].append('x_embedder')
|
||||
|
||||
self.control_lora = LoRASpecialNetwork(
|
||||
text_encoder=sd.text_encoder,
|
||||
unet=sd.unet,
|
||||
lora_dim=self.network_config.linear,
|
||||
multiplier=1.0,
|
||||
alpha=self.network_config.linear_alpha,
|
||||
train_unet=self.train_config.train_unet,
|
||||
train_text_encoder=self.train_config.train_text_encoder,
|
||||
conv_lora_dim=self.network_config.conv,
|
||||
conv_alpha=self.network_config.conv_alpha,
|
||||
is_sdxl=self.model_config.is_xl or self.model_config.is_ssd,
|
||||
is_v2=self.model_config.is_v2,
|
||||
is_v3=self.model_config.is_v3,
|
||||
is_pixart=self.model_config.is_pixart,
|
||||
is_auraflow=self.model_config.is_auraflow,
|
||||
is_flux=self.model_config.is_flux,
|
||||
is_lumina2=self.model_config.is_lumina2,
|
||||
is_ssd=self.model_config.is_ssd,
|
||||
is_vega=self.model_config.is_vega,
|
||||
dropout=self.network_config.dropout,
|
||||
use_text_encoder_1=self.model_config.use_text_encoder_1,
|
||||
use_text_encoder_2=self.model_config.use_text_encoder_2,
|
||||
use_bias=False,
|
||||
is_lorm=False,
|
||||
network_config=self.network_config,
|
||||
network_type=self.network_config.type,
|
||||
transformer_only=self.network_config.transformer_only,
|
||||
is_transformer=sd.is_transformer,
|
||||
base_model=sd,
|
||||
**network_kwargs
|
||||
)
|
||||
self.control_lora.force_to(self.device_torch, dtype=torch.float32)
|
||||
self.control_lora._update_torch_multiplier()
|
||||
self.control_lora.apply_to(
|
||||
sd.text_encoder,
|
||||
sd.unet,
|
||||
self.train_config.train_text_encoder,
|
||||
self.train_config.train_unet
|
||||
)
|
||||
self.control_lora.can_merge_in = False
|
||||
self.control_lora.prepare_grad_etc(sd.text_encoder, sd.unet)
|
||||
if self.train_config.gradient_checkpointing:
|
||||
self.control_lora.enable_gradient_checkpointing()
|
||||
|
||||
self.x_embedder = ImgEmbedder.from_model(
|
||||
sd.unet,
|
||||
self,
|
||||
num_control_images=config.num_control_images,
|
||||
has_inpainting_input=config.has_inpainting_input
|
||||
)
|
||||
self.x_embedder.to(self.device_torch)
|
||||
|
||||
def get_params(self):
|
||||
if self.control_lora is not None:
|
||||
config = {
|
||||
'text_encoder_lr': self.train_config.lr,
|
||||
'unet_lr': self.train_config.lr,
|
||||
}
|
||||
sig = inspect.signature(self.control_lora.prepare_optimizer_params)
|
||||
if 'default_lr' in sig.parameters:
|
||||
config['default_lr'] = self.train_config.lr
|
||||
if 'learning_rate' in sig.parameters:
|
||||
config['learning_rate'] = self.train_config.lr
|
||||
params_net = self.control_lora.prepare_optimizer_params(
|
||||
**config
|
||||
)
|
||||
|
||||
# we want only tensors here
|
||||
params = []
|
||||
for p in params_net:
|
||||
if isinstance(p, dict):
|
||||
params += p["params"]
|
||||
elif isinstance(p, torch.Tensor):
|
||||
params.append(p)
|
||||
elif isinstance(p, list):
|
||||
params += p
|
||||
else:
|
||||
params = []
|
||||
|
||||
# make sure the embedder is float32
|
||||
self.x_embedder.to(torch.float32)
|
||||
|
||||
params += list(self.x_embedder.parameters())
|
||||
|
||||
# we need to be able to yield from the list like yield from params
|
||||
|
||||
return params
|
||||
|
||||
def load_weights(self, state_dict, strict=True):
|
||||
lora_sd = {}
|
||||
img_embedder_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
if "x_embedder" in key:
|
||||
new_key = key.replace("transformer.x_embedder.", "")
|
||||
img_embedder_sd[new_key] = value
|
||||
else:
|
||||
lora_sd[key] = value
|
||||
|
||||
# todo process state dict before loading
|
||||
if self.control_lora is not None:
|
||||
self.control_lora.load_weights(lora_sd)
|
||||
# automatically upgrade the x imbedder if more dims are added
|
||||
if self.x_embedder.weight.shape[1] > img_embedder_sd['weight'].shape[1]:
|
||||
print("Upgrading x_embedder from {} to {}".format(
|
||||
img_embedder_sd['weight'].shape[1],
|
||||
self.x_embedder.weight.shape[1]
|
||||
))
|
||||
while img_embedder_sd['weight'].shape[1] < self.x_embedder.weight.shape[1]:
|
||||
img_embedder_sd['weight'] = torch.cat([img_embedder_sd['weight'] ] * 2, dim=1)
|
||||
if img_embedder_sd['weight'].shape[1] > self.x_embedder.weight.shape[1]:
|
||||
img_embedder_sd['weight'] = img_embedder_sd['weight'][:, :self.x_embedder.weight.shape[1]]
|
||||
self.x_embedder.load_state_dict(img_embedder_sd, strict=False)
|
||||
|
||||
def get_state_dict(self):
|
||||
if self.control_lora is not None:
|
||||
lora_sd = self.control_lora.get_state_dict(dtype=torch.float32)
|
||||
else:
|
||||
lora_sd = {}
|
||||
# todo make sure we match loras elseware.
|
||||
img_embedder_sd = self.x_embedder.state_dict()
|
||||
for key, value in img_embedder_sd.items():
|
||||
lora_sd[f"transformer.x_embedder.{key}"] = value
|
||||
return lora_sd
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
33
toolkit/models/decorator.py
Normal file
33
toolkit/models/decorator.py
Normal file
@@ -0,0 +1,33 @@
|
||||
import torch
|
||||
|
||||
|
||||
class Decorator(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
num_tokens: int = 4,
|
||||
token_size: int = 4096,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.weight: torch.nn.Parameter = torch.nn.Parameter(
|
||||
torch.randn(num_tokens, token_size)
|
||||
)
|
||||
# ensure it is float32
|
||||
self.weight.data = self.weight.data.float()
|
||||
|
||||
def forward(self, text_embeds: torch.Tensor, is_unconditional=False) -> torch.Tensor:
|
||||
# make sure the param is float32
|
||||
if self.weight.dtype != text_embeds.dtype:
|
||||
self.weight.data = self.weight.data.float()
|
||||
# expand batch to match text_embeds
|
||||
batch_size = text_embeds.shape[0]
|
||||
decorator_embeds = self.weight.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
if is_unconditional:
|
||||
# zero pad the decorator embeds
|
||||
decorator_embeds = torch.zeros_like(decorator_embeds)
|
||||
|
||||
if decorator_embeds.dtype != text_embeds.dtype:
|
||||
decorator_embeds = decorator_embeds.to(text_embeds.dtype)
|
||||
text_embeds = torch.cat((text_embeds, decorator_embeds), dim=-2)
|
||||
|
||||
return text_embeds
|
||||
367
toolkit/models/diffusion_feature_extraction.py
Normal file
367
toolkit/models/diffusion_feature_extraction.py
Normal file
@@ -0,0 +1,367 @@
|
||||
import torch
|
||||
import os
|
||||
from torch import nn
|
||||
from safetensors.torch import load_file
|
||||
import torch.nn.functional as F
|
||||
from diffusers import AutoencoderTiny
|
||||
from transformers import SiglipImageProcessor, SiglipVisionModel
|
||||
import lpips
|
||||
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
|
||||
|
||||
class ResBlock(nn.Module):
|
||||
def __init__(self, in_channels, out_channels):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1)
|
||||
self.norm1 = nn.GroupNorm(8, out_channels)
|
||||
self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1)
|
||||
self.norm2 = nn.GroupNorm(8, out_channels)
|
||||
self.skip = nn.Conv2d(in_channels, out_channels,
|
||||
1) if in_channels != out_channels else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
identity = self.skip(x)
|
||||
x = self.conv1(x)
|
||||
x = self.norm1(x)
|
||||
x = F.silu(x)
|
||||
x = self.conv2(x)
|
||||
x = self.norm2(x)
|
||||
x = F.silu(x + identity)
|
||||
return x
|
||||
|
||||
|
||||
class DiffusionFeatureExtractor2(nn.Module):
|
||||
def __init__(self, in_channels=32):
|
||||
super().__init__()
|
||||
self.version = 2
|
||||
|
||||
# Path 1: Upsample to 512x512 (1, 64, 512, 512)
|
||||
self.up_path = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 64, 3, padding=1),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(64, 64),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(64, 64),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(64, 64),
|
||||
nn.Conv2d(64, 64, 3, padding=1),
|
||||
])
|
||||
|
||||
# Path 2: Upsample to 256x256 (1, 128, 256, 256)
|
||||
self.path2 = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 128, 3, padding=1),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(128, 128),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(128, 128),
|
||||
nn.Conv2d(128, 128, 3, padding=1),
|
||||
])
|
||||
|
||||
# Path 3: Upsample to 128x128 (1, 256, 128, 128)
|
||||
self.path3 = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 256, 3, padding=1),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(256, 256),
|
||||
nn.Conv2d(256, 256, 3, padding=1)
|
||||
])
|
||||
|
||||
# Path 4: Original size (1, 512, 64, 64)
|
||||
self.path4 = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 512, 3, padding=1),
|
||||
ResBlock(512, 512),
|
||||
ResBlock(512, 512),
|
||||
nn.Conv2d(512, 512, 3, padding=1)
|
||||
])
|
||||
|
||||
# Path 5: Downsample to 32x32 (1, 512, 32, 32)
|
||||
self.path5 = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 512, 3, padding=1),
|
||||
ResBlock(512, 512),
|
||||
nn.AvgPool2d(2),
|
||||
ResBlock(512, 512),
|
||||
nn.Conv2d(512, 512, 3, padding=1)
|
||||
])
|
||||
|
||||
def forward(self, x):
|
||||
outputs = []
|
||||
|
||||
# Path 1: 512x512
|
||||
x1 = x
|
||||
for layer in self.up_path:
|
||||
x1 = layer(x1)
|
||||
outputs.append(x1) # [1, 64, 512, 512]
|
||||
|
||||
# Path 2: 256x256
|
||||
x2 = x
|
||||
for layer in self.path2:
|
||||
x2 = layer(x2)
|
||||
outputs.append(x2) # [1, 128, 256, 256]
|
||||
|
||||
# Path 3: 128x128
|
||||
x3 = x
|
||||
for layer in self.path3:
|
||||
x3 = layer(x3)
|
||||
outputs.append(x3) # [1, 256, 128, 128]
|
||||
|
||||
# Path 4: 64x64
|
||||
x4 = x
|
||||
for layer in self.path4:
|
||||
x4 = layer(x4)
|
||||
outputs.append(x4) # [1, 512, 64, 64]
|
||||
|
||||
# Path 5: 32x32
|
||||
x5 = x
|
||||
for layer in self.path5:
|
||||
x5 = layer(x5)
|
||||
outputs.append(x5) # [1, 512, 32, 32]
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
class DFEBlock(nn.Module):
|
||||
def __init__(self, channels):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(channels, channels, 3, padding=1)
|
||||
self.conv2 = nn.Conv2d(channels, channels, 3, padding=1)
|
||||
self.act = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
x_in = x
|
||||
x = self.conv1(x)
|
||||
x = self.conv2(x)
|
||||
x = self.act(x)
|
||||
x = x + x_in
|
||||
return x
|
||||
|
||||
|
||||
class DiffusionFeatureExtractor(nn.Module):
|
||||
def __init__(self, in_channels=32):
|
||||
super().__init__()
|
||||
self.version = 1
|
||||
num_blocks = 6
|
||||
self.conv_in = nn.Conv2d(in_channels, 512, 1)
|
||||
self.blocks = nn.ModuleList([DFEBlock(512) for _ in range(num_blocks)])
|
||||
self.conv_out = nn.Conv2d(512, 512, 1)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_in(x)
|
||||
for block in self.blocks:
|
||||
x = block(x)
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class DiffusionFeatureExtractor3(nn.Module):
|
||||
def __init__(self, device=torch.device("cuda"), dtype=torch.bfloat16):
|
||||
super().__init__()
|
||||
self.version = 3
|
||||
vae = AutoencoderTiny.from_pretrained(
|
||||
"madebyollin/taef1", torch_dtype=torch.bfloat16)
|
||||
self.vae = vae
|
||||
image_encoder_path = "google/siglip-so400m-patch14-384"
|
||||
try:
|
||||
self.image_processor = SiglipImageProcessor.from_pretrained(
|
||||
image_encoder_path)
|
||||
except EnvironmentError:
|
||||
self.image_processor = SiglipImageProcessor()
|
||||
self.vision_encoder = SiglipVisionModel.from_pretrained(
|
||||
image_encoder_path,
|
||||
ignore_mismatched_sizes=True
|
||||
).to(device, dtype=dtype)
|
||||
|
||||
self.lpips_model = lpips_model = lpips.LPIPS(net='vgg')
|
||||
self.lpips_model = lpips_model.to(device, dtype=torch.float32)
|
||||
self.losses = {}
|
||||
self.log_every = 100
|
||||
self.step = 0
|
||||
|
||||
def get_siglip_features(self, tensors_0_1):
|
||||
dtype = torch.bfloat16
|
||||
device = self.vae.device
|
||||
# resize to 384x384
|
||||
images = F.interpolate(tensors_0_1, size=(384, 384),
|
||||
mode='bicubic', align_corners=False)
|
||||
|
||||
mean = torch.tensor(self.image_processor.image_mean).to(
|
||||
device, dtype=dtype
|
||||
).detach()
|
||||
std = torch.tensor(self.image_processor.image_std).to(
|
||||
device, dtype=dtype
|
||||
).detach()
|
||||
# tensors_0_1 = torch.clip((255. * tensors_0_1), 0, 255).round() / 255.0
|
||||
clip_image = (
|
||||
images - mean.view([1, 3, 1, 1])) / std.view([1, 3, 1, 1])
|
||||
id_embeds = self.vision_encoder(
|
||||
clip_image,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
|
||||
last_hidden_state = id_embeds['last_hidden_state']
|
||||
return last_hidden_state
|
||||
|
||||
def get_lpips_features(self, tensors_0_1):
|
||||
device = self.vae.device
|
||||
tensors_n1p1 = (tensors_0_1 * 2) - 1
|
||||
def get_lpips_features(img): # -1 to 1
|
||||
in0_input = self.lpips_model.scaling_layer(img)
|
||||
outs0 = self.lpips_model.net.forward(in0_input)
|
||||
|
||||
feats0 = {}
|
||||
|
||||
feats_list = []
|
||||
for kk in range(self.lpips_model.L):
|
||||
feats0[kk] = lpips.normalize_tensor(outs0[kk])
|
||||
feats_list.append(feats0[kk])
|
||||
|
||||
# 512 in
|
||||
# vgg
|
||||
# 0 torch.Size([1, 64, 512, 512])
|
||||
# 1 torch.Size([1, 128, 256, 256])
|
||||
# 2 torch.Size([1, 256, 128, 128])
|
||||
# 3 torch.Size([1, 512, 64, 64])
|
||||
# 4 torch.Size([1, 512, 32, 32])
|
||||
|
||||
return feats_list
|
||||
|
||||
# do lpips
|
||||
lpips_feat_list = [x for x in get_lpips_features(
|
||||
tensors_n1p1.to(device, dtype=torch.float32))]
|
||||
|
||||
return lpips_feat_list
|
||||
|
||||
|
||||
def forward(
|
||||
self,
|
||||
noise,
|
||||
noise_pred,
|
||||
noisy_latents,
|
||||
timesteps,
|
||||
batch: DataLoaderBatchDTO,
|
||||
scheduler: CustomFlowMatchEulerDiscreteScheduler,
|
||||
# lpips_weight=1.0,
|
||||
lpips_weight=10.0,
|
||||
clip_weight=0.1,
|
||||
pixel_weight=0.1
|
||||
):
|
||||
dtype = torch.bfloat16
|
||||
device = self.vae.device
|
||||
|
||||
# first we step the scheduler from current timestep to the very end for a full denoise
|
||||
# bs = noise_pred.shape[0]
|
||||
# noise_pred_chunks = torch.chunk(noise_pred, bs)
|
||||
# timestep_chunks = torch.chunk(timesteps, bs)
|
||||
# noisy_latent_chunks = torch.chunk(noisy_latents, bs)
|
||||
# stepped_chunks = []
|
||||
# for idx in range(bs):
|
||||
# model_output = noise_pred_chunks[idx]
|
||||
# timestep = timestep_chunks[idx]
|
||||
# scheduler._step_index = None
|
||||
# scheduler._init_step_index(timestep)
|
||||
# sample = noisy_latent_chunks[idx].to(torch.float32)
|
||||
|
||||
# sigma = scheduler.sigmas[scheduler.step_index]
|
||||
# sigma_next = scheduler.sigmas[-1] # use last sigma for final step
|
||||
# prev_sample = sample + (sigma_next - sigma) * model_output
|
||||
# stepped_chunks.append(prev_sample)
|
||||
|
||||
# stepped_latents = torch.cat(stepped_chunks, dim=0)
|
||||
|
||||
stepped_latents = noise - noise_pred
|
||||
|
||||
latents = stepped_latents.to(self.vae.device, dtype=self.vae.dtype)
|
||||
|
||||
latents = (
|
||||
latents / self.vae.config['scaling_factor']) + self.vae.config['shift_factor']
|
||||
tensors_n1p1 = self.vae.decode(latents).sample # -1 to 1
|
||||
|
||||
pred_images = (tensors_n1p1 + 1) / 2 # 0 to 1
|
||||
|
||||
lpips_feat_list_pred = self.get_lpips_features(pred_images.float())
|
||||
|
||||
total_loss = 0
|
||||
|
||||
with torch.no_grad():
|
||||
target_img = batch.tensor.to(device, dtype=dtype)
|
||||
# go from -1 to 1 to 0 to 1
|
||||
target_img = (target_img + 1) / 2
|
||||
lpips_feat_list_target = self.get_lpips_features(target_img.float())
|
||||
if clip_weight > 0:
|
||||
target_clip_output = self.get_siglip_features(target_img).detach()
|
||||
if clip_weight > 0:
|
||||
pred_clip_output = self.get_siglip_features(pred_images)
|
||||
clip_loss = torch.nn.functional.mse_loss(
|
||||
pred_clip_output.float(), target_clip_output.float()
|
||||
) * clip_weight
|
||||
|
||||
if 'clip_loss' not in self.losses:
|
||||
self.losses['clip_loss'] = clip_loss.item()
|
||||
else:
|
||||
self.losses['clip_loss'] += clip_loss.item()
|
||||
|
||||
total_loss += clip_loss
|
||||
|
||||
skip_lpips_layers = []
|
||||
|
||||
lpips_loss = 0
|
||||
for idx, lpips_feat in enumerate(lpips_feat_list_pred):
|
||||
if idx in skip_lpips_layers:
|
||||
continue
|
||||
lpips_loss += torch.nn.functional.mse_loss(
|
||||
lpips_feat.float(), lpips_feat_list_target[idx].float()
|
||||
) * lpips_weight
|
||||
|
||||
if f'lpips_loss_{idx}' not in self.losses:
|
||||
self.losses[f'lpips_loss_{idx}'] = lpips_loss.item()
|
||||
else:
|
||||
self.losses[f'lpips_loss_{idx}'] += lpips_loss.item()
|
||||
|
||||
total_loss += lpips_loss
|
||||
|
||||
# mse_loss = torch.nn.functional.mse_loss(
|
||||
# stepped_latents.float(), batch.latents.float()
|
||||
# ) * pixel_weight
|
||||
|
||||
# if 'pixel_loss' not in self.losses:
|
||||
# self.losses['pixel_loss'] = mse_loss.item()
|
||||
# else:
|
||||
# self.losses['pixel_loss'] += mse_loss.item()
|
||||
|
||||
if self.step % self.log_every == 0 and self.step > 0:
|
||||
print(f"DFE losses:")
|
||||
for key in self.losses:
|
||||
self.losses[key] /= self.log_every
|
||||
# print in 2.000e-01 format
|
||||
print(f" - {key}: {self.losses[key]:.3e}")
|
||||
self.losses[key] = 0.0
|
||||
|
||||
# total_loss += mse_loss
|
||||
self.step += 1
|
||||
|
||||
return total_loss
|
||||
|
||||
|
||||
def load_dfe(model_path) -> DiffusionFeatureExtractor:
|
||||
if model_path == "v3":
|
||||
dfe = DiffusionFeatureExtractor3()
|
||||
dfe.eval()
|
||||
return dfe
|
||||
if not os.path.exists(model_path):
|
||||
raise FileNotFoundError(f"Model file not found: {model_path}")
|
||||
# if it ende with safetensors
|
||||
if model_path.endswith('.safetensors'):
|
||||
state_dict = load_file(model_path)
|
||||
else:
|
||||
state_dict = torch.load(model_path, weights_only=True)
|
||||
if 'model_state_dict' in state_dict:
|
||||
state_dict = state_dict['model_state_dict']
|
||||
|
||||
if 'conv_in.weight' in state_dict:
|
||||
dfe = DiffusionFeatureExtractor()
|
||||
else:
|
||||
dfe = DiffusionFeatureExtractor2()
|
||||
|
||||
dfe.load_state_dict(state_dict)
|
||||
dfe.eval()
|
||||
return dfe
|
||||
993
toolkit/models/flex2.py
Normal file
993
toolkit/models/flex2.py
Normal file
@@ -0,0 +1,993 @@
|
||||
from typing import List, Optional, Union
|
||||
from diffusers import FluxPipeline
|
||||
import inspect
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.loaders import FluxLoraLoaderMixin, TextualInversionLoaderMixin
|
||||
from diffusers.utils import (
|
||||
USE_PEFT_BACKEND,
|
||||
is_torch_xla_available,
|
||||
logging,
|
||||
replace_example_docstring,
|
||||
scale_lora_layers,
|
||||
unscale_lora_layers,
|
||||
)
|
||||
from transformers import AutoModel, AutoTokenizer
|
||||
from transformers import (
|
||||
CLIPImageProcessor,
|
||||
CLIPTextModel,
|
||||
CLIPTokenizer,
|
||||
CLIPVisionModelWithProjection
|
||||
)
|
||||
|
||||
|
||||
|
||||
from diffusers.image_processor import PipelineImageInput, VaeImageProcessor
|
||||
from diffusers.loaders import FluxIPAdapterMixin, FluxLoraLoaderMixin, FromSingleFileMixin, TextualInversionLoaderMixin
|
||||
from diffusers.models import AutoencoderKL, FluxTransformer2DModel
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```py
|
||||
>>> import torch
|
||||
>>> from diffusers import Flex2Pipeline
|
||||
|
||||
>>> pipe = Flex2Pipeline.from_pretrained("black-forest-labs/FLUX.1-schnell", torch_dtype=torch.bfloat16)
|
||||
>>> pipe.to("cuda")
|
||||
>>> prompt = "A cat holding a sign that says hello world"
|
||||
>>> # Depending on the variant being used, the pipeline call will slightly vary.
|
||||
>>> # Refer to the pipeline documentation for more details.
|
||||
>>> image = pipe(prompt, num_inference_steps=4, guidance_scale=0.0).images[0]
|
||||
>>> image.save("flux.png")
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
|
||||
XLA_AVAILABLE = True
|
||||
else:
|
||||
XLA_AVAILABLE = False
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
|
||||
def calculate_shift(
|
||||
image_seq_len,
|
||||
base_seq_len: int = 256,
|
||||
max_seq_len: int = 4096,
|
||||
base_shift: float = 0.5,
|
||||
max_shift: float = 1.16,
|
||||
):
|
||||
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
|
||||
b = base_shift - m * base_seq_len
|
||||
mu = image_seq_len * m + b
|
||||
return mu
|
||||
|
||||
|
||||
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
|
||||
def retrieve_timesteps(
|
||||
scheduler,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
device: Optional[Union[str, torch.device]] = None,
|
||||
timesteps: Optional[List[int]] = None,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
|
||||
custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
|
||||
|
||||
Args:
|
||||
scheduler (`SchedulerMixin`):
|
||||
The scheduler to get timesteps from.
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
|
||||
must be `None`.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
timesteps (`List[int]`, *optional*):
|
||||
Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
|
||||
`num_inference_steps` and `sigmas` must be `None`.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
|
||||
`num_inference_steps` and `timesteps` must be `None`.
|
||||
|
||||
Returns:
|
||||
`Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
|
||||
second element is the number of inference steps.
|
||||
"""
|
||||
if timesteps is not None and sigmas is not None:
|
||||
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
|
||||
if timesteps is not None:
|
||||
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
if not accepts_timesteps:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" timestep schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
elif sigmas is not None:
|
||||
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
if not accept_sigmas:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" sigmas schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
else:
|
||||
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
return timesteps, num_inference_steps
|
||||
|
||||
|
||||
class Flex2Pipeline(
|
||||
DiffusionPipeline,
|
||||
FluxLoraLoaderMixin,
|
||||
FromSingleFileMixin,
|
||||
TextualInversionLoaderMixin,
|
||||
FluxIPAdapterMixin,
|
||||
):
|
||||
r"""
|
||||
The Flux pipeline for text-to-image generation.
|
||||
|
||||
Reference: https://blackforestlabs.ai/announcing-black-forest-labs/
|
||||
|
||||
Args:
|
||||
transformer ([`FluxTransformer2DModel`]):
|
||||
Conditional Transformer (MMDiT) architecture to denoise the encoded image latents.
|
||||
scheduler ([`FlowMatchEulerDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
vae ([`AutoencoderKL`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
|
||||
text_encoder ([`CLIPTextModel`]):
|
||||
[CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically
|
||||
the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant.
|
||||
text_encoder_2 ([`T5EncoderModel`]):
|
||||
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
|
||||
the [google/t5-v1_1-xxl](https://huggingface.co/google/t5-v1_1-xxl) variant.
|
||||
tokenizer (`CLIPTokenizer`):
|
||||
Tokenizer of class
|
||||
[CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
|
||||
tokenizer_2 (`T5TokenizerFast`):
|
||||
Second Tokenizer of class
|
||||
[T5TokenizerFast](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5TokenizerFast).
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->text_encoder_2->image_encoder->transformer->vae"
|
||||
_optional_components = ["image_encoder", "feature_extractor"]
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
scheduler: FlowMatchEulerDiscreteScheduler,
|
||||
vae: AutoencoderKL,
|
||||
text_encoder: CLIPTextModel,
|
||||
tokenizer: CLIPTokenizer,
|
||||
text_encoder_2: AutoModel,
|
||||
tokenizer_2: AutoTokenizer,
|
||||
transformer: FluxTransformer2DModel,
|
||||
image_encoder: CLIPVisionModelWithProjection = None,
|
||||
feature_extractor: CLIPImageProcessor = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
text_encoder_2=text_encoder_2,
|
||||
tokenizer=tokenizer,
|
||||
tokenizer_2=tokenizer_2,
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
image_encoder=image_encoder,
|
||||
feature_extractor=feature_extractor,
|
||||
)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) if getattr(self, "vae", None) else 8
|
||||
# Flux latents are turned into 2x2 patches and packed. This means the latent width and height has to be divisible
|
||||
# by the patch size. So the vae scale factor is multiplied by the patch size to account for this
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor * 2)
|
||||
self.tokenizer_max_length = (
|
||||
self.tokenizer.model_max_length if hasattr(self, "tokenizer") and self.tokenizer is not None else 77
|
||||
)
|
||||
self.default_sample_size = 128
|
||||
self.system_prompt = "You are an assistant designed to generate superior images with the superior degree of image-text alignment based on textual prompts or user prompts. <Prompt Start> "
|
||||
|
||||
# determine length of system prompt
|
||||
self.system_prompt_length = self.tokenizer_2(
|
||||
[self.system_prompt],
|
||||
padding="longest",
|
||||
return_tensors="pt",
|
||||
).input_ids[0].shape[0]
|
||||
|
||||
def _get_clip_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
num_images_per_prompt: int = 1,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
if isinstance(self, TextualInversionLoaderMixin):
|
||||
prompt = self.maybe_convert_prompt(prompt, self.tokenizer)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.tokenizer_max_length,
|
||||
truncation=True,
|
||||
return_overflowing_tokens=False,
|
||||
return_length=False,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.tokenizer_max_length - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because CLIP can only handle sequences up to"
|
||||
f" {self.tokenizer_max_length} tokens: {removed_text}"
|
||||
)
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), output_hidden_states=False)
|
||||
|
||||
# Use pooled output of CLIPTextModel
|
||||
prompt_embeds = prompt_embeds.pooler_output
|
||||
prompt_embeds = prompt_embeds.to(dtype=self.text_encoder.dtype, device=device)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, -1)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def _get_llm_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
num_images_per_prompt: int = 1,
|
||||
max_sequence_length: int = 512,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
if isinstance(self, TextualInversionLoaderMixin):
|
||||
prompt = self.maybe_convert_prompt(prompt, self.tokenizer_2)
|
||||
|
||||
text_inputs = self.tokenizer_2(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length + self.system_prompt_length,
|
||||
truncation=True,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
text_input_ids = text_inputs.input_ids.to(device)
|
||||
prompt_attention_mask = text_inputs.attention_mask.to(device)
|
||||
untruncated_ids = self.tokenizer_2(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer_2.batch_decode(untruncated_ids[:, self.tokenizer_max_length - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because `max_sequence_length` is set to "
|
||||
f" {max_sequence_length + self.system_prompt_length} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
prompt_embeds = self.text_encoder_2(
|
||||
text_input_ids,
|
||||
attention_mask=prompt_attention_mask,
|
||||
output_hidden_states=True
|
||||
)
|
||||
prompt_embeds = prompt_embeds.hidden_states[-1]
|
||||
|
||||
# remove the system prompt from the input and attention mask
|
||||
prompt_embeds = prompt_embeds[:, self.system_prompt_length:]
|
||||
prompt_attention_mask = prompt_attention_mask[:, self.system_prompt_length:]
|
||||
|
||||
dtype = self.text_encoder_2.dtype
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
|
||||
# duplicate text embeddings and attention mask for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
prompt_2: Union[str, List[str]],
|
||||
device: Optional[torch.device] = None,
|
||||
num_images_per_prompt: int = 1,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
max_sequence_length: int = 512,
|
||||
lora_scale: Optional[float] = None,
|
||||
):
|
||||
r"""
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
prompt to be encoded
|
||||
prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to be sent to the `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||
used in all text-encoders
|
||||
device: (`torch.device`):
|
||||
torch device
|
||||
num_images_per_prompt (`int`):
|
||||
number of images that should be generated per prompt
|
||||
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||
lora_scale (`float`, *optional*):
|
||||
A lora scale that will be applied to all LoRA layers of the text encoder if LoRA layers are loaded.
|
||||
"""
|
||||
device = device or self._execution_device
|
||||
|
||||
# set lora scale so that monkey patched LoRA
|
||||
# function of text encoder can correctly access it
|
||||
if lora_scale is not None and isinstance(self, FluxLoraLoaderMixin):
|
||||
self._lora_scale = lora_scale
|
||||
|
||||
# dynamically adjust the LoRA scale
|
||||
if self.text_encoder is not None and USE_PEFT_BACKEND:
|
||||
scale_lora_layers(self.text_encoder, lora_scale)
|
||||
if self.text_encoder_2 is not None and USE_PEFT_BACKEND:
|
||||
scale_lora_layers(self.text_encoder_2, lora_scale)
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_2 = prompt_2 or prompt
|
||||
prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2
|
||||
|
||||
# We only use the pooled prompt output from the CLIPTextModel
|
||||
pooled_prompt_embeds = self._get_clip_prompt_embeds(
|
||||
prompt=prompt,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
)
|
||||
prompt_embeds = self._get_llm_prompt_embeds(
|
||||
prompt=prompt_2,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
)
|
||||
|
||||
if self.text_encoder is not None:
|
||||
if isinstance(self, FluxLoraLoaderMixin) and USE_PEFT_BACKEND:
|
||||
# Retrieve the original scale by scaling back the LoRA layers
|
||||
unscale_lora_layers(self.text_encoder, lora_scale)
|
||||
|
||||
if self.text_encoder_2 is not None:
|
||||
if isinstance(self, FluxLoraLoaderMixin) and USE_PEFT_BACKEND:
|
||||
# Retrieve the original scale by scaling back the LoRA layers
|
||||
unscale_lora_layers(self.text_encoder_2, lora_scale)
|
||||
|
||||
dtype = self.text_encoder.dtype if self.text_encoder is not None else self.transformer.dtype
|
||||
text_ids = torch.zeros(prompt_embeds.shape[1], 3).to(device=device, dtype=dtype)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds, text_ids
|
||||
|
||||
def encode_image(self, image, device, num_images_per_prompt):
|
||||
dtype = next(self.image_encoder.parameters()).dtype
|
||||
|
||||
if not isinstance(image, torch.Tensor):
|
||||
image = self.feature_extractor(image, return_tensors="pt").pixel_values
|
||||
|
||||
image = image.to(device=device, dtype=dtype)
|
||||
image_embeds = self.image_encoder(image).image_embeds
|
||||
image_embeds = image_embeds.repeat_interleave(num_images_per_prompt, dim=0)
|
||||
return image_embeds
|
||||
|
||||
def prepare_ip_adapter_image_embeds(
|
||||
self, ip_adapter_image, ip_adapter_image_embeds, device, num_images_per_prompt
|
||||
):
|
||||
image_embeds = []
|
||||
if ip_adapter_image_embeds is None:
|
||||
if not isinstance(ip_adapter_image, list):
|
||||
ip_adapter_image = [ip_adapter_image]
|
||||
|
||||
if len(ip_adapter_image) != len(self.transformer.encoder_hid_proj.image_projection_layers):
|
||||
raise ValueError(
|
||||
f"`ip_adapter_image` must have same length as the number of IP Adapters. Got {len(ip_adapter_image)} images and {len(self.transformer.encoder_hid_proj.image_projection_layers)} IP Adapters."
|
||||
)
|
||||
|
||||
for single_ip_adapter_image, image_proj_layer in zip(
|
||||
ip_adapter_image, self.transformer.encoder_hid_proj.image_projection_layers
|
||||
):
|
||||
single_image_embeds = self.encode_image(single_ip_adapter_image, device, 1)
|
||||
|
||||
image_embeds.append(single_image_embeds[None, :])
|
||||
else:
|
||||
for single_image_embeds in ip_adapter_image_embeds:
|
||||
image_embeds.append(single_image_embeds)
|
||||
|
||||
ip_adapter_image_embeds = []
|
||||
for i, single_image_embeds in enumerate(image_embeds):
|
||||
single_image_embeds = torch.cat([single_image_embeds] * num_images_per_prompt, dim=0)
|
||||
single_image_embeds = single_image_embeds.to(device=device)
|
||||
ip_adapter_image_embeds.append(single_image_embeds)
|
||||
|
||||
return ip_adapter_image_embeds
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
negative_prompt=None,
|
||||
negative_prompt_2=None,
|
||||
prompt_embeds=None,
|
||||
negative_prompt_embeds=None,
|
||||
pooled_prompt_embeds=None,
|
||||
negative_pooled_prompt_embeds=None,
|
||||
callback_on_step_end_tensor_inputs=None,
|
||||
max_sequence_length=None,
|
||||
):
|
||||
if height % (self.vae_scale_factor * 2) != 0 or width % (self.vae_scale_factor * 2) != 0:
|
||||
logger.warning(
|
||||
f"`height` and `width` have to be divisible by {self.vae_scale_factor * 2} but are {height} and {width}. Dimensions will be resized accordingly"
|
||||
)
|
||||
|
||||
if callback_on_step_end_tensor_inputs is not None and not all(
|
||||
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
|
||||
)
|
||||
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt_2 is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt_2`: {prompt_2} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt is None and prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
|
||||
)
|
||||
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
elif prompt_2 is not None and (not isinstance(prompt_2, str) and not isinstance(prompt_2, list)):
|
||||
raise ValueError(f"`prompt_2` has to be of type `str` or `list` but is {type(prompt_2)}")
|
||||
|
||||
if negative_prompt is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
|
||||
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
|
||||
)
|
||||
elif negative_prompt_2 is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `negative_prompt_2`: {negative_prompt_2} and `negative_prompt_embeds`:"
|
||||
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
|
||||
)
|
||||
|
||||
if prompt_embeds is not None and negative_prompt_embeds is not None:
|
||||
if prompt_embeds.shape != negative_prompt_embeds.shape:
|
||||
raise ValueError(
|
||||
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
|
||||
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
|
||||
f" {negative_prompt_embeds.shape}."
|
||||
)
|
||||
|
||||
if prompt_embeds is not None and pooled_prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"If `prompt_embeds` are provided, `pooled_prompt_embeds` also have to be passed. Make sure to generate `pooled_prompt_embeds` from the same text encoder that was used to generate `prompt_embeds`."
|
||||
)
|
||||
if negative_prompt_embeds is not None and negative_pooled_prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"If `negative_prompt_embeds` are provided, `negative_pooled_prompt_embeds` also have to be passed. Make sure to generate `negative_pooled_prompt_embeds` from the same text encoder that was used to generate `negative_prompt_embeds`."
|
||||
)
|
||||
|
||||
if max_sequence_length is not None and max_sequence_length > 512:
|
||||
raise ValueError(f"`max_sequence_length` cannot be greater than 512 but is {max_sequence_length}")
|
||||
|
||||
@staticmethod
|
||||
def _prepare_latent_image_ids(batch_size, height, width, device, dtype):
|
||||
latent_image_ids = torch.zeros(height, width, 3)
|
||||
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None]
|
||||
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :]
|
||||
|
||||
latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape
|
||||
|
||||
latent_image_ids = latent_image_ids.reshape(
|
||||
latent_image_id_height * latent_image_id_width, latent_image_id_channels
|
||||
)
|
||||
|
||||
return latent_image_ids.to(device=device, dtype=dtype)
|
||||
|
||||
@staticmethod
|
||||
def _pack_latents(latents, batch_size, num_channels_latents, height, width):
|
||||
latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2)
|
||||
latents = latents.permute(0, 2, 4, 1, 3, 5)
|
||||
latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4)
|
||||
|
||||
return latents
|
||||
|
||||
@staticmethod
|
||||
def _unpack_latents(latents, height, width, vae_scale_factor):
|
||||
batch_size, num_patches, channels = latents.shape
|
||||
|
||||
# VAE applies 8x compression on images but we must also account for packing which requires
|
||||
# latent height and width to be divisible by 2.
|
||||
height = 2 * (int(height) // (vae_scale_factor * 2))
|
||||
width = 2 * (int(width) // (vae_scale_factor * 2))
|
||||
|
||||
latents = latents.view(batch_size, height // 2, width // 2, channels // 4, 2, 2)
|
||||
latents = latents.permute(0, 3, 1, 4, 2, 5)
|
||||
|
||||
latents = latents.reshape(batch_size, channels // (2 * 2), height, width)
|
||||
|
||||
return latents
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
r"""
|
||||
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
|
||||
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
|
||||
"""
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def disable_vae_slicing(self):
|
||||
r"""
|
||||
Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_slicing()
|
||||
|
||||
def enable_vae_tiling(self):
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
|
||||
processing larger images.
|
||||
"""
|
||||
self.vae.enable_tiling()
|
||||
|
||||
def disable_vae_tiling(self):
|
||||
r"""
|
||||
Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_tiling()
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
dtype,
|
||||
device,
|
||||
generator,
|
||||
latents=None,
|
||||
):
|
||||
# VAE applies 8x compression on images but we must also account for packing which requires
|
||||
# latent height and width to be divisible by 2.
|
||||
height = 2 * (int(height) // (self.vae_scale_factor * 2))
|
||||
width = 2 * (int(width) // (self.vae_scale_factor * 2))
|
||||
|
||||
shape = (batch_size, num_channels_latents, height, width)
|
||||
|
||||
if latents is not None:
|
||||
latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype)
|
||||
return latents.to(device=device, dtype=dtype), latent_image_ids
|
||||
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
latents = self._pack_latents(latents, batch_size, num_channels_latents, height, width)
|
||||
|
||||
latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype)
|
||||
|
||||
return latents, latent_image_ids
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
@property
|
||||
def joint_attention_kwargs(self):
|
||||
return self._joint_attention_kwargs
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def current_timestep(self):
|
||||
return self._current_timestep
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt: Union[str, List[str]] = None,
|
||||
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
true_cfg_scale: float = 1.0,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 28,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
guidance_scale: float = 3.5,
|
||||
num_images_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
ip_adapter_image: Optional[PipelineImageInput] = None,
|
||||
ip_adapter_image_embeds: Optional[List[torch.Tensor]] = None,
|
||||
negative_ip_adapter_image: Optional[PipelineImageInput] = None,
|
||||
negative_ip_adapter_image_embeds: Optional[List[torch.Tensor]] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 512,
|
||||
):
|
||||
r"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||
will be used instead.
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `true_cfg_scale` is
|
||||
not greater than `1`).
|
||||
negative_prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation to be sent to `tokenizer_2` and
|
||||
`text_encoder_2`. If not defined, `negative_prompt` is used in all the text-encoders.
|
||||
true_cfg_scale (`float`, *optional*, defaults to 1.0):
|
||||
When > 1.0 and a provided `negative_prompt`, enables true classifier-free guidance.
|
||||
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The height in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The width in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in
|
||||
their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed
|
||||
will be used.
|
||||
guidance_scale (`float`, *optional*, defaults to 7.0):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||
usually at the expense of lower image quality.
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||
to make generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor will ge generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||
ip_adapter_image: (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters.
|
||||
ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*):
|
||||
Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of
|
||||
IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not
|
||||
provided, embeddings are computed from the `ip_adapter_image` input argument.
|
||||
negative_ip_adapter_image:
|
||||
(`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters.
|
||||
negative_ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*):
|
||||
Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of
|
||||
IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not
|
||||
provided, embeddings are computed from the `ip_adapter_image` input argument.
|
||||
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||
argument.
|
||||
negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, pooled negative_prompt_embeds will be generated from `negative_prompt`
|
||||
input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generate image. Choose between
|
||||
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.flux.FluxPipelineOutput`] instead of a plain tuple.
|
||||
joint_attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict`
|
||||
is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated
|
||||
images.
|
||||
"""
|
||||
|
||||
height = height or self.default_sample_size * self.vae_scale_factor
|
||||
width = width or self.default_sample_size * self.vae_scale_factor
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
negative_prompt=negative_prompt,
|
||||
negative_prompt_2=negative_prompt_2,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
|
||||
max_sequence_length=max_sequence_length,
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._joint_attention_kwargs = joint_attention_kwargs
|
||||
self._current_timestep = None
|
||||
self._interrupt = False
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
lora_scale = (
|
||||
self.joint_attention_kwargs.get("scale", None) if self.joint_attention_kwargs is not None else None
|
||||
)
|
||||
has_neg_prompt = negative_prompt is not None or (
|
||||
negative_prompt_embeds is not None and negative_pooled_prompt_embeds is not None
|
||||
)
|
||||
do_true_cfg = true_cfg_scale > 1 and has_neg_prompt
|
||||
(
|
||||
prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
text_ids,
|
||||
) = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
prompt_2=prompt_2,
|
||||
prompt_embeds=prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
lora_scale=lora_scale,
|
||||
)
|
||||
if do_true_cfg:
|
||||
(
|
||||
negative_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
_,
|
||||
) = self.encode_prompt(
|
||||
prompt=negative_prompt,
|
||||
prompt_2=negative_prompt_2,
|
||||
prompt_embeds=negative_prompt_embeds,
|
||||
pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
lora_scale=lora_scale,
|
||||
)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels // 4
|
||||
latents, latent_image_ids = self.prepare_latents(
|
||||
batch_size * num_images_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 5. Prepare timesteps
|
||||
sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas
|
||||
image_seq_len = latents.shape[1]
|
||||
mu = calculate_shift(
|
||||
image_seq_len,
|
||||
self.scheduler.config.get("base_image_seq_len", 256),
|
||||
self.scheduler.config.get("max_image_seq_len", 4096),
|
||||
self.scheduler.config.get("base_shift", 0.5),
|
||||
self.scheduler.config.get("max_shift", 1.16),
|
||||
)
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
sigmas=sigmas,
|
||||
mu=mu,
|
||||
)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# handle guidance
|
||||
if self.transformer.config.guidance_embeds:
|
||||
guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32)
|
||||
guidance = guidance.expand(latents.shape[0])
|
||||
else:
|
||||
guidance = None
|
||||
|
||||
if (ip_adapter_image is not None or ip_adapter_image_embeds is not None) and (
|
||||
negative_ip_adapter_image is None and negative_ip_adapter_image_embeds is None
|
||||
):
|
||||
negative_ip_adapter_image = np.zeros((width, height, 3), dtype=np.uint8)
|
||||
elif (ip_adapter_image is None and ip_adapter_image_embeds is None) and (
|
||||
negative_ip_adapter_image is not None or negative_ip_adapter_image_embeds is not None
|
||||
):
|
||||
ip_adapter_image = np.zeros((width, height, 3), dtype=np.uint8)
|
||||
|
||||
if self.joint_attention_kwargs is None:
|
||||
self._joint_attention_kwargs = {}
|
||||
|
||||
image_embeds = None
|
||||
negative_image_embeds = None
|
||||
if ip_adapter_image is not None or ip_adapter_image_embeds is not None:
|
||||
image_embeds = self.prepare_ip_adapter_image_embeds(
|
||||
ip_adapter_image,
|
||||
ip_adapter_image_embeds,
|
||||
device,
|
||||
batch_size * num_images_per_prompt,
|
||||
)
|
||||
if negative_ip_adapter_image is not None or negative_ip_adapter_image_embeds is not None:
|
||||
negative_image_embeds = self.prepare_ip_adapter_image_embeds(
|
||||
negative_ip_adapter_image,
|
||||
negative_ip_adapter_image_embeds,
|
||||
device,
|
||||
batch_size * num_images_per_prompt,
|
||||
)
|
||||
|
||||
# 6. Denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
self._current_timestep = t
|
||||
if image_embeds is not None:
|
||||
self._joint_attention_kwargs["ip_adapter_image_embeds"] = image_embeds
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latents.shape[0]).to(latents.dtype)
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latents,
|
||||
timestep=timestep / 1000,
|
||||
guidance=guidance,
|
||||
pooled_projections=pooled_prompt_embeds,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
txt_ids=text_ids,
|
||||
img_ids=latent_image_ids,
|
||||
joint_attention_kwargs=self.joint_attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if do_true_cfg:
|
||||
if negative_image_embeds is not None:
|
||||
self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds
|
||||
neg_noise_pred = self.transformer(
|
||||
hidden_states=latents,
|
||||
timestep=timestep / 1000,
|
||||
guidance=guidance,
|
||||
pooled_projections=negative_pooled_prompt_embeds,
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
txt_ids=text_ids,
|
||||
img_ids=latent_image_ids,
|
||||
joint_attention_kwargs=self.joint_attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents_dtype = latents.dtype
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
|
||||
if latents.dtype != latents_dtype:
|
||||
if torch.backends.mps.is_available():
|
||||
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
|
||||
latents = latents.to(latents_dtype)
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if XLA_AVAILABLE:
|
||||
xm.mark_step()
|
||||
|
||||
self._current_timestep = None
|
||||
|
||||
if output_type == "latent":
|
||||
image = latents
|
||||
else:
|
||||
latents = self._unpack_latents(latents, height, width, self.vae_scale_factor)
|
||||
latents = (latents / self.vae.config.scaling_factor) + self.vae.config.shift_factor
|
||||
image = self.vae.decode(latents, return_dict=False)[0]
|
||||
image = self.image_processor.postprocess(image, output_type=output_type)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (image,)
|
||||
|
||||
return FluxPipelineOutput(images=image)
|
||||
176
toolkit/models/flux.py
Normal file
176
toolkit/models/flux.py
Normal file
@@ -0,0 +1,176 @@
|
||||
|
||||
# forward that bypasses the guidance embedding so it can be avoided during training.
|
||||
from functools import partial
|
||||
from typing import Optional
|
||||
import torch
|
||||
from diffusers import FluxTransformer2DModel
|
||||
|
||||
|
||||
def guidance_embed_bypass_forward(self, timestep, guidance, pooled_projection):
|
||||
timesteps_proj = self.time_proj(timestep)
|
||||
timesteps_emb = self.timestep_embedder(
|
||||
timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D)
|
||||
pooled_projections = self.text_embedder(pooled_projection)
|
||||
conditioning = timesteps_emb + pooled_projections
|
||||
return conditioning
|
||||
|
||||
# bypass the forward function
|
||||
|
||||
|
||||
def bypass_flux_guidance(transformer):
|
||||
if hasattr(transformer.time_text_embed, '_bfg_orig_forward'):
|
||||
return
|
||||
# dont bypass if it doesnt have the guidance embedding
|
||||
if not hasattr(transformer.time_text_embed, 'guidance_embedder'):
|
||||
return
|
||||
transformer.time_text_embed._bfg_orig_forward = transformer.time_text_embed.forward
|
||||
transformer.time_text_embed.forward = partial(
|
||||
guidance_embed_bypass_forward, transformer.time_text_embed
|
||||
)
|
||||
|
||||
# restore the forward function
|
||||
|
||||
|
||||
def restore_flux_guidance(transformer):
|
||||
if not hasattr(transformer.time_text_embed, '_bfg_orig_forward'):
|
||||
return
|
||||
transformer.time_text_embed.forward = transformer.time_text_embed._bfg_orig_forward
|
||||
del transformer.time_text_embed._bfg_orig_forward
|
||||
|
||||
def new_device_to(self: FluxTransformer2DModel, *args, **kwargs):
|
||||
# Store original device if provided in args or kwargs
|
||||
device_in_kwargs = 'device' in kwargs
|
||||
device_in_args = any(isinstance(arg, (str, torch.device)) for arg in args)
|
||||
|
||||
device = None
|
||||
# Remove device from kwargs if present
|
||||
if device_in_kwargs:
|
||||
device = kwargs['device']
|
||||
del kwargs['device']
|
||||
|
||||
# Only filter args if we detected a device argument
|
||||
if device_in_args:
|
||||
args = list(args)
|
||||
for idx, arg in enumerate(args):
|
||||
if isinstance(arg, (str, torch.device)):
|
||||
device = arg
|
||||
del args[idx]
|
||||
|
||||
self.pos_embed = self.pos_embed.to(device, *args, **kwargs)
|
||||
self.time_text_embed = self.time_text_embed.to(device, *args, **kwargs)
|
||||
self.context_embedder = self.context_embedder.to(device, *args, **kwargs)
|
||||
self.x_embedder = self.x_embedder.to(device, *args, **kwargs)
|
||||
for block in self.transformer_blocks:
|
||||
block.to(block._split_device, *args, **kwargs)
|
||||
for block in self.single_transformer_blocks:
|
||||
block.to(block._split_device, *args, **kwargs)
|
||||
|
||||
self.norm_out = self.norm_out.to(device, *args, **kwargs)
|
||||
self.proj_out = self.proj_out.to(device, *args, **kwargs)
|
||||
|
||||
|
||||
|
||||
return self
|
||||
|
||||
|
||||
|
||||
|
||||
def split_gpu_double_block_forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: torch.FloatTensor,
|
||||
temb: torch.FloatTensor,
|
||||
image_rotary_emb=None,
|
||||
joint_attention_kwargs=None,
|
||||
):
|
||||
if hidden_states.device != self._split_device:
|
||||
hidden_states = hidden_states.to(self._split_device)
|
||||
if encoder_hidden_states.device != self._split_device:
|
||||
encoder_hidden_states = encoder_hidden_states.to(self._split_device)
|
||||
if temb.device != self._split_device:
|
||||
temb = temb.to(self._split_device)
|
||||
if image_rotary_emb is not None and image_rotary_emb[0].device != self._split_device:
|
||||
# is a tuple of tensors
|
||||
image_rotary_emb = tuple([t.to(self._split_device) for t in image_rotary_emb])
|
||||
return self._pre_gpu_split_forward(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
|
||||
|
||||
|
||||
def split_gpu_single_block_forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
temb: torch.FloatTensor,
|
||||
image_rotary_emb=None,
|
||||
joint_attention_kwargs=None,
|
||||
**kwargs
|
||||
):
|
||||
if hidden_states.device != self._split_device:
|
||||
hidden_states = hidden_states.to(device=self._split_device)
|
||||
if temb.device != self._split_device:
|
||||
temb = temb.to(device=self._split_device)
|
||||
if image_rotary_emb is not None and image_rotary_emb[0].device != self._split_device:
|
||||
# is a tuple of tensors
|
||||
image_rotary_emb = tuple([t.to(self._split_device) for t in image_rotary_emb])
|
||||
|
||||
hidden_state_out = self._pre_gpu_split_forward(hidden_states, temb, image_rotary_emb, joint_attention_kwargs, **kwargs)
|
||||
if hasattr(self, "_split_output_device"):
|
||||
return hidden_state_out.to(self._split_output_device)
|
||||
return hidden_state_out
|
||||
|
||||
|
||||
def add_model_gpu_splitter_to_flux(
|
||||
transformer: FluxTransformer2DModel,
|
||||
# ~ 5 billion for all other params
|
||||
other_module_params: Optional[int] = 5e9,
|
||||
# since they are not trainable, multiply by smaller number
|
||||
other_module_param_count_scale: Optional[float] = 0.3
|
||||
):
|
||||
gpu_id_list = [i for i in range(torch.cuda.device_count())]
|
||||
|
||||
# if len(gpu_id_list) > 2:
|
||||
# raise ValueError("Cannot split to more than 2 GPUs currently.")
|
||||
other_module_params *= other_module_param_count_scale
|
||||
|
||||
# since we are not tuning the
|
||||
total_params = sum(p.numel() for p in transformer.parameters()) + other_module_params
|
||||
|
||||
params_per_gpu = total_params / len(gpu_id_list)
|
||||
|
||||
current_gpu_idx = 0
|
||||
# text encoders, vae, and some non block layers will all be on gpu 0
|
||||
current_gpu_params = other_module_params
|
||||
|
||||
for double_block in transformer.transformer_blocks:
|
||||
device = torch.device(f"cuda:{current_gpu_idx}")
|
||||
double_block._pre_gpu_split_forward = double_block.forward
|
||||
double_block.forward = partial(
|
||||
split_gpu_double_block_forward, double_block)
|
||||
double_block._split_device = device
|
||||
# add the params to the current gpu
|
||||
current_gpu_params += sum(p.numel() for p in double_block.parameters())
|
||||
# if the current gpu params are greater than the params per gpu, move to next gpu
|
||||
if current_gpu_params > params_per_gpu:
|
||||
current_gpu_idx += 1
|
||||
current_gpu_params = 0
|
||||
if current_gpu_idx >= len(gpu_id_list):
|
||||
current_gpu_idx = gpu_id_list[-1]
|
||||
|
||||
for single_block in transformer.single_transformer_blocks:
|
||||
device = torch.device(f"cuda:{current_gpu_idx}")
|
||||
single_block._pre_gpu_split_forward = single_block.forward
|
||||
single_block.forward = partial(
|
||||
split_gpu_single_block_forward, single_block)
|
||||
single_block._split_device = device
|
||||
# add the params to the current gpu
|
||||
current_gpu_params += sum(p.numel() for p in single_block.parameters())
|
||||
# if the current gpu params are greater than the params per gpu, move to next gpu
|
||||
if current_gpu_params > params_per_gpu:
|
||||
current_gpu_idx += 1
|
||||
current_gpu_params = 0
|
||||
if current_gpu_idx >= len(gpu_id_list):
|
||||
current_gpu_idx = gpu_id_list[-1]
|
||||
|
||||
# add output device to last layer
|
||||
transformer.single_transformer_blocks[-1]._split_output_device = torch.device("cuda:0")
|
||||
|
||||
transformer._pre_gpu_split_to = transformer.to
|
||||
transformer.to = partial(new_device_to, transformer)
|
||||
94
toolkit/models/flux_sage_attn.py
Normal file
94
toolkit/models/flux_sage_attn.py
Normal file
@@ -0,0 +1,94 @@
|
||||
from typing import Optional
|
||||
from diffusers.models.attention_processor import Attention
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class FluxSageAttnProcessor2_0:
|
||||
"""Attention processor used typically in processing the SD3-like self-attention projections."""
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("FluxAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: torch.FloatTensor = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
from sageattention import sageattn
|
||||
|
||||
batch_size, _, _ = hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
|
||||
# `sample` projections.
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(hidden_states)
|
||||
value = attn.to_v(hidden_states)
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
# the attention in FluxSingleTransformerBlock does not use `encoder_hidden_states`
|
||||
if encoder_hidden_states is not None:
|
||||
# `context` projections.
|
||||
encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states)
|
||||
|
||||
encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
|
||||
if attn.norm_added_q is not None:
|
||||
encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj)
|
||||
if attn.norm_added_k is not None:
|
||||
encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj)
|
||||
|
||||
# attention
|
||||
query = torch.cat([encoder_hidden_states_query_proj, query], dim=2)
|
||||
key = torch.cat([encoder_hidden_states_key_proj, key], dim=2)
|
||||
value = torch.cat([encoder_hidden_states_value_proj, value], dim=2)
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
from diffusers.models.embeddings import apply_rotary_emb
|
||||
|
||||
query = apply_rotary_emb(query, image_rotary_emb)
|
||||
key = apply_rotary_emb(key, image_rotary_emb)
|
||||
|
||||
hidden_states = sageattn(query, key, value, dropout_p=0.0, is_causal=False)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
if encoder_hidden_states is not None:
|
||||
encoder_hidden_states, hidden_states = (
|
||||
hidden_states[:, : encoder_hidden_states.shape[1]],
|
||||
hidden_states[:, encoder_hidden_states.shape[1] :],
|
||||
)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
else:
|
||||
return hidden_states
|
||||
364
toolkit/models/ilora.py
Normal file
364
toolkit/models/ilora.py
Normal file
@@ -0,0 +1,364 @@
|
||||
import math
|
||||
import weakref
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import TYPE_CHECKING, List, Dict, Any
|
||||
from toolkit.models.clip_fusion import ZipperBlock
|
||||
from toolkit.models.zipper_resampler import ZipperModule, ZipperResampler
|
||||
import sys
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from ipadapter.ip_adapter.resampler import Resampler
|
||||
from collections import OrderedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRAModule
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, hidden_dim, dropout=0.1, use_residual=True):
|
||||
super().__init__()
|
||||
if use_residual:
|
||||
assert in_dim == out_dim
|
||||
self.layernorm = nn.LayerNorm(in_dim)
|
||||
self.fc1 = nn.Linear(in_dim, hidden_dim)
|
||||
self.fc2 = nn.Linear(hidden_dim, out_dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.use_residual = use_residual
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.layernorm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
x = self.dropout(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
class LoRAGenerator(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int = 768, # projection dimension
|
||||
hidden_size: int = 768,
|
||||
head_size: int = 512,
|
||||
num_heads: int = 1,
|
||||
num_mlp_layers: int = 1,
|
||||
output_size: int = 768,
|
||||
dropout: float = 0.0
|
||||
):
|
||||
super().__init__()
|
||||
self.input_size = input_size
|
||||
self.num_heads = num_heads
|
||||
self.simple = False
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
if self.simple:
|
||||
self.head = nn.Linear(input_size, head_size, bias=False)
|
||||
else:
|
||||
self.lin_in = nn.Linear(input_size, hidden_size)
|
||||
|
||||
self.mlp_blocks = nn.Sequential(*[
|
||||
MLP(hidden_size, hidden_size, hidden_size, dropout=dropout, use_residual=True) for _ in range(num_mlp_layers)
|
||||
])
|
||||
self.head = nn.Linear(hidden_size, head_size, bias=False)
|
||||
self.norm = nn.LayerNorm(head_size)
|
||||
|
||||
if num_heads == 1:
|
||||
self.output = nn.Linear(head_size, self.output_size)
|
||||
# for each output block. multiply weights by 0.01
|
||||
with torch.no_grad():
|
||||
self.output.weight.data *= 0.01
|
||||
else:
|
||||
head_output_size = output_size // num_heads
|
||||
self.outputs = nn.ModuleList([nn.Linear(head_size, head_output_size) for _ in range(num_heads)])
|
||||
# for each output block. multiply weights by 0.01
|
||||
with torch.no_grad():
|
||||
for output in self.outputs:
|
||||
output.weight.data *= 0.01
|
||||
|
||||
# allow get device
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
def forward(self, embedding):
|
||||
if len(embedding.shape) == 2:
|
||||
embedding = embedding.unsqueeze(1)
|
||||
|
||||
x = embedding
|
||||
|
||||
if not self.simple:
|
||||
x = self.lin_in(embedding)
|
||||
x = self.mlp_blocks(x)
|
||||
x = self.head(x)
|
||||
x = self.norm(x)
|
||||
|
||||
if self.num_heads == 1:
|
||||
x = self.output(x)
|
||||
else:
|
||||
out_chunks = torch.chunk(x, self.num_heads, dim=1)
|
||||
x = []
|
||||
for out_layer, chunk in zip(self.outputs, out_chunks):
|
||||
x.append(out_layer(chunk))
|
||||
x = torch.cat(x, dim=-1)
|
||||
|
||||
return x.squeeze(1)
|
||||
|
||||
|
||||
class InstantLoRAMidModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
index: int,
|
||||
lora_module: 'LoRAModule',
|
||||
instant_lora_module: 'InstantLoRAModule',
|
||||
up_shape: list = None,
|
||||
down_shape: list = None,
|
||||
):
|
||||
super(InstantLoRAMidModule, self).__init__()
|
||||
self.up_shape = up_shape
|
||||
self.down_shape = down_shape
|
||||
self.index = index
|
||||
self.lora_module_ref = weakref.ref(lora_module)
|
||||
self.instant_lora_module_ref = weakref.ref(instant_lora_module)
|
||||
|
||||
self.embed = None
|
||||
|
||||
def down_forward(self, x, *args, **kwargs):
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
if x.dtype != self.embed.dtype:
|
||||
x = x.to(self.embed.dtype)
|
||||
down_size = math.prod(self.down_shape)
|
||||
down_weight = self.embed[:, :down_size]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if down_weight.shape[0] * 2 == batch_size:
|
||||
down_weight = torch.cat([down_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(down_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
weight_chunk = weight_chunk.view(self.down_shape)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
org_module = self.lora_module_ref().orig_module_ref()
|
||||
stride = org_module.stride
|
||||
padding = org_module.padding
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding, stride=stride)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
def up_forward(self, x, *args, **kwargs):
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
if x.dtype != self.embed.dtype:
|
||||
x = x.to(self.embed.dtype)
|
||||
up_size = math.prod(self.up_shape)
|
||||
up_weight = self.embed[:, -up_size:]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if up_weight.shape[0] * 2 == batch_size:
|
||||
up_weight = torch.cat([up_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(up_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
weight_chunk = weight_chunk.view(self.up_shape)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
padding = 0
|
||||
if weight_chunk.shape[-1] == 3:
|
||||
padding = 1
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
|
||||
|
||||
class InstantLoRAModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
vision_hidden_size: int,
|
||||
vision_tokens: int,
|
||||
head_dim: int,
|
||||
num_heads: int, # number of heads in the resampler
|
||||
sd: 'StableDiffusion',
|
||||
config=None
|
||||
):
|
||||
super(InstantLoRAModule, self).__init__()
|
||||
# self.linear = torch.nn.Linear(2, 1)
|
||||
self.sd_ref = weakref.ref(sd)
|
||||
self.dim = sd.network.lora_dim
|
||||
self.vision_hidden_size = vision_hidden_size
|
||||
self.vision_tokens = vision_tokens
|
||||
self.head_dim = head_dim
|
||||
self.num_heads = num_heads
|
||||
|
||||
# stores the projection vector. Grabbed by modules
|
||||
self.img_embeds: List[torch.Tensor] = None
|
||||
|
||||
# disable merging in. It is slower on inference
|
||||
self.sd_ref().network.can_merge_in = False
|
||||
|
||||
self.ilora_modules = torch.nn.ModuleList()
|
||||
|
||||
lora_modules = self.sd_ref().network.get_all_modules()
|
||||
|
||||
output_size = 0
|
||||
|
||||
self.embed_lengths = []
|
||||
self.weight_mapping = []
|
||||
|
||||
for idx, lora_module in enumerate(lora_modules):
|
||||
module_dict = lora_module.state_dict()
|
||||
down_shape = list(module_dict['lora_down.weight'].shape)
|
||||
up_shape = list(module_dict['lora_up.weight'].shape)
|
||||
|
||||
self.weight_mapping.append([lora_module.lora_name, [down_shape, up_shape]])
|
||||
|
||||
module_size = math.prod(down_shape) + math.prod(up_shape)
|
||||
output_size += module_size
|
||||
self.embed_lengths.append(module_size)
|
||||
|
||||
|
||||
# add a new mid module that will take the original forward and add a vector to it
|
||||
# this will be used to add the vector to the original forward
|
||||
instant_module = InstantLoRAMidModule(
|
||||
idx,
|
||||
lora_module,
|
||||
self,
|
||||
up_shape=up_shape,
|
||||
down_shape=down_shape
|
||||
)
|
||||
|
||||
self.ilora_modules.append(instant_module)
|
||||
|
||||
# replace the LoRA forwards
|
||||
lora_module.lora_down.forward = instant_module.down_forward
|
||||
lora_module.lora_up.forward = instant_module.up_forward
|
||||
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
number_formatted_output_size = "{:,}".format(output_size)
|
||||
|
||||
print(f" ILORA output size: {number_formatted_output_size}")
|
||||
|
||||
# if not evenly divisible, error
|
||||
if self.output_size % self.num_heads != 0:
|
||||
raise ValueError("Output size must be divisible by the number of heads")
|
||||
|
||||
self.head_output_size = self.output_size // self.num_heads
|
||||
|
||||
if vision_tokens > 1:
|
||||
self.resampler = Resampler(
|
||||
dim=vision_hidden_size,
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=12,
|
||||
num_queries=num_heads, # output tokens
|
||||
embedding_dim=vision_hidden_size,
|
||||
max_seq_len=vision_tokens,
|
||||
output_dim=head_dim,
|
||||
apply_pos_emb=True, # this is new
|
||||
ff_mult=4
|
||||
)
|
||||
|
||||
self.proj_module = LoRAGenerator(
|
||||
input_size=head_dim,
|
||||
hidden_size=head_dim,
|
||||
head_size=head_dim,
|
||||
num_mlp_layers=1,
|
||||
num_heads=self.num_heads,
|
||||
output_size=self.output_size,
|
||||
)
|
||||
|
||||
self.migrate_weight_mapping()
|
||||
|
||||
def migrate_weight_mapping(self):
|
||||
return
|
||||
# # changes the names of the modules to common ones
|
||||
# keymap = self.sd_ref().network.get_keymap()
|
||||
# save_keymap = {}
|
||||
# if keymap is not None:
|
||||
# for ldm_key, diffusers_key in keymap.items():
|
||||
# # invert them
|
||||
# save_keymap[diffusers_key] = ldm_key
|
||||
#
|
||||
# new_keymap = {}
|
||||
# for key, value in self.weight_mapping:
|
||||
# if key in save_keymap:
|
||||
# new_keymap[save_keymap[key]] = value
|
||||
# else:
|
||||
# print(f"Key {key} not found in keymap")
|
||||
# new_keymap[key] = value
|
||||
# self.weight_mapping = new_keymap
|
||||
# else:
|
||||
# print("No keymap found. Using default names")
|
||||
# return
|
||||
|
||||
|
||||
def forward(self, img_embeds):
|
||||
# expand token rank if only rank 2
|
||||
if len(img_embeds.shape) == 2:
|
||||
img_embeds = img_embeds.unsqueeze(1)
|
||||
|
||||
# resample the image embeddings
|
||||
img_embeds = self.resampler(img_embeds)
|
||||
img_embeds = self.proj_module(img_embeds)
|
||||
if len(img_embeds.shape) == 3:
|
||||
# merge the heads
|
||||
img_embeds = img_embeds.mean(dim=1)
|
||||
|
||||
self.img_embeds = []
|
||||
# get all the slices
|
||||
start = 0
|
||||
for length in self.embed_lengths:
|
||||
self.img_embeds.append(img_embeds[:, start:start+length])
|
||||
start += length
|
||||
|
||||
|
||||
def get_additional_save_metadata(self) -> Dict[str, Any]:
|
||||
# save the weight mapping
|
||||
return {
|
||||
"weight_mapping": self.weight_mapping,
|
||||
"num_heads": self.num_heads,
|
||||
"vision_hidden_size": self.vision_hidden_size,
|
||||
"head_dim": self.head_dim,
|
||||
"vision_tokens": self.vision_tokens,
|
||||
"output_size": self.output_size,
|
||||
}
|
||||
|
||||
419
toolkit/models/ilora2.py
Normal file
419
toolkit/models/ilora2.py
Normal file
@@ -0,0 +1,419 @@
|
||||
import math
|
||||
import weakref
|
||||
|
||||
from toolkit.config_modules import AdapterConfig
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import TYPE_CHECKING, List, Dict, Any
|
||||
from toolkit.models.clip_fusion import ZipperBlock
|
||||
from toolkit.models.zipper_resampler import ZipperModule, ZipperResampler
|
||||
import sys
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from ipadapter.ip_adapter.resampler import Resampler
|
||||
from collections import OrderedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRAModule
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, hidden_dim, dropout=0.1, use_residual=True):
|
||||
super().__init__()
|
||||
if use_residual:
|
||||
assert in_dim == out_dim
|
||||
self.layernorm = nn.LayerNorm(in_dim)
|
||||
self.fc1 = nn.Linear(in_dim, hidden_dim)
|
||||
self.fc2 = nn.Linear(hidden_dim, out_dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.use_residual = use_residual
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.layernorm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
x = self.dropout(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
|
||||
class LoRAGenerator(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int = 768, # projection dimension
|
||||
hidden_size: int = 768,
|
||||
head_size: int = 512,
|
||||
num_heads: int = 1,
|
||||
num_mlp_layers: int = 1,
|
||||
output_size: int = 768,
|
||||
dropout: float = 0.0
|
||||
):
|
||||
super().__init__()
|
||||
self.input_size = input_size
|
||||
self.num_heads = num_heads
|
||||
self.simple = False
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
if self.simple:
|
||||
self.head = nn.Linear(input_size, head_size, bias=False)
|
||||
else:
|
||||
self.lin_in = nn.Linear(input_size, hidden_size)
|
||||
|
||||
self.mlp_blocks = nn.Sequential(*[
|
||||
MLP(hidden_size, hidden_size, hidden_size, dropout=dropout, use_residual=True) for _ in
|
||||
range(num_mlp_layers)
|
||||
])
|
||||
self.head = nn.Linear(hidden_size, head_size, bias=False)
|
||||
self.norm = nn.LayerNorm(head_size)
|
||||
|
||||
if num_heads == 1:
|
||||
self.output = nn.Linear(head_size, self.output_size)
|
||||
# for each output block. multiply weights by 0.01
|
||||
with torch.no_grad():
|
||||
self.output.weight.data *= 0.01
|
||||
else:
|
||||
head_output_size = output_size // num_heads
|
||||
self.outputs = nn.ModuleList([nn.Linear(head_size, head_output_size) for _ in range(num_heads)])
|
||||
# for each output block. multiply weights by 0.01
|
||||
with torch.no_grad():
|
||||
for output in self.outputs:
|
||||
output.weight.data *= 0.01
|
||||
|
||||
# allow get device
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
def forward(self, embedding):
|
||||
if len(embedding.shape) == 2:
|
||||
embedding = embedding.unsqueeze(1)
|
||||
|
||||
x = embedding
|
||||
|
||||
if not self.simple:
|
||||
x = self.lin_in(embedding)
|
||||
x = self.mlp_blocks(x)
|
||||
x = self.head(x)
|
||||
x = self.norm(x)
|
||||
|
||||
if self.num_heads == 1:
|
||||
x = self.output(x)
|
||||
else:
|
||||
out_chunks = torch.chunk(x, self.num_heads, dim=1)
|
||||
x = []
|
||||
for out_layer, chunk in zip(self.outputs, out_chunks):
|
||||
x.append(out_layer(chunk))
|
||||
x = torch.cat(x, dim=-1)
|
||||
|
||||
return x.squeeze(1)
|
||||
|
||||
|
||||
class InstantLoRAMidModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
index: int,
|
||||
lora_module: 'LoRAModule',
|
||||
instant_lora_module: 'InstantLoRAModule',
|
||||
up_shape: list = None,
|
||||
down_shape: list = None,
|
||||
):
|
||||
super(InstantLoRAMidModule, self).__init__()
|
||||
self.up_shape = up_shape
|
||||
self.down_shape = down_shape
|
||||
self.index = index
|
||||
self.lora_module_ref = weakref.ref(lora_module)
|
||||
self.instant_lora_module_ref = weakref.ref(instant_lora_module)
|
||||
|
||||
self.do_up = instant_lora_module.config.ilora_up
|
||||
self.do_down = instant_lora_module.config.ilora_down
|
||||
self.do_mid = instant_lora_module.config.ilora_mid
|
||||
|
||||
self.down_dim = self.down_shape[1] if self.do_down else 0
|
||||
self.mid_dim = self.up_shape[1] if self.do_mid else 0
|
||||
self.out_dim = self.up_shape[0] if self.do_up else 0
|
||||
|
||||
self.embed = None
|
||||
|
||||
def down_forward(self, x, *args, **kwargs):
|
||||
if not self.do_down:
|
||||
return self.lora_module_ref().lora_down.orig_forward(x, *args, **kwargs)
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
down_weight = self.embed[:, :self.down_dim]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if down_weight.shape[0] * 2 == batch_size:
|
||||
down_weight = torch.cat([down_weight] * 2, dim=0)
|
||||
|
||||
try:
|
||||
if len(x.shape) == 4:
|
||||
# conv
|
||||
down_weight = down_weight.view(batch_size, -1, 1, 1)
|
||||
if x.shape[1] != down_weight.shape[1]:
|
||||
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
|
||||
elif len(x.shape) == 2:
|
||||
down_weight = down_weight.view(batch_size, -1)
|
||||
if x.shape[1] != down_weight.shape[1]:
|
||||
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
|
||||
else:
|
||||
down_weight = down_weight.view(batch_size, 1, -1)
|
||||
if x.shape[2] != down_weight.shape[2]:
|
||||
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
|
||||
x = x * down_weight
|
||||
x = self.lora_module_ref().lora_down.orig_forward(x, *args, **kwargs)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
|
||||
|
||||
return x
|
||||
|
||||
def up_forward(self, x, *args, **kwargs):
|
||||
# do mid here
|
||||
x = self.mid_forward(x, *args, **kwargs)
|
||||
if not self.do_up:
|
||||
return self.lora_module_ref().lora_up.orig_forward(x, *args, **kwargs)
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
up_weight = self.embed[:, -self.out_dim:]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if up_weight.shape[0] * 2 == batch_size:
|
||||
up_weight = torch.cat([up_weight] * 2, dim=0)
|
||||
|
||||
try:
|
||||
if len(x.shape) == 4:
|
||||
# conv
|
||||
up_weight = up_weight.view(batch_size, -1, 1, 1)
|
||||
elif len(x.shape) == 2:
|
||||
up_weight = up_weight.view(batch_size, -1)
|
||||
else:
|
||||
up_weight = up_weight.view(batch_size, 1, -1)
|
||||
x = self.lora_module_ref().lora_up.orig_forward(x, *args, **kwargs)
|
||||
x = x * up_weight
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise ValueError(f"Up weight shape not understood: {up_weight.shape} {x.shape}")
|
||||
|
||||
return x
|
||||
|
||||
def mid_forward(self, x, *args, **kwargs):
|
||||
if not self.do_mid:
|
||||
return self.lora_module_ref().lora_down.orig_forward(x, *args, **kwargs)
|
||||
batch_size = x.shape[0]
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
mid_weight = self.embed[:, self.down_dim:self.down_dim + self.mid_dim * self.mid_dim]
|
||||
|
||||
# unconditional
|
||||
if mid_weight.shape[0] * 2 == batch_size:
|
||||
mid_weight = torch.cat([mid_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(mid_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
if len(x_chunk.shape) == 4:
|
||||
# conv
|
||||
weight_chunk = weight_chunk.view(self.mid_dim, self.mid_dim, 1, 1)
|
||||
else:
|
||||
weight_chunk = weight_chunk.view(self.mid_dim, self.mid_dim)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
padding = 0
|
||||
if weight_chunk.shape[-1] == 3:
|
||||
padding = 1
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
class InstantLoRAModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
vision_hidden_size: int,
|
||||
vision_tokens: int,
|
||||
head_dim: int,
|
||||
num_heads: int, # number of heads in the resampler
|
||||
sd: 'StableDiffusion',
|
||||
config: AdapterConfig
|
||||
):
|
||||
super(InstantLoRAModule, self).__init__()
|
||||
# self.linear = torch.nn.Linear(2, 1)
|
||||
self.sd_ref = weakref.ref(sd)
|
||||
self.dim = sd.network.lora_dim
|
||||
self.vision_hidden_size = vision_hidden_size
|
||||
self.vision_tokens = vision_tokens
|
||||
self.head_dim = head_dim
|
||||
self.num_heads = num_heads
|
||||
|
||||
self.config: AdapterConfig = config
|
||||
|
||||
# stores the projection vector. Grabbed by modules
|
||||
self.img_embeds: List[torch.Tensor] = None
|
||||
|
||||
# disable merging in. It is slower on inference
|
||||
self.sd_ref().network.can_merge_in = False
|
||||
|
||||
self.ilora_modules = torch.nn.ModuleList()
|
||||
|
||||
lora_modules = self.sd_ref().network.get_all_modules()
|
||||
|
||||
output_size = 0
|
||||
|
||||
self.embed_lengths = []
|
||||
self.weight_mapping = []
|
||||
|
||||
for idx, lora_module in enumerate(lora_modules):
|
||||
module_dict = lora_module.state_dict()
|
||||
down_shape = list(module_dict['lora_down.weight'].shape)
|
||||
up_shape = list(module_dict['lora_up.weight'].shape)
|
||||
|
||||
self.weight_mapping.append([lora_module.lora_name, [down_shape, up_shape]])
|
||||
|
||||
#
|
||||
# module_size = math.prod(down_shape) + math.prod(up_shape)
|
||||
|
||||
# conv weight shape is (out_channels, in_channels, kernel_size, kernel_size)
|
||||
# linear weight shape is (out_features, in_features)
|
||||
|
||||
# just doing in dim and out dim
|
||||
in_dim = down_shape[1] if self.config.ilora_down else 0
|
||||
mid_dim = down_shape[0] * down_shape[0] if self.config.ilora_mid else 0
|
||||
out_dim = up_shape[0] if self.config.ilora_up else 0
|
||||
module_size = in_dim + mid_dim + out_dim
|
||||
|
||||
output_size += module_size
|
||||
self.embed_lengths.append(module_size)
|
||||
|
||||
# add a new mid module that will take the original forward and add a vector to it
|
||||
# this will be used to add the vector to the original forward
|
||||
instant_module = InstantLoRAMidModule(
|
||||
idx,
|
||||
lora_module,
|
||||
self,
|
||||
up_shape=up_shape,
|
||||
down_shape=down_shape
|
||||
)
|
||||
|
||||
self.ilora_modules.append(instant_module)
|
||||
|
||||
# replace the LoRA forwards
|
||||
lora_module.lora_down.orig_forward = lora_module.lora_down.forward
|
||||
lora_module.lora_down.forward = instant_module.down_forward
|
||||
lora_module.lora_up.orig_forward = lora_module.lora_up.forward
|
||||
lora_module.lora_up.forward = instant_module.up_forward
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
number_formatted_output_size = "{:,}".format(output_size)
|
||||
|
||||
print(f" ILORA output size: {number_formatted_output_size}")
|
||||
|
||||
# if not evenly divisible, error
|
||||
if self.output_size % self.num_heads != 0:
|
||||
raise ValueError("Output size must be divisible by the number of heads")
|
||||
|
||||
self.head_output_size = self.output_size // self.num_heads
|
||||
|
||||
if vision_tokens > 1:
|
||||
self.resampler = Resampler(
|
||||
dim=vision_hidden_size,
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=12,
|
||||
num_queries=num_heads, # output tokens
|
||||
embedding_dim=vision_hidden_size,
|
||||
max_seq_len=vision_tokens,
|
||||
output_dim=head_dim,
|
||||
apply_pos_emb=True, # this is new
|
||||
ff_mult=4
|
||||
)
|
||||
|
||||
self.proj_module = LoRAGenerator(
|
||||
input_size=head_dim,
|
||||
hidden_size=head_dim,
|
||||
head_size=head_dim,
|
||||
num_mlp_layers=1,
|
||||
num_heads=self.num_heads,
|
||||
output_size=self.output_size,
|
||||
)
|
||||
|
||||
self.migrate_weight_mapping()
|
||||
|
||||
def migrate_weight_mapping(self):
|
||||
return
|
||||
# # changes the names of the modules to common ones
|
||||
# keymap = self.sd_ref().network.get_keymap()
|
||||
# save_keymap = {}
|
||||
# if keymap is not None:
|
||||
# for ldm_key, diffusers_key in keymap.items():
|
||||
# # invert them
|
||||
# save_keymap[diffusers_key] = ldm_key
|
||||
#
|
||||
# new_keymap = {}
|
||||
# for key, value in self.weight_mapping:
|
||||
# if key in save_keymap:
|
||||
# new_keymap[save_keymap[key]] = value
|
||||
# else:
|
||||
# print(f"Key {key} not found in keymap")
|
||||
# new_keymap[key] = value
|
||||
# self.weight_mapping = new_keymap
|
||||
# else:
|
||||
# print("No keymap found. Using default names")
|
||||
# return
|
||||
|
||||
def forward(self, img_embeds):
|
||||
# expand token rank if only rank 2
|
||||
if len(img_embeds.shape) == 2:
|
||||
img_embeds = img_embeds.unsqueeze(1)
|
||||
|
||||
# resample the image embeddings
|
||||
img_embeds = self.resampler(img_embeds)
|
||||
img_embeds = self.proj_module(img_embeds)
|
||||
if len(img_embeds.shape) == 3:
|
||||
# merge the heads
|
||||
img_embeds = img_embeds.mean(dim=1)
|
||||
|
||||
self.img_embeds = []
|
||||
# get all the slices
|
||||
start = 0
|
||||
for length in self.embed_lengths:
|
||||
self.img_embeds.append(img_embeds[:, start:start + length])
|
||||
start += length
|
||||
|
||||
def get_additional_save_metadata(self) -> Dict[str, Any]:
|
||||
# save the weight mapping
|
||||
return {
|
||||
"weight_mapping": self.weight_mapping,
|
||||
"num_heads": self.num_heads,
|
||||
"vision_hidden_size": self.vision_hidden_size,
|
||||
"head_dim": self.head_dim,
|
||||
"vision_tokens": self.vision_tokens,
|
||||
"output_size": self.output_size,
|
||||
"do_up": self.config.ilora_up,
|
||||
"do_mid": self.config.ilora_mid,
|
||||
"do_down": self.config.ilora_down,
|
||||
}
|
||||
191
toolkit/models/llm_adapter.py
Normal file
191
toolkit/models/llm_adapter.py
Normal file
@@ -0,0 +1,191 @@
|
||||
from functools import partial
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import weakref
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union, TYPE_CHECKING
|
||||
|
||||
from diffusers.models.transformers.transformer_flux import FluxTransformerBlock
|
||||
from transformers import AutoModel, AutoTokenizer, Qwen2Model, LlamaModel, Qwen2Tokenizer, LlamaTokenizer
|
||||
|
||||
from toolkit import train_tools
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from diffusers import Transformer2DModel
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, PixArtSigmaPipeline
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
|
||||
LLM = Union[Qwen2Model, LlamaModel]
|
||||
LLMTokenizer = Union[Qwen2Tokenizer, LlamaTokenizer]
|
||||
|
||||
|
||||
def new_context_embedder_forward(self, x):
|
||||
if self._adapter_ref().is_active:
|
||||
x = self._context_embedder_ref()(x)
|
||||
else:
|
||||
x = self._orig_forward(x)
|
||||
return x
|
||||
|
||||
def new_block_forward(
|
||||
self: FluxTransformerBlock,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
if self._adapter_ref().is_active:
|
||||
return self._new_block_ref()(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
|
||||
else:
|
||||
return self._orig_forward(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
|
||||
|
||||
|
||||
class LLMAdapter(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'CustomAdapter',
|
||||
sd: 'StableDiffusion',
|
||||
llm: LLM,
|
||||
tokenizer: LLMTokenizer,
|
||||
num_cloned_blocks: int = 0,
|
||||
):
|
||||
super(LLMAdapter, self).__init__()
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.llm_ref: weakref.ref = weakref.ref(llm)
|
||||
self.tokenizer_ref: weakref.ref = weakref.ref(tokenizer)
|
||||
self.num_cloned_blocks = num_cloned_blocks
|
||||
self.apply_embedding_mask = False
|
||||
# make sure we can pad
|
||||
if tokenizer.pad_token is None:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
|
||||
# self.system_prompt = ""
|
||||
self.system_prompt = "You are an assistant designed to generate superior images with the superior degree of image-text alignment based on textual prompts or user prompts. <Prompt Start> "
|
||||
|
||||
# determine length of system prompt
|
||||
sys_prompt_tokenized = tokenizer(
|
||||
[self.system_prompt],
|
||||
padding="longest",
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
sys_prompt_tokenized_ids = sys_prompt_tokenized.input_ids[0]
|
||||
|
||||
self.system_prompt_length = sys_prompt_tokenized_ids.shape[0]
|
||||
|
||||
print(f"System prompt length: {self.system_prompt_length}")
|
||||
|
||||
self.hidden_size = llm.config.hidden_size
|
||||
|
||||
blocks = []
|
||||
|
||||
if sd.is_flux:
|
||||
self.apply_embedding_mask = True
|
||||
self.context_embedder = nn.Linear(
|
||||
self.hidden_size, sd.unet.inner_dim)
|
||||
self.sequence_length = 512
|
||||
sd.unet.context_embedder._orig_forward = sd.unet.context_embedder.forward
|
||||
sd.unet.context_embedder.forward = partial(
|
||||
new_context_embedder_forward, sd.unet.context_embedder)
|
||||
sd.unet.context_embedder._context_embedder_ref = weakref.ref(self.context_embedder)
|
||||
# add a is active property to the context embedder
|
||||
sd.unet.context_embedder._adapter_ref = self.adapter_ref
|
||||
|
||||
for idx in range(self.num_cloned_blocks):
|
||||
block = FluxTransformerBlock(
|
||||
dim=sd.unet.inner_dim,
|
||||
num_attention_heads=24,
|
||||
attention_head_dim=128,
|
||||
)
|
||||
# patch it in case it is quantized
|
||||
patch_dequantization_on_save(sd.unet.transformer_blocks[idx])
|
||||
state_dict = sd.unet.transformer_blocks[idx].state_dict()
|
||||
for key, value in state_dict.items():
|
||||
block.state_dict()[key].copy_(value)
|
||||
blocks.append(block)
|
||||
orig_block = sd.unet.transformer_blocks[idx]
|
||||
orig_block._orig_forward = orig_block.forward
|
||||
orig_block.forward = partial(
|
||||
new_block_forward, orig_block)
|
||||
orig_block._new_block_ref = weakref.ref(block)
|
||||
orig_block._adapter_ref = self.adapter_ref
|
||||
|
||||
elif sd.is_lumina2:
|
||||
self.context_embedder = nn.Linear(
|
||||
self.hidden_size, sd.unet.hidden_size)
|
||||
self.sequence_length = 256
|
||||
else:
|
||||
raise ValueError(
|
||||
"llm adapter currently only supports flux or lumina2")
|
||||
|
||||
self.blocks = nn.ModuleList(blocks)
|
||||
|
||||
def _get_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
max_sequence_length: int = 256,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
tokenizer = self.tokenizer_ref()
|
||||
text_encoder = self.llm_ref()
|
||||
device = text_encoder.device
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
text_inputs = tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length + self.system_prompt_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
text_input_ids = text_inputs.input_ids.to(device)
|
||||
prompt_attention_mask = text_inputs.attention_mask.to(device)
|
||||
|
||||
# remove the system prompt from the input and attention mask
|
||||
|
||||
prompt_embeds = text_encoder(
|
||||
text_input_ids, attention_mask=prompt_attention_mask, output_hidden_states=True
|
||||
)
|
||||
prompt_embeds = prompt_embeds.hidden_states[-1]
|
||||
|
||||
prompt_embeds = prompt_embeds[:, self.system_prompt_length:]
|
||||
prompt_attention_mask = prompt_attention_mask[:, self.system_prompt_length:]
|
||||
|
||||
dtype = text_encoder.dtype
|
||||
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
# make a getter to see if is active
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
def encode_text(self, prompt):
|
||||
|
||||
prompt = prompt if isinstance(prompt, list) else [prompt]
|
||||
|
||||
prompt = [self.system_prompt + p for p in prompt]
|
||||
# prompt = [self.system_prompt + p for p in prompt]
|
||||
|
||||
prompt_embeds, prompt_attention_mask = self._get_prompt_embeds(
|
||||
prompt=prompt,
|
||||
max_sequence_length=self.sequence_length,
|
||||
)
|
||||
|
||||
prompt_embeds = PromptEmbeds(
|
||||
prompt_embeds,
|
||||
attention_mask=prompt_attention_mask,
|
||||
).detach()
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def forward(self, input):
|
||||
return input
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user