Compare commits
780 Commits
sdxl
...
qwen_image
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
59ff4efae5 | ||
|
|
aa99784b89 | ||
|
|
bf2700f7be | ||
|
|
38d3814be7 | ||
|
|
83deaec417 | ||
|
|
d2bbe1872c | ||
|
|
b3e666daf4 | ||
|
|
6fffadfc0e | ||
|
|
280aca685f | ||
|
|
1029fa8743 | ||
|
|
8ea2cf00f6 | ||
|
|
ca7bfa414b | ||
|
|
1c96b95617 | ||
|
|
3413fa537f | ||
|
|
be71cc75ce | ||
|
|
e12bb21780 | ||
|
|
3ff4430e84 | ||
|
|
5501521c9f | ||
|
|
85bad57df3 | ||
|
|
259d68d440 | ||
|
|
69ee99b6e1 | ||
|
|
77b10d884d | ||
|
|
4ad18f3d00 | ||
|
|
f0105c33a7 | ||
|
|
ccd449ec49 | ||
|
|
bb6db3d635 | ||
|
|
4c4a10d439 | ||
|
|
14ccf2f3ce | ||
|
|
5d8922fca2 | ||
|
|
1755e58dd9 | ||
|
|
6bb3aed9a2 | ||
|
|
74b4d2d291 | ||
|
|
23327d5659 | ||
|
|
93202c7a2b | ||
|
|
9da8b5408e | ||
|
|
9dfb614755 | ||
|
|
ef1d60ba34 | ||
|
|
75f688766d | ||
|
|
a558d5b68f | ||
|
|
1d1199b15b | ||
|
|
f453e28ea3 | ||
|
|
ca7c5c950b | ||
|
|
e55116d8c9 | ||
|
|
99705ec8be | ||
|
|
ed8d14225f | ||
|
|
b717586ee2 | ||
|
|
cefa2ca5fe | ||
|
|
3f518d9951 | ||
|
|
77dc38a574 | ||
|
|
0d89c44624 | ||
|
|
3e14a674ac | ||
|
|
523c159579 | ||
|
|
c5eb763342 | ||
|
|
ca5cf827a1 | ||
|
|
b1bff66d52 | ||
|
|
a77ba5a089 | ||
|
|
8610c6ed7f | ||
|
|
e25d2feddf | ||
|
|
f500b9f240 | ||
|
|
3916e67455 | ||
|
|
e5ed450dc7 | ||
|
|
1930c3edea | ||
|
|
ef5149180c | ||
|
|
998e8b6537 | ||
|
|
755f0e207c | ||
|
|
2e84b3d5b1 | ||
|
|
7ab44ae0cd | ||
|
|
47002b067f | ||
|
|
8537a8557f | ||
|
|
6e2beef8dd | ||
|
|
611969ec1f | ||
|
|
bbb57de6ec | ||
|
|
5906a76666 | ||
|
|
57a81bc0db | ||
|
|
843be31138 | ||
|
|
8fb01e96e4 | ||
|
|
01a3c8a9b1 | ||
|
|
4f91cb7148 | ||
|
|
446b0b6989 | ||
|
|
60ef2f1df7 | ||
|
|
8d9c47316a | ||
|
|
84c6edca7e | ||
|
|
24cd94929e | ||
|
|
19ea8ecc38 | ||
|
|
18513ec866 | ||
|
|
5e733764aa | ||
|
|
03bc431279 | ||
|
|
f3eb1dff42 | ||
|
|
ba1274d99e | ||
|
|
8602470952 | ||
|
|
4586eb5392 | ||
|
|
989ebfaa11 | ||
|
|
ff617fdaea | ||
|
|
595a6f1735 | ||
|
|
1cc663a664 | ||
|
|
11f2eee53a | ||
|
|
1c2b7298dd | ||
|
|
c0314ba325 | ||
|
|
cbf04b8d53 | ||
|
|
0946a66576 | ||
|
|
3f0ae99d48 | ||
|
|
fc83eb7691 | ||
|
|
cf11f128b9 | ||
|
|
5e86139e0a | ||
|
|
c5d6b74fea | ||
|
|
ba5196dd4a | ||
|
|
ffb5fe0667 | ||
|
|
f8fb3b9c45 | ||
|
|
d5c547da43 | ||
|
|
f19f7f9486 | ||
|
|
7317ed58af | ||
|
|
517bc294fa | ||
|
|
97e101522c | ||
|
|
eefa93f16e | ||
|
|
22cdfadab6 | ||
|
|
adc31ec77d | ||
|
|
82b90b902e | ||
|
|
85f4b47e79 | ||
|
|
e20a869dc1 | ||
|
|
12fa109910 | ||
|
|
7d76165dcf | ||
|
|
b6d25fcd10 | ||
|
|
ffaf2f154a | ||
|
|
34f4c14cd6 | ||
|
|
79bb9be92b | ||
|
|
79499fa795 | ||
|
|
48e11cf843 | ||
|
|
7045a01375 | ||
|
|
fca7fd6c38 | ||
|
|
e5181d23cd | ||
|
|
4f896c0d8a | ||
|
|
01101be196 | ||
|
|
6174ba474e | ||
|
|
64130189ce | ||
|
|
66a41e49d9 | ||
|
|
1210050ead | ||
|
|
25e150b370 | ||
|
|
43cb5603ad | ||
|
|
d9700bdb99 | ||
|
|
5890e67a46 | ||
|
|
2b4c525489 | ||
|
|
88b3fbae37 | ||
|
|
8ff85ba14f | ||
|
|
80f73ce9c0 | ||
|
|
9f42944056 | ||
|
|
add83df5cc | ||
|
|
12e3095d8a | ||
|
|
77001ee77f | ||
|
|
d455e76c4f | ||
|
|
1628884254 | ||
|
|
9c422ac14f | ||
|
|
bfe29e2151 | ||
|
|
bd2de5b74e | ||
|
|
970fac19a5 | ||
|
|
5f312cd46b | ||
|
|
c90615f8bb | ||
|
|
5961ef6c9f | ||
|
|
fd6026ab73 | ||
|
|
79c87701e7 | ||
|
|
fecc64e646 | ||
|
|
d5a64006b5 | ||
|
|
0f99fce004 | ||
|
|
c12036df95 | ||
|
|
68018c908e | ||
|
|
524bd2edfc | ||
|
|
89c0f688db | ||
|
|
1e0bff653c | ||
|
|
3a5ea2c742 | ||
|
|
f80cf99f40 | ||
|
|
594e166ca3 | ||
|
|
ca3ce0f34c | ||
|
|
6fb44db6a0 | ||
|
|
cd37ccfc2e | ||
|
|
4a43589666 | ||
|
|
059155174a | ||
|
|
9794416a5d | ||
|
|
d8bdc03256 | ||
|
|
96ba2fd129 | ||
|
|
615b0d0e94 | ||
|
|
a8680c75eb | ||
|
|
38ad5a4644 | ||
|
|
6c8b5ab606 | ||
|
|
7c21eac1b3 | ||
|
|
2b901cca39 | ||
|
|
ead23cee88 | ||
|
|
ab59ca5091 | ||
|
|
eddd3c1611 | ||
|
|
b0d0466efd | ||
|
|
ac1ee559c5 | ||
|
|
77763a3e5c | ||
|
|
a42c5a1de5 | ||
|
|
3d131fb27a | ||
|
|
5ea19b6292 | ||
|
|
58861005a5 | ||
|
|
c083a0e5ea | ||
|
|
860d892214 | ||
|
|
b94d7aafea | ||
|
|
3c95f87a90 | ||
|
|
1d5f387f54 | ||
|
|
5365200da1 | ||
|
|
e9e30104d3 | ||
|
|
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 | ||
|
|
c2a4b8e058 | ||
|
|
792a5e37e2 | ||
|
|
d7e55b6ad4 | ||
|
|
fbec68681d | ||
|
|
6280284d8b | ||
|
|
ad50921c41 | ||
|
|
e47006ed70 | ||
|
|
4f9cdd916a | ||
|
|
7782caa468 | ||
|
|
fa6d91ba76 | ||
|
|
1ee62562a4 | ||
|
|
dc8448d958 | ||
|
|
a8b3b8b8da | ||
|
|
93ea955d7c | ||
|
|
8a9e8f708f | ||
|
|
d35733ac06 | ||
|
|
ceaf1d9454 | ||
|
|
7d707b2fe6 | ||
|
|
a899ec91c8 | ||
|
|
436a09430e | ||
|
|
3097865203 | ||
|
|
b84e3260cb | ||
|
|
48a9bac22d | ||
|
|
298001439a | ||
|
|
6f3e0d5af2 | ||
|
|
0a79ac9604 | ||
|
|
9636194c09 | ||
|
|
d742792ee4 | ||
|
|
002279cec3 | ||
|
|
73c8b50975 | ||
|
|
34eb563d55 | ||
|
|
dc36bbb3c8 | ||
|
|
9905a1e205 | ||
|
|
0e9fc42816 | ||
|
|
d46112a354 | ||
|
|
07bf7bd7de | ||
|
|
da6302ada8 | ||
|
|
a05459afaf | ||
|
|
b1a22d0b3e | ||
|
|
7909b50d24 | ||
|
|
38e441a29c | ||
|
|
4e3b2c2569 | ||
|
|
239addba51 | ||
|
|
63ceffae24 | ||
|
|
f4c90bb589 | ||
|
|
bb1d3793e3 | ||
|
|
1d3de678aa | ||
|
|
b1cfafa0c6 | ||
|
|
cac8754399 | ||
|
|
f73402473b | ||
|
|
579650eaf8 | ||
|
|
320e109c5f | ||
|
|
560251a24f | ||
|
|
085787b799 | ||
|
|
8d9450ad7c | ||
|
|
8509da60cb | ||
|
|
c5d49ba661 | ||
|
|
76c764af49 | ||
|
|
abf7cd221d | ||
|
|
e5153d87c9 | ||
|
|
830e87cb87 | ||
|
|
19255cdc7c | ||
|
|
0f105690cc | ||
|
|
61badf85a7 | ||
|
|
181f237a7b | ||
|
|
c698837241 | ||
|
|
27f343fc08 | ||
|
|
3eb3535683 | ||
|
|
17e4fe40d7 | ||
|
|
569d7464d5 | ||
|
|
4e945917df | ||
|
|
ae70200d3c | ||
|
|
d8d1e6fd1e | ||
|
|
257da9493d | ||
|
|
b5a2669b74 | ||
|
|
d74dd636ee | ||
|
|
e8583860ad | ||
|
|
3d387103cd | ||
|
|
083cefa78c | ||
|
|
b5ec8e4eb1 | ||
|
|
708b07adb7 | ||
|
|
a437aed45f | ||
|
|
34bfeba229 | ||
|
|
41a3f63b72 | ||
|
|
626ed2939a | ||
|
|
2128ac1e08 | ||
|
|
be804c9cf5 | ||
|
|
408c50ead1 | ||
|
|
4ed03a8d92 | ||
|
|
b01ab5d375 | ||
|
|
cb91b0d6da | ||
|
|
ce4f9fe02a | ||
|
|
92a086d5a5 | ||
|
|
3feb663a51 | ||
|
|
436bf0c6a3 | ||
|
|
f84500159c | ||
|
|
64a5441832 | ||
|
|
a4c3507a62 | ||
|
|
fa8fc32c0a | ||
|
|
22ed539321 | ||
|
|
7cd6945082 | ||
|
|
2a40937b4f | ||
|
|
4ca819a05e | ||
|
|
addf024630 | ||
|
|
33267e117c | ||
|
|
d401348c2e | ||
|
|
836fee47a6 | ||
|
|
14ff51ceb4 | ||
|
|
714854ee86 | ||
|
|
bd758ff203 | ||
|
|
a008d9e63b | ||
|
|
b79ced3e10 | ||
|
|
bee0b6a235 | ||
|
|
2ecb5cf024 | ||
|
|
fab7c2b04a | ||
|
|
e866c75638 | ||
|
|
71da78c8af | ||
|
|
c446f768ea | ||
|
|
cc49786ee9 | ||
|
|
9b164a8688 | ||
|
|
6bd3851058 | ||
|
|
fd338e67bb | ||
|
|
8105c05c12 | ||
|
|
2cb27c3f57 | ||
|
|
3367ab6b2c | ||
|
|
24f46ea7d6 | ||
|
|
5bef2985b5 | ||
|
|
aeaca13d69 | ||
|
|
b408f9f3eb | ||
|
|
7157c316af | ||
|
|
e2c547f6c2 | ||
|
|
7b770bc305 | ||
|
|
f200cf36c5 | ||
|
|
d298240cec | ||
|
|
2e6c55c720 | ||
|
|
36ba08d3fa | ||
|
|
e8667f856f | ||
|
|
bef5551ea5 | ||
|
|
b77b9acc0b | ||
|
|
c6675e2801 | ||
|
|
90eedb78bf | ||
|
|
80e2f4a2a4 | ||
|
|
d51c4ca704 | ||
|
|
c7ec132d5d | ||
|
|
ed9607e8da | ||
|
|
d44f8ac508 | ||
|
|
8d09eb44ec | ||
|
|
55a5fcc7d9 | ||
|
|
e96874241d | ||
|
|
e3be1a1758 | ||
|
|
1a92e97c6d | ||
|
|
355c80df07 | ||
|
|
1487d13191 | ||
|
|
383bad958d | ||
|
|
196b693cf0 | ||
|
|
fd95e7b60c | ||
|
|
379992d89e | ||
|
|
c7054d714f | ||
|
|
67dfd9ced0 | ||
|
|
1a7e346b41 | ||
|
|
df48f0a843 | ||
|
|
fbc8a87a05 | ||
|
|
bf90740b59 | ||
|
|
ff2c9f3d04 | ||
|
|
8bd536df7e | ||
|
|
64fbd4c92a | ||
|
|
8c90fa86c6 | ||
|
|
7e4e660663 | ||
|
|
b865ac8b24 | ||
|
|
66c6f0f6f7 | ||
|
|
75ec5d9292 | ||
|
|
1a25b275c8 | ||
|
|
2bf3e529ce | ||
|
|
f53fd08690 | ||
|
|
8b8d53888d | ||
|
|
63cacf4362 | ||
|
|
c1b1e800df | ||
|
|
7726911562 | ||
|
|
c01673f1b5 | ||
|
|
c35b78f0d4 | ||
|
|
8ba1b11557 | ||
|
|
1e50b39442 | ||
|
|
5fc2bb5d9c | ||
|
|
c7640b0865 | ||
|
|
b2e2e4bf47 | ||
|
|
596e57a6a6 | ||
|
|
6ab8b8b0f1 |
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.
|
||||
13
.gitignore
vendored
13
.gitignore
vendored
@@ -161,6 +161,7 @@ cython_debug/
|
||||
|
||||
/env.sh
|
||||
/models
|
||||
/datasets
|
||||
/custom/*
|
||||
!/custom/.gitkeep
|
||||
/.tmp
|
||||
@@ -170,4 +171,14 @@ cython_debug/
|
||||
!/config/examples
|
||||
!/config/_PUT_YOUR_CONFIGS_HERE).txt
|
||||
/output/*
|
||||
!/output/.gitkeep
|
||||
!/output/.gitkeep
|
||||
/extensions/*
|
||||
!/extensions/example
|
||||
/temp
|
||||
/wandb
|
||||
.vscode/settings.json
|
||||
.DS_Store
|
||||
._.DS_Store
|
||||
aitk_db.db
|
||||
/notes.md
|
||||
/data
|
||||
6
.gitmodules
vendored
6
.gitmodules
vendored
@@ -1,6 +0,0 @@
|
||||
[submodule "repositories/sd-scripts"]
|
||||
path = repositories/sd-scripts
|
||||
url = https://github.com/kohya-ss/sd-scripts.git
|
||||
[submodule "repositories/leco"]
|
||||
path = repositories/leco
|
||||
url = https://github.com/p1atdev/LECO
|
||||
|
||||
56
.vscode/launch.json
vendored
Normal file
56
.vscode/launch.json
vendored
Normal file
@@ -0,0 +1,56 @@
|
||||
{
|
||||
"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": "Run current config (cuda:1)",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${workspaceFolder}/run.py",
|
||||
"args": [
|
||||
"${file}"
|
||||
],
|
||||
"env": {
|
||||
"CUDA_LAUNCH_BLOCKING": "1",
|
||||
"DEBUG_TOOLKIT": "1",
|
||||
"CUDA_VISIBLE_DEVICES": "1"
|
||||
},
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false
|
||||
},
|
||||
{
|
||||
"name": "Python: Debug Current File",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${file}",
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false
|
||||
},
|
||||
{
|
||||
"name": "Python: Debug Current File (cuda:1)",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${file}",
|
||||
"console": "integratedTerminal",
|
||||
"env": {
|
||||
"CUDA_LAUNCH_BLOCKING": "1",
|
||||
"CUDA_VISIBLE_DEVICES": "1"
|
||||
},
|
||||
"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.
|
||||
489
README.md
489
README.md
@@ -1,16 +1,120 @@
|
||||
# AI Toolkit by Ostris
|
||||
|
||||
## IMPORTANT NOTE - READ THIS
|
||||
This is an active WIP repo that is not ready for others to use. And definitely not ready for non developers to use.
|
||||
I am making major breaking changes and pushing straight to master until I have it in a planned state. I have big changes
|
||||
planned for config files and the general structure. I may change how training works entirely. You are welcome to use
|
||||
but keep that in mind. If more people start to use it, I will follow better branch checkout standards, but for now
|
||||
this is my personal active experiment.
|
||||
AI Toolkit is an all in one training suite for diffusion models. I try to support all the latest models on consumer grade hardware. Image and video models. It can be run as a GUI or CLI. It is designed to be easy to use but still have every feature imaginable.
|
||||
|
||||
## Support My Work
|
||||
|
||||
If you enjoy my projects or use them commercially, please consider sponsoring me. Every bit helps! 💖
|
||||
|
||||
[Sponsor on GitHub](https://github.com/orgs/ostris) | [Support on Patreon](https://www.patreon.com/ostris) | [Donate on PayPal](https://www.paypal.com/donate/?hosted_button_id=9GEFUKC8T9R9W)
|
||||
|
||||
### Current Sponsors
|
||||
|
||||
All of these people / organizations are the ones who selflessly make this project possible. Thank you!!
|
||||
|
||||
_Last updated: 2025-08-08 17:01 UTC_
|
||||
|
||||
<p align="center">
|
||||
<a href="https://x.com/NuxZoe" target="_blank" rel="noopener noreferrer"><img src="https://pbs.twimg.com/profile_images/1919488160125616128/QAZXTMEj_400x400.png" alt="a16z" width="200" height="200" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<a href="https://github.com/replicate" target="_blank" rel="noopener noreferrer"><img src="https://avatars.githubusercontent.com/u/60410876?v=4" alt="Replicate" width="200" height="200" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<a href="https://github.com/huggingface" target="_blank" rel="noopener noreferrer"><img src="https://avatars.githubusercontent.com/u/25720743?v=4" alt="Hugging Face" width="200" height="200" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<a href="https://github.com/josephrocca" target="_blank" rel="noopener noreferrer"><img src="https://avatars.githubusercontent.com/u/1167575?u=92d92921b4cb5c8c7e225663fed53c4b41897736&v=4" alt="josephrocca" width="200" height="200" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/162524101/81a72689c3754ac5b9e38612ce5ce914/eyJ3IjoyMDB9/1.png?token-hash=JHRjAxd2XxV1aXIUijj-l65pfTnLoefYSvgNPAsw2lI%3D" alt="Prasanth Veerina" width="200" height="200" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<a href="https://github.com/weights-ai" target="_blank" rel="noopener noreferrer"><img src="https://avatars.githubusercontent.com/u/185568492?v=4" alt="Weights" width="200" height="200" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
</p>
|
||||
<hr style="width:100%;border:none;height:2px;background:#ddd;margin:30px 0;">
|
||||
<p align="center">
|
||||
<img src="https://c8.patreon.com/4/200/93304/J" alt="Joseph Rocca" width="150" height="150" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/161471720/dd330b4036d44a5985ed5985c12a5def/eyJ3IjoyMDB9/1.jpeg?token-hash=k1f4Vv7TevzYa9tqlzAjsogYmkZs8nrXQohPCDGJGkc%3D" alt="Vladimir Sotnikov" width="150" height="150" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/33158543/C" alt="clement Delangue" width="150" height="150" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/8654302/b0f5ebedc62a47c4b56222693e1254e9/eyJ3IjoyMDB9/2.jpeg?token-hash=suI7_QjKUgWpdPuJPaIkElkTrXfItHlL8ZHLPT-w_d4%3D" alt="Misch Strotz" width="150" height="150" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/120239481/49b1ce70d3d24704b8ec34de24ec8f55/eyJ3IjoyMDB9/1.jpeg?token-hash=o0y1JqSXqtGvVXnxb06HMXjQXs6OII9yMMx5WyyUqT4%3D" alt="nitish PNR" width="150" height="150" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
</p>
|
||||
<hr style="width:100%;border:none;height:2px;background:#ddd;margin:30px 0;">
|
||||
<p align="center">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/2298192/1228b69bd7d7481baf3103315183250d/eyJ3IjoyMDB9/1.jpg?token-hash=opN1e4r4Nnvqbtr8R9HI8eyf9m5F50CiHDOdHzb4UcA%3D" alt="Mohamed Oumoumad" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/548524/S" alt="Steve Hanff" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/152118848/3b15a43d71714552b5ed1c9f84e66adf/eyJ3IjoyMDB9/1.png?token-hash=MKf3sWHz0MFPm_OAFjdsNvxoBfN5B5l54mn1ORdlRy8%3D" alt="Kristjan Retter" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/83319230/M" alt="Miguel Lara" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/8449560/P" alt="Patron" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<a href="https://x.com/NuxZoe" target="_blank" rel="noopener noreferrer"><img src="https://pbs.twimg.com/profile_images/1916482710069014528/RDLnPRSg_400x400.jpg" alt="tungsten" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/169502989/220069e79ce745b29237e94c22a729df/eyJ3IjoyMDB9/1.png?token-hash=E8E2JOqx66k2zMtYUw8Gy57dw-gVqA6OPpdCmWFFSFw%3D" alt="Timothy Bielec" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/34200989/58ae95ebda0640c8b7a91b4fa31357aa/eyJ3IjoyMDB9/1.jpeg?token-hash=4mVDM1kCYGauYa33zLG14_g0oj9_UjDK_-Qp4zk42GE%3D" alt="Noah Miller" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/27288932/6c35d2d961ee4e14a7a368c990791315/eyJ3IjoyMDB9/1.jpeg?token-hash=TGIto_PGEG2NEKNyqwzEnRStOkhrjb3QlMhHA3raKJY%3D" alt="David Garrido" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<a href="https://x.com/RalFingerLP" target="_blank" rel="noopener noreferrer"><img src="https://pbs.twimg.com/profile_images/919595465041162241/ZU7X3T5k_400x400.jpg" alt="RalFinger" width="100" height="100" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
</p>
|
||||
<hr style="width:100%;border:none;height:2px;background:#ddd;margin:30px 0;">
|
||||
<p align="center">
|
||||
<a href="http://www.ir-ltd.net" target="_blank" rel="noopener noreferrer"><img src="https://pbs.twimg.com/profile_images/1602579392198283264/6Tm2GYus_400x400.jpg" alt="IR-Entertainment Ltd" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/9547341/bb35d9a222fd460e862e960ba3eacbaf/eyJ3IjoyMDB9/1.jpeg?token-hash=Q2XGDvkCbiONeWNxBCTeTMOcuwTjOaJ8Z-CAf5xq3Hs%3D" alt="Travis Harrington" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/98811435/3a3632d1795b4c2b9f8f0270f2f6a650/eyJ3IjoyMDB9/1.jpeg?token-hash=657rzuJ0bZavMRZW3XZ-xQGqm3Vk6FkMZgFJVMCOPdk%3D" alt="EmmanuelMr18" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/81275465/1e4148fe9c47452b838949d02dd9a70f/eyJ3IjoyMDB9/1.jpeg?token-hash=YAX1ucxybpCIujUCXfdwzUQkttIn3c7pfi59uaFPSwM%3D" alt="Aaron Amortegui" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/155963250/6f8fd7075c3b4247bfeb054ba49172d6/eyJ3IjoyMDB9/1.png?token-hash=z81EHmdU2cqSrwa9vJmZTV3h0LG-z9Qakhxq34FrYT4%3D" alt="Un Defined" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/45562978/0de33cf52ec642ae8a2f612cddec4ca6/eyJ3IjoyMDB9/1.jpeg?token-hash=aD4debMD5ZQjqTII6s4zYSgVK2-bdQt9p3eipi0bENs%3D" alt="Jack English" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/27791680/J" alt="Jean-Tristan Marin" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/570742/4ceb33453a5a4745b430a216aba9280f/eyJ3IjoyMDB9/1.jpg?token-hash=nPcJ2zj3sloND9jvbnbYnob2vMXRnXdRuujthqDLWlU%3D" alt="Al H" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/82763/f99cc484361d4b9d94fe4f0814ada303/eyJ3IjoyMDB9/1.jpeg?token-hash=A3JWlBNL0b24FFWb-FCRDAyhs-OAxg-zrhfBXP_axuU%3D" alt="Doron Adler" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/103077711/bb215761cc004e80bd9cec7d4bcd636d/eyJ3IjoyMDB9/2.jpeg?token-hash=3U8kdZSUpnmeYIDVK4zK9TTXFpnAud_zOwBRXx18018%3D" alt="John Dopamine" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/99036356/7ae9c4d80e604e739b68cca12ee2ed01/eyJ3IjoyMDB9/3.png?token-hash=ZhsBMoTOZjJ-Y6h5NOmU5MT-vDb2fjK46JDlpEehkVQ%3D" alt="Noctre" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/141098579/1a9f0a1249d447a7a0df718a57343912/eyJ3IjoyMDB9/2.png?token-hash=_n-AQmPgY0FP9zCGTIEsr5ka4Y7YuaMkt3qL26ZqGg8%3D" alt="The Local Lab" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/93348210/5c650f32a0bc481d80900d2674528777/eyJ3IjoyMDB9/1.jpeg?token-hash=0jiknRw3jXqYWW6En8bNfuHgVDj4LI_rL7lSS4-_xlo%3D" alt="Armin Behjati" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/134129880/680c7e14cd1a4d1a9face921fb010f88/eyJ3IjoyMDB9/1.png?token-hash=5fqqHE6DCTbt7gDQL7VRcWkV71jF7FvWcLhpYl5aMXA%3D" alt="Bharat Prabhakar" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/70218846/C" alt="Cosmosis" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/30931983/54ab4e4ceab946e79a6418d205f9ed51/eyJ3IjoyMDB9/1.png?token-hash=j2phDrgd6IWuqKqNIDbq9fR2B3fMF-GUCQSdETS1w5Y%3D" alt="HestoySeghuro ." width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/4105384/J" alt="Jack Blakely" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/4541423/S" alt="Sören " width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<a href="https://www.youtube.com/@happyme7055" target="_blank" rel="noopener noreferrer"><img src="https://yt3.googleusercontent.com/ytc/AIdro_mFqhIRk99SoEWY2gvSvVp6u1SkCGMkRqYQ1OlBBeoOVp8=s160-c-k-c0x00ffffff-no-rj" alt="Marcus Rass" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<img src="https://c8.patreon.com/4/200/53077895/M" alt="Marc" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/157407541/bb9d80cffdab4334ad78366060561520/eyJ3IjoyMDB9/2.png?token-hash=WYz-U_9zabhHstOT5UIa5jBaoFwrwwqyWxWEzIR2m_c%3D" alt="Tokio Studio srl IT10640050968" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/44568304/a9d83a0e786b41b4bdada150f7c9271c/eyJ3IjoyMDB9/1.jpeg?token-hash=FtxnwrSrknQUQKvDRv2rqPceX2EF23eLq4pNQYM_fmw%3D" alt="Albert Bukoski" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/5048649/B" alt="Ben Ward" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/111904990/08b1cf65be6a4de091c9b73b693b3468/eyJ3IjoyMDB9/1.png?token-hash=_Odz6RD3CxtubEHbUxYujcjw6zAajbo3w8TRz249VBA%3D" alt="Brian Smith" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/494309/J" alt="Julian Tsependa" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/5602036/K" alt="Kelevra" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/159203973/36c817f941ac4fa18103a4b8c0cb9cae/eyJ3IjoyMDB9/1.png?token-hash=zkt72HW3EoiIEAn3LSk9gJPBsXfuTVcc4rRBS3CeR8w%3D" alt="Marko jak" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/24653779/R" alt="RayHell" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/76566911/6485eaf5ec6249a7b524ee0b979372f0/eyJ3IjoyMDB9/1.jpeg?token-hash=mwCSkTelDBaengG32NkN0lVl5mRjB-cwo6-a47wnOsU%3D" alt="the biitz" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/32633822/1ab5612efe80417cbebfe91e871fc052/eyJ3IjoyMDB9/1.png?token-hash=pOS_IU3b3RL5-iL96A3Xqoj2bQ-dDo4RUkBylcMED_s%3D" alt="Zack Abrams" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/97985240/3d1d0e6905d045aba713e8132cab4a30/eyJ3IjoyMDB9/1.png?token-hash=fRavvbO_yqWKA_OsJb5DzjfKZ1Yt-TG-ihMoeVBvlcM%3D" alt="עומר מכלוף" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<a href="https://github.com/julien-blanchon" target="_blank" rel="noopener noreferrer"><img src="https://avatars.githubusercontent.com/u/11278197?v=4" alt="Blanchon" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/11198131/e696d9647feb4318bcf16243c2425805/eyJ3IjoyMDB9/1.jpeg?token-hash=c2c2p1SaiX86iXAigvGRvzm4jDHvIFCg298A49nIfUM%3D" alt="Nicholas Agranoff" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/785333/bdb9ede5765d42e5a2021a86eebf0d8f/eyJ3IjoyMDB9/2.jpg?token-hash=l_rajMhxTm6wFFPn7YdoKBxeUqhdRXKdy6_8SGCuNsE%3D" alt="Sapjes " width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/2446176/S" alt="Scott VanKirk" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/83034/W" alt="william tatum" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/138787189/2b5662dcb638466282ac758e3ac651b4/eyJ3IjoyMDB9/1.png?token-hash=zwj7MScO18vhDxhKt6s5q4gdeNJM3xCLuhSt8zlqlZs%3D" alt="Антон Антонио" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/30530914/T" alt="Techer " width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/25209707/36ae876d662d4d85aaf162b6d67d31e7/eyJ3IjoyMDB9/1.png?token-hash=Zows_A6uqlY5jClhfr4Y3QfMnDKVkS3mbxNHUDkVejo%3D" alt="fjioq8" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/46680573/ee3d99c04a674dd5a8e1ecfb926db6a2/eyJ3IjoyMDB9/1.jpeg?token-hash=cgD4EXyfZMPnXIrcqWQ5jGqzRUfqjPafb9yWfZUPB4Q%3D" alt="Neil Murray" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://ostris.com/wp-content/uploads/2025/08/supporter_default.jpg" alt="Joakim Sällström" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/63510241/A" alt="Andrew Park" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<a href="https://github.com/Spikhalskiy" target="_blank" rel="noopener noreferrer"><img src="https://avatars.githubusercontent.com/u/532108?u=2464983638afea8caf4cd9f0e4a7bc3e6a63bb0a&v=4" alt="Dmitry Spikhalsky" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<img src="https://c8.patreon.com/4/200/88567307/E" alt="el Chavo" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/117569999/55f75c57f95343e58402529cec852b26/eyJ3IjoyMDB9/1.jpeg?token-hash=squblHZH4-eMs3gI46Uqu1oTOK9sQ-0gcsFdZcB9xQg%3D" alt="James Thompson" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/66157709/6fe70df085e24464995a1a9293a53760/eyJ3IjoyMDB9/1.jpeg?token-hash=eqe0wvg6JfbRUGMKpL_x3YPI5Ppf18aUUJe2EzADU-g%3D" alt="Joey Santana" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://ostris.com/wp-content/uploads/2025/08/supporter_default.jpg" alt="Heikki Rinkinen" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/6175608/B" alt="Bobbie " width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<a href="https://github.com/Slartibart23" target="_blank" rel="noopener noreferrer"><img src="https://avatars.githubusercontent.com/u/133593860?u=31217adb2522fb295805824ffa7e14e8f0fca6fa&v=4" alt="Slarti" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;"></a>
|
||||
<img src="https://ostris.com/wp-content/uploads/2025/08/supporter_default.jpg" alt="Tommy Falkowski" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/28533016/e8f6044ccfa7483f87eeaa01c894a773/eyJ3IjoyMDB9/2.png?token-hash=ak-h3JWB50hyenCavcs32AAPw6nNhmH2nBFKpdk5hvM%3D" alt="William Tatum" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://ostris.com/wp-content/uploads/2025/08/supporter_default.jpg" alt="Karol Stępień" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/156564939/17dbfd45c59d4cf29853d710cb0c5d6f/eyJ3IjoyMDB9/1.png?token-hash=e6wXA_S8cgJeEDI9eJK934eB0TiM8mxJm9zW_VH0gDU%3D" alt="Hans Untch" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/59408413/B" alt="ByteC" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/3712451/432e22a355494ec0a1ea1927ff8d452e/eyJ3IjoyMDB9/7.jpeg?token-hash=OpQ9SAfVQ4Un9dSYlGTHuApZo5GlJ797Mo0DtVtMOSc%3D" alt="David Shorey" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/53634141/c1441f6c605344bbaef885d4272977bb/eyJ3IjoyMDB9/1.JPG?token-hash=Aizd6AxQhY3n6TBE5AwCVeSwEBbjALxQmu6xqc08qBo%3D" alt="Jana Spacelight" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/11180426/J" alt="jarrett towe" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/21828017/J" alt="Jim" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/63232055/2300b4ab370341b5b476902c9b8218ee/eyJ3IjoyMDB9/1.png?token-hash=R9Nb4O0aLBRwxT1cGHUMThlvf6A2MD5SO88lpZBdH7M%3D" alt="Marek P" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/9944625/P" alt="Pomoe " width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/25047900/423e4cb73aba457f8f9c6e5582eddaeb/eyJ3IjoyMDB9/1.jpeg?token-hash=81RvQXBbT66usxqtyWum9Ul4oBn3qHK1cM71IvthC-U%3D" alt="Ruairi Robinson" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c10.patreonusercontent.com/4/patreon-media/p/user/178476551/0b9e83efcd234df5a6bea30d59e6c1cd/eyJ3IjoyMDB9/1.png?token-hash=3XoYMrMxk-K6GelM22mE-FwkjFulX9hpIL7QI3wO2jI%3D" alt="Timmy" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://c8.patreon.com/4/200/10876902/T" alt="Tyssel" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
<img src="https://ostris.com/wp-content/uploads/2025/08/supporter_default.jpg" alt="Juan Franco" width="60" height="60" style="border-radius:8px;margin:5px;display: inline-block;">
|
||||
</p>
|
||||
|
||||
---
|
||||
|
||||
|
||||
Report bugs as you find them, but not knowing how to train ML models, setup an environment, or use python is not a bug.
|
||||
I will make all of this more user-friendly eventually
|
||||
|
||||
I will make a better readme later.
|
||||
|
||||
## Installation
|
||||
|
||||
@@ -21,101 +125,352 @@ Requirements:
|
||||
- git
|
||||
|
||||
|
||||
|
||||
Linux:
|
||||
```bash
|
||||
git clone https://github.com/ostris/ai-toolkit.git
|
||||
cd ai-toolkit
|
||||
git submodule update --init --recursive
|
||||
python3 -m venv venv
|
||||
source venv/bin/activate
|
||||
# or source venv/Scripts/activate on windows
|
||||
# install torch first
|
||||
pip3 install --no-cache-dir torch==2.7.0 torchvision==0.22.0 torchaudio==2.7.0 --index-url https://download.pytorch.org/whl/cu126
|
||||
pip3 install -r requirements.txt
|
||||
```
|
||||
|
||||
---
|
||||
Windows:
|
||||
|
||||
## Current Tools
|
||||
|
||||
I have so many hodge podge scripts I am going to be moving over to this that I use in my ML work. But this is what is
|
||||
here so far.
|
||||
|
||||
### LoRA (lierla), LoCON (LyCORIS) extractor
|
||||
|
||||
It is based on the extractor in the [LyCORIS](https://github.com/KohakuBlueleaf/LyCORIS) tool, but adding some QOL features
|
||||
and LoRA (lierla) support. It can do multiple types of extractions in one run.
|
||||
It all runs off a config file, which you can find an example of in `config/examples/extract.example.yml`.
|
||||
Just copy that file, into the `config` folder, and rename it to `whatever_you_want.yml`.
|
||||
Then you can edit the file to your liking. and call it like so:
|
||||
If you are having issues with Windows. I recommend using the easy install script at [https://github.com/Tavris1/AI-Toolkit-Easy-Install](https://github.com/Tavris1/AI-Toolkit-Easy-Install)
|
||||
|
||||
```bash
|
||||
python3 run.py config/whatever_you_want.yml
|
||||
git clone https://github.com/ostris/ai-toolkit.git
|
||||
cd ai-toolkit
|
||||
python -m venv venv
|
||||
.\venv\Scripts\activate
|
||||
pip install --no-cache-dir torch==2.7.0 torchvision==0.22.0 torchaudio==2.7.0 --index-url https://download.pytorch.org/whl/cu126
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
You can also put a full path to a config file, if you want to keep it somewhere else.
|
||||
|
||||
# AI Toolkit UI
|
||||
|
||||
<img src="https://ostris.com/wp-content/uploads/2025/02/toolkit-ui.jpg" alt="AI Toolkit UI" width="100%">
|
||||
|
||||
The AI Toolkit UI is a web interface for the AI Toolkit. It allows you to easily start, stop, and monitor jobs. It also allows you to easily train models with a few clicks. It also allows you to set a token for the UI to prevent unauthorized access so it is mostly safe to run on an exposed server.
|
||||
|
||||
## Running the UI
|
||||
|
||||
Requirements:
|
||||
- Node.js > 18
|
||||
|
||||
The UI does not need to be kept running for the jobs to run. It is only needed to start/stop/monitor jobs. The commands below
|
||||
will install / update the UI and it's dependencies and start the UI.
|
||||
|
||||
```bash
|
||||
python3 run.py "/home/user/whatever_you_want.yml"
|
||||
cd ui
|
||||
npm run build_and_start
|
||||
```
|
||||
|
||||
More notes on how it works are available in the example config file itself. LoRA and LoCON both support
|
||||
extractions of 'fixed', 'threshold', 'ratio', 'quantile'. I'll update what these do and mean later.
|
||||
Most people used fixed, which is traditional fixed dimension extraction.
|
||||
You can now access the UI at `http://localhost:8675` or `http://<your-ip>:8675` if you are running it on a server.
|
||||
|
||||
`process` is an array of different processes to run. You can add a few and mix and match. One LoRA, one LyCON, etc.
|
||||
## Securing the UI
|
||||
|
||||
|
||||
### LoRA Slider Trainer
|
||||
|
||||
This is how I train most of the recent sliders I have on Civitai, you can check them out in my [Civitai profile](https://civitai.com/user/Ostris/models).
|
||||
It is based off the work by [p1atdev/LECO](https://github.com/p1atdev/LECO) and [rohitgandikota/erasing](https://github.com/rohitgandikota/erasing)
|
||||
But has been heavily modified to create sliders rather than erasing concepts. I have a lot more plans on this, but it is
|
||||
very functional as is. It is also very easy to use. Just copy the example config file in `config/examples/train_slider.example.yml`
|
||||
to the `config` folder and rename it to `whatever_you_want.yml`. Then you can edit the file to your liking. and call it like so:
|
||||
If you are hosting the UI on a cloud provider or any network that is not secure, I highly recommend securing it with an auth token.
|
||||
You can do this by setting the environment variable `AI_TOOLKIT_AUTH` to super secure password. This token will be required to access
|
||||
the UI. You can set this when starting the UI like so:
|
||||
|
||||
```bash
|
||||
python3 run.py config/whatever_you_want.yml
|
||||
# Linux
|
||||
AI_TOOLKIT_AUTH=super_secure_password npm run build_and_start
|
||||
|
||||
# Windows
|
||||
set AI_TOOLKIT_AUTH=super_secure_password && npm run build_and_start
|
||||
|
||||
# Windows Powershell
|
||||
$env:AI_TOOLKIT_AUTH="super_secure_password"; npm run build_and_start
|
||||
```
|
||||
|
||||
There is a lot more information in that example file. You can even run the example as is without any modifications to see
|
||||
how it works. It will create a slider that turns all animals into dogs(neg) or cats(pos). Just run it like so:
|
||||
|
||||
## FLUX.1 Training
|
||||
|
||||
### Tutorial
|
||||
|
||||
To get started quickly, check out [@araminta_k](https://x.com/araminta_k) tutorial on [Finetuning Flux Dev on a 3090](https://www.youtube.com/watch?v=HzGW_Kyermg) with 24GB VRAM.
|
||||
|
||||
|
||||
### Requirements
|
||||
You currently need a GPU with **at least 24GB of VRAM** to train FLUX.1. If you are using it as your GPU to control
|
||||
your monitors, you probably need to set the flag `low_vram: true` in the config file under `model:`. This will quantize
|
||||
the model on CPU and should allow it to train with monitors attached. Users have gotten it to work on Windows with WSL,
|
||||
but there are some reports of a bug when running on windows natively.
|
||||
I have only tested on linux for now. This is still extremely experimental
|
||||
and a lot of quantizing and tricks had to happen to get it to fit on 24GB at all.
|
||||
|
||||
### FLUX.1-dev
|
||||
|
||||
FLUX.1-dev has a non-commercial license. 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.
|
||||
|
||||
1. Sign into HF and accept the model access here [black-forest-labs/FLUX.1-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev)
|
||||
2. Make a file named `.env` in the root on this folder
|
||||
3. [Get a READ key from huggingface](https://huggingface.co/settings/tokens/new?) and add it to the `.env` file like so `HF_TOKEN=your_key_here`
|
||||
|
||||
### FLUX.1-schnell
|
||||
|
||||
FLUX.1-schnell is Apache 2.0. Anything trained on it can be licensed however you want and it does not require a HF_TOKEN to train.
|
||||
However, it does require a special adapter to train with it, [ostris/FLUX.1-schnell-training-adapter](https://huggingface.co/ostris/FLUX.1-schnell-training-adapter).
|
||||
It is also highly experimental. For best overall quality, training on FLUX.1-dev is recommended.
|
||||
|
||||
To use it, You just need to add the assistant to the `model` section of your config file like so:
|
||||
|
||||
```yaml
|
||||
model:
|
||||
name_or_path: "black-forest-labs/FLUX.1-schnell"
|
||||
assistant_lora_path: "ostris/FLUX.1-schnell-training-adapter"
|
||||
is_flux: true
|
||||
quantize: true
|
||||
```
|
||||
|
||||
You also need to adjust your sample steps since schnell does not require as many
|
||||
|
||||
```yaml
|
||||
sample:
|
||||
guidance_scale: 1 # schnell does not do guidance
|
||||
sample_steps: 4 # 1 - 4 works well
|
||||
```
|
||||
|
||||
### Training
|
||||
1. Copy the example config file located at `config/examples/train_lora_flux_24gb.yaml` (`config/examples/train_lora_flux_schnell_24gb.yaml` for schnell) to the `config` folder and rename it to `whatever_you_want.yml`
|
||||
2. Edit the file following the comments in the file
|
||||
3. Run the file like so `python run.py config/whatever_you_want.yml`
|
||||
|
||||
A folder with the name and the training folder from the config file will be created when you start. It will have all
|
||||
checkpoints and images in it. You can stop the training at any time using ctrl+c and when you resume, it will pick back up
|
||||
from the last checkpoint.
|
||||
|
||||
IMPORTANT. If you press crtl+c while it is saving, it will likely corrupt that checkpoint. So wait until it is done saving
|
||||
|
||||
### Need help?
|
||||
|
||||
Please do not open a bug report unless it is a bug in the code. You are welcome to [Join my Discord](https://discord.gg/VXmU2f5WEU)
|
||||
and ask for help there. However, please refrain from PMing me directly with general question or support. Ask in the discord
|
||||
and I will answer when I can.
|
||||
|
||||
## Gradio UI
|
||||
|
||||
To get started training locally with a with a custom UI, once you followed the steps above and `ai-toolkit` is installed:
|
||||
|
||||
```bash
|
||||
python3 run.py config/examples/train_slider.example.yml
|
||||
cd ai-toolkit #in case you are not yet in the ai-toolkit folder
|
||||
huggingface-cli login #provide a `write` token to publish your LoRA at the end
|
||||
python flux_train_ui.py
|
||||
```
|
||||
|
||||
And you will be able to see how it works without configuring anything. No datasets are required for this method.
|
||||
I will post an better tutorial soon.
|
||||
You will instantiate a UI that will let you upload your images, caption them, train and publish your LoRA
|
||||

|
||||
|
||||
|
||||
## Training in RunPod
|
||||
Example RunPod template: **runpod/pytorch:2.2.0-py3.10-cuda12.1.1-devel-ubuntu22.04**
|
||||
> You need a minimum of 24GB VRAM, pick a GPU by your preference.
|
||||
|
||||
#### Example config ($0.5/hr):
|
||||
- 1x A40 (48 GB VRAM)
|
||||
- 19 vCPU 100 GB RAM
|
||||
|
||||
#### Custom overrides (you need some storage to clone FLUX.1, store datasets, store trained models and samples):
|
||||
- ~120 GB Disk
|
||||
- ~120 GB Pod Volume
|
||||
- Start Jupyter Notebook
|
||||
|
||||
### 1. Setup
|
||||
```
|
||||
git clone https://github.com/ostris/ai-toolkit.git
|
||||
cd ai-toolkit
|
||||
git submodule update --init --recursive
|
||||
python -m venv venv
|
||||
source venv/bin/activate
|
||||
pip install torch
|
||||
pip install -r requirements.txt
|
||||
pip install --upgrade accelerate transformers diffusers huggingface_hub #Optional, run it if you run into issues
|
||||
```
|
||||
### 2. Upload your dataset
|
||||
- Create a new folder in the root, name it `dataset` or whatever you like.
|
||||
- Drag and drop your .jpg, .jpeg, or .png images and .txt files inside the newly created dataset folder.
|
||||
|
||||
### 3. Login into Hugging Face with an Access Token
|
||||
- Get a READ token from [here](https://huggingface.co/settings/tokens) and request access to Flux.1-dev model from [here](https://huggingface.co/black-forest-labs/FLUX.1-dev).
|
||||
- Run ```huggingface-cli login``` and paste your token.
|
||||
|
||||
### 4. Training
|
||||
- Copy an example config file located at ```config/examples``` to the config folder and rename it to ```whatever_you_want.yml```.
|
||||
- Edit the config following the comments in the file.
|
||||
- Change ```folder_path: "/path/to/images/folder"``` to your dataset path like ```folder_path: "/workspace/ai-toolkit/your-dataset"```.
|
||||
- Run the file: ```python run.py config/whatever_you_want.yml```.
|
||||
|
||||
### Screenshot from RunPod
|
||||
<img width="1728" alt="RunPod Training Screenshot" src="https://github.com/user-attachments/assets/53a1b8ef-92fa-4481-81a7-bde45a14a7b5">
|
||||
|
||||
## Training in Modal
|
||||
|
||||
### 1. Setup
|
||||
#### ai-toolkit:
|
||||
```
|
||||
git clone https://github.com/ostris/ai-toolkit.git
|
||||
cd ai-toolkit
|
||||
git submodule update --init --recursive
|
||||
python -m venv venv
|
||||
source venv/bin/activate
|
||||
pip install torch
|
||||
pip install -r requirements.txt
|
||||
pip install --upgrade accelerate transformers diffusers huggingface_hub #Optional, run it if you run into issues
|
||||
```
|
||||
#### Modal:
|
||||
- Run `pip install modal` to install the modal Python package.
|
||||
- Run `modal setup` to authenticate (if this doesn’t work, try `python -m modal setup`).
|
||||
|
||||
#### Hugging Face:
|
||||
- Get a READ token from [here](https://huggingface.co/settings/tokens) and request access to Flux.1-dev model from [here](https://huggingface.co/black-forest-labs/FLUX.1-dev).
|
||||
- Run `huggingface-cli login` and paste your token.
|
||||
|
||||
### 2. Upload your dataset
|
||||
- Drag and drop your dataset folder containing the .jpg, .jpeg, or .png images and .txt files in `ai-toolkit`.
|
||||
|
||||
### 3. Configs
|
||||
- Copy an example config file located at ```config/examples/modal``` to the `config` folder and rename it to ```whatever_you_want.yml```.
|
||||
- Edit the config following the comments in the file, **<ins>be careful and follow the example `/root/ai-toolkit` paths</ins>**.
|
||||
|
||||
### 4. Edit run_modal.py
|
||||
- Set your entire local `ai-toolkit` path at `code_mount = modal.Mount.from_local_dir` like:
|
||||
|
||||
```
|
||||
code_mount = modal.Mount.from_local_dir("/Users/username/ai-toolkit", remote_path="/root/ai-toolkit")
|
||||
```
|
||||
- Choose a `GPU` and `Timeout` in `@app.function` _(default is A100 40GB and 2 hour timeout)_.
|
||||
|
||||
### 5. Training
|
||||
- Run the config file in your terminal: `modal run run_modal.py --config-file-list-str=/root/ai-toolkit/config/whatever_you_want.yml`.
|
||||
- You can monitor your training in your local terminal, or on [modal.com](https://modal.com/).
|
||||
- Models, samples and optimizer will be stored in `Storage > flux-lora-models`.
|
||||
|
||||
### 6. Saving the model
|
||||
- Check contents of the volume by running `modal volume ls flux-lora-models`.
|
||||
- Download the content by running `modal volume get flux-lora-models your-model-name`.
|
||||
- Example: `modal volume get flux-lora-models my_first_flux_lora_v1`.
|
||||
|
||||
### Screenshot from Modal
|
||||
|
||||
<img width="1728" alt="Modal Traning Screenshot" src="https://github.com/user-attachments/assets/7497eb38-0090-49d6-8ad9-9c8ea7b5388b">
|
||||
|
||||
---
|
||||
|
||||
## WIP Tools
|
||||
## Dataset Preparation
|
||||
|
||||
Datasets generally need to be a folder containing images and associated text files. Currently, the only supported
|
||||
formats are jpg, jpeg, and png. Webp currently has issues. The text files should be named the same as the images
|
||||
but with a `.txt` extension. For example `image2.jpg` and `image2.txt`. The text file should contain only the caption.
|
||||
You can add the word `[trigger]` in the caption file and if you have `trigger_word` in your config, it will be automatically
|
||||
replaced.
|
||||
|
||||
Images are never upscaled but they are downscaled and placed in buckets for batching. **You do not need to crop/resize your images**.
|
||||
The loader will automatically resize them and can handle varying aspect ratios.
|
||||
|
||||
|
||||
### VAE (Variational Auto Encoder) Trainer
|
||||
## Training Specific Layers
|
||||
|
||||
This works, but is not ready for others to use and therefore does not have an example config.
|
||||
I am still working on it. I will update this when it is ready.
|
||||
I am adding a lot of features for criteria that I have used in my image enlargement work. A Critic (discriminator),
|
||||
content loss, style loss, and a few more. If you don't know, the VAE
|
||||
for stable diffusion (yes even the MSE one, and SDXL), are horrible at smaller faces and it holds SD back. I will fix this.
|
||||
I'll post more about this later with better examples later, but here is a quick test of a run through with various VAEs.
|
||||
Just went in and out. It is much worse on smaller faces than shown here.
|
||||
To train specific layers with LoRA, you can use the `only_if_contains` network kwargs. For instance, if you want to train only the 2 layers
|
||||
used by The Last Ben, [mentioned in this post](https://x.com/__TheBen/status/1829554120270987740), you can adjust your
|
||||
network kwargs like so:
|
||||
|
||||
<img src="https://raw.githubusercontent.com/ostris/ai-toolkit/main/assets/VAE_test1.jpg" width="768" height="auto">
|
||||
```yaml
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 128
|
||||
linear_alpha: 128
|
||||
network_kwargs:
|
||||
only_if_contains:
|
||||
- "transformer.single_transformer_blocks.7.proj_out"
|
||||
- "transformer.single_transformer_blocks.20.proj_out"
|
||||
```
|
||||
|
||||
---
|
||||
The naming conventions of the layers are in diffusers format, so checking the state dict of a model will reveal
|
||||
the suffix of the name of the layers you want to train. You can also use this method to only train specific groups of weights.
|
||||
For instance to only train the `single_transformer` for FLUX.1, you can use the following:
|
||||
|
||||
## TODO
|
||||
- [X] Add proper regs on sliders
|
||||
- [ ] Add SDXL support (base model only for now)
|
||||
- [ ] Add plain erasing
|
||||
- [ ] Make Textual inversion network trainer (network that spits out TI embeddings)
|
||||
```yaml
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 128
|
||||
linear_alpha: 128
|
||||
network_kwargs:
|
||||
only_if_contains:
|
||||
- "transformer.single_transformer_blocks."
|
||||
```
|
||||
|
||||
---
|
||||
You can also exclude layers by their names by using `ignore_if_contains` network kwarg. So to exclude all the single transformer blocks,
|
||||
|
||||
## Change Log
|
||||
#### 2021-07-30
|
||||
Added "anchors" to the slider trainer. This allows you to set a prompt that will be used as a
|
||||
regularizer. You can set the network multiplier to force spread consistency at high weights
|
||||
|
||||
```yaml
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 128
|
||||
linear_alpha: 128
|
||||
network_kwargs:
|
||||
ignore_if_contains:
|
||||
- "transformer.single_transformer_blocks."
|
||||
```
|
||||
|
||||
`ignore_if_contains` takes priority over `only_if_contains`. So if a weight is covered by both,
|
||||
if will be ignored.
|
||||
|
||||
## LoKr Training
|
||||
|
||||
To learn more about LoKr, read more about it at [KohakuBlueleaf/LyCORIS](https://github.com/KohakuBlueleaf/LyCORIS/blob/main/docs/Guidelines.md). To train a LoKr model, you can adjust the network type in the config file like so:
|
||||
|
||||
```yaml
|
||||
network:
|
||||
type: "lokr"
|
||||
lokr_full_rank: true
|
||||
lokr_factor: 8
|
||||
```
|
||||
|
||||
Everything else should work the same including layer targeting.
|
||||
|
||||
|
||||
## Updates
|
||||
|
||||
Only larger updates are listed here. There are usually smaller daily updated that are omitted.
|
||||
|
||||
### Jul 17, 2025
|
||||
- Make it easy to add control images to the samples in the ui
|
||||
|
||||
### Jul 11, 2025
|
||||
- Added better video config settings to the UI for video models.
|
||||
- Added Wan I2V training to the UI
|
||||
|
||||
### June 29, 2025
|
||||
- Fixed issue where Kontext forced sizes on sampling
|
||||
|
||||
### June 26, 2025
|
||||
- Added support for FLUX.1 Kontext training
|
||||
- added support for instruction dataset training
|
||||
|
||||
### June 25, 2025
|
||||
- Added support for OmniGen2 training
|
||||
-
|
||||
### June 17, 2025
|
||||
- Performance optimizations for batch preparation
|
||||
- Added some docs via a popup for items in the simple ui explaining what settings do. Still a WIP
|
||||
|
||||
### June 16, 2025
|
||||
- Hide control images in the UI when viewing datasets
|
||||
- WIP on mean flow loss
|
||||
|
||||
### June 12, 2025
|
||||
- Fixed issue that resulted in blank captions in the dataloader
|
||||
|
||||
### June 10, 2025
|
||||
- Decided to keep track up updates in the readme
|
||||
- Added support for SDXL in the UI
|
||||
- Added support for SD 1.5 in the UI
|
||||
- Fixed UI Wan 2.1 14b name bug
|
||||
- Added support for for conv training in the UI for models that support it
|
||||
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"
|
||||
21
build_and_push_docker_dev
Normal file
21
build_and_push_docker_dev
Normal file
@@ -0,0 +1,21 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
VERSION=dev
|
||||
GIT_COMMIT=dev
|
||||
|
||||
echo "Docker builds from the repo, not this dir. Make sure changes are pushed to the repo."
|
||||
echo "Building version: $VERSION"
|
||||
# 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
|
||||
|
||||
# Push both tags
|
||||
echo "Pushing images to Docker Hub..."
|
||||
docker push ostris/aitoolkit:$VERSION
|
||||
|
||||
echo "Successfully built and pushed ostris/aitoolkit:$VERSION"
|
||||
60
config/examples/generate.example.yaml
Normal file
60
config/examples/generate.example.yaml
Normal file
@@ -0,0 +1,60 @@
|
||||
---
|
||||
|
||||
job: generate # tells the runner what to do
|
||||
config:
|
||||
name: "generate" # this is not really used anywhere currently but required by runner
|
||||
process:
|
||||
# process 1
|
||||
- type: to_folder # process images to a folder
|
||||
output_folder: "output/gen"
|
||||
device: cuda:0 # cpu, cuda:0, etc
|
||||
generate:
|
||||
# these are your defaults you can override most of them with flags
|
||||
sampler: "ddpm" # ignored for now, will add later though ddpm is used regardless for now
|
||||
width: 1024
|
||||
height: 1024
|
||||
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime"
|
||||
seed: -1 # -1 is random
|
||||
guidance_scale: 7
|
||||
sample_steps: 20
|
||||
ext: ".png" # .png, .jpg, .jpeg, .webp
|
||||
|
||||
# here ate the flags you can use for prompts. Always start with
|
||||
# your prompt first then add these flags after. You can use as many
|
||||
# like
|
||||
# photo of a baseball --n painting, ugly --w 1024 --h 1024 --seed 42 --cfg 7 --steps 20
|
||||
# we will try to support all sd-scripts flags where we can
|
||||
|
||||
# FROM SD-SCRIPTS
|
||||
# --n Treat everything until the next option as a negative prompt.
|
||||
# --w Specify the width of the generated image.
|
||||
# --h Specify the height of the generated image.
|
||||
# --d Specify the seed for the generated image.
|
||||
# --l Specify the CFG scale for the generated image.
|
||||
# --s Specify the number of steps during generation.
|
||||
|
||||
# OURS and some QOL additions
|
||||
# --p2 Prompt for the second text encoder (SDXL only)
|
||||
# --n2 Negative prompt for the second text encoder (SDXL only)
|
||||
# --gr Specify the guidance rescale for the generated image (SDXL only)
|
||||
# --seed Specify the seed for the generated image same as --d
|
||||
# --cfg Specify the CFG scale for the generated image same as --l
|
||||
# --steps Specify the number of steps during generation same as --s
|
||||
|
||||
prompt_file: false # if true a txt file will be created next to images with prompt strings used
|
||||
# prompts can also be a path to a text file with one prompt per line
|
||||
# prompts: "/path/to/prompts.txt"
|
||||
prompts:
|
||||
- "photo of batman"
|
||||
- "photo of superman"
|
||||
- "photo of spiderman"
|
||||
- "photo of a superhero --n batman superman spiderman"
|
||||
|
||||
model:
|
||||
# huggingface name, relative prom project path, or absolute path to .safetensors or .ckpt
|
||||
# name_or_path: "runwayml/stable-diffusion-v1-5"
|
||||
name_or_path: "/mnt/Models/stable-diffusion/models/stable-diffusion/Ostris/Ostris_Real_v1.safetensors"
|
||||
is_v2: false # for v2 models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
is_xl: false # for SDXL models
|
||||
dtype: bf16
|
||||
48
config/examples/mod_lora_scale.yaml
Normal file
48
config/examples/mod_lora_scale.yaml
Normal file
@@ -0,0 +1,48 @@
|
||||
---
|
||||
job: mod
|
||||
config:
|
||||
name: name_of_your_model_v1
|
||||
process:
|
||||
- type: rescale_lora
|
||||
# path to your current lora model
|
||||
input_path: "/path/to/lora/lora.safetensors"
|
||||
# output path for your new lora model, can be the same as input_path to replace
|
||||
output_path: "/path/to/lora/output_lora_v1.safetensors"
|
||||
# replaces meta with the meta below (plus minimum meta fields)
|
||||
# if false, we will leave the meta alone except for updating hashes (sd-script hashes)
|
||||
replace_meta: true
|
||||
# how to adjust, we can scale the up_down weights or the alpha
|
||||
# up_down is the default and probably the best, they will both net the same outputs
|
||||
# would only affect rare NaN cases and maybe merging with old merge tools
|
||||
scale_target: 'up_down'
|
||||
# precision to save, fp16 is the default and standard
|
||||
save_dtype: fp16
|
||||
# current_weight is the ideal weight you use as a multiplier when using the lora
|
||||
# IE in automatic1111 <lora:my_lora:6.0> the 6.0 is the current_weight
|
||||
# you can do negatives here too if you want to flip the lora
|
||||
current_weight: 6.0
|
||||
# target_weight is the ideal weight you use as a multiplier when using the lora
|
||||
# instead of the one above. IE in automatic1111 instead of using <lora:my_lora:6.0>
|
||||
# we want to use <lora:my_lora:1.0> so 1.0 is the target_weight
|
||||
target_weight: 1.0
|
||||
|
||||
# base model for the lora
|
||||
# this is just used to add meta so automatic111 knows which model it is for
|
||||
# assume v1.5 if these are not set
|
||||
is_xl: false
|
||||
is_v2: false
|
||||
meta:
|
||||
# this is only used if you set replace_meta to true above
|
||||
name: "[name]" # [name] gets replaced with the name above
|
||||
description: A short description of your lora
|
||||
trigger_words:
|
||||
- put
|
||||
- trigger
|
||||
- words
|
||||
- here
|
||||
version: '0.1'
|
||||
creator:
|
||||
name: Your Name
|
||||
email: your@email.com
|
||||
website: https://yourwebsite.com
|
||||
any: All meta data above is arbitrary, it can be whatever you want.
|
||||
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'
|
||||
112
config/examples/train_flex_redux.yaml
Normal file
112
config/examples/train_flex_redux.yaml
Normal file
@@ -0,0 +1,112 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flex_redux_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
|
||||
adapter:
|
||||
type: "redux"
|
||||
# you can finetune an existing adapter or start from scratch. Set to null to start from scratch
|
||||
name_or_path: '/local/path/to/redux_adapter_to_finetune.safetensors'
|
||||
# name_or_path: null
|
||||
# image_encoder_path: 'google/siglip-so400m-patch14-384' # Flux.1 redux adapter
|
||||
image_encoder_path: 'google/siglip2-so400m-patch16-512' # Flex.1 512 redux adapter
|
||||
# image_encoder_arch: 'siglip' # for Flux.1
|
||||
image_encoder_arch: 'siglip2'
|
||||
# You need a control input for each sample. Best to do squares for both images
|
||||
test_img_path:
|
||||
- "/path/to/x_01.jpg"
|
||||
- "/path/to/x_02.jpg"
|
||||
- "/path/to/x_03.jpg"
|
||||
- "/path/to/x_04.jpg"
|
||||
- "/path/to/x_05.jpg"
|
||||
- "/path/to/x_06.jpg"
|
||||
- "/path/to/x_07.jpg"
|
||||
- "/path/to/x_08.jpg"
|
||||
- "/path/to/x_09.jpg"
|
||||
- "/path/to/x_10.jpg"
|
||||
clip_layer: 'last_hidden_state'
|
||||
train: true
|
||||
save:
|
||||
dtype: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4
|
||||
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"
|
||||
# clip_image_path is directory containting your control images. They must have filename as their train image. (extension does not matter)
|
||||
# for normal redux, we are just recreating the same image, so you can use the same folder path above
|
||||
clip_image_path: "/path/to/control/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
resolution: [ 512, 768, 1024 ] # flex enjoys multiple resolutions
|
||||
train:
|
||||
# this is what I used for the 24GB card, but feel free to adjust
|
||||
# total batch size is 6 here
|
||||
batch_size: 3
|
||||
gradient_accumulation: 2
|
||||
|
||||
# captions are not needed for this training, we cache a blank proompt and rely on the vision encoder
|
||||
unload_text_encoder: true
|
||||
|
||||
loss_type: "mse"
|
||||
train_unet: true
|
||||
train_text_encoder: false
|
||||
steps: 4000000 # I set this very high and stop when I like the results
|
||||
content_or_style: balanced # content, style, balanced
|
||||
gradient_checkpointing: true
|
||||
noise_scheduler: "flowmatch" # or "ddpm", "lms", "euler_a"
|
||||
timestep_type: "flux_shift"
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
|
||||
# this is for Flex.1, comment this out for FLUX.1-dev
|
||||
bypass_guidance_embedding: true
|
||||
|
||||
dtype: bf16
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
model:
|
||||
name_or_path: "ostris/Flex.1-alpha"
|
||||
is_flux: true
|
||||
quantize: true
|
||||
text_encoder_bits: 8
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
# I leave half blank to test prompt and unprompted
|
||||
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"
|
||||
- ""
|
||||
- ""
|
||||
- ""
|
||||
- ""
|
||||
- ""
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
network_multiplier: 1.0
|
||||
|
||||
# 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'
|
||||
104
config/examples/train_lora_chroma_24gb.yaml
Normal file
104
config/examples/train_lora_chroma_24gb.yaml
Normal file
@@ -0,0 +1,104 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_chroma_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 ] # chroma enjoys multiple resolutions
|
||||
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 chroma
|
||||
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 chroma, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# Download the whichever model you prefer from the Chroma repo
|
||||
# https://huggingface.co/lodestones/Chroma/tree/main
|
||||
# point to it here.
|
||||
# name_or_path: "/path/to/chroma/chroma-unlocked-vVERSION.safetensors"
|
||||
|
||||
# using lodestones/Chroma will automatically use the latest version
|
||||
name_or_path: "lodestones/Chroma"
|
||||
|
||||
# # You can also select a version of Chroma like so
|
||||
# name_or_path: "lodestones/Chroma/v28"
|
||||
|
||||
arch: "chroma"
|
||||
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: "" # negative prompt, optional
|
||||
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'
|
||||
165
config/examples/train_lora_flex2_24gb.yaml
Normal file
165
config/examples/train_lora_flex2_24gb.yaml
Normal file
@@ -0,0 +1,165 @@
|
||||
# Note, Flex2 is a highly experimental WIP model. Finetuning a model with built in controls and inpainting has not
|
||||
# been done before, so you will be experimenting with me on how to do it. This is my recommended setup, but this is highly
|
||||
# subject to change as we learn more about how Flex2 works.
|
||||
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flex2_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"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
# Flex2 is trained with controls and inpainting. If you want the model to truely understand how the
|
||||
# controls function with your dataset, it is a good idea to keep doing controls during training.
|
||||
# this will automatically generate the controls for you before training. The current script is not
|
||||
# fully optimized so this could be rather slow for large datasets, but it caches them to disk so it
|
||||
# only needs to be done once. If you want to skip this step, you can set the controls to [] and it will
|
||||
controls:
|
||||
- "depth"
|
||||
- "line"
|
||||
- "pose"
|
||||
- "inpaint"
|
||||
|
||||
# you can make custom inpainting images as well. These images must be webp or png format with an alpha.
|
||||
# just erase the part of the image you want to inpaint and save it as a webp or png. Again, erase your
|
||||
# train target. So the person if training a person. The automatic controls above with inpaint will
|
||||
# just run a background remover mask and erase the foreground, which works well for subjects.
|
||||
|
||||
# inpaint_path: "/my/impaint/images"
|
||||
|
||||
# you can also specify existing control image pairs. It can handle multiple groups and will randomly
|
||||
# select one for each step.
|
||||
|
||||
# control_path:
|
||||
# - "/my/custom/control/images"
|
||||
# - "/my/custom/control/images2"
|
||||
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
resolution: [ 512, 768, 1024 ] # flex2 enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
# IMPORTANT! For Flex2, you must bypass the guidance embedder during training
|
||||
bypass_guidance_embedding: true
|
||||
|
||||
steps: 3000 # 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 flex2
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
# shift works well for training fast and learning composition and style.
|
||||
# for just subject, you may want to change this to sigmoid
|
||||
timestep_type: 'shift' # 'linear', 'sigmoid', 'shift'
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
|
||||
optimizer_params:
|
||||
weight_decay: 1e-5
|
||||
# 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. Defaults off
|
||||
ema_config:
|
||||
use_ema: false
|
||||
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.2-preview"
|
||||
arch: "flex2"
|
||||
quantize: true # run 8bit mixed precision
|
||||
quantize_te: true
|
||||
|
||||
# you can pass special training infor for controls to the model here
|
||||
# percentages are decimal based so 0.0 is 0% and 1.0 is 100% of the time.
|
||||
model_kwargs:
|
||||
# inverts the inpainting mask, good to learn outpainting as well, recommended 0.0 for characters
|
||||
invert_inpaint_mask_chance: 0.5
|
||||
# this will do a normal t2i training step without inpaint when dropped out. REcommended if you want
|
||||
# your lora to be able to inference with and without inpainting.
|
||||
inpaint_dropout: 0.5
|
||||
# randomly drops out the control image. Dropout recvommended if your want it to work without controls as well.
|
||||
control_dropout: 0.5
|
||||
# does a random inpaint blob. Usually a good idea to keep. Without it, the model will learn to always 100%
|
||||
# fill the inpaint area with your subject. This is not always a good thing.
|
||||
inpaint_random_chance: 0.5
|
||||
# generates random inpaint blobs if you did not provide an inpaint image for your dataset. Inpaint breaks down fast
|
||||
# if you are not training with it. Controls are a little more robust and can be left out,
|
||||
# but when in doubt, always leave this on
|
||||
do_random_inpainting: false
|
||||
# does random blurring of the inpaint mask. Helps prevent weird edge artifacts for real workd inpainting. Leave on.
|
||||
random_blur_mask: true
|
||||
# applies a small amount of random dialition and restriction to the inpaint mask. Helps with edge artifacts.
|
||||
# Leave on.
|
||||
random_dialate_mask: 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!'"\
|
||||
|
||||
# you can use a single inpaint or single control image on your samples.
|
||||
# for controls, the ctrl_idx is 1, the images can be any name and image format.
|
||||
# use either a pose/line/depth image or whatever you are training with. An example is
|
||||
# - "photo of [trigger] --ctrl_idx 1 --ctrl_img /path/to/control/image.jpg"
|
||||
|
||||
# for an inpainting image, it must be png/webp. Erase the part of the image you want to inpaint
|
||||
# IMPORTANT! the inpaint images must be ctrl_idx 0 and have .inpaint.{ext} in the name for this to work right.
|
||||
# - "photo of [trigger] --ctrl_idx 0 --ctrl_img /path/to/inpaint/image.inpaint.png"
|
||||
|
||||
- "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 flex2
|
||||
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_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'
|
||||
106
config/examples/train_lora_flux_kontext_24gb.yaml
Normal file
106
config/examples/train_lora_flux_kontext_24gb.yaml
Normal file
@@ -0,0 +1,106 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flux_kontext_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"
|
||||
# control path is the input images for kontext for a paired dataset. These are the source images you want to change.
|
||||
# You can comment this out and only use normal images if you don't have a paired dataset.
|
||||
# Control images need to match the filenames on the folder path but in
|
||||
# a different folder. These do not need captions.
|
||||
control_path: "/path/to/control/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
|
||||
# Kontext runs images in at 2x the latent size. It may OOM at 1024 resolution with 24GB vram.
|
||||
resolution: [ 512, 768 ] # flux enjoys multiple resolutions
|
||||
# resolution: [ 512, 768, 1024 ]
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 3000 # 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
|
||||
timestep_type: "weighted" # sigmoid, linear, or weighted.
|
||||
# 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.
|
||||
|
||||
# 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. This model is gated.
|
||||
# visit https://huggingface.co/black-forest-labs/FLUX.1-Kontext-dev to accept the terms and conditions
|
||||
# and then you can use this model.
|
||||
name_or_path: "black-forest-labs/FLUX.1-Kontext-dev"
|
||||
arch: "flux_kontext"
|
||||
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
|
||||
# the --ctrl_img path is the one loaded to apply the kontext editing to
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "make the person smile --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
- "give the person an afro --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
- "turn this image into a cartoon --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
- "put this person in an action film --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
- "make this person a rapper in a rap music video --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
- "make the person smile --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
- "give the person an afro --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
- "turn this image into a cartoon --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
- "put this person in an action film --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
- "make this person a rapper in a rap music video --ctrl_img /path/to/control/folder/person1.jpg"
|
||||
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'
|
||||
112
config/examples/train_lora_hidream_48.yaml
Normal file
112
config/examples/train_lora_hidream_48.yaml
Normal file
@@ -0,0 +1,112 @@
|
||||
# HiDream training is still highly experimental. The settings here will take ~35.2GB of vram to train.
|
||||
# It is not possible to train on a single 24GB card yet, but I am working on it. If you have more VRAM
|
||||
# I highly recommend first disabling quantization on the model itself if you can. You can leave the TEs quantized.
|
||||
# HiDream has a mixture of experts that may take special training considerations that I do not
|
||||
# have implemented properly. The current implementation seems to work well for LoRA training, but
|
||||
# may not be effective for longer training runs. The implementation could change in future updates
|
||||
# so your results may vary when this happens.
|
||||
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_hidream_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
|
||||
network_kwargs:
|
||||
# it is probably best to ignore the mixture of experts since only 2 are active each block. It works activating it, but I wouldnt.
|
||||
# proper training of it is not fully implemented
|
||||
ignore_if_contains:
|
||||
- "ff_i.experts"
|
||||
- "ff_i.gate"
|
||||
save:
|
||||
dtype: bfloat16 # 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"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
resolution: [ 512, 768, 1024 ] # hidream enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 3000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # wont work with hidream
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
timestep_type: shift # sigmoid, shift, linear
|
||||
optimizer: "adamw8bit"
|
||||
lr: 2e-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. Defaults off
|
||||
ema_config:
|
||||
use_ema: false
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for hidream, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# the transformer will get grabbed from this hf repo
|
||||
# warning ONLY train on Full. The dev and fast models are distilled and will break
|
||||
name_or_path: "HiDream-ai/HiDream-I1-Full"
|
||||
# the extras will be grabbed from this hf repo. (text encoder, vae)
|
||||
extras_name_or_path: "HiDream-ai/HiDream-I1-Full"
|
||||
arch: "hidream"
|
||||
# both need to be quantized to train on 48GB currently
|
||||
quantize: true
|
||||
quantize_te: true
|
||||
model_kwargs:
|
||||
# llama is a gated model, It defaults to unsloth version, but you can set the llama path here
|
||||
llama_model_path: "unsloth/Meta-Llama-3.1-8B-Instruct"
|
||||
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'
|
||||
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'
|
||||
94
config/examples/train_lora_omnigen2_24gb.yaml
Normal file
94
config/examples/train_lora_omnigen2_24gb.yaml
Normal file
@@ -0,0 +1,94 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_omnigen2_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 ] # omnigen2 should work with multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 3000 # 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 omnigen2
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
timestep_type: 'sigmoid' # sigmoid, linear, shift
|
||||
# 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.
|
||||
# ema_config:
|
||||
# use_ema: true
|
||||
# ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for omnigen2, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
name_or_path: "OmniGen2/OmniGen2
|
||||
arch: "omnigen2"
|
||||
quantize_te: true # quantize_only te
|
||||
# quantize: true # quantize transformer
|
||||
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: "" # negative prompt, optional
|
||||
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'
|
||||
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'
|
||||
@@ -7,7 +7,7 @@ job: train
|
||||
config:
|
||||
# the name will be used to create a folder in the output folder
|
||||
# it will also replace any [name] token in the rest of this config
|
||||
name: pet_slider_v1
|
||||
name: detail_slider_v1
|
||||
# folder will be created with name above in folder below
|
||||
# it can be relative to the project root or absolute
|
||||
training_folder: "output/LoRA"
|
||||
@@ -23,9 +23,8 @@ config:
|
||||
# network type lierla is traditional LoRA that works everywhere, only linear layers
|
||||
type: "lierla"
|
||||
# rank / dim of the network. Bigger is not always better. Especially for sliders. 8 is good
|
||||
rank: 8
|
||||
alpha: 1.0 # just leave it
|
||||
|
||||
linear: 8
|
||||
linear_alpha: 4 # Do about half of rank
|
||||
# training config
|
||||
train:
|
||||
# this is also used in sampling. Stick with ddpm unless you know what you are doing
|
||||
@@ -33,14 +32,17 @@ config:
|
||||
# how many steps to train. More is not always better. I rarely go over 1000
|
||||
steps: 500
|
||||
# I have had good results with 4e-4 to 1e-4 at 500 steps
|
||||
lr: 1e-4
|
||||
lr: 2e-4
|
||||
# enables gradient checkpoint, saves vram, leave it on
|
||||
gradient_checkpointing: true
|
||||
# train the unet. I recommend leaving this true
|
||||
train_unet: true
|
||||
# train the text encoder. I don't recommend this unless you have a special use case
|
||||
# for sliders we are adjusting representation of the concept (unet),
|
||||
# not the description of it (text encoder)
|
||||
train_text_encoder: false
|
||||
|
||||
# same as from sd-scripts, not fully tested but should speed up training
|
||||
min_snr_gamma: 5.0
|
||||
# just leave unless you know what you are doing
|
||||
# also supports "dadaptation" but set lr to 1 if you use that,
|
||||
# but it learns too fast and I don't recommend it
|
||||
@@ -51,14 +53,17 @@ config:
|
||||
# while training. Just leave it
|
||||
max_denoising_steps: 40
|
||||
# works great at 1. I do 1 even with my 4090.
|
||||
# higher may not work right with newer single batch stacking code anyway
|
||||
batch_size: 1
|
||||
# bf16 works best if your GPU supports it (modern)
|
||||
dtype: bf16 # fp32, bf16, fp16
|
||||
# if you have it, use it. It is faster and better
|
||||
xformers: true
|
||||
# torch 2.0 doesnt need xformers anymore, only use if you have lower version
|
||||
# xformers: true
|
||||
# I don't recommend using unless you are trying to make a darker lora. Then do 0.1 MAX
|
||||
# although, the way we train sliders is comparative, so it probably won't work anyway
|
||||
noise_offset: 0.0
|
||||
# noise_offset: 0.0357 # SDXL was trained with offset of 0.0357. So use that when training on SDXL
|
||||
|
||||
# the model to train the LoRA network on
|
||||
model:
|
||||
@@ -66,11 +71,17 @@ config:
|
||||
name_or_path: "runwayml/stable-diffusion-v1-5"
|
||||
is_v2: false # for v2 models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
# has some issues with the dual text encoder and the way we train sliders
|
||||
# it works bit weights need to probably be higher to see it.
|
||||
is_xl: false # for SDXL models
|
||||
|
||||
# saving config
|
||||
save:
|
||||
dtype: float16 # precision to save. I recommend float16
|
||||
save_every: 50 # save every this many steps
|
||||
# this will remove step counts more than this number
|
||||
# allows you to save more often in case of a crash without filling up your drive
|
||||
max_step_saves_to_keep: 2
|
||||
|
||||
# sampling config
|
||||
sample:
|
||||
@@ -88,21 +99,22 @@ config:
|
||||
# --m [number] # network multiplier. LoRA weight. -3 for the negative slide, 3 for the positive
|
||||
# slide are good tests. will inherit sample.network_multiplier if not set
|
||||
# --n [string] # negative prompt, will inherit sample.neg if not set
|
||||
|
||||
# Only 75 tokens allowed currently
|
||||
prompts: # our example is an animal slider, neg: dog, pos: cat
|
||||
- "a golden retriever --m -5"
|
||||
- "a golden retriever --m -3"
|
||||
- "a golden retriever --m 3"
|
||||
- "a golden retriever --m 5"
|
||||
- "calico cat --m -5"
|
||||
- "calico cat --m -3"
|
||||
- "calico cat --m 3"
|
||||
- "calico cat --m 5"
|
||||
- "an elephant --m -5"
|
||||
- "an elephant --m -3"
|
||||
- "an elephant --m 3"
|
||||
- "an elephant --m 5"
|
||||
# I like to do a wide positive and negative spread so I can see a good range and stop
|
||||
# early if the network is braking down
|
||||
prompts:
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -5"
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -3"
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 3"
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 5"
|
||||
- "a golden retriever sitting on a leather couch, --m -5"
|
||||
- "a golden retriever sitting on a leather couch --m -3"
|
||||
- "a golden retriever sitting on a leather couch --m 3"
|
||||
- "a golden retriever sitting on a leather couch --m 5"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -5"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -3"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 3"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 5"
|
||||
# negative prompt used on all prompts above as default if they don't have one
|
||||
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime, monochrome"
|
||||
# seed for sampling. 42 is the answer for everything
|
||||
@@ -131,11 +143,16 @@ config:
|
||||
# resolutions to train on. [ width, height ]. This is less important for sliders
|
||||
# as we are not teaching the model anything it doesn't already know
|
||||
# but must be a size it understands [ 512, 512 ] for sd_v1.5 and [ 768, 768 ] for sd_v2.1
|
||||
# and [ 1024, 1024 ] for sd_xl
|
||||
# you can do as many as you want here
|
||||
resolutions:
|
||||
- [ 512, 512 ]
|
||||
# - [ 512, 768 ]
|
||||
# - [ 768, 768 ]
|
||||
# slider training uses 4 combined steps for a single round. This will do it in one gradient
|
||||
# step. It is highly optimized and shouldn't take anymore vram than doing without it,
|
||||
# since we break down batches for gradient accumulation now. so just leave it on.
|
||||
batch_full_slide: true
|
||||
# These are the concepts to train on. You can do as many as you want here,
|
||||
# but they can conflict outweigh each other. Other than experimenting, I recommend
|
||||
# just doing one for good results
|
||||
@@ -146,7 +163,9 @@ config:
|
||||
# a keyword necessarily but what the model understands the concept to represent.
|
||||
# "person" will affect men, women, children, etc but will not affect cats, dogs, etc
|
||||
# it is the models base general understanding of the concept and everything it represents
|
||||
- target_class: "animal"
|
||||
# you can leave it blank to affect everything. In this example, we are adjusting
|
||||
# detail, so we will leave it blank to affect everything
|
||||
- target_class: ""
|
||||
# positive is the prompt for the positive side of the slider.
|
||||
# It is the concept that will be excited and amplified in the model when we slide the slider
|
||||
# to the positive side and forgotten / inverted when we slide
|
||||
@@ -154,33 +173,48 @@ config:
|
||||
# the prompt. You want it to be the extreme of what you want to train on. For example,
|
||||
# if you want to train on fat people, you would use "an extremely fat, morbidly obese person"
|
||||
# as the prompt. Not just "fat person"
|
||||
positive: "cat"
|
||||
# max 75 tokens for now
|
||||
positive: "high detail, 8k, intricate, detailed, high resolution, high res, high quality"
|
||||
# negative is the prompt for the negative side of the slider and works the same as positive
|
||||
# it does not necessarily work the same as a negative prompt when generating images
|
||||
negative: "dog"
|
||||
# these need to be polar opposites.
|
||||
# max 76 tokens for now
|
||||
negative: "blurry, boring, fuzzy, low detail, low resolution, low res, low quality"
|
||||
# the loss for this target is multiplied by this number.
|
||||
# if you are doing more than one target it may be good to set less important ones
|
||||
# to a lower number like 0.1 so they dont outweigh the primary target
|
||||
# to a lower number like 0.1 so they don't outweigh the primary target
|
||||
weight: 1.0
|
||||
# shuffle the prompts split by the comma. We will run every combination randomly
|
||||
# this will make the LoRA more robust. You probably want this on unless prompt order
|
||||
# is important for some reason
|
||||
shuffle: true
|
||||
|
||||
# anchors are prompts that wer try to hold on to while training the slider
|
||||
# you want these to generate an image very similar to the target_class
|
||||
# without directly overlapping it. For example, if you are training on a person smiling,
|
||||
# you would use "a person with a face mask" as an anchor. It is a person, the image is the same
|
||||
# regardless if they are smiling or not
|
||||
anchors:
|
||||
# only positive prompt for now
|
||||
- prompt: "a woman"
|
||||
neg_prompt: "animal"
|
||||
# the multiplier applied to the LoRA when this is run.
|
||||
# higher will give it more weight but also help keep the lora from collapsing
|
||||
multiplier: 8.0
|
||||
- prompt: "a man"
|
||||
neg_prompt: "animal"
|
||||
multiplier: 8.0
|
||||
- prompt: "a person"
|
||||
neg_prompt: "animal"
|
||||
multiplier: 8.0
|
||||
|
||||
# anchors are prompts that we will try to hold on to while training the slider
|
||||
# these are NOT necessary and can prevent the slider from converging if not done right
|
||||
# leave them off if you are having issues, but they can help lock the network
|
||||
# on certain concepts to help prevent catastrophic forgetting
|
||||
# you want these to generate an image that is not your target_class, but close to it
|
||||
# is fine as long as it does not directly overlap it.
|
||||
# For example, if you are training on a person smiling,
|
||||
# you could use "a person with a face mask" as an anchor. It is a person, the image is the same
|
||||
# regardless if they are smiling or not, however, the closer the concept is to the target_class
|
||||
# the less the multiplier needs to be. Keep multipliers less than 1.0 for anchors usually
|
||||
# for close concepts, you want to be closer to 0.1 or 0.2
|
||||
# these will slow down training. I am leaving them off for the demo
|
||||
|
||||
# anchors:
|
||||
# - prompt: "a woman"
|
||||
# neg_prompt: "animal"
|
||||
# # the multiplier applied to the LoRA when this is run.
|
||||
# # higher will give it more weight but also help keep the lora from collapsing
|
||||
# multiplier: 1.0
|
||||
# - prompt: "a man"
|
||||
# neg_prompt: "animal"
|
||||
# multiplier: 1.0
|
||||
# - prompt: "a person"
|
||||
# neg_prompt: "animal"
|
||||
# multiplier: 1.0
|
||||
|
||||
# You can put any information you want here, and it will be saved in the model.
|
||||
# The below is an example, but you can put your grocery list in it if you want.
|
||||
|
||||
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]
|
||||
83
docker/Dockerfile
Normal file
83
docker/Dockerfile
Normal file
@@ -0,0 +1,83 @@
|
||||
FROM nvidia/cuda:12.8.1-devel-ubuntu22.04
|
||||
|
||||
LABEL authors="jaret"
|
||||
|
||||
# Set noninteractive to avoid timezone prompts
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# ref https://en.wikipedia.org/wiki/CUDA
|
||||
ENV TORCH_CUDA_ARCH_LIST="8.0 8.6 8.9 9.0 10.0 12.0"
|
||||
|
||||
# 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 \
|
||||
openssh-client \
|
||||
openssh-server \
|
||||
openssl \
|
||||
rsync \
|
||||
unzip \
|
||||
&& 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 --pre --no-cache-dir torch torchvision torchaudio --index-url https://download.pytorch.org/whl/nightly/cu128
|
||||
|
||||
# Fix cache busting by moving CACHEBUST to right before git clone
|
||||
ARG CACHEBUST=1234
|
||||
ARG GIT_COMMIT=main
|
||||
RUN echo "Cache bust: ${CACHEBUST}" && \
|
||||
git clone https://github.com/ostris/ai-toolkit.git && \
|
||||
cd ai-toolkit && \
|
||||
git checkout ${GIT_COMMIT}
|
||||
|
||||
WORKDIR /app/ai-toolkit
|
||||
|
||||
# Install Python dependencies
|
||||
RUN pip install --no-cache-dir -r requirements.txt && \
|
||||
pip install --pre --no-cache-dir torch torchvision torchaudio --index-url https://download.pytorch.org/whl/nightly/cu128 --force && \
|
||||
pip install setuptools==69.5.1 --no-cache-dir
|
||||
|
||||
# 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
|
||||
|
||||
WORKDIR /
|
||||
|
||||
COPY docker/start.sh /start.sh
|
||||
RUN chmod +x /start.sh
|
||||
|
||||
CMD ["/start.sh"]
|
||||
70
docker/start.sh
Normal file
70
docker/start.sh
Normal file
@@ -0,0 +1,70 @@
|
||||
#!/bin/bash
|
||||
set -e # Exit the script if any statement returns a non-true return value
|
||||
|
||||
# ref https://github.com/runpod/containers/blob/main/container-template/start.sh
|
||||
|
||||
# ---------------------------------------------------------------------------- #
|
||||
# Function Definitions #
|
||||
# ---------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
# Setup ssh
|
||||
setup_ssh() {
|
||||
if [[ $PUBLIC_KEY ]]; then
|
||||
echo "Setting up SSH..."
|
||||
mkdir -p ~/.ssh
|
||||
echo "$PUBLIC_KEY" >> ~/.ssh/authorized_keys
|
||||
chmod 700 -R ~/.ssh
|
||||
|
||||
if [ ! -f /etc/ssh/ssh_host_rsa_key ]; then
|
||||
ssh-keygen -t rsa -f /etc/ssh/ssh_host_rsa_key -q -N ''
|
||||
echo "RSA key fingerprint:"
|
||||
ssh-keygen -lf /etc/ssh/ssh_host_rsa_key.pub
|
||||
fi
|
||||
|
||||
if [ ! -f /etc/ssh/ssh_host_dsa_key ]; then
|
||||
ssh-keygen -t dsa -f /etc/ssh/ssh_host_dsa_key -q -N ''
|
||||
echo "DSA key fingerprint:"
|
||||
ssh-keygen -lf /etc/ssh/ssh_host_dsa_key.pub
|
||||
fi
|
||||
|
||||
if [ ! -f /etc/ssh/ssh_host_ecdsa_key ]; then
|
||||
ssh-keygen -t ecdsa -f /etc/ssh/ssh_host_ecdsa_key -q -N ''
|
||||
echo "ECDSA key fingerprint:"
|
||||
ssh-keygen -lf /etc/ssh/ssh_host_ecdsa_key.pub
|
||||
fi
|
||||
|
||||
if [ ! -f /etc/ssh/ssh_host_ed25519_key ]; then
|
||||
ssh-keygen -t ed25519 -f /etc/ssh/ssh_host_ed25519_key -q -N ''
|
||||
echo "ED25519 key fingerprint:"
|
||||
ssh-keygen -lf /etc/ssh/ssh_host_ed25519_key.pub
|
||||
fi
|
||||
|
||||
service ssh start
|
||||
|
||||
echo "SSH host keys:"
|
||||
for key in /etc/ssh/*.pub; do
|
||||
echo "Key: $key"
|
||||
ssh-keygen -lf $key
|
||||
done
|
||||
fi
|
||||
}
|
||||
|
||||
# Export env vars
|
||||
export_env_vars() {
|
||||
echo "Exporting environment variables..."
|
||||
printenv | grep -E '^RUNPOD_|^PATH=|^_=' | awk -F = '{ print "export " $1 "=\"" $2 "\"" }' >> /etc/rp_environment
|
||||
echo 'source /etc/rp_environment' >> ~/.bashrc
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------- #
|
||||
# Main Program #
|
||||
# ---------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
echo "Pod Started"
|
||||
|
||||
setup_ssh
|
||||
export_env_vars
|
||||
echo "Starting AI Toolkit UI..."
|
||||
cd /app/ai-toolkit/ui && npm run start
|
||||
129
extensions/example/ExampleMergeModels.py
Normal file
129
extensions/example/ExampleMergeModels.py
Normal file
@@ -0,0 +1,129 @@
|
||||
import torch
|
||||
import gc
|
||||
from collections import OrderedDict
|
||||
from typing import TYPE_CHECKING
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from toolkit.config_modules import ModelConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from tqdm import tqdm
|
||||
|
||||
# Type check imports. Prevents circular imports
|
||||
if TYPE_CHECKING:
|
||||
from jobs import ExtensionJob
|
||||
|
||||
|
||||
# extend standard config classes to add weight
|
||||
class ModelInputConfig(ModelConfig):
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.weight = kwargs.get('weight', 1.0)
|
||||
# overwrite default dtype unless user specifies otherwise
|
||||
# float 32 will give up better precision on the merging functions
|
||||
self.dtype: str = kwargs.get('dtype', 'float32')
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
# this is our main class process
|
||||
class ExampleMergeModels(BaseExtensionProcess):
|
||||
def __init__(
|
||||
self,
|
||||
process_id: int,
|
||||
job: 'ExtensionJob',
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
# this is the setup process, do not do process intensive stuff here, just variable setup and
|
||||
# checking requirements. This is called before the run() function
|
||||
# no loading models or anything like that, it is just for setting up the process
|
||||
# all of your process intensive stuff should be done in the run() function
|
||||
# config will have everything from the process item in the config file
|
||||
|
||||
# convince methods exist on BaseProcess to get config values
|
||||
# if required is set to true and the value is not found it will throw an error
|
||||
# you can pass a default value to get_conf() as well if it was not in the config file
|
||||
# as well as a type to cast the value to
|
||||
self.save_path = self.get_conf('save_path', required=True)
|
||||
self.save_dtype = self.get_conf('save_dtype', default='float16', as_type=get_torch_dtype)
|
||||
self.device = self.get_conf('device', default='cpu', as_type=torch.device)
|
||||
|
||||
# build models to merge list
|
||||
models_to_merge = self.get_conf('models_to_merge', required=True, as_type=list)
|
||||
# build list of ModelInputConfig objects. I find it is a good idea to make a class for each config
|
||||
# this way you can add methods to it and it is easier to read and code. There are a lot of
|
||||
# inbuilt config classes located in toolkit.config_modules as well
|
||||
self.models_to_merge = [ModelInputConfig(**model) for model in models_to_merge]
|
||||
# setup is complete. Don't load anything else here, just setup variables and stuff
|
||||
|
||||
# this is the entire run process be sure to call super().run() first
|
||||
def run(self):
|
||||
# always call first
|
||||
super().run()
|
||||
print(f"Running process: {self.__class__.__name__}")
|
||||
|
||||
# let's adjust our weights first to normalize them so the total is 1.0
|
||||
total_weight = sum([model.weight for model in self.models_to_merge])
|
||||
weight_adjust = 1.0 / total_weight
|
||||
for model in self.models_to_merge:
|
||||
model.weight *= weight_adjust
|
||||
|
||||
output_model: StableDiffusion = None
|
||||
# let's do the merge, it is a good idea to use tqdm to show progress
|
||||
for model_config in tqdm(self.models_to_merge, desc="Merging models"):
|
||||
# setup model class with our helper class
|
||||
sd_model = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=model_config,
|
||||
dtype="float32"
|
||||
)
|
||||
# load the model
|
||||
sd_model.load_model()
|
||||
|
||||
# adjust the weight of the text encoder
|
||||
if isinstance(sd_model.text_encoder, list):
|
||||
# sdxl model
|
||||
for text_encoder in sd_model.text_encoder:
|
||||
for key, value in text_encoder.state_dict().items():
|
||||
value *= model_config.weight
|
||||
else:
|
||||
# normal model
|
||||
for key, value in sd_model.text_encoder.state_dict().items():
|
||||
value *= model_config.weight
|
||||
# adjust the weights of the unet
|
||||
for key, value in sd_model.unet.state_dict().items():
|
||||
value *= model_config.weight
|
||||
|
||||
if output_model is None:
|
||||
# use this one as the base
|
||||
output_model = sd_model
|
||||
else:
|
||||
# merge the models
|
||||
# text encoder
|
||||
if isinstance(output_model.text_encoder, list):
|
||||
# sdxl model
|
||||
for i, text_encoder in enumerate(output_model.text_encoder):
|
||||
for key, value in text_encoder.state_dict().items():
|
||||
value += sd_model.text_encoder[i].state_dict()[key]
|
||||
else:
|
||||
# normal model
|
||||
for key, value in output_model.text_encoder.state_dict().items():
|
||||
value += sd_model.text_encoder.state_dict()[key]
|
||||
# unet
|
||||
for key, value in output_model.unet.state_dict().items():
|
||||
value += sd_model.unet.state_dict()[key]
|
||||
|
||||
# remove the model to free memory
|
||||
del sd_model
|
||||
flush()
|
||||
|
||||
# merge loop is done, let's save the model
|
||||
print(f"Saving merged model to {self.save_path}")
|
||||
output_model.save(self.save_path, meta=self.meta, save_dtype=self.save_dtype)
|
||||
print(f"Saved merged model to {self.save_path}")
|
||||
# do cleanup here
|
||||
del output_model
|
||||
flush()
|
||||
25
extensions/example/__init__.py
Normal file
25
extensions/example/__init__.py
Normal file
@@ -0,0 +1,25 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# We make a subclass of Extension
|
||||
class ExampleMergeExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "example_merge_extension"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Example Merge Extension"
|
||||
|
||||
# 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 .ExampleMergeModels import ExampleMergeModels
|
||||
return ExampleMergeModels
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
ExampleMergeExtension
|
||||
]
|
||||
48
extensions/example/config/config.example.yaml
Normal file
48
extensions/example/config/config.example.yaml
Normal file
@@ -0,0 +1,48 @@
|
||||
---
|
||||
# Always include at least one example config file to show how to use your extension.
|
||||
# use plenty of comments so users know how to use it and what everything does
|
||||
|
||||
# all extensions will use this job name
|
||||
job: extension
|
||||
config:
|
||||
name: 'my_awesome_merge'
|
||||
process:
|
||||
# Put your example processes here. This will be passed
|
||||
# to your extension process in the config argument.
|
||||
# the type MUST match your extension uid
|
||||
- type: "example_merge_extension"
|
||||
# save path for the merged model
|
||||
save_path: "output/merge/[name].safetensors"
|
||||
# save type
|
||||
dtype: fp16
|
||||
# device to run it on
|
||||
device: cuda:0
|
||||
# input models can only be SD1.x and SD2.x models for this example (currently)
|
||||
models_to_merge:
|
||||
# weights are relative, total weights will be normalized
|
||||
# for example. If you have 2 models with weight 1.0, they will
|
||||
# both be weighted 0.5. If you have 1 model with weight 1.0 and
|
||||
# another with weight 2.0, the first will be weighted 1/3 and the
|
||||
# second will be weighted 2/3
|
||||
- name_or_path: "input/model1.safetensors"
|
||||
weight: 1.0
|
||||
- name_or_path: "input/model2.safetensors"
|
||||
weight: 1.0
|
||||
- name_or_path: "input/model3.safetensors"
|
||||
weight: 0.3
|
||||
- name_or_path: "input/model4.safetensors"
|
||||
weight: 1.0
|
||||
|
||||
|
||||
# you can put any information you want here, and it will be saved in the model
|
||||
# the below is an example. I recommend doing trigger words at a minimum
|
||||
# in the metadata. The software will include this plus some other information
|
||||
meta:
|
||||
name: "[name]" # [name] gets replaced with the name above
|
||||
description: A short description of your model
|
||||
version: '0.1'
|
||||
creator:
|
||||
name: Your Name
|
||||
email: your@email.com
|
||||
website: https://yourwebsite.com
|
||||
any: All meta data above is arbitrary, it can be whatever you want.
|
||||
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()
|
||||
102
extensions_built_in/advanced_generator/PureLoraGenerator.py
Normal file
102
extensions_built_in/advanced_generator/PureLoraGenerator.py
Normal file
@@ -0,0 +1,102 @@
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, SampleConfig, LoRMConfig
|
||||
from toolkit.lorm import ExtractMode, convert_diffusers_unet_to_lorm
|
||||
from toolkit.sd_device_states_presets import get_train_sd_device_state_preset
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class PureLoraGenerator(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.device = self.get_conf('device', 'cuda')
|
||||
self.device_torch = torch.device(self.device)
|
||||
self.model_config = ModelConfig(**self.get_conf('model', required=True))
|
||||
self.generate_config = SampleConfig(**self.get_conf('sample', required=True))
|
||||
self.dtype = self.get_conf('dtype', 'float16')
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
lorm_config = self.get_conf('lorm', None)
|
||||
self.lorm_config = LoRMConfig(**lorm_config) if lorm_config is not None else None
|
||||
|
||||
self.device_state_preset = get_train_sd_device_state_preset(
|
||||
device=torch.device(self.device),
|
||||
)
|
||||
|
||||
self.progress_bar = None
|
||||
self.sd = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=self.model_config,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
with torch.no_grad():
|
||||
self.sd.load_model()
|
||||
self.sd.unet.eval()
|
||||
self.sd.unet.to(self.device_torch)
|
||||
if isinstance(self.sd.text_encoder, list):
|
||||
for te in self.sd.text_encoder:
|
||||
te.eval()
|
||||
te.to(self.device_torch)
|
||||
else:
|
||||
self.sd.text_encoder.eval()
|
||||
self.sd.to(self.device_torch)
|
||||
|
||||
print(f"Converting to LoRM UNet")
|
||||
# replace the unet with LoRMUnet
|
||||
convert_diffusers_unet_to_lorm(
|
||||
self.sd.unet,
|
||||
config=self.lorm_config,
|
||||
)
|
||||
|
||||
sample_folder = os.path.join(self.output_folder)
|
||||
gen_img_config_list = []
|
||||
|
||||
sample_config = self.generate_config
|
||||
start_seed = sample_config.seed
|
||||
current_seed = start_seed
|
||||
for i in range(len(sample_config.prompts)):
|
||||
if sample_config.walk_seed:
|
||||
current_seed = start_seed + i
|
||||
|
||||
filename = f"[time]_[count].{self.generate_config.ext}"
|
||||
output_path = os.path.join(sample_folder, filename)
|
||||
prompt = sample_config.prompts[i]
|
||||
extra_args = {}
|
||||
gen_img_config_list.append(GenerateImageConfig(
|
||||
prompt=prompt, # it will autoparse the prompt
|
||||
width=sample_config.width,
|
||||
height=sample_config.height,
|
||||
negative_prompt=sample_config.neg,
|
||||
seed=current_seed,
|
||||
guidance_scale=sample_config.guidance_scale,
|
||||
guidance_rescale=sample_config.guidance_rescale,
|
||||
num_inference_steps=sample_config.sample_steps,
|
||||
network_multiplier=sample_config.network_multiplier,
|
||||
output_path=output_path,
|
||||
output_ext=sample_config.ext,
|
||||
adapter_conditioning_scale=sample_config.adapter_conditioning_scale,
|
||||
**extra_args
|
||||
))
|
||||
|
||||
# send to be generated
|
||||
self.sd.generate_images(gen_img_config_list, sampler=sample_config.sampler)
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
212
extensions_built_in/advanced_generator/ReferenceGenerator.py
Normal file
212
extensions_built_in/advanced_generator/ReferenceGenerator.py
Normal file
@@ -0,0 +1,212 @@
|
||||
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 torch.utils.data import DataLoader
|
||||
from diffusers import StableDiffusionXLAdapterPipeline, StableDiffusionAdapterPipeline
|
||||
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
|
||||
|
||||
|
||||
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.t2i_adapter_path = kwargs.get('t2i_adapter_path', None)
|
||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||
self.sample_steps = kwargs.get('sample_steps', 20)
|
||||
self.prompt_2 = kwargs.get('prompt_2', None)
|
||||
self.neg_2 = kwargs.get('neg_2', None)
|
||||
self.prompts = kwargs.get('prompts', None)
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
self.ext = kwargs.get('ext', 'png')
|
||||
self.adapter_conditioning_scale = kwargs.get('adapter_conditioning_scale', 1.0)
|
||||
if kwargs.get('shuffle', False):
|
||||
# shuffle the prompts
|
||||
random.shuffle(self.prompts)
|
||||
|
||||
|
||||
class ReferenceGenerator(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.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 run(self):
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
device = torch.device(self.device)
|
||||
|
||||
if self.generate_config.t2i_adapter_path is not None:
|
||||
self.adapter = T2IAdapter.from_pretrained(
|
||||
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)
|
||||
|
||||
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)
|
||||
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
|
||||
|
||||
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]
|
||||
output_path = os.path.join(self.output_folder, img_filename)
|
||||
output_caption_path = os.path.join(self.output_folder, img_filename_no_ext + '.txt')
|
||||
output_depth_path = os.path.join(self.output_folder, img_filename_no_ext + '.depth.png')
|
||||
|
||||
caption = batch.get_caption_list()[0]
|
||||
|
||||
img: torch.Tensor = batch.tensor.clone()
|
||||
# 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)
|
||||
|
||||
width, height = image.size
|
||||
min_res = min(width, height)
|
||||
|
||||
if self.generate_config.walk_seed:
|
||||
seed = seed + 1
|
||||
|
||||
if self.generate_config.seed == -1:
|
||||
# random
|
||||
seed = random.randint(0, 1000000)
|
||||
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
# generate depth map
|
||||
image = midas_depth(
|
||||
image,
|
||||
detect_resolution=min_res, # do 512 ?
|
||||
image_resolution=min_res
|
||||
)
|
||||
|
||||
# image.save(output_depth_path)
|
||||
|
||||
gen_images = pipe(
|
||||
prompt=caption,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
image=image,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
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
|
||||
with open(output_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()
|
||||
59
extensions_built_in/advanced_generator/__init__.py
Normal file
59
extensions_built_in/advanced_generator/__init__.py
Normal file
@@ -0,0 +1,59 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class AdvancedReferenceGeneratorExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "reference_generator"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Reference Generator"
|
||||
|
||||
# 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 .ReferenceGenerator import ReferenceGenerator
|
||||
return ReferenceGenerator
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class PureLoraGenerator(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "pure_lora_generator"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Pure LoRA Generator"
|
||||
|
||||
# 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 .PureLoraGenerator import PureLoraGenerator
|
||||
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, Img2ImgGeneratorExtension
|
||||
]
|
||||
@@ -0,0 +1,91 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
name: test_v1
|
||||
process:
|
||||
- type: 'textual_inversion_trainer'
|
||||
training_folder: "out/TI"
|
||||
device: cuda:0
|
||||
# for tensorboard logging
|
||||
log_dir: "out/.tensorboard"
|
||||
embedding:
|
||||
trigger: "your_trigger_here"
|
||||
tokens: 12
|
||||
init_words: "man with short brown hair"
|
||||
save_format: "safetensors" # 'safetensors' or 'pt'
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 100 # save every this many steps
|
||||
max_step_saves_to_keep: 5 # only affects step counts
|
||||
datasets:
|
||||
- folder_path: "/path/to/dataset"
|
||||
caption_ext: "txt"
|
||||
default_caption: "[trigger]"
|
||||
buckets: true
|
||||
resolution: 512
|
||||
train:
|
||||
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
|
||||
steps: 3000
|
||||
weight_jitter: 0.0
|
||||
lr: 5e-5
|
||||
train_unet: false
|
||||
gradient_checkpointing: true
|
||||
train_text_encoder: false
|
||||
optimizer: "adamw"
|
||||
# optimizer: "prodigy"
|
||||
optimizer_params:
|
||||
weight_decay: 1e-2
|
||||
lr_scheduler: "constant"
|
||||
max_denoising_steps: 1000
|
||||
batch_size: 4
|
||||
dtype: bf16
|
||||
xformers: true
|
||||
min_snr_gamma: 5.0
|
||||
# skip_first_sample: true
|
||||
noise_offset: 0.0 # not needed for this
|
||||
model:
|
||||
# objective reality v2
|
||||
name_or_path: "https://civitai.com/models/128453?modelVersionId=142465"
|
||||
is_v2: false # for v2 models
|
||||
is_xl: false # for SDXL models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
sample:
|
||||
sampler: "ddpm" # must match train.noise_scheduler
|
||||
sample_every: 100 # sample every this many steps
|
||||
width: 512
|
||||
height: 512
|
||||
prompts:
|
||||
- "photo of [trigger] laughing"
|
||||
- "photo of [trigger] smiling"
|
||||
- "[trigger] close up"
|
||||
- "dark scene [trigger] frozen"
|
||||
- "[trigger] nighttime"
|
||||
- "a painting of [trigger]"
|
||||
- "a drawing of [trigger]"
|
||||
- "a cartoon of [trigger]"
|
||||
- "[trigger] pixar style"
|
||||
- "[trigger] costume"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: false
|
||||
guidance_scale: 7
|
||||
sample_steps: 20
|
||||
network_multiplier: 1.0
|
||||
|
||||
logging:
|
||||
log_every: 10 # log every this many steps
|
||||
use_wandb: false # not supported yet
|
||||
verbose: false
|
||||
|
||||
# You can put any information you want here, and it will be saved in the model.
|
||||
# The below is an example, but you can put your grocery list in it if you want.
|
||||
# It is saved in the model so be aware of that. The software will include this
|
||||
# plus some other information for you automatically
|
||||
meta:
|
||||
# [name] gets replaced with the name above
|
||||
name: "[name]"
|
||||
# version: '1.0'
|
||||
# creator:
|
||||
# name: Your Name
|
||||
# email: your@gmail.com
|
||||
# website: https://your.website
|
||||
151
extensions_built_in/concept_replacer/ConceptReplacer.py
Normal file
151
extensions_built_in/concept_replacer/ConceptReplacer.py
Normal file
@@ -0,0 +1,151 @@
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from torch.utils.data import DataLoader
|
||||
from toolkit.prompt_utils import concat_prompt_embeds, split_prompt_embeds
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, BlankNetwork
|
||||
from toolkit.train_tools import get_torch_dtype, apply_snr_weight
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseSDTrainProcess
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class ConceptReplacementConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.concept: str = kwargs.get('concept', '')
|
||||
self.replacement: str = kwargs.get('replacement', '')
|
||||
|
||||
|
||||
class ConceptReplacer(BaseSDTrainProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
|
||||
super().__init__(process_id, job, config, **kwargs)
|
||||
replacement_list = self.config.get('replacements', [])
|
||||
self.replacement_list = [ConceptReplacementConfig(**x) for x in replacement_list]
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
self.sd.vae.eval()
|
||||
self.sd.vae.to(self.device_torch)
|
||||
|
||||
# textual inversion
|
||||
if self.embedding is not None:
|
||||
# set text encoder to train. Not sure if this is necessary but diffusers example did it
|
||||
self.sd.text_encoder.train()
|
||||
|
||||
def hook_train_loop(self, batch):
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
noisy_latents, noise, timesteps, conditioned_prompts, imgs = self.process_general_training_batch(batch)
|
||||
network_weight_list = batch.get_network_weight_list()
|
||||
|
||||
# have a blank network so we can wrap it in a context and set multipliers without checking every time
|
||||
if self.network is not None:
|
||||
network = self.network
|
||||
else:
|
||||
network = BlankNetwork()
|
||||
|
||||
batch_replacement_list = []
|
||||
# get a random replacement for each prompt
|
||||
for prompt in conditioned_prompts:
|
||||
replacement = random.choice(self.replacement_list)
|
||||
batch_replacement_list.append(replacement)
|
||||
|
||||
# build out prompts
|
||||
concept_prompts = []
|
||||
replacement_prompts = []
|
||||
for idx, replacement in enumerate(batch_replacement_list):
|
||||
prompt = conditioned_prompts[idx]
|
||||
|
||||
# insert shuffled concept at beginning and end of prompt
|
||||
shuffled_concept = [x.strip() for x in replacement.concept.split(',')]
|
||||
random.shuffle(shuffled_concept)
|
||||
shuffled_concept = ', '.join(shuffled_concept)
|
||||
concept_prompts.append(f"{shuffled_concept}, {prompt}, {shuffled_concept}")
|
||||
|
||||
# insert replacement at beginning and end of prompt
|
||||
shuffled_replacement = [x.strip() for x in replacement.replacement.split(',')]
|
||||
random.shuffle(shuffled_replacement)
|
||||
shuffled_replacement = ', '.join(shuffled_replacement)
|
||||
replacement_prompts.append(f"{shuffled_replacement}, {prompt}, {shuffled_replacement}")
|
||||
|
||||
# predict the replacement without network
|
||||
conditional_embeds = self.sd.encode_prompt(replacement_prompts).to(self.device_torch, dtype=dtype)
|
||||
|
||||
replacement_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents.to(self.device_torch, dtype=dtype),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
)
|
||||
|
||||
del conditional_embeds
|
||||
replacement_pred = replacement_pred.detach()
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
flush()
|
||||
|
||||
# text encoding
|
||||
grad_on_text_encoder = False
|
||||
if self.train_config.train_text_encoder:
|
||||
grad_on_text_encoder = True
|
||||
|
||||
if self.embedding:
|
||||
grad_on_text_encoder = True
|
||||
|
||||
# set the weights
|
||||
network.multiplier = network_weight_list
|
||||
|
||||
# activate network if it exits
|
||||
with network:
|
||||
with torch.set_grad_enabled(grad_on_text_encoder):
|
||||
# embed the prompts
|
||||
conditional_embeds = self.sd.encode_prompt(concept_prompts).to(self.device_torch, dtype=dtype)
|
||||
if not grad_on_text_encoder:
|
||||
# detach the embeddings
|
||||
conditional_embeds = conditional_embeds.detach()
|
||||
self.optimizer.zero_grad()
|
||||
flush()
|
||||
|
||||
noise_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents.to(self.device_torch, dtype=dtype),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
)
|
||||
|
||||
loss = torch.nn.functional.mse_loss(noise_pred.float(), replacement_pred.float(), reduction="none")
|
||||
loss = loss.mean([1, 2, 3])
|
||||
|
||||
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
|
||||
# add min_snr_gamma
|
||||
loss = apply_snr_weight(loss, timesteps, self.sd.noise_scheduler, self.train_config.min_snr_gamma)
|
||||
|
||||
loss = loss.mean()
|
||||
|
||||
# back propagate loss to free ram
|
||||
loss.backward()
|
||||
flush()
|
||||
|
||||
# apply gradients
|
||||
self.optimizer.step()
|
||||
self.optimizer.zero_grad()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
if self.embedding is not None:
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
self.embedding.restore_embeddings()
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': loss.item()}
|
||||
)
|
||||
# reset network multiplier
|
||||
network.multiplier = 1.0
|
||||
|
||||
return loss_dict
|
||||
26
extensions_built_in/concept_replacer/__init__.py
Normal file
26
extensions_built_in/concept_replacer/__init__.py
Normal file
@@ -0,0 +1,26 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class ConceptReplacerExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "concept_replacer"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Concept Replacer"
|
||||
|
||||
# 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 .ConceptReplacer import ConceptReplacer
|
||||
return ConceptReplacer
|
||||
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
ConceptReplacerExtension,
|
||||
]
|
||||
@@ -0,0 +1,91 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
name: test_v1
|
||||
process:
|
||||
- type: 'textual_inversion_trainer'
|
||||
training_folder: "out/TI"
|
||||
device: cuda:0
|
||||
# for tensorboard logging
|
||||
log_dir: "out/.tensorboard"
|
||||
embedding:
|
||||
trigger: "your_trigger_here"
|
||||
tokens: 12
|
||||
init_words: "man with short brown hair"
|
||||
save_format: "safetensors" # 'safetensors' or 'pt'
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 100 # save every this many steps
|
||||
max_step_saves_to_keep: 5 # only affects step counts
|
||||
datasets:
|
||||
- folder_path: "/path/to/dataset"
|
||||
caption_ext: "txt"
|
||||
default_caption: "[trigger]"
|
||||
buckets: true
|
||||
resolution: 512
|
||||
train:
|
||||
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
|
||||
steps: 3000
|
||||
weight_jitter: 0.0
|
||||
lr: 5e-5
|
||||
train_unet: false
|
||||
gradient_checkpointing: true
|
||||
train_text_encoder: false
|
||||
optimizer: "adamw"
|
||||
# optimizer: "prodigy"
|
||||
optimizer_params:
|
||||
weight_decay: 1e-2
|
||||
lr_scheduler: "constant"
|
||||
max_denoising_steps: 1000
|
||||
batch_size: 4
|
||||
dtype: bf16
|
||||
xformers: true
|
||||
min_snr_gamma: 5.0
|
||||
# skip_first_sample: true
|
||||
noise_offset: 0.0 # not needed for this
|
||||
model:
|
||||
# objective reality v2
|
||||
name_or_path: "https://civitai.com/models/128453?modelVersionId=142465"
|
||||
is_v2: false # for v2 models
|
||||
is_xl: false # for SDXL models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
sample:
|
||||
sampler: "ddpm" # must match train.noise_scheduler
|
||||
sample_every: 100 # sample every this many steps
|
||||
width: 512
|
||||
height: 512
|
||||
prompts:
|
||||
- "photo of [trigger] laughing"
|
||||
- "photo of [trigger] smiling"
|
||||
- "[trigger] close up"
|
||||
- "dark scene [trigger] frozen"
|
||||
- "[trigger] nighttime"
|
||||
- "a painting of [trigger]"
|
||||
- "a drawing of [trigger]"
|
||||
- "a cartoon of [trigger]"
|
||||
- "[trigger] pixar style"
|
||||
- "[trigger] costume"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: false
|
||||
guidance_scale: 7
|
||||
sample_steps: 20
|
||||
network_multiplier: 1.0
|
||||
|
||||
logging:
|
||||
log_every: 10 # log every this many steps
|
||||
use_wandb: false # not supported yet
|
||||
verbose: false
|
||||
|
||||
# You can put any information you want here, and it will be saved in the model.
|
||||
# The below is an example, but you can put your grocery list in it if you want.
|
||||
# It is saved in the model so be aware of that. The software will include this
|
||||
# plus some other information for you automatically
|
||||
meta:
|
||||
# [name] gets replaced with the name above
|
||||
name: "[name]"
|
||||
# version: '1.0'
|
||||
# creator:
|
||||
# name: Your Name
|
||||
# email: your@gmail.com
|
||||
# website: https://your.website
|
||||
20
extensions_built_in/dataset_tools/DatasetTools.py
Normal file
20
extensions_built_in/dataset_tools/DatasetTools.py
Normal file
@@ -0,0 +1,20 @@
|
||||
from collections import OrderedDict
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseExtensionProcess
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class DatasetTools(BaseExtensionProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
|
||||
raise NotImplementedError("This extension is not yet implemented")
|
||||
196
extensions_built_in/dataset_tools/SuperTagger.py
Normal file
196
extensions_built_in/dataset_tools/SuperTagger.py
Normal file
@@ -0,0 +1,196 @@
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
import gc
|
||||
import traceback
|
||||
import torch
|
||||
from PIL import Image, ImageOps
|
||||
from tqdm import tqdm
|
||||
|
||||
from .tools.dataset_tools_config_modules import RAW_DIR, TRAIN_DIR, Step, ImgInfo
|
||||
from .tools.fuyu_utils import FuyuImageProcessor
|
||||
from .tools.image_tools import load_image, ImageProcessor, resize_to_max
|
||||
from .tools.llava_utils import LLaVAImageProcessor
|
||||
from .tools.caption import default_long_prompt, default_short_prompt, default_replacements
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from .tools.sync_tools import get_img_paths
|
||||
|
||||
img_ext = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
VERSION = 2
|
||||
|
||||
|
||||
class SuperTagger(BaseExtensionProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
parent_dir = config.get('parent_dir', None)
|
||||
self.dataset_paths: list[str] = config.get('dataset_paths', [])
|
||||
self.device = config.get('device', 'cuda')
|
||||
self.steps: list[Step] = config.get('steps', [])
|
||||
self.caption_method = config.get('caption_method', 'llava:default')
|
||||
self.caption_prompt = config.get('caption_prompt', default_long_prompt)
|
||||
self.caption_short_prompt = config.get('caption_short_prompt', default_short_prompt)
|
||||
self.force_reprocess_img = config.get('force_reprocess_img', False)
|
||||
self.caption_replacements = config.get('caption_replacements', default_replacements)
|
||||
self.caption_short_replacements = config.get('caption_short_replacements', default_replacements)
|
||||
self.master_dataset_dict = OrderedDict()
|
||||
self.dataset_master_config_file = config.get('dataset_master_config_file', None)
|
||||
if parent_dir is not None and len(self.dataset_paths) == 0:
|
||||
# find all folders in the patent_dataset_path
|
||||
self.dataset_paths = [
|
||||
os.path.join(parent_dir, folder)
|
||||
for folder in os.listdir(parent_dir)
|
||||
if os.path.isdir(os.path.join(parent_dir, folder))
|
||||
]
|
||||
else:
|
||||
# make sure they exist
|
||||
for dataset_path in self.dataset_paths:
|
||||
if not os.path.exists(dataset_path):
|
||||
raise ValueError(f"Dataset path does not exist: {dataset_path}")
|
||||
|
||||
print(f"Found {len(self.dataset_paths)} dataset paths")
|
||||
|
||||
self.image_processor: ImageProcessor = self.get_image_processor()
|
||||
|
||||
def get_image_processor(self):
|
||||
if self.caption_method.startswith('llava'):
|
||||
return LLaVAImageProcessor(device=self.device)
|
||||
elif self.caption_method.startswith('fuyu'):
|
||||
return FuyuImageProcessor(device=self.device)
|
||||
else:
|
||||
raise ValueError(f"Unknown caption method: {self.caption_method}")
|
||||
|
||||
def process_image(self, img_path: str):
|
||||
root_img_dir = os.path.dirname(os.path.dirname(img_path))
|
||||
filename = os.path.basename(img_path)
|
||||
filename_no_ext = os.path.splitext(filename)[0]
|
||||
train_dir = os.path.join(root_img_dir, TRAIN_DIR)
|
||||
train_img_path = os.path.join(train_dir, filename)
|
||||
json_path = os.path.join(train_dir, f"{filename_no_ext}.json")
|
||||
|
||||
# check if json exists, if it does load it as image info
|
||||
if os.path.exists(json_path):
|
||||
with open(json_path, 'r') as f:
|
||||
img_info = ImgInfo(**json.load(f))
|
||||
else:
|
||||
img_info = ImgInfo()
|
||||
|
||||
# always send steps first in case other processes need them
|
||||
img_info.add_steps(copy.deepcopy(self.steps))
|
||||
img_info.set_version(VERSION)
|
||||
img_info.set_caption_method(self.caption_method)
|
||||
|
||||
image: Image = None
|
||||
caption_image: Image = None
|
||||
|
||||
did_update_image = False
|
||||
|
||||
# trigger reprocess of steps
|
||||
if self.force_reprocess_img:
|
||||
img_info.trigger_image_reprocess()
|
||||
|
||||
# set the image as updated if it does not exist on disk
|
||||
if not os.path.exists(train_img_path):
|
||||
did_update_image = True
|
||||
image = load_image(img_path)
|
||||
if img_info.force_image_process:
|
||||
did_update_image = True
|
||||
image = load_image(img_path)
|
||||
|
||||
# go through the needed steps
|
||||
for step in copy.deepcopy(img_info.state.steps_to_complete):
|
||||
if step == 'caption':
|
||||
# load image
|
||||
if image is None:
|
||||
image = load_image(img_path)
|
||||
if caption_image is None:
|
||||
caption_image = resize_to_max(image, 1024, 1024)
|
||||
|
||||
if not self.image_processor.is_loaded:
|
||||
print('Loading Model. Takes a while, especially the first time')
|
||||
self.image_processor.load_model()
|
||||
|
||||
img_info.caption = self.image_processor.generate_caption(
|
||||
image=caption_image,
|
||||
prompt=self.caption_prompt,
|
||||
replacements=self.caption_replacements
|
||||
)
|
||||
img_info.mark_step_complete(step)
|
||||
elif step == 'caption_short':
|
||||
# load image
|
||||
if image is None:
|
||||
image = load_image(img_path)
|
||||
|
||||
if caption_image is None:
|
||||
caption_image = resize_to_max(image, 1024, 1024)
|
||||
|
||||
if not self.image_processor.is_loaded:
|
||||
print('Loading Model. Takes a while, especially the first time')
|
||||
self.image_processor.load_model()
|
||||
img_info.caption_short = self.image_processor.generate_caption(
|
||||
image=caption_image,
|
||||
prompt=self.caption_short_prompt,
|
||||
replacements=self.caption_short_replacements
|
||||
)
|
||||
img_info.mark_step_complete(step)
|
||||
elif step == 'contrast_stretch':
|
||||
# load image
|
||||
if image is None:
|
||||
image = load_image(img_path)
|
||||
image = ImageOps.autocontrast(image, cutoff=(0.1, 0), preserve_tone=True)
|
||||
did_update_image = True
|
||||
img_info.mark_step_complete(step)
|
||||
else:
|
||||
raise ValueError(f"Unknown step: {step}")
|
||||
|
||||
os.makedirs(os.path.dirname(train_img_path), exist_ok=True)
|
||||
if did_update_image:
|
||||
image.save(train_img_path)
|
||||
|
||||
if img_info.is_dirty:
|
||||
with open(json_path, 'w') as f:
|
||||
json.dump(img_info.to_dict(), f, indent=4)
|
||||
|
||||
if self.dataset_master_config_file:
|
||||
# add to master dict
|
||||
self.master_dataset_dict[train_img_path] = img_info.to_dict()
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
imgs_to_process = []
|
||||
# find all images
|
||||
for dataset_path in self.dataset_paths:
|
||||
raw_dir = os.path.join(dataset_path, RAW_DIR)
|
||||
raw_image_paths = get_img_paths(raw_dir)
|
||||
for raw_image_path in raw_image_paths:
|
||||
imgs_to_process.append(raw_image_path)
|
||||
|
||||
if len(imgs_to_process) == 0:
|
||||
print(f"No images to process")
|
||||
else:
|
||||
print(f"Found {len(imgs_to_process)} to process")
|
||||
|
||||
for img_path in tqdm(imgs_to_process, desc="Processing images"):
|
||||
try:
|
||||
self.process_image(img_path)
|
||||
except Exception:
|
||||
# print full stack trace
|
||||
print(traceback.format_exc())
|
||||
continue
|
||||
# self.process_image(img_path)
|
||||
|
||||
if self.dataset_master_config_file is not None:
|
||||
# save it as json
|
||||
with open(self.dataset_master_config_file, 'w') as f:
|
||||
json.dump(self.master_dataset_dict, f, indent=4)
|
||||
|
||||
del self.image_processor
|
||||
flush()
|
||||
131
extensions_built_in/dataset_tools/SyncFromCollection.py
Normal file
131
extensions_built_in/dataset_tools/SyncFromCollection.py
Normal file
@@ -0,0 +1,131 @@
|
||||
import os
|
||||
import shutil
|
||||
from collections import OrderedDict
|
||||
import gc
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from .tools.dataset_tools_config_modules import DatasetSyncCollectionConfig, RAW_DIR, NEW_DIR
|
||||
from .tools.sync_tools import get_unsplash_images, get_pexels_images, get_local_image_file_names, download_image, \
|
||||
get_img_paths
|
||||
from jobs.process import BaseExtensionProcess
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class SyncFromCollection(BaseExtensionProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
|
||||
self.min_width = config.get('min_width', 1024)
|
||||
self.min_height = config.get('min_height', 1024)
|
||||
|
||||
# add our min_width and min_height to each dataset config if they don't exist
|
||||
for dataset_config in config.get('dataset_sync', []):
|
||||
if 'min_width' not in dataset_config:
|
||||
dataset_config['min_width'] = self.min_width
|
||||
if 'min_height' not in dataset_config:
|
||||
dataset_config['min_height'] = self.min_height
|
||||
|
||||
self.dataset_configs: List[DatasetSyncCollectionConfig] = [
|
||||
DatasetSyncCollectionConfig(**dataset_config)
|
||||
for dataset_config in config.get('dataset_sync', [])
|
||||
]
|
||||
print(f"Found {len(self.dataset_configs)} dataset configs")
|
||||
|
||||
def move_new_images(self, root_dir: str):
|
||||
raw_dir = os.path.join(root_dir, RAW_DIR)
|
||||
new_dir = os.path.join(root_dir, NEW_DIR)
|
||||
new_images = get_img_paths(new_dir)
|
||||
|
||||
for img_path in new_images:
|
||||
# move to raw
|
||||
new_path = os.path.join(raw_dir, os.path.basename(img_path))
|
||||
shutil.move(img_path, new_path)
|
||||
|
||||
# remove new dir
|
||||
shutil.rmtree(new_dir)
|
||||
|
||||
def sync_dataset(self, config: DatasetSyncCollectionConfig):
|
||||
if config.host == 'unsplash':
|
||||
get_images = get_unsplash_images
|
||||
elif config.host == 'pexels':
|
||||
get_images = get_pexels_images
|
||||
else:
|
||||
raise ValueError(f"Unknown host: {config.host}")
|
||||
|
||||
results = {
|
||||
'num_downloaded': 0,
|
||||
'num_skipped': 0,
|
||||
'bad': 0,
|
||||
'total': 0,
|
||||
}
|
||||
|
||||
photos = get_images(config)
|
||||
raw_dir = os.path.join(config.directory, RAW_DIR)
|
||||
new_dir = os.path.join(config.directory, NEW_DIR)
|
||||
raw_images = get_local_image_file_names(raw_dir)
|
||||
new_images = get_local_image_file_names(new_dir)
|
||||
|
||||
for photo in tqdm(photos, desc=f"{config.host}-{config.collection_id}"):
|
||||
try:
|
||||
if photo.filename not in raw_images and photo.filename not in new_images:
|
||||
download_image(photo, new_dir, min_width=self.min_width, min_height=self.min_height)
|
||||
results['num_downloaded'] += 1
|
||||
else:
|
||||
results['num_skipped'] += 1
|
||||
except Exception as e:
|
||||
print(f" - BAD({photo.id}): {e}")
|
||||
results['bad'] += 1
|
||||
continue
|
||||
results['total'] += 1
|
||||
|
||||
return results
|
||||
|
||||
def print_results(self, results):
|
||||
print(
|
||||
f" - new:{results['num_downloaded']}, old:{results['num_skipped']}, bad:{results['bad']} total:{results['total']}")
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print(f"Syncing {len(self.dataset_configs)} datasets")
|
||||
all_results = None
|
||||
failed_datasets = []
|
||||
for dataset_config in tqdm(self.dataset_configs, desc="Syncing datasets", leave=True):
|
||||
try:
|
||||
results = self.sync_dataset(dataset_config)
|
||||
if all_results is None:
|
||||
all_results = {**results}
|
||||
else:
|
||||
for key, value in results.items():
|
||||
all_results[key] += value
|
||||
|
||||
self.print_results(results)
|
||||
except Exception as e:
|
||||
print(f" - FAILED: {e}")
|
||||
if 'response' in e.__dict__:
|
||||
error = f"{e.response.status_code}: {e.response.text}"
|
||||
print(f" - {error}")
|
||||
failed_datasets.append({'dataset': dataset_config, 'error': error})
|
||||
else:
|
||||
failed_datasets.append({'dataset': dataset_config, 'error': str(e)})
|
||||
continue
|
||||
|
||||
print("Moving new images to raw")
|
||||
for dataset_config in self.dataset_configs:
|
||||
self.move_new_images(dataset_config.directory)
|
||||
|
||||
print("Done syncing datasets")
|
||||
self.print_results(all_results)
|
||||
|
||||
if len(failed_datasets) > 0:
|
||||
print(f"Failed to sync {len(failed_datasets)} datasets")
|
||||
for failed in failed_datasets:
|
||||
print(f" - {failed['dataset'].host}-{failed['dataset'].collection_id}")
|
||||
print(f" - ERR: {failed['error']}")
|
||||
43
extensions_built_in/dataset_tools/__init__.py
Normal file
43
extensions_built_in/dataset_tools/__init__.py
Normal file
@@ -0,0 +1,43 @@
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
class DatasetToolsExtension(Extension):
|
||||
uid = "dataset_tools"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Dataset Tools"
|
||||
|
||||
# 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 .DatasetTools import DatasetTools
|
||||
return DatasetTools
|
||||
|
||||
|
||||
class SyncFromCollectionExtension(Extension):
|
||||
uid = "sync_from_collection"
|
||||
name = "Sync from Collection"
|
||||
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .SyncFromCollection import SyncFromCollection
|
||||
return SyncFromCollection
|
||||
|
||||
|
||||
class SuperTaggerExtension(Extension):
|
||||
uid = "super_tagger"
|
||||
name = "Super Tagger"
|
||||
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .SuperTagger import SuperTagger
|
||||
return SuperTagger
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
SyncFromCollectionExtension, DatasetToolsExtension, SuperTaggerExtension
|
||||
]
|
||||
53
extensions_built_in/dataset_tools/tools/caption.py
Normal file
53
extensions_built_in/dataset_tools/tools/caption.py
Normal file
@@ -0,0 +1,53 @@
|
||||
|
||||
caption_manipulation_steps = ['caption', 'caption_short']
|
||||
|
||||
default_long_prompt = 'caption this image. describe every single thing in the image in detail. Do not include any unnecessary words in your description for the sake of good grammar. I want many short statements that serve the single purpose of giving the most thorough description if items as possible in the smallest, comma separated way possible. be sure to describe people\'s moods, clothing, the environment, lighting, colors, and everything.'
|
||||
default_short_prompt = 'caption this image in less than ten words'
|
||||
|
||||
default_replacements = [
|
||||
("the image features", ""),
|
||||
("the image shows", ""),
|
||||
("the image depicts", ""),
|
||||
("the image is", ""),
|
||||
("in this image", ""),
|
||||
("in the image", ""),
|
||||
]
|
||||
|
||||
|
||||
def clean_caption(cap, replacements=None):
|
||||
if replacements is None:
|
||||
replacements = default_replacements
|
||||
|
||||
# remove any newlines
|
||||
cap = cap.replace("\n", ", ")
|
||||
cap = cap.replace("\r", ", ")
|
||||
cap = cap.replace(".", ",")
|
||||
cap = cap.replace("\"", "")
|
||||
|
||||
# remove unicode characters
|
||||
cap = cap.encode('ascii', 'ignore').decode('ascii')
|
||||
|
||||
# make lowercase
|
||||
cap = cap.lower()
|
||||
# remove any extra spaces
|
||||
cap = " ".join(cap.split())
|
||||
|
||||
for replacement in replacements:
|
||||
if replacement[0].startswith('*'):
|
||||
# we are removing all text if it starts with this and the rest matches
|
||||
search_text = replacement[0][1:]
|
||||
if cap.startswith(search_text):
|
||||
cap = ""
|
||||
else:
|
||||
cap = cap.replace(replacement[0].lower(), replacement[1].lower())
|
||||
|
||||
cap_list = cap.split(",")
|
||||
# trim whitespace
|
||||
cap_list = [c.strip() for c in cap_list]
|
||||
# remove empty strings
|
||||
cap_list = [c for c in cap_list if c != ""]
|
||||
# remove duplicates
|
||||
cap_list = list(dict.fromkeys(cap_list))
|
||||
# join back together
|
||||
cap = ", ".join(cap_list)
|
||||
return cap
|
||||
@@ -0,0 +1,187 @@
|
||||
import json
|
||||
from typing import Literal, Type, TYPE_CHECKING
|
||||
|
||||
Host: Type = Literal['unsplash', 'pexels']
|
||||
|
||||
RAW_DIR = "raw"
|
||||
NEW_DIR = "_tmp"
|
||||
TRAIN_DIR = "train"
|
||||
DEPTH_DIR = "depth"
|
||||
|
||||
from .image_tools import Step, img_manipulation_steps
|
||||
from .caption import caption_manipulation_steps
|
||||
|
||||
|
||||
class DatasetSyncCollectionConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.host: Host = kwargs.get('host', None)
|
||||
self.collection_id: str = kwargs.get('collection_id', None)
|
||||
self.directory: str = kwargs.get('directory', None)
|
||||
self.api_key: str = kwargs.get('api_key', None)
|
||||
self.min_width: int = kwargs.get('min_width', 1024)
|
||||
self.min_height: int = kwargs.get('min_height', 1024)
|
||||
|
||||
if self.host is None:
|
||||
raise ValueError("host is required")
|
||||
if self.collection_id is None:
|
||||
raise ValueError("collection_id is required")
|
||||
if self.directory is None:
|
||||
raise ValueError("directory is required")
|
||||
if self.api_key is None:
|
||||
raise ValueError(f"api_key is required: {self.host}:{self.collection_id}")
|
||||
|
||||
|
||||
class ImageState:
|
||||
def __init__(self, **kwargs):
|
||||
self.steps_complete: list[Step] = kwargs.get('steps_complete', [])
|
||||
self.steps_to_complete: list[Step] = kwargs.get('steps_to_complete', [])
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
'steps_complete': self.steps_complete
|
||||
}
|
||||
|
||||
|
||||
class Rect:
|
||||
def __init__(self, **kwargs):
|
||||
self.x = kwargs.get('x', 0)
|
||||
self.y = kwargs.get('y', 0)
|
||||
self.width = kwargs.get('width', 0)
|
||||
self.height = kwargs.get('height', 0)
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
'x': self.x,
|
||||
'y': self.y,
|
||||
'width': self.width,
|
||||
'height': self.height
|
||||
}
|
||||
|
||||
|
||||
class ImgInfo:
|
||||
def __init__(self, **kwargs):
|
||||
self.version: int = kwargs.get('version', None)
|
||||
self.caption: str = kwargs.get('caption', None)
|
||||
self.caption_short: str = kwargs.get('caption_short', None)
|
||||
self.poi = [Rect(**poi) for poi in kwargs.get('poi', [])]
|
||||
self.state = ImageState(**kwargs.get('state', {}))
|
||||
self.caption_method = kwargs.get('caption_method', None)
|
||||
self.other_captions = kwargs.get('other_captions', {})
|
||||
self._upgrade_state()
|
||||
self.force_image_process: bool = False
|
||||
self._requested_steps: list[Step] = []
|
||||
|
||||
self.is_dirty: bool = False
|
||||
|
||||
def _upgrade_state(self):
|
||||
# upgrades older states
|
||||
if self.caption is not None and 'caption' not in self.state.steps_complete:
|
||||
self.mark_step_complete('caption')
|
||||
self.is_dirty = True
|
||||
if self.caption_short is not None and 'caption_short' not in self.state.steps_complete:
|
||||
self.mark_step_complete('caption_short')
|
||||
self.is_dirty = True
|
||||
if self.caption_method is None and self.caption is not None:
|
||||
# added caption method in version 2. Was all llava before that
|
||||
self.caption_method = 'llava:default'
|
||||
self.is_dirty = True
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
'version': self.version,
|
||||
'caption_method': self.caption_method,
|
||||
'caption': self.caption,
|
||||
'caption_short': self.caption_short,
|
||||
'poi': [poi.to_dict() for poi in self.poi],
|
||||
'state': self.state.to_dict(),
|
||||
'other_captions': self.other_captions
|
||||
}
|
||||
|
||||
def mark_step_complete(self, step: Step):
|
||||
if step not in self.state.steps_complete:
|
||||
self.state.steps_complete.append(step)
|
||||
if step in self.state.steps_to_complete:
|
||||
self.state.steps_to_complete.remove(step)
|
||||
self.is_dirty = True
|
||||
|
||||
def add_step(self, step: Step):
|
||||
if step not in self.state.steps_to_complete and step not in self.state.steps_complete:
|
||||
self.state.steps_to_complete.append(step)
|
||||
|
||||
def trigger_image_reprocess(self):
|
||||
if self._requested_steps is None:
|
||||
raise Exception("Must call add_steps before trigger_image_reprocess")
|
||||
steps = self._requested_steps
|
||||
# remove all image manipulationf from steps_to_complete
|
||||
for step in img_manipulation_steps:
|
||||
if step in self.state.steps_to_complete:
|
||||
self.state.steps_to_complete.remove(step)
|
||||
if step in self.state.steps_complete:
|
||||
self.state.steps_complete.remove(step)
|
||||
self.force_image_process = True
|
||||
self.is_dirty = True
|
||||
# we want to keep the order passed in process file
|
||||
for step in steps:
|
||||
if step in img_manipulation_steps:
|
||||
self.add_step(step)
|
||||
|
||||
def add_steps(self, steps: list[Step]):
|
||||
self._requested_steps = [step for step in steps]
|
||||
for stage in steps:
|
||||
self.add_step(stage)
|
||||
|
||||
# update steps if we have any img processes not complete, we have to reprocess them all
|
||||
# if any steps_to_complete are in img_manipulation_steps
|
||||
|
||||
is_manipulating_image = any([step in img_manipulation_steps for step in self.state.steps_to_complete])
|
||||
order_has_changed = False
|
||||
|
||||
if not is_manipulating_image:
|
||||
# check to see if order has changed. No need to if already redoing it. Will detect if ones are removed
|
||||
target_img_manipulation_order = [step for step in steps if step in img_manipulation_steps]
|
||||
current_img_manipulation_order = [step for step in self.state.steps_complete if
|
||||
step in img_manipulation_steps]
|
||||
if target_img_manipulation_order != current_img_manipulation_order:
|
||||
order_has_changed = True
|
||||
|
||||
if is_manipulating_image or order_has_changed:
|
||||
self.trigger_image_reprocess()
|
||||
|
||||
def set_caption_method(self, method: str):
|
||||
if self._requested_steps is None:
|
||||
raise Exception("Must call add_steps before set_caption_method")
|
||||
if self.caption_method != method:
|
||||
self.is_dirty = True
|
||||
# move previous caption method to other_captions
|
||||
if self.caption_method is not None and self.caption is not None or self.caption_short is not None:
|
||||
self.other_captions[self.caption_method] = {
|
||||
'caption': self.caption,
|
||||
'caption_short': self.caption_short,
|
||||
}
|
||||
self.caption_method = method
|
||||
self.caption = None
|
||||
self.caption_short = None
|
||||
# see if we have a caption from the new method
|
||||
if method in self.other_captions:
|
||||
self.caption = self.other_captions[method].get('caption', None)
|
||||
self.caption_short = self.other_captions[method].get('caption_short', None)
|
||||
else:
|
||||
self.trigger_new_caption()
|
||||
|
||||
def trigger_new_caption(self):
|
||||
self.caption = None
|
||||
self.caption_short = None
|
||||
self.is_dirty = True
|
||||
# check to see if we have any steps in the complete list and move them to the to_complete list
|
||||
for step in self.state.steps_complete:
|
||||
if step in caption_manipulation_steps:
|
||||
self.state.steps_complete.remove(step)
|
||||
self.state.steps_to_complete.append(step)
|
||||
|
||||
def to_json(self):
|
||||
return json.dumps(self.to_dict())
|
||||
|
||||
def set_version(self, version: int):
|
||||
if self.version != version:
|
||||
self.is_dirty = True
|
||||
self.version = version
|
||||
66
extensions_built_in/dataset_tools/tools/fuyu_utils.py
Normal file
66
extensions_built_in/dataset_tools/tools/fuyu_utils.py
Normal file
@@ -0,0 +1,66 @@
|
||||
from transformers import CLIPImageProcessor, BitsAndBytesConfig, AutoTokenizer
|
||||
|
||||
from .caption import default_long_prompt, default_short_prompt, default_replacements, clean_caption
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
|
||||
class FuyuImageProcessor:
|
||||
def __init__(self, device='cuda'):
|
||||
from transformers import FuyuProcessor, FuyuForCausalLM
|
||||
self.device = device
|
||||
self.model: FuyuForCausalLM = None
|
||||
self.processor: FuyuProcessor = None
|
||||
self.dtype = torch.bfloat16
|
||||
self.tokenizer: AutoTokenizer
|
||||
self.is_loaded = False
|
||||
|
||||
def load_model(self):
|
||||
from transformers import FuyuProcessor, FuyuForCausalLM
|
||||
model_path = "adept/fuyu-8b"
|
||||
kwargs = {"device_map": self.device}
|
||||
kwargs['load_in_4bit'] = True
|
||||
kwargs['quantization_config'] = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_compute_dtype=self.dtype,
|
||||
bnb_4bit_use_double_quant=True,
|
||||
bnb_4bit_quant_type='nf4'
|
||||
)
|
||||
self.processor = FuyuProcessor.from_pretrained(model_path)
|
||||
self.model = FuyuForCausalLM.from_pretrained(model_path, low_cpu_mem_usage=True, **kwargs)
|
||||
self.is_loaded = True
|
||||
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(model_path)
|
||||
self.model = FuyuForCausalLM.from_pretrained(model_path, torch_dtype=self.dtype, **kwargs)
|
||||
self.processor = FuyuProcessor(image_processor=FuyuImageProcessor(), tokenizer=self.tokenizer)
|
||||
|
||||
def generate_caption(
|
||||
self, image: Image,
|
||||
prompt: str = default_long_prompt,
|
||||
replacements=default_replacements,
|
||||
max_new_tokens=512
|
||||
):
|
||||
# prepare inputs for the model
|
||||
# text_prompt = f"{prompt}\n"
|
||||
|
||||
# image = image.convert('RGB')
|
||||
model_inputs = self.processor(text=prompt, images=[image])
|
||||
model_inputs = {k: v.to(dtype=self.dtype if torch.is_floating_point(v) else v.dtype, device=self.device) for k, v in
|
||||
model_inputs.items()}
|
||||
|
||||
generation_output = self.model.generate(**model_inputs, max_new_tokens=max_new_tokens)
|
||||
prompt_len = model_inputs["input_ids"].shape[-1]
|
||||
output = self.tokenizer.decode(generation_output[0][prompt_len:], skip_special_tokens=True)
|
||||
output = clean_caption(output, replacements=replacements)
|
||||
return output
|
||||
|
||||
# inputs = self.processor(text=text_prompt, images=image, return_tensors="pt")
|
||||
# for k, v in inputs.items():
|
||||
# inputs[k] = v.to(self.device)
|
||||
|
||||
# # autoregressively generate text
|
||||
# generation_output = self.model.generate(**inputs, max_new_tokens=max_new_tokens)
|
||||
# generation_text = self.processor.batch_decode(generation_output[:, -max_new_tokens:], skip_special_tokens=True)
|
||||
# output = generation_text[0]
|
||||
#
|
||||
# return clean_caption(output, replacements=replacements)
|
||||
49
extensions_built_in/dataset_tools/tools/image_tools.py
Normal file
49
extensions_built_in/dataset_tools/tools/image_tools.py
Normal file
@@ -0,0 +1,49 @@
|
||||
from typing import Literal, Type, TYPE_CHECKING, Union
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
Step: Type = Literal['caption', 'caption_short', 'create_mask', 'contrast_stretch']
|
||||
|
||||
img_manipulation_steps = ['contrast_stretch']
|
||||
|
||||
img_ext = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .llava_utils import LLaVAImageProcessor
|
||||
from .fuyu_utils import FuyuImageProcessor
|
||||
|
||||
ImageProcessor = Union['LLaVAImageProcessor', 'FuyuImageProcessor']
|
||||
|
||||
|
||||
def pil_to_cv2(image):
|
||||
"""Convert a PIL image to a cv2 image."""
|
||||
return cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
|
||||
|
||||
|
||||
def cv2_to_pil(image):
|
||||
"""Convert a cv2 image to a PIL image."""
|
||||
return Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))
|
||||
|
||||
|
||||
def load_image(img_path: str):
|
||||
image = Image.open(img_path).convert('RGB')
|
||||
try:
|
||||
# transpose with exif data
|
||||
image = ImageOps.exif_transpose(image)
|
||||
except Exception as e:
|
||||
pass
|
||||
return image
|
||||
|
||||
|
||||
def resize_to_max(image, max_width=1024, max_height=1024):
|
||||
width, height = image.size
|
||||
if width <= max_width and height <= max_height:
|
||||
return image
|
||||
|
||||
scale = min(max_width / width, max_height / height)
|
||||
width = int(width * scale)
|
||||
height = int(height * scale)
|
||||
|
||||
return image.resize((width, height), Image.LANCZOS)
|
||||
85
extensions_built_in/dataset_tools/tools/llava_utils.py
Normal file
85
extensions_built_in/dataset_tools/tools/llava_utils.py
Normal file
@@ -0,0 +1,85 @@
|
||||
|
||||
from .caption import default_long_prompt, default_short_prompt, default_replacements, clean_caption
|
||||
|
||||
import torch
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
from transformers import AutoTokenizer, BitsAndBytesConfig, CLIPImageProcessor
|
||||
|
||||
img_ext = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
|
||||
|
||||
class LLaVAImageProcessor:
|
||||
def __init__(self, device='cuda'):
|
||||
try:
|
||||
from llava.model import LlavaLlamaForCausalLM
|
||||
except ImportError:
|
||||
# print("You need to manually install llava -> pip install --no-deps git+https://github.com/haotian-liu/LLaVA.git")
|
||||
print(
|
||||
"You need to manually install llava -> pip install --no-deps git+https://github.com/haotian-liu/LLaVA.git")
|
||||
raise
|
||||
self.device = device
|
||||
self.model: LlavaLlamaForCausalLM = None
|
||||
self.tokenizer: AutoTokenizer = None
|
||||
self.image_processor: CLIPImageProcessor = None
|
||||
self.is_loaded = False
|
||||
|
||||
def load_model(self):
|
||||
from llava.model import LlavaLlamaForCausalLM
|
||||
|
||||
model_path = "4bit/llava-v1.5-13b-3GB"
|
||||
# kwargs = {"device_map": "auto"}
|
||||
kwargs = {"device_map": self.device}
|
||||
kwargs['load_in_4bit'] = True
|
||||
kwargs['quantization_config'] = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_compute_dtype=torch.float16,
|
||||
bnb_4bit_use_double_quant=True,
|
||||
bnb_4bit_quant_type='nf4'
|
||||
)
|
||||
self.model = LlavaLlamaForCausalLM.from_pretrained(model_path, low_cpu_mem_usage=True, **kwargs)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(model_path, use_fast=False)
|
||||
vision_tower = self.model.get_vision_tower()
|
||||
if not vision_tower.is_loaded:
|
||||
vision_tower.load_model()
|
||||
vision_tower.to(device=self.device)
|
||||
self.image_processor = vision_tower.image_processor
|
||||
self.is_loaded = True
|
||||
|
||||
def generate_caption(
|
||||
self, image:
|
||||
Image, prompt: str = default_long_prompt,
|
||||
replacements=default_replacements,
|
||||
max_new_tokens=512
|
||||
):
|
||||
from llava.conversation import conv_templates, SeparatorStyle
|
||||
from llava.utils import disable_torch_init
|
||||
from llava.constants import IMAGE_TOKEN_INDEX, DEFAULT_IMAGE_TOKEN, DEFAULT_IM_START_TOKEN, DEFAULT_IM_END_TOKEN
|
||||
from llava.mm_utils import tokenizer_image_token, KeywordsStoppingCriteria
|
||||
# question = "how many dogs are in the picture?"
|
||||
disable_torch_init()
|
||||
conv_mode = "llava_v0"
|
||||
conv = conv_templates[conv_mode].copy()
|
||||
roles = conv.roles
|
||||
image_tensor = self.image_processor.preprocess([image], return_tensors='pt')['pixel_values'].half().cuda()
|
||||
|
||||
inp = f"{roles[0]}: {prompt}"
|
||||
inp = DEFAULT_IM_START_TOKEN + DEFAULT_IMAGE_TOKEN + DEFAULT_IM_END_TOKEN + '\n' + inp
|
||||
conv.append_message(conv.roles[0], inp)
|
||||
conv.append_message(conv.roles[1], None)
|
||||
raw_prompt = conv.get_prompt()
|
||||
input_ids = tokenizer_image_token(raw_prompt, self.tokenizer, IMAGE_TOKEN_INDEX,
|
||||
return_tensors='pt').unsqueeze(0).cuda()
|
||||
stop_str = conv.sep if conv.sep_style != SeparatorStyle.TWO else conv.sep2
|
||||
keywords = [stop_str]
|
||||
stopping_criteria = KeywordsStoppingCriteria(keywords, self.tokenizer, input_ids)
|
||||
with torch.inference_mode():
|
||||
output_ids = self.model.generate(
|
||||
input_ids, images=image_tensor, do_sample=True, temperature=0.1,
|
||||
max_new_tokens=max_new_tokens, use_cache=True, stopping_criteria=[stopping_criteria],
|
||||
top_p=0.8
|
||||
)
|
||||
outputs = self.tokenizer.decode(output_ids[0, input_ids.shape[1]:]).strip()
|
||||
conv.messages[-1][-1] = outputs
|
||||
output = outputs.rsplit('</s>', 1)[0]
|
||||
return clean_caption(output, replacements=replacements)
|
||||
279
extensions_built_in/dataset_tools/tools/sync_tools.py
Normal file
279
extensions_built_in/dataset_tools/tools/sync_tools.py
Normal file
@@ -0,0 +1,279 @@
|
||||
import os
|
||||
import requests
|
||||
import tqdm
|
||||
from typing import List, Optional, TYPE_CHECKING
|
||||
|
||||
|
||||
def img_root_path(img_id: str):
|
||||
return os.path.dirname(os.path.dirname(img_id))
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .dataset_tools_config_modules import DatasetSyncCollectionConfig
|
||||
|
||||
img_exts = ['.jpg', '.jpeg', '.webp', '.png']
|
||||
|
||||
class Photo:
|
||||
def __init__(
|
||||
self,
|
||||
id,
|
||||
host,
|
||||
width,
|
||||
height,
|
||||
url,
|
||||
filename
|
||||
):
|
||||
self.id = str(id)
|
||||
self.host = host
|
||||
self.width = width
|
||||
self.height = height
|
||||
self.url = url
|
||||
self.filename = filename
|
||||
|
||||
|
||||
def get_desired_size(img_width: int, img_height: int, min_width: int, min_height: int):
|
||||
if img_width > img_height:
|
||||
scale = min_height / img_height
|
||||
else:
|
||||
scale = min_width / img_width
|
||||
|
||||
new_width = int(img_width * scale)
|
||||
new_height = int(img_height * scale)
|
||||
|
||||
return new_width, new_height
|
||||
|
||||
|
||||
def get_pexels_images(config: 'DatasetSyncCollectionConfig') -> List[Photo]:
|
||||
all_images = []
|
||||
next_page = f"https://api.pexels.com/v1/collections/{config.collection_id}?page=1&per_page=80&type=photos"
|
||||
|
||||
while True:
|
||||
response = requests.get(next_page, headers={
|
||||
"Authorization": f"{config.api_key}"
|
||||
})
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
all_images.extend(data['media'])
|
||||
if 'next_page' in data and data['next_page']:
|
||||
next_page = data['next_page']
|
||||
else:
|
||||
break
|
||||
|
||||
photos = []
|
||||
for image in all_images:
|
||||
new_width, new_height = get_desired_size(image['width'], image['height'], config.min_width, config.min_height)
|
||||
url = f"{image['src']['original']}?auto=compress&cs=tinysrgb&h={new_height}&w={new_width}"
|
||||
filename = os.path.basename(image['src']['original'])
|
||||
|
||||
photos.append(Photo(
|
||||
id=image['id'],
|
||||
host="pexels",
|
||||
width=image['width'],
|
||||
height=image['height'],
|
||||
url=url,
|
||||
filename=filename
|
||||
))
|
||||
|
||||
return photos
|
||||
|
||||
|
||||
def get_unsplash_images(config: 'DatasetSyncCollectionConfig') -> List[Photo]:
|
||||
headers = {
|
||||
# "Authorization": f"Client-ID {UNSPLASH_ACCESS_KEY}"
|
||||
"Authorization": f"Client-ID {config.api_key}"
|
||||
}
|
||||
# headers['Authorization'] = f"Bearer {token}"
|
||||
|
||||
url = f"https://api.unsplash.com/collections/{config.collection_id}/photos?page=1&per_page=30"
|
||||
response = requests.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
res_headers = response.headers
|
||||
# parse the link header to get the next page
|
||||
# 'Link': '<https://api.unsplash.com/collections/mIPWwLdfct8/photos?page=82>; rel="last", <https://api.unsplash.com/collections/mIPWwLdfct8/photos?page=2>; rel="next"'
|
||||
has_next_page = False
|
||||
if 'Link' in res_headers:
|
||||
has_next_page = True
|
||||
link_header = res_headers['Link']
|
||||
link_header = link_header.split(',')
|
||||
link_header = [link.strip() for link in link_header]
|
||||
link_header = [link.split(';') for link in link_header]
|
||||
link_header = [[link[0].strip('<>'), link[1].strip().strip('"')] for link in link_header]
|
||||
link_header = {link[1]: link[0] for link in link_header}
|
||||
|
||||
# get page number from last url
|
||||
last_page = link_header['rel="last']
|
||||
last_page = last_page.split('?')[1]
|
||||
last_page = last_page.split('&')
|
||||
last_page = [param.split('=') for param in last_page]
|
||||
last_page = {param[0]: param[1] for param in last_page}
|
||||
last_page = int(last_page['page'])
|
||||
|
||||
all_images = response.json()
|
||||
|
||||
if has_next_page:
|
||||
# assume we start on page 1, so we don't need to get it again
|
||||
for page in tqdm.tqdm(range(2, last_page + 1)):
|
||||
url = f"https://api.unsplash.com/collections/{config.collection_id}/photos?page={page}&per_page=30"
|
||||
response = requests.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
all_images.extend(response.json())
|
||||
|
||||
photos = []
|
||||
for image in all_images:
|
||||
new_width, new_height = get_desired_size(image['width'], image['height'], config.min_width, config.min_height)
|
||||
url = f"{image['urls']['raw']}&w={new_width}"
|
||||
filename = f"{image['id']}.jpg"
|
||||
|
||||
photos.append(Photo(
|
||||
id=image['id'],
|
||||
host="unsplash",
|
||||
width=image['width'],
|
||||
height=image['height'],
|
||||
url=url,
|
||||
filename=filename
|
||||
))
|
||||
|
||||
return photos
|
||||
|
||||
|
||||
def get_img_paths(dir_path: str):
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
local_files = os.listdir(dir_path)
|
||||
# remove non image files
|
||||
local_files = [file for file in local_files if os.path.splitext(file)[1].lower() in img_exts]
|
||||
# make full path
|
||||
local_files = [os.path.join(dir_path, file) for file in local_files]
|
||||
return local_files
|
||||
|
||||
|
||||
def get_local_image_ids(dir_path: str):
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
local_files = get_img_paths(dir_path)
|
||||
# assuming local files are named after Unsplash IDs, e.g., 'abc123.jpg'
|
||||
return set([os.path.basename(file).split('.')[0] for file in local_files])
|
||||
|
||||
|
||||
def get_local_image_file_names(dir_path: str):
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
local_files = get_img_paths(dir_path)
|
||||
# assuming local files are named after Unsplash IDs, e.g., 'abc123.jpg'
|
||||
return set([os.path.basename(file) for file in local_files])
|
||||
|
||||
|
||||
def download_image(photo: Photo, dir_path: str, min_width: int = 1024, min_height: int = 1024):
|
||||
img_width = photo.width
|
||||
img_height = photo.height
|
||||
|
||||
if img_width < min_width or img_height < min_height:
|
||||
raise ValueError(f"Skipping {photo.id} because it is too small: {img_width}x{img_height}")
|
||||
|
||||
img_response = requests.get(photo.url)
|
||||
img_response.raise_for_status()
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
|
||||
filename = os.path.join(dir_path, photo.filename)
|
||||
with open(filename, 'wb') as file:
|
||||
file.write(img_response.content)
|
||||
|
||||
|
||||
def update_caption(img_path: str):
|
||||
# if the caption is a txt file, convert it to a json file
|
||||
filename_no_ext = os.path.splitext(os.path.basename(img_path))[0]
|
||||
# see if it exists
|
||||
if os.path.exists(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.json")):
|
||||
# todo add poi and what not
|
||||
return # we have a json file
|
||||
caption = ""
|
||||
# see if txt file exists
|
||||
if os.path.exists(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.txt")):
|
||||
# read it
|
||||
with open(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.txt"), 'r') as file:
|
||||
caption = file.read()
|
||||
# write json file
|
||||
with open(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.json"), 'w') as file:
|
||||
file.write(f'{{"caption": "{caption}"}}')
|
||||
|
||||
# delete txt file
|
||||
os.remove(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.txt"))
|
||||
|
||||
|
||||
# def equalize_img(img_path: str):
|
||||
# input_path = img_path
|
||||
# output_path = os.path.join(img_root_path(img_path), COLOR_CORRECTED_DIR, os.path.basename(img_path))
|
||||
# os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
# process_img(
|
||||
# img_path=input_path,
|
||||
# output_path=output_path,
|
||||
# equalize=True,
|
||||
# max_size=2056,
|
||||
# white_balance=False,
|
||||
# gamma_correction=False,
|
||||
# strength=0.6,
|
||||
# )
|
||||
|
||||
|
||||
# def annotate_depth(img_path: str):
|
||||
# # make fake args
|
||||
# args = argparse.Namespace()
|
||||
# args.annotator = "midas"
|
||||
# args.res = 1024
|
||||
#
|
||||
# img = cv2.imread(img_path)
|
||||
# img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
#
|
||||
# output = annotate(img, args)
|
||||
#
|
||||
# output = output.astype('uint8')
|
||||
# output = cv2.cvtColor(output, cv2.COLOR_RGB2BGR)
|
||||
#
|
||||
# os.makedirs(os.path.dirname(img_path), exist_ok=True)
|
||||
# output_path = os.path.join(img_root_path(img_path), DEPTH_DIR, os.path.basename(img_path))
|
||||
#
|
||||
# cv2.imwrite(output_path, output)
|
||||
|
||||
|
||||
# def invert_depth(img_path: str):
|
||||
# img = cv2.imread(img_path)
|
||||
# img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
# # invert the colors
|
||||
# img = cv2.bitwise_not(img)
|
||||
#
|
||||
# os.makedirs(os.path.dirname(img_path), exist_ok=True)
|
||||
# output_path = os.path.join(img_root_path(img_path), INVERTED_DEPTH_DIR, os.path.basename(img_path))
|
||||
# cv2.imwrite(output_path, img)
|
||||
|
||||
|
||||
#
|
||||
# # update our list of raw images
|
||||
# raw_images = get_img_paths(raw_dir)
|
||||
#
|
||||
# # update raw captions
|
||||
# for image_id in tqdm.tqdm(raw_images, desc="Updating raw captions"):
|
||||
# update_caption(image_id)
|
||||
#
|
||||
# # equalize images
|
||||
# for img_path in tqdm.tqdm(raw_images, desc="Equalizing images"):
|
||||
# if img_path not in eq_images:
|
||||
# equalize_img(img_path)
|
||||
#
|
||||
# # update our list of eq images
|
||||
# eq_images = get_img_paths(eq_dir)
|
||||
# # update eq captions
|
||||
# for image_id in tqdm.tqdm(eq_images, desc="Updating eq captions"):
|
||||
# update_caption(image_id)
|
||||
#
|
||||
# # annotate depth
|
||||
# depth_dir = os.path.join(root_dir, DEPTH_DIR)
|
||||
# depth_images = get_img_paths(depth_dir)
|
||||
# for img_path in tqdm.tqdm(eq_images, desc="Annotating depth"):
|
||||
# if img_path not in depth_images:
|
||||
# annotate_depth(img_path)
|
||||
#
|
||||
# depth_images = get_img_paths(depth_dir)
|
||||
#
|
||||
# # invert depth
|
||||
# inv_depth_dir = os.path.join(root_dir, INVERTED_DEPTH_DIR)
|
||||
# inv_depth_images = get_img_paths(inv_depth_dir)
|
||||
# for img_path in tqdm.tqdm(depth_images, desc="Inverting depth"):
|
||||
# if img_path not in inv_depth_images:
|
||||
# invert_depth(img_path)
|
||||
22
extensions_built_in/diffusion_models/__init__.py
Normal file
22
extensions_built_in/diffusion_models/__init__.py
Normal file
@@ -0,0 +1,22 @@
|
||||
from .chroma import ChromaModel
|
||||
from .hidream import HidreamModel, HidreamE1Model
|
||||
from .f_light import FLiteModel
|
||||
from .omnigen2 import OmniGen2Model
|
||||
from .flux_kontext import FluxKontextModel
|
||||
from .wan22 import Wan225bModel, Wan2214bModel, Wan2214bI2VModel
|
||||
from .qwen_image import QwenImageModel, QwenImageEditModel
|
||||
|
||||
AI_TOOLKIT_MODELS = [
|
||||
# put a list of models here
|
||||
ChromaModel,
|
||||
HidreamModel,
|
||||
HidreamE1Model,
|
||||
FLiteModel,
|
||||
OmniGen2Model,
|
||||
FluxKontextModel,
|
||||
Wan225bModel,
|
||||
Wan2214bI2VModel,
|
||||
Wan2214bModel,
|
||||
QwenImageModel,
|
||||
QwenImageEditModel,
|
||||
]
|
||||
1
extensions_built_in/diffusion_models/chroma/__init__.py
Normal file
1
extensions_built_in/diffusion_models/chroma/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
from .chroma_model import ChromaModel
|
||||
450
extensions_built_in/diffusion_models/chroma/chroma_model.py
Normal file
450
extensions_built_in/diffusion_models/chroma/chroma_model.py
Normal file
@@ -0,0 +1,450 @@
|
||||
import os
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||
from PIL import Image
|
||||
from toolkit.models.base_model import BaseModel
|
||||
from toolkit.basic import flush
|
||||
from diffusers import AutoencoderKL
|
||||
# from toolkit.pixel_shuffle_encoder import AutoencoderPixelMixer
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
from toolkit.accelerator import unwrap_model
|
||||
from optimum.quanto import freeze, QTensor
|
||||
from toolkit.util.quantize import quantize, get_qtype
|
||||
from transformers import T5TokenizerFast, T5EncoderModel, CLIPTextModel, CLIPTokenizer
|
||||
from .pipeline import ChromaPipeline
|
||||
from einops import rearrange, repeat
|
||||
import random
|
||||
import torch.nn.functional as F
|
||||
from .src.model import Chroma, chroma_params
|
||||
from safetensors.torch import load_file, save_file
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
import huggingface_hub
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
|
||||
scheduler_config = {
|
||||
"base_image_seq_len": 256,
|
||||
"base_shift": 0.5,
|
||||
"max_image_seq_len": 4096,
|
||||
"max_shift": 1.15,
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 3.0,
|
||||
"use_dynamic_shifting": True
|
||||
}
|
||||
|
||||
class FakeConfig:
|
||||
# for diffusers compatability
|
||||
def __init__(self):
|
||||
self.attention_head_dim = 128
|
||||
self.guidance_embeds = True
|
||||
self.in_channels = 64
|
||||
self.joint_attention_dim = 4096
|
||||
self.num_attention_heads = 24
|
||||
self.num_layers = 19
|
||||
self.num_single_layers = 38
|
||||
self.patch_size = 1
|
||||
|
||||
class FakeCLIP(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.dtype = torch.bfloat16
|
||||
self.device = 'cuda'
|
||||
self.text_model = None
|
||||
self.tokenizer = None
|
||||
self.model_max_length = 77
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
return torch.zeros(1, 1, 1).to(self.device)
|
||||
|
||||
|
||||
class ChromaModel(BaseModel):
|
||||
arch = "chroma"
|
||||
|
||||
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 = ['Chroma']
|
||||
|
||||
# static method to get the noise scheduler
|
||||
@staticmethod
|
||||
def get_train_scheduler():
|
||||
return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
|
||||
|
||||
def get_bucket_divisibility(self):
|
||||
# return the bucket divisibility for the model
|
||||
return 32
|
||||
|
||||
def load_model(self):
|
||||
dtype = self.torch_dtype
|
||||
|
||||
# will be updated if we detect a existing checkpoint in training folder
|
||||
model_path = self.model_config.name_or_path
|
||||
|
||||
if model_path == "lodestones/Chroma":
|
||||
print("Looking for latest Chroma checkpoint")
|
||||
# get the latest checkpoint
|
||||
files_list = huggingface_hub.list_repo_files(model_path)
|
||||
print(files_list)
|
||||
latest_version = 28 # current latest version at time of writing
|
||||
while True:
|
||||
if f"chroma-unlocked-v{latest_version}.safetensors" not in files_list:
|
||||
latest_version -= 1
|
||||
break
|
||||
else:
|
||||
latest_version += 1
|
||||
print(f"Using latest Chroma version: v{latest_version}")
|
||||
|
||||
# make sure we have it
|
||||
model_path = huggingface_hub.hf_hub_download(
|
||||
repo_id=model_path,
|
||||
filename=f"chroma-unlocked-v{latest_version}.safetensors",
|
||||
)
|
||||
elif model_path.startswith("lodestones/Chroma/v"):
|
||||
# get the version number
|
||||
version = model_path.split("/")[-1].split("v")[-1]
|
||||
print(f"Using Chroma version: v{version}")
|
||||
# make sure we have it
|
||||
model_path = huggingface_hub.hf_hub_download(
|
||||
repo_id='lodestones/Chroma',
|
||||
filename=f"chroma-unlocked-v{version}.safetensors",
|
||||
)
|
||||
else:
|
||||
# check if the model path is a local file
|
||||
if os.path.exists(model_path):
|
||||
print(f"Using local model: {model_path}")
|
||||
else:
|
||||
raise ValueError(f"Model path {model_path} does not exist")
|
||||
|
||||
# extras_path = 'black-forest-labs/FLUX.1-schnell'
|
||||
# schnell model is gated now, use flex instead
|
||||
extras_path = 'ostris/Flex.1-alpha'
|
||||
|
||||
self.print_and_status_update("Loading transformer")
|
||||
|
||||
chroma_state_dict = load_file(model_path, 'cpu')
|
||||
|
||||
# determine number of double and single blocks
|
||||
double_blocks = 0
|
||||
single_blocks = 0
|
||||
for key in chroma_state_dict.keys():
|
||||
if "double_blocks" in key:
|
||||
block_num = int(key.split(".")[1]) + 1
|
||||
if block_num > double_blocks:
|
||||
double_blocks = block_num
|
||||
elif "single_blocks" in key:
|
||||
block_num = int(key.split(".")[1]) + 1
|
||||
if block_num > single_blocks:
|
||||
single_blocks = block_num
|
||||
print(f"Double Blocks: {double_blocks}")
|
||||
print(f"Single Blocks: {single_blocks}")
|
||||
|
||||
chroma_params.depth = double_blocks
|
||||
chroma_params.depth_single_blocks = single_blocks
|
||||
transformer = Chroma(chroma_params)
|
||||
|
||||
# add dtype, not sure why it doesnt have it
|
||||
transformer.dtype = dtype
|
||||
# load the state dict into the model
|
||||
transformer.load_state_dict(chroma_state_dict)
|
||||
|
||||
transformer.to(self.quantize_device, dtype=dtype)
|
||||
|
||||
transformer.config = FakeConfig()
|
||||
transformer.config.num_layers = double_blocks
|
||||
transformer.config.num_single_layers = single_blocks
|
||||
|
||||
if self.model_config.quantize:
|
||||
# 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,
|
||||
**self.model_config.quantize_kwargs)
|
||||
freeze(transformer)
|
||||
transformer.to(self.device_torch)
|
||||
else:
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Loading T5")
|
||||
tokenizer_2 = T5TokenizerFast.from_pretrained(
|
||||
extras_path, subfolder="tokenizer_2", torch_dtype=dtype
|
||||
)
|
||||
text_encoder_2 = T5EncoderModel.from_pretrained(
|
||||
extras_path, subfolder="text_encoder_2", torch_dtype=dtype
|
||||
)
|
||||
text_encoder_2.to(self.device_torch, dtype=dtype)
|
||||
flush()
|
||||
|
||||
if self.model_config.quantize_te:
|
||||
self.print_and_status_update("Quantizing T5")
|
||||
quantize(text_encoder_2, weights=get_qtype(
|
||||
self.model_config.qtype))
|
||||
freeze(text_encoder_2)
|
||||
flush()
|
||||
|
||||
# self.print_and_status_update("Loading CLIP")
|
||||
text_encoder = FakeCLIP()
|
||||
tokenizer = FakeCLIP()
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
|
||||
self.noise_scheduler = ChromaModel.get_train_scheduler()
|
||||
|
||||
self.print_and_status_update("Loading VAE")
|
||||
vae = AutoencoderKL.from_pretrained(
|
||||
extras_path,
|
||||
subfolder="vae",
|
||||
torch_dtype=dtype
|
||||
)
|
||||
vae = vae.to(self.device_torch, dtype=dtype)
|
||||
|
||||
self.print_and_status_update("Making pipe")
|
||||
|
||||
pipe: ChromaPipeline = ChromaPipeline(
|
||||
scheduler=self.noise_scheduler,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
text_encoder_2=None,
|
||||
tokenizer_2=tokenizer_2,
|
||||
vae=vae,
|
||||
transformer=None,
|
||||
)
|
||||
# for quantization, it works best to do these after making the pipe
|
||||
pipe.text_encoder_2 = text_encoder_2
|
||||
pipe.transformer = transformer
|
||||
|
||||
self.print_and_status_update("Preparing Model")
|
||||
|
||||
text_encoder = [pipe.text_encoder, pipe.text_encoder_2]
|
||||
tokenizer = [pipe.tokenizer, pipe.tokenizer_2]
|
||||
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
|
||||
flush()
|
||||
# just to make sure everything is on the right device and dtype
|
||||
text_encoder[0].to(self.device_torch)
|
||||
text_encoder[0].requires_grad_(False)
|
||||
text_encoder[0].eval()
|
||||
text_encoder[1].to(self.device_torch)
|
||||
text_encoder[1].requires_grad_(False)
|
||||
text_encoder[1].eval()
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
flush()
|
||||
|
||||
# save it to the model class
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder # list of text encoders
|
||||
self.tokenizer = tokenizer # list of tokenizers
|
||||
self.model = pipe.transformer
|
||||
self.pipeline = pipe
|
||||
self.print_and_status_update("Model Loaded")
|
||||
|
||||
def get_generation_pipeline(self):
|
||||
scheduler = ChromaModel.get_train_scheduler()
|
||||
pipeline = ChromaPipeline(
|
||||
scheduler=scheduler,
|
||||
text_encoder=unwrap_model(self.text_encoder[0]),
|
||||
tokenizer=self.tokenizer[0],
|
||||
text_encoder_2=unwrap_model(self.text_encoder[1]),
|
||||
tokenizer_2=self.tokenizer[1],
|
||||
vae=unwrap_model(self.vae),
|
||||
transformer=unwrap_model(self.transformer)
|
||||
)
|
||||
|
||||
# pipeline = pipeline.to(self.device_torch)
|
||||
|
||||
return pipeline
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: ChromaPipeline,
|
||||
gen_config: GenerateImageConfig,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
unconditional_embeds: PromptEmbeds,
|
||||
generator: torch.Generator,
|
||||
extra: dict,
|
||||
):
|
||||
|
||||
extra['negative_prompt_embeds'] = unconditional_embeds.text_embeds
|
||||
extra['negative_prompt_attn_mask'] = unconditional_embeds.attention_mask
|
||||
|
||||
img = pipeline(
|
||||
prompt_embeds=conditional_embeds.text_embeds,
|
||||
prompt_attn_mask=conditional_embeds.attention_mask,
|
||||
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
|
||||
):
|
||||
with torch.no_grad():
|
||||
bs, c, h, w = latent_model_input.shape
|
||||
latent_model_input_packed = rearrange(
|
||||
latent_model_input,
|
||||
"b c (h ph) (w pw) -> b (h w) (c ph pw)",
|
||||
ph=2,
|
||||
pw=2
|
||||
)
|
||||
|
||||
img_ids = torch.zeros(h // 2, w // 2, 3)
|
||||
img_ids[..., 1] = img_ids[..., 1] + torch.arange(h // 2)[:, None]
|
||||
img_ids[..., 2] = img_ids[..., 2] + torch.arange(w // 2)[None, :]
|
||||
img_ids = repeat(img_ids, "h w c -> b (h w) c",
|
||||
b=bs).to(self.device_torch)
|
||||
|
||||
txt_ids = torch.zeros(
|
||||
bs, text_embeddings.text_embeds.shape[1], 3).to(self.device_torch)
|
||||
|
||||
guidance = torch.full([1], 0, device=self.device_torch, dtype=torch.float32)
|
||||
guidance = guidance.expand(latent_model_input_packed.shape[0])
|
||||
|
||||
cast_dtype = self.unet.dtype
|
||||
|
||||
noise_pred = self.unet(
|
||||
img=latent_model_input_packed.to(
|
||||
self.device_torch, cast_dtype
|
||||
),
|
||||
img_ids=img_ids,
|
||||
txt=text_embeddings.text_embeds.to(
|
||||
self.device_torch, cast_dtype
|
||||
),
|
||||
txt_ids=txt_ids,
|
||||
txt_mask=text_embeddings.attention_mask.to(
|
||||
self.device_torch, cast_dtype
|
||||
),
|
||||
timesteps=timestep / 1000,
|
||||
guidance=guidance
|
||||
)
|
||||
|
||||
if isinstance(noise_pred, QTensor):
|
||||
noise_pred = noise_pred.dequantize()
|
||||
|
||||
noise_pred = rearrange(
|
||||
noise_pred,
|
||||
"b (h w) (c ph pw) -> b c (h ph) (w pw)",
|
||||
h=latent_model_input.shape[2] // 2,
|
||||
w=latent_model_input.shape[3] // 2,
|
||||
ph=2,
|
||||
pw=2,
|
||||
c=self.vae.config.latent_channels
|
||||
)
|
||||
|
||||
return noise_pred
|
||||
|
||||
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
|
||||
if isinstance(prompt, str):
|
||||
prompts = [prompt]
|
||||
else:
|
||||
prompts = prompt
|
||||
if self.pipeline.text_encoder.device != self.device_torch:
|
||||
self.pipeline.text_encoder.to(self.device_torch)
|
||||
|
||||
max_length = 512
|
||||
|
||||
device = self.text_encoder[1].device
|
||||
dtype = self.text_encoder[1].dtype
|
||||
|
||||
# T5
|
||||
text_inputs = self.tokenizer[1](
|
||||
prompts,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
|
||||
prompt_embeds = self.text_encoder[1](text_input_ids.to(device), output_hidden_states=False)[0]
|
||||
|
||||
dtype = self.text_encoder[1].dtype
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
prompt_attention_mask = text_inputs["attention_mask"]
|
||||
|
||||
pe = PromptEmbeds(
|
||||
prompt_embeds
|
||||
)
|
||||
pe.attention_mask = prompt_attention_mask
|
||||
return pe
|
||||
|
||||
def get_model_has_grad(self):
|
||||
# return from a weight if it has grad
|
||||
return self.model.final_layer.linear.weight.requires_grad
|
||||
|
||||
def get_te_has_grad(self):
|
||||
# return from a weight if it has grad
|
||||
return self.text_encoder[1].encoder.block[0].layer[0].SelfAttention.q.weight.requires_grad
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
if not output_path.endswith(".safetensors"):
|
||||
output_path = output_path + ".safetensors"
|
||||
# only save the unet
|
||||
transformer: Chroma = unwrap_model(self.model)
|
||||
state_dict = transformer.state_dict()
|
||||
save_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
if isinstance(v, QTensor):
|
||||
v = v.dequantize()
|
||||
save_dict[k] = v.clone().to('cpu', dtype=save_dtype)
|
||||
|
||||
meta = get_meta_for_safetensors(meta, name='chroma')
|
||||
save_file(save_dict, output_path, metadata=meta)
|
||||
|
||||
def get_loss_target(self, *args, **kwargs):
|
||||
noise = kwargs.get('noise')
|
||||
batch = kwargs.get('batch')
|
||||
return (noise - batch.latents).detach()
|
||||
|
||||
def convert_lora_weights_before_save(self, state_dict):
|
||||
# currently starte with transformer. but needs to start with diffusion_model. for comfyui
|
||||
new_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
new_key = key.replace("transformer.", "diffusion_model.")
|
||||
new_sd[new_key] = value
|
||||
return new_sd
|
||||
|
||||
def convert_lora_weights_before_load(self, state_dict):
|
||||
# saved as diffusion_model. but needs to be transformer. for ai-toolkit
|
||||
new_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
new_key = key.replace("diffusion_model.", "transformer.")
|
||||
new_sd[new_key] = value
|
||||
return new_sd
|
||||
|
||||
def get_base_model_version(self):
|
||||
return "chroma"
|
||||
197
extensions_built_in/diffusion_models/chroma/pipeline.py
Normal file
197
extensions_built_in/diffusion_models/chroma/pipeline.py
Normal file
@@ -0,0 +1,197 @@
|
||||
from typing import Union, List, Optional, Dict, Any, Callable
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import FluxPipeline
|
||||
from diffusers.pipelines.flux.pipeline_flux import calculate_shift, retrieve_timesteps
|
||||
from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput
|
||||
from diffusers.utils import is_torch_xla_available
|
||||
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
|
||||
XLA_AVAILABLE = True
|
||||
else:
|
||||
XLA_AVAILABLE = False
|
||||
|
||||
|
||||
class ChromaPipeline(FluxPipeline):
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 28,
|
||||
timesteps: List[int] = None,
|
||||
guidance_scale: float = 7.0,
|
||||
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,
|
||||
prompt_attn_mask: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_attn_mask: 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,
|
||||
):
|
||||
|
||||
height = height or self.default_sample_size * self.vae_scale_factor
|
||||
width = width or self.default_sample_size * self.vae_scale_factor
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._joint_attention_kwargs = joint_attention_kwargs
|
||||
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
|
||||
if isinstance(device, str):
|
||||
device = torch.device(device)
|
||||
|
||||
text_ids = torch.zeros(batch_size, prompt_embeds.shape[1], 3).to(device=device, dtype=torch.bfloat16)
|
||||
if guidance_scale > 1.00001:
|
||||
negative_text_ids = torch.zeros(batch_size, negative_prompt_embeds.shape[1], 3).to(device=device, dtype=torch.bfloat16)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
num_channels_latents = 64 // 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,
|
||||
)
|
||||
|
||||
# extend img ids to match batch size
|
||||
latent_image_ids = latent_image_ids.unsqueeze(0)
|
||||
latent_image_ids = torch.cat([latent_image_ids] * batch_size, dim=0)
|
||||
|
||||
# 5. Prepare timesteps
|
||||
sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps)
|
||||
image_seq_len = latents.shape[1]
|
||||
mu = calculate_shift(
|
||||
image_seq_len,
|
||||
self.scheduler.config.base_image_seq_len,
|
||||
self.scheduler.config.max_image_seq_len,
|
||||
self.scheduler.config.base_shift,
|
||||
self.scheduler.config.max_shift,
|
||||
)
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
timesteps,
|
||||
sigmas,
|
||||
mu=mu,
|
||||
)
|
||||
num_warmup_steps = max(
|
||||
len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
guidance = torch.full([1], 0, device=device, dtype=torch.float32)
|
||||
guidance = guidance.expand(latents.shape[0])
|
||||
|
||||
# 6. Denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latents.shape[0]).to(latents.dtype)
|
||||
|
||||
# handle guidance
|
||||
|
||||
noise_pred_text = self.transformer(
|
||||
img=latents,
|
||||
img_ids=latent_image_ids,
|
||||
txt=prompt_embeds,
|
||||
txt_ids=text_ids,
|
||||
txt_mask=prompt_attn_mask, # todo add this
|
||||
timesteps=timestep / 1000,
|
||||
guidance=guidance
|
||||
)
|
||||
|
||||
if guidance_scale > 1.00001:
|
||||
noise_pred_uncond = self.transformer(
|
||||
img=latents,
|
||||
img_ids=latent_image_ids,
|
||||
txt=negative_prompt_embeds,
|
||||
txt_ids=negative_text_ids,
|
||||
txt_mask=negative_prompt_attn_mask, # todo add this
|
||||
timesteps=timestep / 1000,
|
||||
guidance=guidance
|
||||
)
|
||||
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * \
|
||||
(noise_pred_text - noise_pred_uncond)
|
||||
|
||||
else:
|
||||
noise_pred = noise_pred_text
|
||||
|
||||
# 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()
|
||||
|
||||
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)
|
||||
@@ -0,0 +1 @@
|
||||
# This is taken and slightly modified from https://github.com/lodestone-rock/flow/tree/master/src/models/chroma
|
||||
505
extensions_built_in/diffusion_models/chroma/src/layers.py
Normal file
505
extensions_built_in/diffusion_models/chroma/src/layers.py
Normal file
@@ -0,0 +1,505 @@
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from torch import Tensor, nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .math import attention, rope
|
||||
|
||||
|
||||
class EmbedND(nn.Module):
|
||||
def __init__(self, dim: int, theta: int, axes_dim: list[int]):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.theta = theta
|
||||
self.axes_dim = axes_dim
|
||||
|
||||
def forward(self, ids: Tensor) -> Tensor:
|
||||
n_axes = ids.shape[-1]
|
||||
emb = torch.cat(
|
||||
[rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)],
|
||||
dim=-3,
|
||||
)
|
||||
|
||||
return emb.unsqueeze(1)
|
||||
|
||||
|
||||
def timestep_embedding(t: Tensor, dim, max_period=10000, time_factor: float = 1000.0):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param t: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
:param dim: the dimension of the output.
|
||||
:param max_period: controls the minimum frequency of the embeddings.
|
||||
:return: an (N, D) Tensor of positional embeddings.
|
||||
"""
|
||||
t = time_factor * t
|
||||
half = dim // 2
|
||||
freqs = torch.exp(
|
||||
-math.log(max_period)
|
||||
* torch.arange(start=0, end=half, dtype=torch.float32)
|
||||
/ half
|
||||
).to(t.device)
|
||||
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
if torch.is_floating_point(t):
|
||||
embedding = embedding.to(t)
|
||||
return embedding
|
||||
|
||||
|
||||
class MLPEmbedder(nn.Module):
|
||||
def __init__(self, in_dim: int, hidden_dim: int):
|
||||
super().__init__()
|
||||
self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True)
|
||||
self.silu = nn.SiLU()
|
||||
self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
# Get the device of the module (assumes all parameters are on the same device)
|
||||
return next(self.parameters()).device
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return self.out_layer(self.silu(self.in_layer(x)))
|
||||
|
||||
|
||||
class RMSNorm(torch.nn.Module):
|
||||
def __init__(self, dim: int, use_compiled: bool = False):
|
||||
super().__init__()
|
||||
self.scale = nn.Parameter(torch.ones(dim))
|
||||
self.use_compiled = use_compiled
|
||||
|
||||
def _forward(self, x: Tensor):
|
||||
x_dtype = x.dtype
|
||||
x = x.float()
|
||||
rrms = torch.rsqrt(torch.mean(x**2, dim=-1, keepdim=True) + 1e-6)
|
||||
return (x * rrms).to(dtype=x_dtype) * self.scale
|
||||
|
||||
def forward(self, x: Tensor):
|
||||
return F.rms_norm(x, self.scale.shape, weight=self.scale, eps=1e-6)
|
||||
# if self.use_compiled:
|
||||
# return torch.compile(self._forward)(x)
|
||||
# else:
|
||||
# return self._forward(x)
|
||||
|
||||
|
||||
def distribute_modulations(tensor: torch.Tensor):
|
||||
"""
|
||||
Distributes slices of the tensor into the block_dict as ModulationOut objects.
|
||||
|
||||
Args:
|
||||
tensor (torch.Tensor): Input tensor with shape [batch_size, vectors, dim].
|
||||
"""
|
||||
batch_size, vectors, dim = tensor.shape
|
||||
|
||||
block_dict = {}
|
||||
|
||||
# HARD CODED VALUES! lookup table for the generated vectors
|
||||
# TODO: move this into chroma config!
|
||||
# Add 38 single mod blocks
|
||||
for i in range(38):
|
||||
key = f"single_blocks.{i}.modulation.lin"
|
||||
block_dict[key] = None
|
||||
|
||||
# Add 19 image double blocks
|
||||
for i in range(19):
|
||||
key = f"double_blocks.{i}.img_mod.lin"
|
||||
block_dict[key] = None
|
||||
|
||||
# Add 19 text double blocks
|
||||
for i in range(19):
|
||||
key = f"double_blocks.{i}.txt_mod.lin"
|
||||
block_dict[key] = None
|
||||
|
||||
# Add the final layer
|
||||
block_dict["final_layer.adaLN_modulation.1"] = None
|
||||
# 6.2b version
|
||||
block_dict["lite_double_blocks.4.img_mod.lin"] = None
|
||||
block_dict["lite_double_blocks.4.txt_mod.lin"] = None
|
||||
|
||||
idx = 0 # Index to keep track of the vector slices
|
||||
|
||||
for key in block_dict.keys():
|
||||
if "single_blocks" in key:
|
||||
# Single block: 1 ModulationOut
|
||||
block_dict[key] = ModulationOut(
|
||||
shift=tensor[:, idx : idx + 1, :],
|
||||
scale=tensor[:, idx + 1 : idx + 2, :],
|
||||
gate=tensor[:, idx + 2 : idx + 3, :],
|
||||
)
|
||||
idx += 3 # Advance by 3 vectors
|
||||
|
||||
elif "img_mod" in key:
|
||||
# Double block: List of 2 ModulationOut
|
||||
double_block = []
|
||||
for _ in range(2): # Create 2 ModulationOut objects
|
||||
double_block.append(
|
||||
ModulationOut(
|
||||
shift=tensor[:, idx : idx + 1, :],
|
||||
scale=tensor[:, idx + 1 : idx + 2, :],
|
||||
gate=tensor[:, idx + 2 : idx + 3, :],
|
||||
)
|
||||
)
|
||||
idx += 3 # Advance by 3 vectors per ModulationOut
|
||||
block_dict[key] = double_block
|
||||
|
||||
elif "txt_mod" in key:
|
||||
# Double block: List of 2 ModulationOut
|
||||
double_block = []
|
||||
for _ in range(2): # Create 2 ModulationOut objects
|
||||
double_block.append(
|
||||
ModulationOut(
|
||||
shift=tensor[:, idx : idx + 1, :],
|
||||
scale=tensor[:, idx + 1 : idx + 2, :],
|
||||
gate=tensor[:, idx + 2 : idx + 3, :],
|
||||
)
|
||||
)
|
||||
idx += 3 # Advance by 3 vectors per ModulationOut
|
||||
block_dict[key] = double_block
|
||||
|
||||
elif "final_layer" in key:
|
||||
# Final layer: 1 ModulationOut
|
||||
block_dict[key] = [
|
||||
tensor[:, idx : idx + 1, :],
|
||||
tensor[:, idx + 1 : idx + 2, :],
|
||||
]
|
||||
idx += 2 # Advance by 3 vectors
|
||||
|
||||
return block_dict
|
||||
|
||||
|
||||
class Approximator(nn.Module):
|
||||
def __init__(self, in_dim: int, out_dim: int, hidden_dim: int, n_layers=4):
|
||||
super().__init__()
|
||||
self.in_proj = nn.Linear(in_dim, hidden_dim, bias=True)
|
||||
self.layers = nn.ModuleList(
|
||||
[MLPEmbedder(hidden_dim, hidden_dim) for x in range(n_layers)]
|
||||
)
|
||||
self.norms = nn.ModuleList([RMSNorm(hidden_dim) for x in range(n_layers)])
|
||||
self.out_proj = nn.Linear(hidden_dim, out_dim)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
# Get the device of the module (assumes all parameters are on the same device)
|
||||
return next(self.parameters()).device
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
x = self.in_proj(x)
|
||||
|
||||
for layer, norms in zip(self.layers, self.norms):
|
||||
x = x + layer(norms(x))
|
||||
|
||||
x = self.out_proj(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class QKNorm(torch.nn.Module):
|
||||
def __init__(self, dim: int, use_compiled: bool = False):
|
||||
super().__init__()
|
||||
self.query_norm = RMSNorm(dim, use_compiled=use_compiled)
|
||||
self.key_norm = RMSNorm(dim, use_compiled=use_compiled)
|
||||
self.use_compiled = use_compiled
|
||||
|
||||
def forward(self, q: Tensor, k: Tensor, v: Tensor) -> tuple[Tensor, Tensor]:
|
||||
q = self.query_norm(q)
|
||||
k = self.key_norm(k)
|
||||
return q.to(v), k.to(v)
|
||||
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int = 8,
|
||||
qkv_bias: bool = False,
|
||||
use_compiled: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.norm = QKNorm(head_dim, use_compiled=use_compiled)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.use_compiled = use_compiled
|
||||
|
||||
def forward(self, x: Tensor, pe: Tensor) -> Tensor:
|
||||
qkv = self.qkv(x)
|
||||
q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
q, k = self.norm(q, k, v)
|
||||
x = attention(q, k, v, pe=pe)
|
||||
x = self.proj(x)
|
||||
return x
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModulationOut:
|
||||
shift: Tensor
|
||||
scale: Tensor
|
||||
gate: Tensor
|
||||
|
||||
|
||||
def _modulation_shift_scale_fn(x, scale, shift):
|
||||
return (1 + scale) * x + shift
|
||||
|
||||
|
||||
def _modulation_gate_fn(x, gate, gate_params):
|
||||
return x + gate * gate_params
|
||||
|
||||
|
||||
class DoubleStreamBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
mlp_ratio: float,
|
||||
qkv_bias: bool = False,
|
||||
use_compiled: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
self.num_heads = num_heads
|
||||
self.hidden_size = hidden_size
|
||||
self.img_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.img_attn = SelfAttention(
|
||||
dim=hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
use_compiled=use_compiled,
|
||||
)
|
||||
|
||||
self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.img_mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
|
||||
nn.GELU(approximate="tanh"),
|
||||
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
|
||||
)
|
||||
|
||||
self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.txt_attn = SelfAttention(
|
||||
dim=hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
use_compiled=use_compiled,
|
||||
)
|
||||
|
||||
self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.txt_mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
|
||||
nn.GELU(approximate="tanh"),
|
||||
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
|
||||
)
|
||||
self.use_compiled = use_compiled
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
# Get the device of the module (assumes all parameters are on the same device)
|
||||
return next(self.parameters()).device
|
||||
|
||||
def modulation_shift_scale_fn(self, x, scale, shift):
|
||||
if self.use_compiled:
|
||||
return torch.compile(_modulation_shift_scale_fn)(x, scale, shift)
|
||||
else:
|
||||
return _modulation_shift_scale_fn(x, scale, shift)
|
||||
|
||||
def modulation_gate_fn(self, x, gate, gate_params):
|
||||
if self.use_compiled:
|
||||
return torch.compile(_modulation_gate_fn)(x, gate, gate_params)
|
||||
else:
|
||||
return _modulation_gate_fn(x, gate, gate_params)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: Tensor,
|
||||
txt: Tensor,
|
||||
pe: Tensor,
|
||||
distill_vec: list[ModulationOut],
|
||||
mask: Tensor,
|
||||
) -> tuple[Tensor, Tensor]:
|
||||
(img_mod1, img_mod2), (txt_mod1, txt_mod2) = distill_vec
|
||||
|
||||
# prepare image for attention
|
||||
img_modulated = self.img_norm1(img)
|
||||
# replaced with compiled fn
|
||||
# img_modulated = (1 + img_mod1.scale) * img_modulated + img_mod1.shift
|
||||
img_modulated = self.modulation_shift_scale_fn(
|
||||
img_modulated, img_mod1.scale, img_mod1.shift
|
||||
)
|
||||
img_qkv = self.img_attn.qkv(img_modulated)
|
||||
img_q, img_k, img_v = rearrange(
|
||||
img_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads
|
||||
)
|
||||
img_q, img_k = self.img_attn.norm(img_q, img_k, img_v)
|
||||
|
||||
# prepare txt for attention
|
||||
txt_modulated = self.txt_norm1(txt)
|
||||
# replaced with compiled fn
|
||||
# txt_modulated = (1 + txt_mod1.scale) * txt_modulated + txt_mod1.shift
|
||||
txt_modulated = self.modulation_shift_scale_fn(
|
||||
txt_modulated, txt_mod1.scale, txt_mod1.shift
|
||||
)
|
||||
txt_qkv = self.txt_attn.qkv(txt_modulated)
|
||||
txt_q, txt_k, txt_v = rearrange(
|
||||
txt_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads
|
||||
)
|
||||
txt_q, txt_k = self.txt_attn.norm(txt_q, txt_k, txt_v)
|
||||
|
||||
# run actual attention
|
||||
q = torch.cat((txt_q, img_q), dim=2)
|
||||
k = torch.cat((txt_k, img_k), dim=2)
|
||||
v = torch.cat((txt_v, img_v), dim=2)
|
||||
|
||||
attn = attention(q, k, v, pe=pe, mask=mask)
|
||||
txt_attn, img_attn = attn[:, : txt.shape[1]], attn[:, txt.shape[1] :]
|
||||
|
||||
# calculate the img bloks
|
||||
# replaced with compiled fn
|
||||
# img = img + img_mod1.gate * self.img_attn.proj(img_attn)
|
||||
# img = img + img_mod2.gate * self.img_mlp((1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
|
||||
img = self.modulation_gate_fn(img, img_mod1.gate, self.img_attn.proj(img_attn))
|
||||
img = self.modulation_gate_fn(
|
||||
img,
|
||||
img_mod2.gate,
|
||||
self.img_mlp(
|
||||
self.modulation_shift_scale_fn(
|
||||
self.img_norm2(img), img_mod2.scale, img_mod2.shift
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
# calculate the txt bloks
|
||||
# replaced with compiled fn
|
||||
# txt = txt + txt_mod1.gate * self.txt_attn.proj(txt_attn)
|
||||
# txt = txt + txt_mod2.gate * self.txt_mlp((1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
|
||||
txt = self.modulation_gate_fn(txt, txt_mod1.gate, self.txt_attn.proj(txt_attn))
|
||||
txt = self.modulation_gate_fn(
|
||||
txt,
|
||||
txt_mod2.gate,
|
||||
self.txt_mlp(
|
||||
self.modulation_shift_scale_fn(
|
||||
self.txt_norm2(txt), txt_mod2.scale, txt_mod2.shift
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
return img, txt
|
||||
|
||||
|
||||
class SingleStreamBlock(nn.Module):
|
||||
"""
|
||||
A DiT block with parallel linear layers as described in
|
||||
https://arxiv.org/abs/2302.05442 and adapted modulation interface.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
qk_scale: float | None = None,
|
||||
use_compiled: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.hidden_dim = hidden_size
|
||||
self.num_heads = num_heads
|
||||
head_dim = hidden_size // num_heads
|
||||
self.scale = qk_scale or head_dim**-0.5
|
||||
|
||||
self.mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
# qkv and mlp_in
|
||||
self.linear1 = nn.Linear(hidden_size, hidden_size * 3 + self.mlp_hidden_dim)
|
||||
# proj and mlp_out
|
||||
self.linear2 = nn.Linear(hidden_size + self.mlp_hidden_dim, hidden_size)
|
||||
|
||||
self.norm = QKNorm(head_dim, use_compiled=use_compiled)
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
|
||||
self.mlp_act = nn.GELU(approximate="tanh")
|
||||
self.use_compiled = use_compiled
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
# Get the device of the module (assumes all parameters are on the same device)
|
||||
return next(self.parameters()).device
|
||||
|
||||
def modulation_shift_scale_fn(self, x, scale, shift):
|
||||
if self.use_compiled:
|
||||
return torch.compile(_modulation_shift_scale_fn)(x, scale, shift)
|
||||
else:
|
||||
return _modulation_shift_scale_fn(x, scale, shift)
|
||||
|
||||
def modulation_gate_fn(self, x, gate, gate_params):
|
||||
if self.use_compiled:
|
||||
return torch.compile(_modulation_gate_fn)(x, gate, gate_params)
|
||||
else:
|
||||
return _modulation_gate_fn(x, gate, gate_params)
|
||||
|
||||
def forward(
|
||||
self, x: Tensor, pe: Tensor, distill_vec: list[ModulationOut], mask: Tensor
|
||||
) -> Tensor:
|
||||
mod = distill_vec
|
||||
# replaced with compiled fn
|
||||
# x_mod = (1 + mod.scale) * self.pre_norm(x) + mod.shift
|
||||
x_mod = self.modulation_shift_scale_fn(self.pre_norm(x), mod.scale, mod.shift)
|
||||
qkv, mlp = torch.split(
|
||||
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
|
||||
)
|
||||
|
||||
q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
q, k = self.norm(q, k, v)
|
||||
|
||||
# compute attention
|
||||
attn = attention(q, k, v, pe=pe, mask=mask)
|
||||
# compute activation in mlp stream, cat again and run second linear layer
|
||||
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
|
||||
# replaced with compiled fn
|
||||
# return x + mod.gate * output
|
||||
return self.modulation_gate_fn(x, mod.gate, output)
|
||||
|
||||
|
||||
class LastLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
patch_size: int,
|
||||
out_channels: int,
|
||||
use_compiled: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.linear = nn.Linear(
|
||||
hidden_size, patch_size * patch_size * out_channels, bias=True
|
||||
)
|
||||
self.use_compiled = use_compiled
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
# Get the device of the module (assumes all parameters are on the same device)
|
||||
return next(self.parameters()).device
|
||||
|
||||
def modulation_shift_scale_fn(self, x, scale, shift):
|
||||
if self.use_compiled:
|
||||
return torch.compile(_modulation_shift_scale_fn)(x, scale, shift)
|
||||
else:
|
||||
return _modulation_shift_scale_fn(x, scale, shift)
|
||||
|
||||
def forward(self, x: Tensor, distill_vec: list[Tensor]) -> Tensor:
|
||||
shift, scale = distill_vec
|
||||
shift = shift.squeeze(1)
|
||||
scale = scale.squeeze(1)
|
||||
# replaced with compiled fn
|
||||
# x = (1 + scale[:, None, :]) * self.norm_final(x) + shift[:, None, :]
|
||||
x = self.modulation_shift_scale_fn(
|
||||
self.norm_final(x), scale[:, None, :], shift[:, None, :]
|
||||
)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
51
extensions_built_in/diffusion_models/chroma/src/math.py
Normal file
51
extensions_built_in/diffusion_models/chroma/src/math.py
Normal file
@@ -0,0 +1,51 @@
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from torch import Tensor
|
||||
|
||||
# Flash-Attention 2 (optional)
|
||||
try:
|
||||
from flash_attn.flash_attn_interface import flash_attn_func # type: ignore
|
||||
_HAS_FLASH = True
|
||||
except (ImportError, ModuleNotFoundError):
|
||||
_HAS_FLASH = False
|
||||
|
||||
|
||||
def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, mask: Tensor) -> Tensor:
|
||||
q, k = apply_rope(q, k, pe)
|
||||
|
||||
# mask should have shape [B, H, L, D]
|
||||
if _HAS_FLASH and mask is None and q.is_cuda:
|
||||
x = flash_attn_func(
|
||||
rearrange(q, "B H L D -> B L H D").contiguous(),
|
||||
rearrange(k, "B H L D -> B L H D").contiguous(),
|
||||
rearrange(v, "B H L D -> B L H D").contiguous(),
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
)
|
||||
x = rearrange(x, "B L H D -> B H L D")
|
||||
else:
|
||||
x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
|
||||
x = rearrange(x, "B H L D -> B L (H D)")
|
||||
return x
|
||||
|
||||
|
||||
def rope(pos: Tensor, dim: int, theta: int) -> Tensor:
|
||||
assert dim % 2 == 0
|
||||
scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim
|
||||
omega = 1.0 / (theta**scale)
|
||||
out = torch.einsum("...n,d->...nd", pos, omega)
|
||||
out = torch.stack(
|
||||
[torch.cos(out), -torch.sin(out), torch.sin(out), torch.cos(out)], dim=-1
|
||||
)
|
||||
out = rearrange(out, "b n d (i j) -> b n d i j", i=2, j=2)
|
||||
return out.float()
|
||||
|
||||
|
||||
def apply_rope(xq: Tensor, xk: Tensor, freqs_cis: Tensor) -> tuple[Tensor, Tensor]:
|
||||
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
|
||||
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
|
||||
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
|
||||
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
|
||||
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
|
||||
276
extensions_built_in/diffusion_models/chroma/src/model.py
Normal file
276
extensions_built_in/diffusion_models/chroma/src/model.py
Normal file
@@ -0,0 +1,276 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
from torch import Tensor, nn
|
||||
import torch.utils.checkpoint as ckpt
|
||||
|
||||
from .layers import (
|
||||
DoubleStreamBlock,
|
||||
EmbedND,
|
||||
LastLayer,
|
||||
SingleStreamBlock,
|
||||
timestep_embedding,
|
||||
Approximator,
|
||||
distribute_modulations,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChromaParams:
|
||||
in_channels: int
|
||||
context_in_dim: int
|
||||
hidden_size: int
|
||||
mlp_ratio: float
|
||||
num_heads: int
|
||||
depth: int
|
||||
depth_single_blocks: int
|
||||
axes_dim: list[int]
|
||||
theta: int
|
||||
qkv_bias: bool
|
||||
guidance_embed: bool
|
||||
approximator_in_dim: int
|
||||
approximator_depth: int
|
||||
approximator_hidden_size: int
|
||||
_use_compiled: bool
|
||||
|
||||
|
||||
chroma_params = ChromaParams(
|
||||
in_channels=64,
|
||||
context_in_dim=4096,
|
||||
hidden_size=3072,
|
||||
mlp_ratio=4.0,
|
||||
num_heads=24,
|
||||
depth=19,
|
||||
depth_single_blocks=38,
|
||||
axes_dim=[16, 56, 56],
|
||||
theta=10_000,
|
||||
qkv_bias=True,
|
||||
guidance_embed=True,
|
||||
approximator_in_dim=64,
|
||||
approximator_depth=5,
|
||||
approximator_hidden_size=5120,
|
||||
_use_compiled=False,
|
||||
)
|
||||
|
||||
|
||||
def modify_mask_to_attend_padding(mask, max_seq_length, num_extra_padding=8):
|
||||
"""
|
||||
Modifies attention mask to allow attention to a few extra padding tokens.
|
||||
|
||||
Args:
|
||||
mask: Original attention mask (1 for tokens to attend to, 0 for masked tokens)
|
||||
max_seq_length: Maximum sequence length of the model
|
||||
num_extra_padding: Number of padding tokens to unmask
|
||||
|
||||
Returns:
|
||||
Modified mask
|
||||
"""
|
||||
# Get the actual sequence length from the mask
|
||||
seq_length = mask.sum(dim=-1)
|
||||
batch_size = mask.shape[0]
|
||||
|
||||
modified_mask = mask.clone()
|
||||
|
||||
for i in range(batch_size):
|
||||
current_seq_len = int(seq_length[i].item())
|
||||
|
||||
# Only add extra padding tokens if there's room
|
||||
if current_seq_len < max_seq_length:
|
||||
# Calculate how many padding tokens we can unmask
|
||||
available_padding = max_seq_length - current_seq_len
|
||||
tokens_to_unmask = min(num_extra_padding, available_padding)
|
||||
|
||||
# Unmask the specified number of padding tokens right after the sequence
|
||||
modified_mask[i, current_seq_len : current_seq_len + tokens_to_unmask] = 1
|
||||
|
||||
return modified_mask
|
||||
|
||||
|
||||
class Chroma(nn.Module):
|
||||
"""
|
||||
Transformer model for flow matching on sequences.
|
||||
"""
|
||||
|
||||
def __init__(self, params: ChromaParams):
|
||||
super().__init__()
|
||||
self.params = params
|
||||
self.in_channels = params.in_channels
|
||||
self.out_channels = self.in_channels
|
||||
self.gradient_checkpointing = False
|
||||
if params.hidden_size % params.num_heads != 0:
|
||||
raise ValueError(
|
||||
f"Hidden size {params.hidden_size} must be divisible by num_heads {params.num_heads}"
|
||||
)
|
||||
pe_dim = params.hidden_size // params.num_heads
|
||||
if sum(params.axes_dim) != pe_dim:
|
||||
raise ValueError(
|
||||
f"Got {params.axes_dim} but expected positional dim {pe_dim}"
|
||||
)
|
||||
self.hidden_size = params.hidden_size
|
||||
self.num_heads = params.num_heads
|
||||
self.pe_embedder = EmbedND(
|
||||
dim=pe_dim, theta=params.theta, axes_dim=params.axes_dim
|
||||
)
|
||||
self.img_in = nn.Linear(self.in_channels, self.hidden_size, bias=True)
|
||||
|
||||
# TODO: need proper mapping for this approximator output!
|
||||
# currently the mapping is hardcoded in distribute_modulations function
|
||||
self.distilled_guidance_layer = Approximator(
|
||||
params.approximator_in_dim,
|
||||
self.hidden_size,
|
||||
params.approximator_hidden_size,
|
||||
params.approximator_depth,
|
||||
)
|
||||
self.txt_in = nn.Linear(params.context_in_dim, self.hidden_size)
|
||||
|
||||
self.double_blocks = nn.ModuleList(
|
||||
[
|
||||
DoubleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.num_heads,
|
||||
mlp_ratio=params.mlp_ratio,
|
||||
qkv_bias=params.qkv_bias,
|
||||
use_compiled=params._use_compiled,
|
||||
)
|
||||
for _ in range(params.depth)
|
||||
]
|
||||
)
|
||||
|
||||
self.single_blocks = nn.ModuleList(
|
||||
[
|
||||
SingleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.num_heads,
|
||||
mlp_ratio=params.mlp_ratio,
|
||||
use_compiled=params._use_compiled,
|
||||
)
|
||||
for _ in range(params.depth_single_blocks)
|
||||
]
|
||||
)
|
||||
|
||||
self.final_layer = LastLayer(
|
||||
self.hidden_size,
|
||||
1,
|
||||
self.out_channels,
|
||||
use_compiled=params._use_compiled,
|
||||
)
|
||||
|
||||
# TODO: move this hardcoded value to config
|
||||
self.mod_index_length = 344
|
||||
# self.mod_index = torch.tensor(list(range(self.mod_index_length)), device=0)
|
||||
self.register_buffer(
|
||||
"mod_index",
|
||||
torch.tensor(list(range(self.mod_index_length)), device="cpu"),
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
# Get the device of the module (assumes all parameters are on the same device)
|
||||
return next(self.parameters()).device
|
||||
|
||||
def enable_gradient_checkpointing(self, enable: bool = True):
|
||||
self.gradient_checkpointing = enable
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: Tensor,
|
||||
img_ids: Tensor,
|
||||
txt: Tensor,
|
||||
txt_ids: Tensor,
|
||||
txt_mask: Tensor,
|
||||
timesteps: Tensor,
|
||||
guidance: Tensor,
|
||||
attn_padding: int = 1,
|
||||
) -> Tensor:
|
||||
if img.ndim != 3 or txt.ndim != 3:
|
||||
raise ValueError("Input img and txt tensors must have 3 dimensions.")
|
||||
|
||||
# running on sequences img
|
||||
img = self.img_in(img)
|
||||
txt = self.txt_in(txt)
|
||||
|
||||
# TODO:
|
||||
# need to fix grad accumulation issue here for now it's in no grad mode
|
||||
# besides, i don't want to wash out the PFP that's trained on this model weights anyway
|
||||
# the fan out operation here is deleting the backward graph
|
||||
# alternatively doing forward pass for every block manually is doable but slow
|
||||
# custom backward probably be better
|
||||
with torch.no_grad():
|
||||
distill_timestep = timestep_embedding(timesteps, 16)
|
||||
# TODO: need to add toggle to omit this from schnell but that's not a priority
|
||||
distil_guidance = timestep_embedding(guidance, 16)
|
||||
# get all modulation index
|
||||
modulation_index = timestep_embedding(self.mod_index, 32)
|
||||
# we need to broadcast the modulation index here so each batch has all of the index
|
||||
modulation_index = modulation_index.unsqueeze(0).repeat(img.shape[0], 1, 1)
|
||||
# and we need to broadcast timestep and guidance along too
|
||||
timestep_guidance = (
|
||||
torch.cat([distill_timestep, distil_guidance], dim=1)
|
||||
.unsqueeze(1)
|
||||
.repeat(1, self.mod_index_length, 1)
|
||||
)
|
||||
# then and only then we could concatenate it together
|
||||
input_vec = torch.cat([timestep_guidance, modulation_index], dim=-1)
|
||||
mod_vectors = self.distilled_guidance_layer(input_vec.requires_grad_(True))
|
||||
mod_vectors_dict = distribute_modulations(mod_vectors)
|
||||
|
||||
ids = torch.cat((txt_ids, img_ids), dim=1)
|
||||
pe = self.pe_embedder(ids)
|
||||
|
||||
# compute mask
|
||||
# assume max seq length from the batched input
|
||||
|
||||
max_len = txt.shape[1]
|
||||
|
||||
# mask
|
||||
with torch.no_grad():
|
||||
txt_mask_w_padding = modify_mask_to_attend_padding(
|
||||
txt_mask, max_len, attn_padding
|
||||
)
|
||||
txt_img_mask = torch.cat(
|
||||
[
|
||||
txt_mask_w_padding,
|
||||
torch.ones([img.shape[0], img.shape[1]], device=txt_mask.device),
|
||||
],
|
||||
dim=1,
|
||||
)
|
||||
txt_img_mask = txt_img_mask.float().T @ txt_img_mask.float()
|
||||
txt_img_mask = (
|
||||
txt_img_mask[None, None, ...]
|
||||
.repeat(txt.shape[0], self.num_heads, 1, 1)
|
||||
.int()
|
||||
.bool()
|
||||
)
|
||||
# txt_mask_w_padding[txt_mask_w_padding==False] = True
|
||||
|
||||
for i, block in enumerate(self.double_blocks):
|
||||
# the guidance replaced by FFN output
|
||||
img_mod = mod_vectors_dict[f"double_blocks.{i}.img_mod.lin"]
|
||||
txt_mod = mod_vectors_dict[f"double_blocks.{i}.txt_mod.lin"]
|
||||
double_mod = [img_mod, txt_mod]
|
||||
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
img.requires_grad_(True)
|
||||
img, txt = ckpt.checkpoint(
|
||||
block, img, txt, pe, double_mod, txt_img_mask
|
||||
)
|
||||
else:
|
||||
img, txt = block(
|
||||
img=img, txt=txt, pe=pe, distill_vec=double_mod, mask=txt_img_mask
|
||||
)
|
||||
|
||||
img = torch.cat((txt, img), 1)
|
||||
for i, block in enumerate(self.single_blocks):
|
||||
single_mod = mod_vectors_dict[f"single_blocks.{i}.modulation.lin"]
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
img.requires_grad_(True)
|
||||
img = ckpt.checkpoint(block, img, pe, single_mod, txt_img_mask)
|
||||
else:
|
||||
img = block(img, pe=pe, distill_vec=single_mod, mask=txt_img_mask)
|
||||
img = img[:, txt.shape[1] :, ...]
|
||||
final_mod = mod_vectors_dict["final_layer.adaLN_modulation.1"]
|
||||
img = self.final_layer(
|
||||
img, distill_vec=final_mod
|
||||
) # (N, T, patch_size ** 2 * out_channels)
|
||||
return img
|
||||
1
extensions_built_in/diffusion_models/f_light/__init__.py
Normal file
1
extensions_built_in/diffusion_models/f_light/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
from .f_light import FLiteModel
|
||||
295
extensions_built_in/diffusion_models/f_light/f_light.py
Normal file
295
extensions_built_in/diffusion_models/f_light/f_light.py
Normal file
@@ -0,0 +1,295 @@
|
||||
import os
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
import yaml
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||
from PIL import Image
|
||||
from toolkit.models.base_model import BaseModel
|
||||
from toolkit.basic import flush
|
||||
from diffusers import AutoencoderKL
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
from toolkit.accelerator import unwrap_model
|
||||
from optimum.quanto import freeze, QTensor
|
||||
from toolkit.util.quantize import quantize, get_qtype
|
||||
from transformers import T5TokenizerFast, T5EncoderModel
|
||||
from .src import FLitePipeline, DiT
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
|
||||
scheduler_config = {
|
||||
"base_image_seq_len": 256,
|
||||
"base_shift": 0.5,
|
||||
"max_image_seq_len": 4096,
|
||||
"max_shift": 1.15,
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 3.0,
|
||||
"use_dynamic_shifting": True
|
||||
}
|
||||
|
||||
|
||||
class FLiteModel(BaseModel):
|
||||
arch = "f-lite"
|
||||
|
||||
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 = ['DiT']
|
||||
|
||||
# static method to get the noise scheduler
|
||||
@staticmethod
|
||||
def get_train_scheduler():
|
||||
return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
|
||||
|
||||
def get_bucket_divisibility(self):
|
||||
# return the bucket divisibility for the model
|
||||
return 16
|
||||
|
||||
def load_model(self):
|
||||
dtype = self.torch_dtype
|
||||
|
||||
# will be updated if we detect a existing checkpoint in training folder
|
||||
model_path = self.model_config.name_or_path
|
||||
|
||||
extras_path = self.model_config.extras_name_or_path
|
||||
|
||||
self.print_and_status_update("Loading transformer")
|
||||
|
||||
transformer = DiT.from_pretrained(
|
||||
model_path,
|
||||
subfolder="dit_model",
|
||||
torch_dtype=dtype,
|
||||
)
|
||||
|
||||
transformer.to(self.quantize_device, dtype=dtype)
|
||||
|
||||
if self.model_config.quantize:
|
||||
# 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,
|
||||
**self.model_config.quantize_kwargs)
|
||||
freeze(transformer)
|
||||
transformer.to(self.device_torch)
|
||||
else:
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Loading T5")
|
||||
tokenizer = T5TokenizerFast.from_pretrained(
|
||||
extras_path, subfolder="tokenizer", torch_dtype=dtype
|
||||
)
|
||||
text_encoder = T5EncoderModel.from_pretrained(
|
||||
extras_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 T5")
|
||||
quantize(text_encoder, weights=get_qtype(
|
||||
self.model_config.qtype))
|
||||
freeze(text_encoder)
|
||||
flush()
|
||||
|
||||
self.noise_scheduler = FLiteModel.get_train_scheduler()
|
||||
|
||||
self.print_and_status_update("Loading VAE")
|
||||
vae = AutoencoderKL.from_pretrained(
|
||||
extras_path,
|
||||
subfolder="vae",
|
||||
torch_dtype=dtype
|
||||
)
|
||||
vae = vae.to(self.device_torch, dtype=dtype)
|
||||
|
||||
self.print_and_status_update("Making pipe")
|
||||
|
||||
pipe: FLitePipeline = FLitePipeline(
|
||||
text_encoder=None,
|
||||
tokenizer=tokenizer,
|
||||
vae=vae,
|
||||
dit_model=None,
|
||||
)
|
||||
# for quantization, it works best to do these after making the pipe
|
||||
pipe.text_encoder = text_encoder
|
||||
pipe.dit_model = transformer
|
||||
pipe.transformer = transformer
|
||||
pipe.scheduler = self.noise_scheduler,
|
||||
|
||||
self.print_and_status_update("Preparing Model")
|
||||
|
||||
text_encoder = [pipe.text_encoder]
|
||||
tokenizer = [pipe.tokenizer]
|
||||
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
|
||||
flush()
|
||||
# just to make sure everything is on the right device and dtype
|
||||
text_encoder[0].to(self.device_torch)
|
||||
text_encoder[0].requires_grad_(False)
|
||||
text_encoder[0].eval()
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
flush()
|
||||
|
||||
# save it to the model class
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder # list of text encoders
|
||||
self.tokenizer = tokenizer # list of tokenizers
|
||||
self.model = pipe.transformer
|
||||
self.pipeline = pipe
|
||||
self.print_and_status_update("Model Loaded")
|
||||
|
||||
def get_generation_pipeline(self):
|
||||
scheduler = FLiteModel.get_train_scheduler()
|
||||
# it has built in scheduler. Basically euler flowmatching
|
||||
pipeline = FLitePipeline(
|
||||
text_encoder=unwrap_model(self.text_encoder[0]),
|
||||
tokenizer=self.tokenizer[0],
|
||||
vae=unwrap_model(self.vae),
|
||||
dit_model=unwrap_model(self.transformer)
|
||||
)
|
||||
pipeline.transformer = pipeline.dit_model
|
||||
pipeline.scheduler = scheduler
|
||||
|
||||
return pipeline
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: FLitePipeline,
|
||||
gen_config: GenerateImageConfig,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
unconditional_embeds: PromptEmbeds,
|
||||
generator: torch.Generator,
|
||||
extra: dict,
|
||||
):
|
||||
|
||||
extra['negative_prompt_embeds'] = unconditional_embeds.text_embeds
|
||||
|
||||
img = pipeline(
|
||||
prompt_embeds=conditional_embeds.text_embeds,
|
||||
negative_prompt_embeds=unconditional_embeds.text_embeds,
|
||||
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,
|
||||
).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
|
||||
):
|
||||
cast_dtype = self.unet.dtype
|
||||
|
||||
noise_pred = self.unet(
|
||||
latent_model_input.to(
|
||||
self.device_torch, cast_dtype
|
||||
),
|
||||
text_embeddings.text_embeds.to(
|
||||
self.device_torch, cast_dtype
|
||||
),
|
||||
timestep / 1000,
|
||||
)
|
||||
|
||||
if isinstance(noise_pred, QTensor):
|
||||
noise_pred = noise_pred.dequantize()
|
||||
|
||||
return noise_pred
|
||||
|
||||
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
|
||||
if isinstance(prompt, str):
|
||||
prompts = [prompt]
|
||||
else:
|
||||
prompts = prompt
|
||||
if self.pipeline.text_encoder.device != self.device_torch:
|
||||
self.pipeline.text_encoder.to(self.device_torch)
|
||||
|
||||
prompt_embeds, negative_embeds = self.pipeline.encode_prompt(
|
||||
prompt=prompts,
|
||||
negative_prompt=None,
|
||||
device=self.text_encoder[0].device,
|
||||
dtype=self.torch_dtype,
|
||||
)
|
||||
|
||||
pe = PromptEmbeds(prompt_embeds)
|
||||
|
||||
return pe
|
||||
|
||||
def get_model_has_grad(self):
|
||||
# return from a weight if it has grad
|
||||
return False
|
||||
|
||||
def get_te_has_grad(self):
|
||||
# return from a weight if it has grad
|
||||
return False
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
# only save the unet
|
||||
transformer: DiT = unwrap_model(self.model)
|
||||
# diffusers
|
||||
# only save the unet
|
||||
transformer: DiT = unwrap_model(self.transformer)
|
||||
transformer.save_pretrained(
|
||||
save_directory=os.path.join(output_path, 'dit_model'),
|
||||
safe_serialization=True,
|
||||
)
|
||||
# save out meta config
|
||||
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')
|
||||
batch = kwargs.get('batch')
|
||||
# return (noise - batch.latents).detach()
|
||||
return (batch.latents - noise).detach()
|
||||
|
||||
def convert_lora_weights_before_save(self, state_dict):
|
||||
# currently starte with transformer. but needs to start with diffusion_model. for comfyui
|
||||
new_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
new_key = key.replace("transformer.", "diffusion_model.")
|
||||
new_sd[new_key] = value
|
||||
return new_sd
|
||||
|
||||
def convert_lora_weights_before_load(self, state_dict):
|
||||
# saved as diffusion_model. but needs to be transformer. for ai-toolkit
|
||||
new_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
new_key = key.replace("diffusion_model.", "transformer.")
|
||||
new_sd[new_key] = value
|
||||
return new_sd
|
||||
|
||||
def get_base_model_version(self):
|
||||
return "f-lite"
|
||||
|
||||
def get_stepped_pred(self, pred, noise):
|
||||
# just used for DFE support
|
||||
latents = pred + noise
|
||||
return latents
|
||||
@@ -0,0 +1,5 @@
|
||||
from .pipeline import FLitePipeline, FLitePipelineOutput, APGConfig
|
||||
from .model import DiT
|
||||
|
||||
|
||||
__all__ = ["FLitePipeline", "FLitePipelineOutput", "APGConfig", "DiT"]
|
||||
456
extensions_built_in/diffusion_models/f_light/src/model.py
Normal file
456
extensions_built_in/diffusion_models/f_light/src/model.py
Normal file
@@ -0,0 +1,456 @@
|
||||
# originally from https://github.com/fal-ai/f-lite/blob/main/f_lite/model.py but modified slightly
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.utils.accelerate_utils import apply_forward_hook
|
||||
from einops import rearrange
|
||||
from peft import get_peft_model_state_dict, set_peft_model_state_dict
|
||||
from torch import nn
|
||||
|
||||
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
|
||||
device=t.device
|
||||
)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
|
||||
return embedding
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim, eps=1e-6, trainable=False):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
if trainable:
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
else:
|
||||
self.weight = None
|
||||
|
||||
def forward(self, x):
|
||||
x_dtype = x.dtype
|
||||
x = x.float()
|
||||
norm = torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
||||
if self.weight is not None:
|
||||
return (x * norm * self.weight).to(dtype=x_dtype)
|
||||
else:
|
||||
return (x * norm).to(dtype=x_dtype)
|
||||
|
||||
|
||||
class QKNorm(nn.Module):
|
||||
"""Normalizing the query and the key independently, as Flux proposes"""
|
||||
|
||||
def __init__(self, dim, trainable=False):
|
||||
super().__init__()
|
||||
self.query_norm = RMSNorm(dim, trainable=trainable)
|
||||
self.key_norm = RMSNorm(dim, trainable=trainable)
|
||||
|
||||
def forward(self, q, k):
|
||||
q = self.query_norm(q)
|
||||
k = self.key_norm(k)
|
||||
return q, k
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_heads=8,
|
||||
qkv_bias=False,
|
||||
is_self_attn=True,
|
||||
cross_attn_input_size=None,
|
||||
residual_v=False,
|
||||
dynamic_softmax_temperature=False,
|
||||
):
|
||||
super().__init__()
|
||||
assert dim % num_heads == 0
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.scale = self.head_dim**-0.5
|
||||
self.is_self_attn = is_self_attn
|
||||
self.residual_v = residual_v
|
||||
self.dynamic_softmax_temperature = dynamic_softmax_temperature
|
||||
|
||||
if is_self_attn:
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
else:
|
||||
self.q = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.context_kv = nn.Linear(cross_attn_input_size, dim * 2, bias=qkv_bias)
|
||||
|
||||
self.proj = nn.Linear(dim, dim, bias=False)
|
||||
|
||||
if residual_v:
|
||||
self.lambda_param = nn.Parameter(torch.tensor(0.5).reshape(1))
|
||||
|
||||
self.qk_norm = QKNorm(self.head_dim)
|
||||
|
||||
def forward(self, x, context=None, v_0=None, rope=None):
|
||||
if self.is_self_attn:
|
||||
qkv = self.qkv(x)
|
||||
qkv = rearrange(qkv, "b l (k h d) -> k b h l d", k=3, h=self.num_heads)
|
||||
q, k, v = qkv.unbind(0)
|
||||
|
||||
if self.residual_v and v_0 is not None:
|
||||
v = self.lambda_param * v + (1 - self.lambda_param) * v_0
|
||||
|
||||
if rope is not None:
|
||||
# print(q.shape, rope[0].shape, rope[1].shape)
|
||||
q = apply_rotary_emb(q, rope[0], rope[1])
|
||||
k = apply_rotary_emb(k, rope[0], rope[1])
|
||||
|
||||
# https://arxiv.org/abs/2306.08645
|
||||
# https://arxiv.org/abs/2410.01104
|
||||
# ratioonale is that if tokens get larger, categorical distribution get more uniform
|
||||
# so you want to enlargen entropy.
|
||||
|
||||
token_length = q.shape[2]
|
||||
if self.dynamic_softmax_temperature:
|
||||
ratio = math.sqrt(math.log(token_length) / math.log(1040.0)) # 1024 + 16
|
||||
k = k * ratio
|
||||
q, k = self.qk_norm(q, k)
|
||||
|
||||
else:
|
||||
q = rearrange(self.q(x), "b l (h d) -> b h l d", h=self.num_heads)
|
||||
kv = rearrange(
|
||||
self.context_kv(context),
|
||||
"b l (k h d) -> k b h l d",
|
||||
k=2,
|
||||
h=self.num_heads,
|
||||
)
|
||||
k, v = kv.unbind(0)
|
||||
q, k = self.qk_norm(q, k)
|
||||
|
||||
x = F.scaled_dot_product_attention(q, k, v)
|
||||
x = rearrange(x, "b h l d -> b l (h d)")
|
||||
x = self.proj(x)
|
||||
return x, v if self.is_self_attn else None
|
||||
|
||||
|
||||
class DiTBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
cross_attn_input_size,
|
||||
num_heads,
|
||||
mlp_ratio=4.0,
|
||||
qkv_bias=True,
|
||||
residual_v=False,
|
||||
dynamic_softmax_temperature=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.norm1 = RMSNorm(hidden_size, trainable=qkv_bias)
|
||||
self.self_attn = Attention(
|
||||
hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
is_self_attn=True,
|
||||
residual_v=residual_v,
|
||||
dynamic_softmax_temperature=dynamic_softmax_temperature,
|
||||
)
|
||||
|
||||
if cross_attn_input_size is not None:
|
||||
self.norm2 = RMSNorm(hidden_size, trainable=qkv_bias)
|
||||
self.cross_attn = Attention(
|
||||
hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
is_self_attn=False,
|
||||
cross_attn_input_size=cross_attn_input_size,
|
||||
dynamic_softmax_temperature=dynamic_softmax_temperature,
|
||||
)
|
||||
else:
|
||||
self.norm2 = None
|
||||
self.cross_attn = None
|
||||
|
||||
self.norm3 = RMSNorm(hidden_size, trainable=qkv_bias)
|
||||
mlp_hidden = int(hidden_size * mlp_ratio)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, mlp_hidden),
|
||||
nn.GELU(),
|
||||
nn.Linear(mlp_hidden, hidden_size),
|
||||
)
|
||||
|
||||
self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 9 * hidden_size, bias=True))
|
||||
|
||||
self.adaLN_modulation[-1].weight.data.zero_()
|
||||
self.adaLN_modulation[-1].bias.data.zero_()
|
||||
|
||||
# @torch.compile(mode='reduce-overhead')
|
||||
def forward(self, x, context, c, v_0=None, rope=None):
|
||||
(
|
||||
shift_sa,
|
||||
scale_sa,
|
||||
gate_sa,
|
||||
shift_ca,
|
||||
scale_ca,
|
||||
gate_ca,
|
||||
shift_mlp,
|
||||
scale_mlp,
|
||||
gate_mlp,
|
||||
) = self.adaLN_modulation(c).chunk(9, dim=1)
|
||||
|
||||
scale_sa = scale_sa[:, None, :]
|
||||
scale_ca = scale_ca[:, None, :]
|
||||
scale_mlp = scale_mlp[:, None, :]
|
||||
|
||||
shift_sa = shift_sa[:, None, :]
|
||||
shift_ca = shift_ca[:, None, :]
|
||||
shift_mlp = shift_mlp[:, None, :]
|
||||
|
||||
gate_sa = gate_sa[:, None, :]
|
||||
gate_ca = gate_ca[:, None, :]
|
||||
gate_mlp = gate_mlp[:, None, :]
|
||||
|
||||
norm_x = self.norm1(x.clone())
|
||||
norm_x = norm_x * (1 + scale_sa) + shift_sa
|
||||
attn_out, v = self.self_attn(norm_x, v_0=v_0, rope=rope)
|
||||
x = x + attn_out * gate_sa
|
||||
|
||||
if self.norm2 is not None:
|
||||
norm_x = self.norm2(x)
|
||||
norm_x = norm_x * (1 + scale_ca) + shift_ca
|
||||
x = x + self.cross_attn(norm_x, context)[0] * gate_ca
|
||||
|
||||
norm_x = self.norm3(x)
|
||||
norm_x = norm_x * (1 + scale_mlp) + shift_mlp
|
||||
x = x + self.mlp(norm_x) * gate_mlp
|
||||
|
||||
return x, v
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
def __init__(self, patch_size=16, in_channels=3, embed_dim=768):
|
||||
super().__init__()
|
||||
self.patch_proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
|
||||
self.patch_size = patch_size
|
||||
|
||||
def forward(self, x):
|
||||
B, C, H, W = x.shape
|
||||
x = self.patch_proj(x)
|
||||
x = rearrange(x, "b c h w -> b (h w) c")
|
||||
return x
|
||||
|
||||
|
||||
class TwoDimRotary(torch.nn.Module):
|
||||
def __init__(self, dim, base=10000, h=256, w=256):
|
||||
super().__init__()
|
||||
self.inv_freq = torch.FloatTensor([1.0 / (base ** (i / dim)) for i in range(0, dim, 2)])
|
||||
self.h = h
|
||||
self.w = w
|
||||
|
||||
t_h = torch.arange(h, dtype=torch.float32)
|
||||
t_w = torch.arange(w, dtype=torch.float32)
|
||||
|
||||
freqs_h = torch.outer(t_h, self.inv_freq).unsqueeze(1) # h, 1, d / 2
|
||||
freqs_w = torch.outer(t_w, self.inv_freq).unsqueeze(0) # 1, w, d / 2
|
||||
freqs_h = freqs_h.repeat(1, w, 1) # h, w, d / 2
|
||||
freqs_w = freqs_w.repeat(h, 1, 1) # h, w, d / 2
|
||||
freqs_hw = torch.cat([freqs_h, freqs_w], 2) # h, w, d
|
||||
|
||||
self.register_buffer("freqs_hw_cos", freqs_hw.cos())
|
||||
self.register_buffer("freqs_hw_sin", freqs_hw.sin())
|
||||
|
||||
def forward(self, x, height_width=None, extend_with_register_tokens=0):
|
||||
if height_width is not None:
|
||||
this_h, this_w = height_width
|
||||
else:
|
||||
this_hw = x.shape[1]
|
||||
this_h, this_w = int(this_hw**0.5), int(this_hw**0.5)
|
||||
|
||||
cos = self.freqs_hw_cos[0 : this_h, 0 : this_w]
|
||||
sin = self.freqs_hw_sin[0 : this_h, 0 : this_w]
|
||||
|
||||
cos = cos.clone().reshape(this_h * this_w, -1)
|
||||
sin = sin.clone().reshape(this_h * this_w, -1)
|
||||
|
||||
# append N of zero-attn tokens
|
||||
if extend_with_register_tokens > 0:
|
||||
cos = torch.cat(
|
||||
[
|
||||
torch.ones(extend_with_register_tokens, cos.shape[1]).to(cos.device),
|
||||
cos,
|
||||
],
|
||||
0,
|
||||
)
|
||||
sin = torch.cat(
|
||||
[
|
||||
torch.zeros(extend_with_register_tokens, sin.shape[1]).to(sin.device),
|
||||
sin,
|
||||
],
|
||||
0,
|
||||
)
|
||||
|
||||
return cos[None, None, :, :], sin[None, None, :, :] # [1, 1, T + N, Attn-dim]
|
||||
|
||||
|
||||
def apply_rotary_emb(x, cos, sin):
|
||||
orig_dtype = x.dtype
|
||||
x = x.to(dtype=torch.float32)
|
||||
assert x.ndim == 4 # multihead attention
|
||||
d = x.shape[3] // 2
|
||||
x1 = x[..., :d]
|
||||
x2 = x[..., d:]
|
||||
y1 = x1 * cos + x2 * sin
|
||||
y2 = x1 * (-sin) + x2 * cos
|
||||
return torch.cat([y1, y2], 3).to(dtype=orig_dtype)
|
||||
|
||||
|
||||
class DiT(ModelMixin, ConfigMixin, FromOriginalModelMixin, PeftAdapterMixin): # type: ignore[misc]
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
in_channels=4,
|
||||
patch_size=2,
|
||||
hidden_size=1152,
|
||||
depth=28,
|
||||
num_heads=16,
|
||||
mlp_ratio=4.0,
|
||||
cross_attn_input_size=128,
|
||||
residual_v=False,
|
||||
train_bias_and_rms=True,
|
||||
use_rope=True,
|
||||
gradient_checkpoint=False,
|
||||
dynamic_softmax_temperature=False,
|
||||
rope_base=10000,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.patch_embed = PatchEmbed(patch_size, in_channels, hidden_size)
|
||||
|
||||
if use_rope:
|
||||
self.rope = TwoDimRotary(hidden_size // (2 * num_heads), base=rope_base, h=512, w=512)
|
||||
else:
|
||||
self.positional_embedding = nn.Parameter(torch.zeros(1, 2048, hidden_size))
|
||||
|
||||
self.register_tokens = nn.Parameter(torch.randn(1, 16, hidden_size))
|
||||
|
||||
self.time_embed = nn.Sequential(
|
||||
nn.Linear(hidden_size, 4 * hidden_size),
|
||||
nn.SiLU(),
|
||||
nn.Linear(4 * hidden_size, hidden_size),
|
||||
)
|
||||
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
DiTBlock(
|
||||
hidden_size=hidden_size,
|
||||
num_heads=num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
cross_attn_input_size=cross_attn_input_size,
|
||||
residual_v=residual_v,
|
||||
qkv_bias=train_bias_and_rms,
|
||||
dynamic_softmax_temperature=dynamic_softmax_temperature,
|
||||
)
|
||||
for _ in range(depth)
|
||||
]
|
||||
)
|
||||
|
||||
self.final_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
|
||||
self.final_norm = RMSNorm(hidden_size, trainable=train_bias_and_rms)
|
||||
self.final_proj = nn.Linear(hidden_size, patch_size * patch_size * in_channels)
|
||||
nn.init.zeros_(self.final_modulation[-1].weight)
|
||||
nn.init.zeros_(self.final_modulation[-1].bias)
|
||||
nn.init.zeros_(self.final_proj.weight)
|
||||
nn.init.zeros_(self.final_proj.bias)
|
||||
self.paramstatus = {}
|
||||
for n, p in self.named_parameters():
|
||||
self.paramstatus[n] = {
|
||||
"shape": p.shape,
|
||||
"requires_grad": p.requires_grad,
|
||||
}
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def save_lora_weights(self, save_directory):
|
||||
"""Save LoRA weights to a file"""
|
||||
lora_state_dict = get_peft_model_state_dict(self)
|
||||
torch.save(lora_state_dict, f"{save_directory}/lora_weights.pt")
|
||||
|
||||
def load_lora_weights(self, load_directory):
|
||||
"""Load LoRA weights from a file"""
|
||||
lora_state_dict = torch.load(f"{load_directory}/lora_weights.pt")
|
||||
set_peft_model_state_dict(self, lora_state_dict)
|
||||
|
||||
@apply_forward_hook
|
||||
def forward(self, x, context, timesteps):
|
||||
b, c, h, w = x.shape
|
||||
x = self.patch_embed(x) # b, T, d
|
||||
|
||||
x = torch.cat([self.register_tokens.repeat(b, 1, 1), x], 1) # b, T + N, d
|
||||
|
||||
if self.config.use_rope:
|
||||
cos, sin = self.rope(
|
||||
x,
|
||||
extend_with_register_tokens=16,
|
||||
height_width=(h // self.config.patch_size, w // self.config.patch_size),
|
||||
)
|
||||
else:
|
||||
x = x + self.positional_embedding.repeat(b, 1, 1)[:, : x.shape[1], :]
|
||||
cos, sin = None, None
|
||||
|
||||
t_emb = timestep_embedding(timesteps * 1000, self.config.hidden_size).to(x.device, dtype=x.dtype)
|
||||
t_emb = self.time_embed(t_emb)
|
||||
|
||||
v_0 = None
|
||||
|
||||
for _idx, block in enumerate(self.blocks):
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
x, v = self._gradient_checkpointing_func(
|
||||
block,
|
||||
x,
|
||||
context,
|
||||
t_emb,
|
||||
v_0,
|
||||
(cos, sin)
|
||||
)
|
||||
else:
|
||||
x, v = block(x, context, t_emb, v_0, (cos, sin))
|
||||
if v_0 is None:
|
||||
v_0 = v
|
||||
|
||||
x = x[:, 16:, :]
|
||||
final_shift, final_scale = self.final_modulation(t_emb).chunk(2, dim=1)
|
||||
x = self.final_norm(x)
|
||||
x = x * (1 + final_scale[:, None, :]) + final_shift[:, None, :]
|
||||
x = self.final_proj(x)
|
||||
|
||||
x = rearrange(
|
||||
x,
|
||||
"b (h w) (p1 p2 c) -> b c (h p1) (w p2)",
|
||||
h=h // self.config.patch_size,
|
||||
w=w // self.config.patch_size,
|
||||
p1=self.config.patch_size,
|
||||
p2=self.config.patch_size,
|
||||
)
|
||||
return x
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
model = DiT(
|
||||
in_channels=4,
|
||||
patch_size=2,
|
||||
hidden_size=1152,
|
||||
depth=28,
|
||||
num_heads=16,
|
||||
mlp_ratio=4.0,
|
||||
cross_attn_input_size=128,
|
||||
residual_v=False,
|
||||
train_bias_and_rms=True,
|
||||
use_rope=True,
|
||||
).cuda()
|
||||
print(
|
||||
model(
|
||||
torch.randn(1, 4, 64, 64).cuda(),
|
||||
torch.randn(1, 37, 128).cuda(),
|
||||
torch.tensor([1.0]).cuda(),
|
||||
)
|
||||
)
|
||||
308
extensions_built_in/diffusion_models/f_light/src/pipeline.py
Normal file
308
extensions_built_in/diffusion_models/f_light/src/pipeline.py
Normal file
@@ -0,0 +1,308 @@
|
||||
# originally from https://github.com/fal-ai/f-lite/blob/main/f_lite/pipeline.py but modified slightly
|
||||
import logging
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import AutoencoderKL, DiffusionPipeline
|
||||
from diffusers.utils import BaseOutput
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from PIL import Image
|
||||
from torch import FloatTensor
|
||||
from tqdm.auto import tqdm
|
||||
from transformers import T5EncoderModel, T5TokenizerFast
|
||||
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class APGConfig:
|
||||
"""APG (Augmented Parallel Guidance) configuration"""
|
||||
|
||||
enabled: bool = True
|
||||
orthogonal_threshold: float = 0.03
|
||||
|
||||
|
||||
@dataclass
|
||||
class FLitePipelineOutput(BaseOutput):
|
||||
"""
|
||||
Output class for FLitePipeline pipeline.
|
||||
Args:
|
||||
images (`List[PIL.Image.Image]` or `np.ndarray`)
|
||||
List of denoised PIL images of length `batch_size` or numpy array of shape `(batch_size, height, width,
|
||||
num_channels)`. PIL images or numpy array present the denoised images of the diffusion pipeline.
|
||||
"""
|
||||
|
||||
images: Union[List[Image.Image], np.ndarray]
|
||||
|
||||
|
||||
class FLitePipeline(DiffusionPipeline):
|
||||
r"""
|
||||
Pipeline for text-to-image generation using F-Lite model.
|
||||
This model inherits from [`DiffusionPipeline`].
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->dit_model->vae"
|
||||
|
||||
dit_model: torch.nn.Module
|
||||
vae: AutoencoderKL
|
||||
text_encoder: T5EncoderModel
|
||||
tokenizer: T5TokenizerFast
|
||||
_progress_bar_config: Dict[str, Any]
|
||||
|
||||
def __init__(
|
||||
self, dit_model: torch.nn.Module, vae: AutoencoderKL, text_encoder: T5EncoderModel, tokenizer: T5TokenizerFast
|
||||
):
|
||||
super().__init__()
|
||||
# Register all modules for the pipeline
|
||||
# Access DiffusionPipeline's register_modules directly to avoid mypy error
|
||||
DiffusionPipeline.register_modules(
|
||||
self, dit_model=dit_model, vae=vae, text_encoder=text_encoder, tokenizer=tokenizer
|
||||
)
|
||||
|
||||
# Move models to channels last for better performance
|
||||
# AutoencoderKL inherits from torch.nn.Module which has these methods
|
||||
if hasattr(self.vae, "to"):
|
||||
self.vae.to(memory_format=torch.channels_last)
|
||||
if hasattr(self.vae, "requires_grad_"):
|
||||
self.vae.requires_grad_(False)
|
||||
if hasattr(self.text_encoder, "requires_grad_"):
|
||||
self.text_encoder.requires_grad_(False)
|
||||
|
||||
# Constants
|
||||
self.vae_scale_factor = 8
|
||||
self.return_index = -8 # T5 hidden state index to use
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
"""Enable VAE slicing for memory efficiency."""
|
||||
if hasattr(self.vae, "enable_slicing"):
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def enable_vae_tiling(self):
|
||||
"""Enable VAE tiling for memory efficiency."""
|
||||
if hasattr(self.vae, "enable_tiling"):
|
||||
self.vae.enable_tiling()
|
||||
|
||||
def set_progress_bar_config(self, **kwargs):
|
||||
"""Set progress bar configuration."""
|
||||
self._progress_bar_config = kwargs
|
||||
|
||||
def progress_bar(self, iterable=None, **kwargs):
|
||||
"""Create progress bar for iterations."""
|
||||
self._progress_bar_config = getattr(self, "_progress_bar_config", None) or {}
|
||||
config = {**self._progress_bar_config, **kwargs}
|
||||
return tqdm(iterable, **config)
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
max_sequence_length: int = 512,
|
||||
return_index: int = -8,
|
||||
) -> Tuple[FloatTensor, FloatTensor]:
|
||||
"""Encodes the prompt and negative prompt."""
|
||||
if isinstance(prompt, str):
|
||||
prompt = [prompt]
|
||||
device = device or self.text_encoder.device
|
||||
# Text encoder forward pass
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids.to(device)
|
||||
prompt_embeds = self.text_encoder(text_input_ids, return_dict=True, output_hidden_states=True)
|
||||
prompt_embeds_tensor = prompt_embeds.hidden_states[return_index]
|
||||
if return_index != -1:
|
||||
prompt_embeds_tensor = self.text_encoder.encoder.final_layer_norm(prompt_embeds_tensor)
|
||||
prompt_embeds_tensor = self.text_encoder.encoder.dropout(prompt_embeds_tensor)
|
||||
|
||||
dtype = dtype or next(self.text_encoder.parameters()).dtype
|
||||
prompt_embeds_tensor = prompt_embeds_tensor.to(dtype=dtype, device=device)
|
||||
|
||||
# Handle negative prompts
|
||||
if negative_prompt is None:
|
||||
negative_embeds = torch.zeros_like(prompt_embeds_tensor)
|
||||
else:
|
||||
if isinstance(negative_prompt, str):
|
||||
negative_prompt = [negative_prompt]
|
||||
negative_result = self.encode_prompt(
|
||||
prompt=negative_prompt, device=device, dtype=dtype, return_index=return_index
|
||||
)
|
||||
negative_embeds = negative_result[0]
|
||||
|
||||
# Explicitly cast both tensors to FloatTensor for mypy
|
||||
from typing import cast
|
||||
|
||||
prompt_tensor = cast(FloatTensor, prompt_embeds_tensor.to(dtype=dtype))
|
||||
negative_tensor = cast(FloatTensor, negative_embeds.to(dtype=dtype))
|
||||
return (prompt_tensor, negative_tensor)
|
||||
|
||||
def to(self, torch_device=None, torch_dtype=None, silence_dtype_warnings=False):
|
||||
"""Move pipeline components to specified device and dtype."""
|
||||
if hasattr(self, "vae"):
|
||||
self.vae.to(device=torch_device, dtype=torch_dtype)
|
||||
if hasattr(self, "text_encoder"):
|
||||
self.text_encoder.to(device=torch_device, dtype=torch_dtype)
|
||||
if hasattr(self, "dit_model"):
|
||||
self.dit_model.to(device=torch_device, dtype=torch_dtype)
|
||||
return self
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]]=None,
|
||||
prompt_embeds: Optional[FloatTensor] = None,
|
||||
height: Optional[int] = 1024,
|
||||
width: Optional[int] = 1024,
|
||||
num_inference_steps: int = 30,
|
||||
guidance_scale: float = 6.0,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_embeds: Optional[FloatTensor] = None,
|
||||
num_images_per_prompt: int = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
alpha: Optional[float] = None,
|
||||
apg_config: Optional[APGConfig] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""Generate images from text prompt."""
|
||||
# Ensure height and width are not None for calculation
|
||||
if height is None:
|
||||
height = 1024
|
||||
if width is None:
|
||||
width = 1024
|
||||
|
||||
dtype = dtype or next(self.dit_model.parameters()).dtype
|
||||
apg_config = apg_config or APGConfig(enabled=False)
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
# 2. Encode prompts
|
||||
prompt_batch_size = len(prompt) if isinstance(prompt, list) else 1
|
||||
batch_size = prompt_batch_size * num_images_per_prompt
|
||||
|
||||
if prompt_embeds is None or negative_prompt_embeds is None:
|
||||
prompt_embeds, negative_embeds = self.encode_prompt(
|
||||
prompt=prompt, negative_prompt=negative_prompt, device=self.text_encoder.device, dtype=dtype,
|
||||
return_index=self.return_index,
|
||||
)
|
||||
else:
|
||||
negative_embeds = negative_prompt_embeds
|
||||
|
||||
# Repeat embeddings for num_images_per_prompt
|
||||
prompt_embeds = prompt_embeds.repeat_interleave(num_images_per_prompt, dim=0)
|
||||
negative_embeds = negative_embeds.repeat_interleave(num_images_per_prompt, dim=0)
|
||||
|
||||
# 3. Initialize latents
|
||||
latent_height = height // self.vae_scale_factor
|
||||
latent_width = width // self.vae_scale_factor
|
||||
|
||||
if isinstance(generator, list):
|
||||
if len(generator) != batch_size:
|
||||
raise ValueError(f"Got {len(generator)} generators for {batch_size} samples")
|
||||
|
||||
latents = randn_tensor((batch_size, 16, latent_height, latent_width), generator=generator, device=device, dtype=dtype)
|
||||
acc_latents = latents.clone()
|
||||
|
||||
# 4. Calculate alpha if not provided
|
||||
if alpha is None:
|
||||
image_token_size = latent_height * latent_width
|
||||
alpha = 2 * math.sqrt(image_token_size / (64 * 64))
|
||||
|
||||
# 6. Sampling loop
|
||||
self.dit_model.eval()
|
||||
|
||||
# Check if guidance is needed
|
||||
do_classifier_free_guidance = guidance_scale >= 1.0
|
||||
|
||||
for i in self.progress_bar(range(num_inference_steps, 0, -1)):
|
||||
# Calculate timesteps
|
||||
t = i / num_inference_steps
|
||||
t_next = (i - 1) / num_inference_steps
|
||||
# Scale timesteps according to alpha
|
||||
t = t * alpha / (1 + (alpha - 1) * t)
|
||||
t_next = t_next * alpha / (1 + (alpha - 1) * t_next)
|
||||
dt = t - t_next
|
||||
|
||||
# Create tensor with proper device
|
||||
t_tensor = torch.tensor([t] * batch_size, device=device, dtype=dtype)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
# Duplicate latents for both conditional and unconditional inputs
|
||||
latents_input = torch.cat([latents] * 2)
|
||||
# Concatenate negative and positive prompt embeddings
|
||||
context_input = torch.cat([negative_embeds, prompt_embeds])
|
||||
# Duplicate timesteps for the batch
|
||||
t_input = torch.cat([t_tensor] * 2)
|
||||
|
||||
# Get model predictions in a single pass
|
||||
model_outputs = self.dit_model(latents_input, context_input, t_input)
|
||||
|
||||
# Split outputs back into unconditional and conditional predictions
|
||||
uncond_output, cond_output = model_outputs.chunk(2)
|
||||
|
||||
if apg_config.enabled:
|
||||
# Augmented Parallel Guidance
|
||||
dy = cond_output
|
||||
dd = cond_output - uncond_output
|
||||
# Find parallel direction
|
||||
parallel_direction = (dy * dd).sum() / (dy * dy).sum() * dy
|
||||
orthogonal_direction = dd - parallel_direction
|
||||
# Scale orthogonal component
|
||||
orthogonal_std = orthogonal_direction.std()
|
||||
orthogonal_scale = min(1, apg_config.orthogonal_threshold / orthogonal_std)
|
||||
orthogonal_direction = orthogonal_direction * orthogonal_scale
|
||||
model_output = dy + (guidance_scale - 1) * orthogonal_direction
|
||||
else:
|
||||
# Standard classifier-free guidance
|
||||
model_output = uncond_output + guidance_scale * (cond_output - uncond_output)
|
||||
else:
|
||||
# If no guidance needed, just run the model normally
|
||||
model_output = self.dit_model(latents, prompt_embeds, t_tensor)
|
||||
|
||||
# Update latents
|
||||
acc_latents = acc_latents + dt * model_output.to(device)
|
||||
latents = acc_latents.clone()
|
||||
|
||||
# 7. Decode latents
|
||||
# These checks handle the case where mypy doesn't recognize these attributes
|
||||
scaling_factor = getattr(self.vae.config, "scaling_factor", 0.18215) if hasattr(self.vae, "config") else 0.18215
|
||||
shift_factor = getattr(self.vae.config, "shift_factor", 0) if hasattr(self.vae, "config") else 0
|
||||
|
||||
latents = latents / scaling_factor + shift_factor
|
||||
|
||||
vae_dtype = self.vae.dtype if hasattr(self.vae, "dtype") else dtype
|
||||
decoded_images = self.vae.decode(latents.to(vae_dtype)).sample if hasattr(self.vae, "decode") else latents
|
||||
|
||||
# Offload all models
|
||||
try:
|
||||
self.maybe_free_model_hooks()
|
||||
except AttributeError as e:
|
||||
if "OptimizedModule" in str(e):
|
||||
import warnings
|
||||
warnings.warn(
|
||||
"Encountered 'OptimizedModule' error when offloading models. "
|
||||
"This issue might be fixed in the future by: "
|
||||
"https://github.com/huggingface/diffusers/pull/10730"
|
||||
)
|
||||
else:
|
||||
raise
|
||||
|
||||
# 8. Post-process images
|
||||
images = (decoded_images / 2 + 0.5).clamp(0, 1)
|
||||
# Convert to PIL Images
|
||||
images = (images * 255).round().clamp(0, 255).to(torch.uint8).cpu()
|
||||
pil_images = [Image.fromarray(img.permute(1, 2, 0).numpy()) for img in images]
|
||||
|
||||
return FLitePipelineOutput(
|
||||
images=pil_images,
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
from .flux_kontext import FluxKontextModel
|
||||
@@ -0,0 +1,420 @@
|
||||
import os
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
import torch
|
||||
import torchvision
|
||||
import yaml
|
||||
from toolkit import train_tools
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||
from PIL import Image
|
||||
from toolkit.models.base_model import BaseModel
|
||||
from diffusers import FluxTransformer2DModel, AutoencoderKL, FluxKontextPipeline
|
||||
from toolkit.basic import flush
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
from toolkit.models.flux import add_model_gpu_splitter_to_flux, bypass_flux_guidance, restore_flux_guidance
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
from toolkit.accelerator import get_accelerator, unwrap_model
|
||||
from optimum.quanto import freeze, QTensor
|
||||
from toolkit.util.mask import generate_random_mask, random_dialate_mask
|
||||
from toolkit.util.quantize import quantize, get_qtype
|
||||
from transformers import T5TokenizerFast, T5EncoderModel, CLIPTextModel, CLIPTokenizer
|
||||
from einops import rearrange, repeat
|
||||
import random
|
||||
import torch.nn.functional as F
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
|
||||
scheduler_config = {
|
||||
"base_image_seq_len": 256,
|
||||
"base_shift": 0.5,
|
||||
"max_image_seq_len": 4096,
|
||||
"max_shift": 1.15,
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 3.0,
|
||||
"use_dynamic_shifting": True
|
||||
}
|
||||
|
||||
|
||||
|
||||
class FluxKontextModel(BaseModel):
|
||||
arch = "flux_kontext"
|
||||
|
||||
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 = ['FluxTransformer2DModel']
|
||||
|
||||
# static method to get the noise scheduler
|
||||
@staticmethod
|
||||
def get_train_scheduler():
|
||||
return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
|
||||
|
||||
def get_bucket_divisibility(self):
|
||||
return 16
|
||||
|
||||
def load_model(self):
|
||||
dtype = self.torch_dtype
|
||||
self.print_and_status_update("Loading Flux Kontext model")
|
||||
# will be updated if we detect a existing checkpoint in training folder
|
||||
model_path = self.model_config.name_or_path
|
||||
# this is the original path put in the model directory
|
||||
# it is here because for finetuning we only save the transformer usually
|
||||
# so we need this for the VAE, te, etc
|
||||
base_model_path = self.model_config.extras_name_or_path
|
||||
|
||||
transformer_path = model_path
|
||||
transformer_subfolder = 'transformer'
|
||||
if os.path.exists(transformer_path):
|
||||
transformer_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 transformer")
|
||||
transformer = FluxTransformer2DModel.from_pretrained(
|
||||
transformer_path,
|
||||
subfolder=transformer_subfolder,
|
||||
torch_dtype=dtype
|
||||
)
|
||||
transformer.to(self.quantize_device, dtype=dtype)
|
||||
|
||||
if self.model_config.quantize:
|
||||
# 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,
|
||||
**self.model_config.quantize_kwargs)
|
||||
freeze(transformer)
|
||||
transformer.to(self.device_torch)
|
||||
else:
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Loading T5")
|
||||
tokenizer_2 = T5TokenizerFast.from_pretrained(
|
||||
base_model_path, subfolder="tokenizer_2", torch_dtype=dtype
|
||||
)
|
||||
text_encoder_2 = T5EncoderModel.from_pretrained(
|
||||
base_model_path, subfolder="text_encoder_2", torch_dtype=dtype
|
||||
)
|
||||
text_encoder_2.to(self.device_torch, dtype=dtype)
|
||||
flush()
|
||||
|
||||
if self.model_config.quantize_te:
|
||||
self.print_and_status_update("Quantizing T5")
|
||||
quantize(text_encoder_2, weights=get_qtype(
|
||||
self.model_config.qtype))
|
||||
freeze(text_encoder_2)
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Loading CLIP")
|
||||
text_encoder = CLIPTextModel.from_pretrained(
|
||||
base_model_path, subfolder="text_encoder", torch_dtype=dtype)
|
||||
tokenizer = CLIPTokenizer.from_pretrained(
|
||||
base_model_path, subfolder="tokenizer", torch_dtype=dtype)
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
|
||||
self.print_and_status_update("Loading VAE")
|
||||
vae = AutoencoderKL.from_pretrained(
|
||||
base_model_path, subfolder="vae", torch_dtype=dtype)
|
||||
|
||||
self.noise_scheduler = FluxKontextModel.get_train_scheduler()
|
||||
|
||||
self.print_and_status_update("Making pipe")
|
||||
|
||||
pipe: FluxKontextPipeline = FluxKontextPipeline(
|
||||
scheduler=self.noise_scheduler,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
text_encoder_2=None,
|
||||
tokenizer_2=tokenizer_2,
|
||||
vae=vae,
|
||||
transformer=None,
|
||||
)
|
||||
# for quantization, it works best to do these after making the pipe
|
||||
pipe.text_encoder_2 = text_encoder_2
|
||||
pipe.transformer = transformer
|
||||
|
||||
self.print_and_status_update("Preparing Model")
|
||||
|
||||
text_encoder = [pipe.text_encoder, pipe.text_encoder_2]
|
||||
tokenizer = [pipe.tokenizer, pipe.tokenizer_2]
|
||||
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
|
||||
flush()
|
||||
# just to make sure everything is on the right device and dtype
|
||||
text_encoder[0].to(self.device_torch)
|
||||
text_encoder[0].requires_grad_(False)
|
||||
text_encoder[0].eval()
|
||||
text_encoder[1].to(self.device_torch)
|
||||
text_encoder[1].requires_grad_(False)
|
||||
text_encoder[1].eval()
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
flush()
|
||||
|
||||
# save it to the model class
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder # list of text encoders
|
||||
self.tokenizer = tokenizer # list of tokenizers
|
||||
self.model = pipe.transformer
|
||||
self.pipeline = pipe
|
||||
self.print_and_status_update("Model Loaded")
|
||||
|
||||
def get_generation_pipeline(self):
|
||||
scheduler = FluxKontextModel.get_train_scheduler()
|
||||
|
||||
pipeline: FluxKontextPipeline = FluxKontextPipeline(
|
||||
scheduler=scheduler,
|
||||
text_encoder=unwrap_model(self.text_encoder[0]),
|
||||
tokenizer=self.tokenizer[0],
|
||||
text_encoder_2=unwrap_model(self.text_encoder[1]),
|
||||
tokenizer_2=self.tokenizer[1],
|
||||
vae=unwrap_model(self.vae),
|
||||
transformer=unwrap_model(self.transformer)
|
||||
)
|
||||
|
||||
pipeline = pipeline.to(self.device_torch)
|
||||
|
||||
return pipeline
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: FluxKontextPipeline,
|
||||
gen_config: GenerateImageConfig,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
unconditional_embeds: PromptEmbeds,
|
||||
generator: torch.Generator,
|
||||
extra: dict,
|
||||
):
|
||||
if gen_config.ctrl_img is None:
|
||||
raise ValueError(
|
||||
"Control image is required for Flux Kontext model generation."
|
||||
)
|
||||
else:
|
||||
control_img = Image.open(gen_config.ctrl_img)
|
||||
control_img = control_img.convert("RGB")
|
||||
# resize to width and height
|
||||
if control_img.size != (gen_config.width, gen_config.height):
|
||||
control_img = control_img.resize(
|
||||
(gen_config.width, gen_config.height), Image.BILINEAR
|
||||
)
|
||||
gen_config.width = int(gen_config.width // 16 * 16)
|
||||
gen_config.height = int(gen_config.height // 16 * 16)
|
||||
img = pipeline(
|
||||
image=control_img,
|
||||
prompt_embeds=conditional_embeds.text_embeds,
|
||||
pooled_prompt_embeds=conditional_embeds.pooled_embeds,
|
||||
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,
|
||||
max_area=gen_config.height * gen_config.width,
|
||||
_auto_resize=False,
|
||||
**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,
|
||||
guidance_embedding_scale: float,
|
||||
bypass_guidance_embedding: bool,
|
||||
**kwargs
|
||||
):
|
||||
with torch.no_grad():
|
||||
bs, c, h, w = latent_model_input.shape
|
||||
# if we have a control on the channel dimension, put it on the batch for packing
|
||||
has_control = False
|
||||
if latent_model_input.shape[1] == 32:
|
||||
# chunk it and stack it on batch dimension
|
||||
# dont update batch size for img_its
|
||||
lat, control = torch.chunk(latent_model_input, 2, dim=1)
|
||||
latent_model_input = torch.cat([lat, control], dim=0)
|
||||
has_control = True
|
||||
|
||||
latent_model_input_packed = rearrange(
|
||||
latent_model_input,
|
||||
"b c (h ph) (w pw) -> b (h w) (c ph pw)",
|
||||
ph=2,
|
||||
pw=2
|
||||
)
|
||||
|
||||
img_ids = torch.zeros(h // 2, w // 2, 3)
|
||||
img_ids[..., 1] = img_ids[..., 1] + torch.arange(h // 2)[:, None]
|
||||
img_ids[..., 2] = img_ids[..., 2] + torch.arange(w // 2)[None, :]
|
||||
img_ids = repeat(img_ids, "h w c -> b (h w) c",
|
||||
b=bs).to(self.device_torch)
|
||||
|
||||
# handle control image ids
|
||||
if has_control:
|
||||
ctrl_ids = img_ids.clone()
|
||||
ctrl_ids[..., 0] = 1
|
||||
img_ids = torch.cat([img_ids, ctrl_ids], dim=1)
|
||||
|
||||
|
||||
txt_ids = torch.zeros(
|
||||
bs, text_embeddings.text_embeds.shape[1], 3).to(self.device_torch)
|
||||
|
||||
# # handle guidance
|
||||
if self.unet_unwrapped.config.guidance_embeds:
|
||||
if isinstance(guidance_embedding_scale, list):
|
||||
guidance = torch.tensor(
|
||||
guidance_embedding_scale, device=self.device_torch)
|
||||
else:
|
||||
guidance = torch.tensor(
|
||||
[guidance_embedding_scale], device=self.device_torch)
|
||||
# Expand guidance to match original batch_size
|
||||
guidance = guidance.expand(bs)
|
||||
else:
|
||||
guidance = None
|
||||
|
||||
if bypass_guidance_embedding:
|
||||
bypass_flux_guidance(self.unet)
|
||||
|
||||
cast_dtype = self.unet.dtype
|
||||
# changes from orig implementation
|
||||
if txt_ids.ndim == 3:
|
||||
txt_ids = txt_ids[0]
|
||||
if img_ids.ndim == 3:
|
||||
img_ids = img_ids[0]
|
||||
|
||||
latent_size = latent_model_input_packed.shape[1]
|
||||
# move the kontext channels. We have them on batch dimension to here, but need to put them on the latent dimension
|
||||
if has_control:
|
||||
latent, control = torch.chunk(latent_model_input_packed, 2, dim=0)
|
||||
latent_model_input_packed = torch.cat(
|
||||
[latent, control], dim=1
|
||||
)
|
||||
latent_size = latent.shape[1]
|
||||
|
||||
noise_pred = self.unet(
|
||||
hidden_states=latent_model_input_packed.to(
|
||||
self.device_torch, cast_dtype),
|
||||
timestep=timestep / 1000,
|
||||
encoder_hidden_states=text_embeddings.text_embeds.to(
|
||||
self.device_torch, cast_dtype),
|
||||
pooled_projections=text_embeddings.pooled_embeds.to(
|
||||
self.device_torch, cast_dtype),
|
||||
txt_ids=txt_ids,
|
||||
img_ids=img_ids,
|
||||
guidance=guidance,
|
||||
return_dict=False,
|
||||
**kwargs,
|
||||
)[0]
|
||||
|
||||
# remove kontext image conditioning
|
||||
noise_pred = noise_pred[:, :latent_size]
|
||||
|
||||
if isinstance(noise_pred, QTensor):
|
||||
noise_pred = noise_pred.dequantize()
|
||||
|
||||
noise_pred = rearrange(
|
||||
noise_pred,
|
||||
"b (h w) (c ph pw) -> b c (h ph) (w pw)",
|
||||
h=latent_model_input.shape[2] // 2,
|
||||
w=latent_model_input.shape[3] // 2,
|
||||
ph=2,
|
||||
pw=2,
|
||||
c=self.vae.config.latent_channels
|
||||
)
|
||||
|
||||
if bypass_guidance_embedding:
|
||||
restore_flux_guidance(self.unet)
|
||||
|
||||
return noise_pred
|
||||
|
||||
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
|
||||
if self.pipeline.text_encoder.device != self.device_torch:
|
||||
self.pipeline.text_encoder.to(self.device_torch)
|
||||
prompt_embeds, pooled_prompt_embeds = train_tools.encode_prompts_flux(
|
||||
self.tokenizer,
|
||||
self.text_encoder,
|
||||
prompt,
|
||||
max_length=512,
|
||||
)
|
||||
pe = PromptEmbeds(
|
||||
prompt_embeds
|
||||
)
|
||||
pe.pooled_embeds = pooled_prompt_embeds
|
||||
return pe
|
||||
|
||||
def get_model_has_grad(self):
|
||||
# return from a weight if it has grad
|
||||
return self.model.proj_out.weight.requires_grad
|
||||
|
||||
def get_te_has_grad(self):
|
||||
# return from a weight if it has grad
|
||||
return self.text_encoder[1].encoder.block[0].layer[0].SelfAttention.q.weight.requires_grad
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
# only save the unet
|
||||
transformer: FluxTransformer2DModel = 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')
|
||||
batch = kwargs.get('batch')
|
||||
return (noise - batch.latents).detach()
|
||||
|
||||
def condition_noisy_latents(self, latents: torch.Tensor, batch:'DataLoaderBatchDTO'):
|
||||
with torch.no_grad():
|
||||
control_tensor = batch.control_tensor
|
||||
if control_tensor is not None:
|
||||
self.vae.to(self.device_torch)
|
||||
# we are not packed here, so we just need to pass them so we can pack them later
|
||||
control_tensor = control_tensor * 2 - 1
|
||||
control_tensor = control_tensor.to(self.vae_device_torch, dtype=self.torch_dtype)
|
||||
|
||||
# if it is not the size of batch.tensor, (bs,ch,h,w) then we need to resize it
|
||||
if batch.tensor is not None:
|
||||
target_h, target_w = batch.tensor.shape[2], batch.tensor.shape[3]
|
||||
else:
|
||||
# When caching latents, batch.tensor is None. We get the size from the file_items instead.
|
||||
target_h = batch.file_items[0].crop_height
|
||||
target_w = batch.file_items[0].crop_width
|
||||
|
||||
if control_tensor.shape[2] != target_h or control_tensor.shape[3] != target_w:
|
||||
control_tensor = F.interpolate(control_tensor, size=(target_h, target_w), mode='bilinear')
|
||||
|
||||
control_latent = self.encode_images(control_tensor).to(latents.device, latents.dtype)
|
||||
latents = torch.cat((latents, control_latent), dim=1)
|
||||
|
||||
return latents.detach()
|
||||
|
||||
def get_base_model_version(self):
|
||||
return "flux.1_kontext"
|
||||
2
extensions_built_in/diffusion_models/hidream/__init__.py
Normal file
2
extensions_built_in/diffusion_models/hidream/__init__.py
Normal file
@@ -0,0 +1,2 @@
|
||||
from .hidream_model import HidreamModel
|
||||
from .hidream_e1_model import HidreamE1Model
|
||||
189
extensions_built_in/diffusion_models/hidream/hidream_e1_model.py
Normal file
189
extensions_built_in/diffusion_models/hidream/hidream_e1_model.py
Normal file
@@ -0,0 +1,189 @@
|
||||
from .hidream_model import HidreamModel
|
||||
from .src.pipelines.hidream_image.pipeline_hidream_image_editing import (
|
||||
HiDreamImageEditingPipeline,
|
||||
)
|
||||
from .src.schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
from toolkit.accelerator import unwrap_model
|
||||
import torch
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.config_modules import GenerateImageConfig
|
||||
from diffusers.models import HiDreamImageTransformer2DModel
|
||||
|
||||
import torch.nn.functional as F
|
||||
from PIL import Image
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
|
||||
|
||||
class HidreamE1Model(HidreamModel):
|
||||
arch = "hidream_e1"
|
||||
hidream_transformer_class = HiDreamImageTransformer2DModel
|
||||
hidream_pipeline_class = HiDreamImageEditingPipeline
|
||||
|
||||
def get_generation_pipeline(self):
|
||||
scheduler = FlowUniPCMultistepScheduler(
|
||||
num_train_timesteps=1000, shift=3.0, use_dynamic_shifting=False
|
||||
)
|
||||
|
||||
pipeline: HiDreamImageEditingPipeline = HiDreamImageEditingPipeline(
|
||||
scheduler=scheduler,
|
||||
vae=self.vae,
|
||||
text_encoder=self.text_encoder[0],
|
||||
tokenizer=self.tokenizer[0],
|
||||
text_encoder_2=self.text_encoder[1],
|
||||
tokenizer_2=self.tokenizer[1],
|
||||
text_encoder_3=self.text_encoder[2],
|
||||
tokenizer_3=self.tokenizer[2],
|
||||
text_encoder_4=self.text_encoder[3],
|
||||
tokenizer_4=self.tokenizer[3],
|
||||
transformer=unwrap_model(self.model),
|
||||
aggressive_unloading=self.low_vram,
|
||||
)
|
||||
|
||||
pipeline = pipeline.to(self.device_torch)
|
||||
|
||||
return pipeline
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: HiDreamImageEditingPipeline,
|
||||
gen_config: GenerateImageConfig,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
unconditional_embeds: PromptEmbeds,
|
||||
generator: torch.Generator,
|
||||
extra: dict,
|
||||
):
|
||||
if gen_config.ctrl_img is None:
|
||||
raise ValueError(
|
||||
"Control image is required for Flux Kontext model generation."
|
||||
)
|
||||
else:
|
||||
control_img = Image.open(gen_config.ctrl_img)
|
||||
control_img = control_img.convert("RGB")
|
||||
# resize to width and height
|
||||
if control_img.size != (gen_config.width, gen_config.height):
|
||||
control_img = control_img.resize(
|
||||
(gen_config.width, gen_config.height), Image.BILINEAR
|
||||
)
|
||||
img = pipeline(
|
||||
prompt_embeds_t5=conditional_embeds.text_embeds[0],
|
||||
prompt_embeds_llama3=conditional_embeds.text_embeds[1],
|
||||
pooled_prompt_embeds=conditional_embeds.pooled_embeds,
|
||||
negative_prompt_embeds_t5=unconditional_embeds.text_embeds[0],
|
||||
negative_prompt_embeds_llama3=unconditional_embeds.text_embeds[1],
|
||||
negative_pooled_prompt_embeds=unconditional_embeds.pooled_embeds,
|
||||
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,
|
||||
image=control_img,
|
||||
**extra,
|
||||
).images[0]
|
||||
return img
|
||||
|
||||
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
|
||||
self.text_encoder_to(self.device_torch, dtype=self.torch_dtype)
|
||||
max_sequence_length = 128
|
||||
(
|
||||
prompt_embeds_t5,
|
||||
negative_prompt_embeds_t5,
|
||||
prompt_embeds_llama3,
|
||||
negative_prompt_embeds_llama3,
|
||||
pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
) = self.pipeline.encode_prompt(
|
||||
prompt=prompt,
|
||||
prompt_2=prompt,
|
||||
prompt_3=prompt,
|
||||
prompt_4=prompt,
|
||||
device=self.device_torch,
|
||||
dtype=self.torch_dtype,
|
||||
num_images_per_prompt=1,
|
||||
max_sequence_length=max_sequence_length,
|
||||
do_classifier_free_guidance=False,
|
||||
)
|
||||
prompt_embeds = [prompt_embeds_t5, prompt_embeds_llama3]
|
||||
pe = PromptEmbeds([prompt_embeds, pooled_prompt_embeds])
|
||||
return pe
|
||||
|
||||
def condition_noisy_latents(
|
||||
self, latents: torch.Tensor, batch: "DataLoaderBatchDTO"
|
||||
):
|
||||
with torch.no_grad():
|
||||
control_tensor = batch.control_tensor
|
||||
if control_tensor is not None:
|
||||
self.vae.to(self.device_torch)
|
||||
# we are not packed here, so we just need to pass them so we can pack them later
|
||||
control_tensor = control_tensor * 2 - 1
|
||||
control_tensor = control_tensor.to(
|
||||
self.vae_device_torch, dtype=self.torch_dtype
|
||||
)
|
||||
|
||||
# if it is not the size of batch.tensor, (bs,ch,h,w) then we need to resize it
|
||||
if batch.tensor is not None:
|
||||
target_h, target_w = batch.tensor.shape[2], batch.tensor.shape[3]
|
||||
else:
|
||||
# When caching latents, batch.tensor is None. We get the size from the file_items instead.
|
||||
target_h = batch.file_items[0].crop_height
|
||||
target_w = batch.file_items[0].crop_width
|
||||
|
||||
if (
|
||||
control_tensor.shape[2] != target_h
|
||||
or control_tensor.shape[3] != target_w
|
||||
):
|
||||
control_tensor = F.interpolate(
|
||||
control_tensor, size=(target_h, target_w), mode="bilinear"
|
||||
)
|
||||
|
||||
control_latent = self.encode_images(control_tensor).to(
|
||||
latents.device, latents.dtype
|
||||
)
|
||||
latents = torch.cat((latents, control_latent), dim=1)
|
||||
|
||||
return latents.detach()
|
||||
|
||||
def get_noise_prediction(
|
||||
self,
|
||||
latent_model_input: torch.Tensor,
|
||||
timestep: torch.Tensor, # 0 to 1000 scale
|
||||
text_embeddings: PromptEmbeds,
|
||||
**kwargs,
|
||||
):
|
||||
with torch.no_grad():
|
||||
# make sure config is set
|
||||
self.model.config.force_inference_output = True
|
||||
has_control = False
|
||||
lat_size = latent_model_input.shape[-1]
|
||||
if latent_model_input.shape[1] == 32:
|
||||
# chunk it and stack it on batch dimension
|
||||
# dont update batch size for img_its
|
||||
lat, control = torch.chunk(latent_model_input, 2, dim=1)
|
||||
latent_model_input = torch.cat([lat, control], dim=-1)
|
||||
has_control = True
|
||||
|
||||
dtype = self.model.dtype
|
||||
device = self.device_torch
|
||||
|
||||
text_embeds = text_embeddings.text_embeds
|
||||
# run the to for the list
|
||||
text_embeds = [te.to(device, dtype=dtype) for te in text_embeds]
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timesteps=timestep,
|
||||
encoder_hidden_states_t5=text_embeds[0],
|
||||
encoder_hidden_states_llama3=text_embeds[1],
|
||||
pooled_embeds=text_embeddings.pooled_embeds.to(device, dtype=dtype),
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if has_control:
|
||||
noise_pred = -1.0 * noise_pred[..., :lat_size]
|
||||
else:
|
||||
noise_pred = -1.0 * noise_pred
|
||||
|
||||
return noise_pred
|
||||
453
extensions_built_in/diffusion_models/hidream/hidream_model.py
Normal file
453
extensions_built_in/diffusion_models/hidream/hidream_model.py
Normal file
@@ -0,0 +1,453 @@
|
||||
import os
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
import einops
|
||||
import torch
|
||||
import torchvision
|
||||
import yaml
|
||||
from toolkit import train_tools
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||
from PIL import Image
|
||||
from toolkit.models.base_model import BaseModel
|
||||
from diffusers import AutoencoderKL, TorchAoConfig
|
||||
from toolkit.basic import flush
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
from toolkit.models.flux import add_model_gpu_splitter_to_flux, bypass_flux_guidance, restore_flux_guidance
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
from toolkit.accelerator import get_accelerator, unwrap_model
|
||||
from optimum.quanto import freeze, QTensor
|
||||
from toolkit.util.mask import generate_random_mask, random_dialate_mask
|
||||
from toolkit.util.quantize import quantize, get_qtype
|
||||
from transformers import T5TokenizerFast, T5EncoderModel, CLIPTextModel, CLIPTokenizer, TorchAoConfig as TorchAoConfigTransformers
|
||||
from .src.pipelines.hidream_image.pipeline_hidream_image import HiDreamImagePipeline
|
||||
from .src.models.transformers.transformer_hidream_image import HiDreamImageTransformer2DModel
|
||||
from .src.schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
from transformers import LlamaForCausalLM, PreTrainedTokenizerFast
|
||||
from einops import rearrange, repeat
|
||||
import random
|
||||
import torch.nn.functional as F
|
||||
from tqdm import tqdm
|
||||
from transformers import (
|
||||
CLIPTextModelWithProjection,
|
||||
CLIPTokenizer,
|
||||
T5EncoderModel,
|
||||
T5Tokenizer,
|
||||
LlamaForCausalLM,
|
||||
PreTrainedTokenizerFast
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
|
||||
scheduler_config = {
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 3.0
|
||||
}
|
||||
|
||||
# LLAMA_MODEL_NAME = "meta-llama/Meta-Llama-3.1-8B-Instruct"
|
||||
LLAMA_MODEL_PATH = "unsloth/Meta-Llama-3.1-8B-Instruct"
|
||||
BASE_MODEL_PATH = "HiDream-ai/HiDream-I1-Full"
|
||||
|
||||
|
||||
class HidreamModel(BaseModel):
|
||||
arch = "hidream"
|
||||
hidream_transformer_class = HiDreamImageTransformer2DModel
|
||||
hidream_pipeline_class = HiDreamImagePipeline
|
||||
|
||||
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 = ['HiDreamImageTransformer2DModel']
|
||||
|
||||
# static method to get the noise scheduler
|
||||
@staticmethod
|
||||
def get_train_scheduler():
|
||||
return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
|
||||
|
||||
def get_bucket_divisibility(self):
|
||||
return 16
|
||||
|
||||
def load_model(self):
|
||||
dtype = self.torch_dtype
|
||||
# HiDream-ai/HiDream-I1-Full
|
||||
self.print_and_status_update("Loading HiDream model")
|
||||
# will be updated if we detect a existing checkpoint in training folder
|
||||
model_path = self.model_config.name_or_path
|
||||
extras_path = self.model_config.extras_name_or_path
|
||||
|
||||
llama_model_path = self.model_config.model_kwargs.get('llama_model_path', LLAMA_MODEL_PATH)
|
||||
|
||||
scheduler = HidreamModel.get_train_scheduler()
|
||||
|
||||
self.print_and_status_update("Loading llama 8b model")
|
||||
|
||||
tokenizer_4 = PreTrainedTokenizerFast.from_pretrained(
|
||||
llama_model_path,
|
||||
use_fast=False
|
||||
)
|
||||
|
||||
text_encoder_4 = LlamaForCausalLM.from_pretrained(
|
||||
llama_model_path,
|
||||
output_hidden_states=True,
|
||||
output_attentions=True,
|
||||
torch_dtype=torch.bfloat16,
|
||||
)
|
||||
text_encoder_4.to(self.device_torch, dtype=dtype)
|
||||
|
||||
if self.model_config.quantize_te:
|
||||
self.print_and_status_update("Quantizing llama 8b model")
|
||||
quantization_type = get_qtype(self.model_config.qtype_te)
|
||||
quantize(text_encoder_4, weights=quantization_type)
|
||||
freeze(text_encoder_4)
|
||||
|
||||
if self.low_vram:
|
||||
# unload it for now
|
||||
text_encoder_4.to('cpu')
|
||||
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Loading transformer")
|
||||
|
||||
transformer = self.hidream_transformer_class.from_pretrained(
|
||||
model_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype=torch.bfloat16
|
||||
)
|
||||
|
||||
if not self.low_vram:
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
|
||||
if self.model_config.quantize:
|
||||
self.print_and_status_update("Quantizing transformer")
|
||||
quantization_type = get_qtype(self.model_config.qtype)
|
||||
if self.low_vram:
|
||||
# move and quantize only certain pieces at a time.
|
||||
all_blocks = list(transformer.double_stream_blocks) + list(transformer.single_stream_blocks)
|
||||
self.print_and_status_update(" - quantizing transformer blocks")
|
||||
for block in tqdm(all_blocks):
|
||||
block.to(self.device_torch, dtype=dtype)
|
||||
quantize(block, weights=quantization_type)
|
||||
freeze(block)
|
||||
block.to('cpu')
|
||||
# flush()
|
||||
|
||||
self.print_and_status_update(" - quantizing extras")
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
quantize(transformer, weights=quantization_type)
|
||||
freeze(transformer)
|
||||
else:
|
||||
quantize(transformer, weights=quantization_type)
|
||||
freeze(transformer)
|
||||
|
||||
if self.low_vram:
|
||||
# unload it for now
|
||||
transformer.to('cpu')
|
||||
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Loading vae")
|
||||
|
||||
vae = AutoencoderKL.from_pretrained(
|
||||
extras_path,
|
||||
subfolder="vae",
|
||||
torch_dtype=torch.bfloat16
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
|
||||
self.print_and_status_update("Loading clip encoders")
|
||||
|
||||
text_encoder = CLIPTextModelWithProjection.from_pretrained(
|
||||
extras_path,
|
||||
subfolder="text_encoder",
|
||||
torch_dtype=torch.bfloat16
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
tokenizer = CLIPTokenizer.from_pretrained(
|
||||
extras_path,
|
||||
subfolder="tokenizer"
|
||||
)
|
||||
|
||||
text_encoder_2 = CLIPTextModelWithProjection.from_pretrained(
|
||||
extras_path,
|
||||
subfolder="text_encoder_2",
|
||||
torch_dtype=torch.bfloat16
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
tokenizer_2 = CLIPTokenizer.from_pretrained(
|
||||
extras_path,
|
||||
subfolder="tokenizer_2"
|
||||
)
|
||||
|
||||
flush()
|
||||
self.print_and_status_update("Loading T5 encoders")
|
||||
|
||||
text_encoder_3 = T5EncoderModel.from_pretrained(
|
||||
extras_path,
|
||||
subfolder="text_encoder_3",
|
||||
torch_dtype=torch.bfloat16
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
if self.model_config.quantize_te:
|
||||
self.print_and_status_update("Quantizing T5")
|
||||
quantization_type = get_qtype(self.model_config.qtype_te)
|
||||
quantize(text_encoder_3, weights=quantization_type)
|
||||
freeze(text_encoder_3)
|
||||
flush()
|
||||
|
||||
tokenizer_3 = T5Tokenizer.from_pretrained(
|
||||
extras_path,
|
||||
subfolder="tokenizer_3"
|
||||
)
|
||||
flush()
|
||||
|
||||
if self.low_vram:
|
||||
self.print_and_status_update("Moving everything to device")
|
||||
# move it all back
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
vae.to(self.device_torch, dtype=dtype)
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
text_encoder_2.to(self.device_torch, dtype=dtype)
|
||||
text_encoder_4.to(self.device_torch, dtype=dtype)
|
||||
text_encoder_3.to(self.device_torch, dtype=dtype)
|
||||
|
||||
# set to eval mode
|
||||
# transformer.eval()
|
||||
vae.eval()
|
||||
text_encoder.eval()
|
||||
text_encoder_2.eval()
|
||||
text_encoder_4.eval()
|
||||
text_encoder_3.eval()
|
||||
|
||||
pipe = self.hidream_pipeline_class(
|
||||
scheduler=scheduler,
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
text_encoder_2=text_encoder_2,
|
||||
tokenizer_2=tokenizer_2,
|
||||
text_encoder_3=text_encoder_3,
|
||||
tokenizer_3=tokenizer_3,
|
||||
text_encoder_4=text_encoder_4,
|
||||
tokenizer_4=tokenizer_4,
|
||||
transformer=transformer,
|
||||
)
|
||||
|
||||
flush()
|
||||
|
||||
text_encoder_list = [text_encoder, text_encoder_2, text_encoder_3, text_encoder_4]
|
||||
tokenizer_list = [tokenizer, tokenizer_2, tokenizer_3, tokenizer_4]
|
||||
|
||||
for te in text_encoder_list:
|
||||
# set the dtype
|
||||
te.to(self.device_torch, dtype=dtype)
|
||||
# freeze the model
|
||||
freeze(te)
|
||||
# set to eval mode
|
||||
te.eval()
|
||||
# set the requires grad to false
|
||||
te.requires_grad_(False)
|
||||
|
||||
flush()
|
||||
|
||||
# save it to the model class
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder_list # list of text encoders
|
||||
self.tokenizer = tokenizer_list # list of tokenizers
|
||||
self.model = pipe.transformer
|
||||
self.pipeline = pipe
|
||||
self.print_and_status_update("Model Loaded")
|
||||
|
||||
def get_generation_pipeline(self):
|
||||
scheduler = FlowUniPCMultistepScheduler(
|
||||
num_train_timesteps=1000,
|
||||
shift=3.0,
|
||||
use_dynamic_shifting=False
|
||||
)
|
||||
|
||||
pipeline: HiDreamImagePipeline = HiDreamImagePipeline(
|
||||
scheduler=scheduler,
|
||||
vae=self.vae,
|
||||
text_encoder=self.text_encoder[0],
|
||||
tokenizer=self.tokenizer[0],
|
||||
text_encoder_2=self.text_encoder[1],
|
||||
tokenizer_2=self.tokenizer[1],
|
||||
text_encoder_3=self.text_encoder[2],
|
||||
tokenizer_3=self.tokenizer[2],
|
||||
text_encoder_4=self.text_encoder[3],
|
||||
tokenizer_4=self.tokenizer[3],
|
||||
transformer=unwrap_model(self.model),
|
||||
aggressive_unloading=self.low_vram
|
||||
)
|
||||
|
||||
pipeline = pipeline.to(self.device_torch)
|
||||
|
||||
return pipeline
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: HiDreamImagePipeline,
|
||||
gen_config: GenerateImageConfig,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
unconditional_embeds: PromptEmbeds,
|
||||
generator: torch.Generator,
|
||||
extra: dict,
|
||||
):
|
||||
img = pipeline(
|
||||
prompt_embeds=conditional_embeds.text_embeds,
|
||||
pooled_prompt_embeds=conditional_embeds.pooled_embeds,
|
||||
negative_prompt_embeds=unconditional_embeds.text_embeds,
|
||||
negative_pooled_prompt_embeds=unconditional_embeds.pooled_embeds,
|
||||
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
|
||||
):
|
||||
batch_size = latent_model_input.shape[0]
|
||||
with torch.no_grad():
|
||||
if latent_model_input.shape[-2] != latent_model_input.shape[-1]:
|
||||
B, C, H, W = latent_model_input.shape
|
||||
pH, pW = H // self.model.config.patch_size, W // self.model.config.patch_size
|
||||
|
||||
img_sizes = torch.tensor([pH, pW], dtype=torch.int64).reshape(-1)
|
||||
img_ids = torch.zeros(pH, pW, 3)
|
||||
img_ids[..., 1] = img_ids[..., 1] + torch.arange(pH)[:, None]
|
||||
img_ids[..., 2] = img_ids[..., 2] + torch.arange(pW)[None, :]
|
||||
img_ids = img_ids.reshape(pH * pW, -1)
|
||||
img_ids_pad = torch.zeros(self.transformer.max_seq, 3)
|
||||
img_ids_pad[:pH*pW, :] = img_ids
|
||||
|
||||
img_sizes = img_sizes.unsqueeze(0).to(latent_model_input.device)
|
||||
img_sizes = torch.cat([img_sizes] * batch_size, dim=0)
|
||||
img_ids = img_ids_pad.unsqueeze(0).to(latent_model_input.device)
|
||||
img_ids = torch.cat([img_ids] * batch_size, dim=0)
|
||||
else:
|
||||
img_sizes = img_ids = None
|
||||
|
||||
dtype = self.model.dtype
|
||||
device = self.device_torch
|
||||
|
||||
# Pack the latent
|
||||
if latent_model_input.shape[-2] != latent_model_input.shape[-1]:
|
||||
B, C, H, W = latent_model_input.shape
|
||||
patch_size = self.transformer.config.patch_size
|
||||
pH, pW = H // patch_size, W // patch_size
|
||||
out = torch.zeros(
|
||||
(B, C, self.transformer.max_seq, patch_size * patch_size),
|
||||
dtype=latent_model_input.dtype,
|
||||
device=latent_model_input.device
|
||||
)
|
||||
latent_model_input = einops.rearrange(latent_model_input, 'B C (H p1) (W p2) -> B C (H W) (p1 p2)', p1=patch_size, p2=patch_size)
|
||||
out[:, :, 0:pH*pW] = latent_model_input
|
||||
latent_model_input = out
|
||||
|
||||
text_embeds = text_embeddings.text_embeds
|
||||
# run the to for the list
|
||||
text_embeds = [te.to(device, dtype=dtype) for te in text_embeds]
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states = latent_model_input,
|
||||
timesteps = timestep,
|
||||
encoder_hidden_states = text_embeds,
|
||||
pooled_embeds = text_embeddings.pooled_embeds.to(device, dtype=dtype),
|
||||
img_sizes = img_sizes,
|
||||
img_ids = img_ids,
|
||||
return_dict = False,
|
||||
)[0]
|
||||
noise_pred = -noise_pred
|
||||
|
||||
return noise_pred
|
||||
|
||||
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
|
||||
self.text_encoder_to(self.device_torch, dtype=self.torch_dtype)
|
||||
max_sequence_length = 128
|
||||
prompt_embeds, pooled_prompt_embeds = self.pipeline._encode_prompt(
|
||||
prompt = prompt,
|
||||
prompt_2 = prompt,
|
||||
prompt_3 = prompt,
|
||||
prompt_4 = prompt,
|
||||
device = self.device_torch,
|
||||
dtype = self.torch_dtype,
|
||||
num_images_per_prompt = 1,
|
||||
max_sequence_length = max_sequence_length,
|
||||
)
|
||||
pe = PromptEmbeds(
|
||||
[prompt_embeds, pooled_prompt_embeds]
|
||||
)
|
||||
return pe
|
||||
|
||||
def get_model_has_grad(self):
|
||||
# return from a weight if it has grad
|
||||
return self.model.double_stream_blocks[0].block.attn1.to_q.weight.requires_grad
|
||||
|
||||
def get_te_has_grad(self):
|
||||
# assume no one wants to finetune 4 text encoders.
|
||||
return False
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
# only save the unet
|
||||
transformer: HiDreamImageTransformer2DModel = 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')
|
||||
batch = kwargs.get('batch')
|
||||
return (noise - batch.latents).detach()
|
||||
|
||||
def get_transformer_block_names(self) -> Optional[List[str]]:
|
||||
return ['double_stream_blocks', 'single_stream_blocks']
|
||||
|
||||
def convert_lora_weights_before_save(self, state_dict):
|
||||
# currently starte with transformer. but needs to start with diffusion_model. for comfyui
|
||||
new_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
new_key = key.replace("transformer.", "diffusion_model.")
|
||||
new_sd[new_key] = value
|
||||
return new_sd
|
||||
|
||||
def convert_lora_weights_before_load(self, state_dict):
|
||||
# saved as diffusion_model. but needs to be transformer. for ai-toolkit
|
||||
new_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
new_key = key.replace("diffusion_model.", "transformer.")
|
||||
new_sd[new_key] = value
|
||||
return new_sd
|
||||
|
||||
def get_base_model_version(self):
|
||||
return "hidream_i1"
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
from .models.transformers.transformer_hidream_image import HiDreamImageTransformer2DModel
|
||||
from .pipelines.hidream_image.pipeline_hidream_image import HiDreamImagePipeline
|
||||
@@ -0,0 +1,106 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
from typing import Optional
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.utils.torch_utils import maybe_allow_in_graph
|
||||
|
||||
@maybe_allow_in_graph
|
||||
class HiDreamAttention(Attention):
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
heads: int = 8,
|
||||
dim_head: int = 64,
|
||||
upcast_attention: bool = False,
|
||||
upcast_softmax: bool = False,
|
||||
scale_qk: bool = True,
|
||||
eps: float = 1e-5,
|
||||
processor = None,
|
||||
out_dim: int = None,
|
||||
single: bool = False
|
||||
):
|
||||
super(Attention, self).__init__()
|
||||
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
|
||||
self.query_dim = query_dim
|
||||
self.upcast_attention = upcast_attention
|
||||
self.upcast_softmax = upcast_softmax
|
||||
self.out_dim = out_dim if out_dim is not None else query_dim
|
||||
|
||||
self.scale_qk = scale_qk
|
||||
self.scale = dim_head**-0.5 if self.scale_qk else 1.0
|
||||
|
||||
self.heads = out_dim // dim_head if out_dim is not None else heads
|
||||
self.sliceable_head_dim = heads
|
||||
self.single = single
|
||||
|
||||
linear_cls = nn.Linear
|
||||
self.linear_cls = linear_cls
|
||||
self.to_q = linear_cls(query_dim, self.inner_dim)
|
||||
self.to_k = linear_cls(self.inner_dim, self.inner_dim)
|
||||
self.to_v = linear_cls(self.inner_dim, self.inner_dim)
|
||||
self.to_out = linear_cls(self.inner_dim, self.out_dim)
|
||||
self.q_rms_norm = nn.RMSNorm(self.inner_dim, eps)
|
||||
self.k_rms_norm = nn.RMSNorm(self.inner_dim, eps)
|
||||
|
||||
if not single:
|
||||
self.to_q_t = linear_cls(query_dim, self.inner_dim)
|
||||
self.to_k_t = linear_cls(self.inner_dim, self.inner_dim)
|
||||
self.to_v_t = linear_cls(self.inner_dim, self.inner_dim)
|
||||
self.to_out_t = linear_cls(self.inner_dim, self.out_dim)
|
||||
self.q_rms_norm_t = nn.RMSNorm(self.inner_dim, eps)
|
||||
self.k_rms_norm_t = nn.RMSNorm(self.inner_dim, eps)
|
||||
|
||||
self.set_processor(processor)
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
norm_image_tokens: torch.FloatTensor,
|
||||
image_tokens_masks: torch.FloatTensor = None,
|
||||
norm_text_tokens: torch.FloatTensor = None,
|
||||
rope: torch.FloatTensor = None,
|
||||
) -> torch.Tensor:
|
||||
return self.processor(
|
||||
self,
|
||||
image_tokens = norm_image_tokens,
|
||||
image_tokens_masks = image_tokens_masks,
|
||||
text_tokens = norm_text_tokens,
|
||||
rope = rope,
|
||||
)
|
||||
|
||||
class FeedForwardSwiGLU(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
hidden_dim: int,
|
||||
multiple_of: int = 256,
|
||||
ffn_dim_multiplier: Optional[float] = None,
|
||||
):
|
||||
super().__init__()
|
||||
hidden_dim = int(2 * hidden_dim / 3)
|
||||
# custom dim factor multiplier
|
||||
if ffn_dim_multiplier is not None:
|
||||
hidden_dim = int(ffn_dim_multiplier * hidden_dim)
|
||||
hidden_dim = multiple_of * (
|
||||
(hidden_dim + multiple_of - 1) // multiple_of
|
||||
)
|
||||
|
||||
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
|
||||
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
|
||||
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, x):
|
||||
return self.w2(torch.nn.functional.silu(self.w1(x)) * self.w3(x))
|
||||
@@ -0,0 +1,121 @@
|
||||
from typing import Optional
|
||||
import torch
|
||||
from .attention import HiDreamAttention
|
||||
|
||||
# Try to import Flash Attention first
|
||||
flash_attn_available = False
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func
|
||||
USE_FLASH_ATTN3 = True
|
||||
flash_attn_available = True
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn import flash_attn_func
|
||||
USE_FLASH_ATTN3 = False
|
||||
flash_attn_available = True
|
||||
except ImportError:
|
||||
USE_FLASH_ATTN3 = False
|
||||
flash_attn_available = False
|
||||
|
||||
# Copied from https://github.com/black-forest-labs/flux/blob/main/src/flux/math.py
|
||||
def apply_rope(xq: torch.Tensor, xk: torch.Tensor, freqs_cis: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
|
||||
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
|
||||
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
|
||||
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
|
||||
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
|
||||
|
||||
def attention(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor):
|
||||
if flash_attn_available:
|
||||
if USE_FLASH_ATTN3:
|
||||
hidden_states = flash_attn_func(query, key, value, causal=False, deterministic=False)[0]
|
||||
else:
|
||||
hidden_states = flash_attn_func(query, key, value, dropout_p=0., causal=False)
|
||||
else:
|
||||
# Use torch's scaled dot-product attention as fallback
|
||||
# Reshape for torch.nn.functional.scaled_dot_product_attention which expects [batch, heads, seq_len, head_dim]
|
||||
query = query.transpose(1, 2) # [batch, heads, seq_len, head_dim]
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
|
||||
hidden_states = torch.nn.functional.scaled_dot_product_attention(
|
||||
query, key, value,
|
||||
attn_mask=None,
|
||||
dropout_p=0.0,
|
||||
is_causal=False
|
||||
)
|
||||
|
||||
# Restore original shape
|
||||
hidden_states = hidden_states.transpose(1, 2) # [batch, seq_len, heads, head_dim]
|
||||
|
||||
hidden_states = hidden_states.flatten(-2)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
return hidden_states
|
||||
|
||||
class HiDreamAttnProcessor_flashattn:
|
||||
"""Attention processor used typically in processing the SD3-like self-attention projections."""
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: HiDreamAttention,
|
||||
image_tokens: torch.FloatTensor,
|
||||
image_tokens_masks: Optional[torch.FloatTensor] = None,
|
||||
text_tokens: Optional[torch.FloatTensor] = None,
|
||||
rope: torch.FloatTensor = None,
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> torch.FloatTensor:
|
||||
dtype = image_tokens.dtype
|
||||
batch_size = image_tokens.shape[0]
|
||||
|
||||
query_i = attn.q_rms_norm(attn.to_q(image_tokens)).to(dtype=dtype)
|
||||
key_i = attn.k_rms_norm(attn.to_k(image_tokens)).to(dtype=dtype)
|
||||
value_i = attn.to_v(image_tokens)
|
||||
|
||||
inner_dim = key_i.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
query_i = query_i.view(batch_size, -1, attn.heads, head_dim)
|
||||
key_i = key_i.view(batch_size, -1, attn.heads, head_dim)
|
||||
value_i = value_i.view(batch_size, -1, attn.heads, head_dim)
|
||||
if image_tokens_masks is not None:
|
||||
key_i = key_i * image_tokens_masks.view(batch_size, -1, 1, 1)
|
||||
|
||||
if not attn.single:
|
||||
query_t = attn.q_rms_norm_t(attn.to_q_t(text_tokens)).to(dtype=dtype)
|
||||
key_t = attn.k_rms_norm_t(attn.to_k_t(text_tokens)).to(dtype=dtype)
|
||||
value_t = attn.to_v_t(text_tokens)
|
||||
|
||||
query_t = query_t.view(batch_size, -1, attn.heads, head_dim)
|
||||
key_t = key_t.view(batch_size, -1, attn.heads, head_dim)
|
||||
value_t = value_t.view(batch_size, -1, attn.heads, head_dim)
|
||||
|
||||
num_image_tokens = query_i.shape[1]
|
||||
num_text_tokens = query_t.shape[1]
|
||||
query = torch.cat([query_i, query_t], dim=1)
|
||||
key = torch.cat([key_i, key_t], dim=1)
|
||||
value = torch.cat([value_i, value_t], dim=1)
|
||||
else:
|
||||
query = query_i
|
||||
key = key_i
|
||||
value = value_i
|
||||
|
||||
if query.shape[-1] == rope.shape[-3] * 2:
|
||||
query, key = apply_rope(query, key, rope)
|
||||
else:
|
||||
query_1, query_2 = query.chunk(2, dim=-1)
|
||||
key_1, key_2 = key.chunk(2, dim=-1)
|
||||
query_1, key_1 = apply_rope(query_1, key_1, rope)
|
||||
query = torch.cat([query_1, query_2], dim=-1)
|
||||
key = torch.cat([key_1, key_2], dim=-1)
|
||||
|
||||
hidden_states = attention(query, key, value)
|
||||
|
||||
if not attn.single:
|
||||
hidden_states_i, hidden_states_t = torch.split(hidden_states, [num_image_tokens, num_text_tokens], dim=1)
|
||||
hidden_states_i = attn.to_out(hidden_states_i)
|
||||
hidden_states_t = attn.to_out_t(hidden_states_t)
|
||||
return hidden_states_i, hidden_states_t
|
||||
else:
|
||||
hidden_states = attn.to_out(hidden_states)
|
||||
return hidden_states
|
||||
@@ -0,0 +1,114 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
from typing import List
|
||||
from diffusers.models.embeddings import Timesteps, TimestepEmbedding
|
||||
|
||||
# Copied from https://github.com/black-forest-labs/flux/blob/main/src/flux/math.py
|
||||
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
|
||||
assert dim % 2 == 0, "The dimension must be even."
|
||||
|
||||
scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim
|
||||
omega = 1.0 / (theta**scale)
|
||||
|
||||
batch_size, seq_length = pos.shape
|
||||
out = torch.einsum("...n,d->...nd", pos, omega)
|
||||
cos_out = torch.cos(out)
|
||||
sin_out = torch.sin(out)
|
||||
|
||||
stacked_out = torch.stack([cos_out, -sin_out, sin_out, cos_out], dim=-1)
|
||||
out = stacked_out.view(batch_size, -1, dim // 2, 2, 2)
|
||||
return out.float()
|
||||
|
||||
# Copied from https://github.com/black-forest-labs/flux/blob/main/src/flux/modules/layers.py
|
||||
class EmbedND(nn.Module):
|
||||
def __init__(self, theta: int, axes_dim: List[int]):
|
||||
super().__init__()
|
||||
self.theta = theta
|
||||
self.axes_dim = axes_dim
|
||||
|
||||
def forward(self, ids: torch.Tensor) -> torch.Tensor:
|
||||
n_axes = ids.shape[-1]
|
||||
emb = torch.cat(
|
||||
[rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)],
|
||||
dim=-3,
|
||||
)
|
||||
return emb.unsqueeze(2)
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=2,
|
||||
in_channels=4,
|
||||
out_channels=1024,
|
||||
):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.out_channels = out_channels
|
||||
self.proj = nn.Linear(in_channels * patch_size * patch_size, out_channels, bias=True)
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, latent):
|
||||
latent = self.proj(latent)
|
||||
return latent
|
||||
|
||||
class PooledEmbed(nn.Module):
|
||||
def __init__(self, text_emb_dim, hidden_size):
|
||||
super().__init__()
|
||||
self.pooled_embedder = TimestepEmbedding(in_channels=text_emb_dim, time_embed_dim=hidden_size)
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=0.02)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, pooled_embed):
|
||||
return self.pooled_embedder(pooled_embed)
|
||||
|
||||
class TimestepEmbed(nn.Module):
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__()
|
||||
self.time_proj = Timesteps(num_channels=frequency_embedding_size, flip_sin_to_cos=True, downscale_freq_shift=0)
|
||||
self.timestep_embedder = TimestepEmbedding(in_channels=frequency_embedding_size, time_embed_dim=hidden_size)
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=0.02)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, timesteps, wdtype):
|
||||
t_emb = self.time_proj(timesteps).to(dtype=wdtype)
|
||||
t_emb = self.timestep_embedder(t_emb)
|
||||
return t_emb
|
||||
|
||||
class OutEmbed(nn.Module):
|
||||
def __init__(self, hidden_size, patch_size, out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, 2 * hidden_size, bias=True)
|
||||
)
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.zeros_(m.weight)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, x, adaln_input):
|
||||
shift, scale = self.adaLN_modulation(adaln_input).chunk(2, dim=1)
|
||||
x = self.norm_final(x) * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
157
extensions_built_in/diffusion_models/hidream/src/models/moe.py
Normal file
157
extensions_built_in/diffusion_models/hidream/src/models/moe.py
Normal file
@@ -0,0 +1,157 @@
|
||||
import math
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
from .attention import FeedForwardSwiGLU
|
||||
from torch.distributed.nn.functional import all_gather
|
||||
|
||||
_LOAD_BALANCING_LOSS = []
|
||||
def save_load_balancing_loss(loss):
|
||||
global _LOAD_BALANCING_LOSS
|
||||
_LOAD_BALANCING_LOSS.append(loss)
|
||||
|
||||
def clear_load_balancing_loss():
|
||||
global _LOAD_BALANCING_LOSS
|
||||
_LOAD_BALANCING_LOSS.clear()
|
||||
|
||||
def get_load_balancing_loss():
|
||||
global _LOAD_BALANCING_LOSS
|
||||
return _LOAD_BALANCING_LOSS
|
||||
|
||||
def batched_load_balancing_loss():
|
||||
aux_losses_arr = get_load_balancing_loss()
|
||||
alpha = aux_losses_arr[0][-1]
|
||||
Pi = torch.stack([ent[1] for ent in aux_losses_arr], dim=0)
|
||||
fi = torch.stack([ent[2] for ent in aux_losses_arr], dim=0)
|
||||
|
||||
fi_list = all_gather(fi)
|
||||
fi = torch.stack(fi_list, 0).mean(0)
|
||||
|
||||
aux_loss = (Pi * fi).sum(-1).mean() * alpha
|
||||
return aux_loss
|
||||
|
||||
# Modified from https://github.com/deepseek-ai/DeepSeek-V3/blob/main/inference/model.py
|
||||
class MoEGate(nn.Module):
|
||||
def __init__(self, embed_dim, num_routed_experts=4, num_activated_experts=2, aux_loss_alpha=0.01):
|
||||
super().__init__()
|
||||
self.top_k = num_activated_experts
|
||||
self.n_routed_experts = num_routed_experts
|
||||
|
||||
self.scoring_func = 'softmax'
|
||||
self.alpha = aux_loss_alpha
|
||||
self.seq_aux = False
|
||||
|
||||
# topk selection algorithm
|
||||
self.norm_topk_prob = False
|
||||
self.gating_dim = embed_dim
|
||||
self.weight = nn.Parameter(torch.empty((self.n_routed_experts, self.gating_dim)))
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self) -> None:
|
||||
import torch.nn.init as init
|
||||
init.kaiming_uniform_(self.weight, a=math.sqrt(5))
|
||||
|
||||
def forward(self, hidden_states):
|
||||
bsz, seq_len, h = hidden_states.shape
|
||||
# print(bsz, seq_len, h)
|
||||
### compute gating score
|
||||
hidden_states = hidden_states.view(-1, h)
|
||||
logits = F.linear(hidden_states, self.weight, None)
|
||||
if self.scoring_func == 'softmax':
|
||||
scores = logits.softmax(dim=-1)
|
||||
else:
|
||||
raise NotImplementedError(f'insupportable scoring function for MoE gating: {self.scoring_func}')
|
||||
|
||||
### select top-k experts
|
||||
topk_weight, topk_idx = torch.topk(scores, k=self.top_k, dim=-1, sorted=False)
|
||||
|
||||
### norm gate to sum 1
|
||||
if self.top_k > 1 and self.norm_topk_prob:
|
||||
denominator = topk_weight.sum(dim=-1, keepdim=True) + 1e-20
|
||||
topk_weight = topk_weight / denominator
|
||||
# this was in original and memory leaks, not needed
|
||||
|
||||
# ### expert-level computation auxiliary loss
|
||||
# if self.training and self.alpha > 0.0:
|
||||
# scores_for_aux = scores
|
||||
# aux_topk = self.top_k
|
||||
# # always compute aux loss based on the naive greedy topk method
|
||||
# topk_idx_for_aux_loss = topk_idx.view(bsz, -1)
|
||||
# if self.seq_aux:
|
||||
# scores_for_seq_aux = scores_for_aux.view(bsz, seq_len, -1)
|
||||
# ce = torch.zeros(bsz, self.n_routed_experts, device=hidden_states.device)
|
||||
# ce.scatter_add_(1, topk_idx_for_aux_loss, torch.ones(bsz, seq_len * aux_topk, device=hidden_states.device)).div_(seq_len * aux_topk / self.n_routed_experts)
|
||||
# aux_loss = (ce * scores_for_seq_aux.mean(dim = 1)).sum(dim = 1).mean() * self.alpha
|
||||
# else:
|
||||
# mask_ce = F.one_hot(topk_idx_for_aux_loss.view(-1), num_classes=self.n_routed_experts)
|
||||
# ce = mask_ce.float().mean(0)
|
||||
|
||||
# Pi = scores_for_aux.mean(0)
|
||||
# fi = ce * self.n_routed_experts
|
||||
# aux_loss = (Pi * fi).sum() * self.alpha
|
||||
# save_load_balancing_loss((aux_loss, Pi, fi, self.alpha))
|
||||
# else:
|
||||
aux_loss = None
|
||||
return topk_idx, topk_weight, aux_loss
|
||||
|
||||
# Modified from https://github.com/deepseek-ai/DeepSeek-V3/blob/main/inference/model.py
|
||||
class MOEFeedForwardSwiGLU(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
hidden_dim: int,
|
||||
num_routed_experts: int,
|
||||
num_activated_experts: int,
|
||||
):
|
||||
super().__init__()
|
||||
self.shared_experts = FeedForwardSwiGLU(dim, hidden_dim // 2)
|
||||
self.experts = nn.ModuleList([FeedForwardSwiGLU(dim, hidden_dim) for i in range(num_routed_experts)])
|
||||
self.gate = MoEGate(
|
||||
embed_dim = dim,
|
||||
num_routed_experts = num_routed_experts,
|
||||
num_activated_experts = num_activated_experts
|
||||
)
|
||||
self.num_activated_experts = num_activated_experts
|
||||
|
||||
def forward(self, x):
|
||||
wtype = x.dtype
|
||||
identity = x
|
||||
orig_shape = x.shape
|
||||
topk_idx, topk_weight, aux_loss = self.gate(x)
|
||||
x = x.view(-1, x.shape[-1])
|
||||
flat_topk_idx = topk_idx.view(-1)
|
||||
y = self.moe_infer(x, flat_topk_idx, topk_weight.view(-1, 1)).view(*orig_shape)
|
||||
# this was in original and memory leaks, not needed
|
||||
# if self.training:
|
||||
# x = x.repeat_interleave(self.num_activated_experts, dim=0)
|
||||
# y = torch.empty_like(x, dtype=wtype)
|
||||
# for i, expert in enumerate(self.experts):
|
||||
# y[flat_topk_idx == i] = expert(x[flat_topk_idx == i]).to(dtype=wtype)
|
||||
# y = (y.view(*topk_weight.shape, -1) * topk_weight.unsqueeze(-1)).sum(dim=1)
|
||||
# y = y.view(*orig_shape).to(dtype=wtype)
|
||||
# #y = AddAuxiliaryLoss.apply(y, aux_loss)
|
||||
# else:
|
||||
# y = self.moe_infer(x, flat_topk_idx, topk_weight.view(-1, 1)).view(*orig_shape)
|
||||
y = y + self.shared_experts(identity)
|
||||
return y
|
||||
|
||||
# @torch.no_grad()
|
||||
def moe_infer(self, x, flat_expert_indices, flat_expert_weights):
|
||||
expert_cache = torch.zeros_like(x)
|
||||
idxs = flat_expert_indices.argsort()
|
||||
tokens_per_expert = flat_expert_indices.bincount().cpu().numpy().cumsum(0)
|
||||
token_idxs = idxs // self.num_activated_experts
|
||||
for i, end_idx in enumerate(tokens_per_expert):
|
||||
start_idx = 0 if i == 0 else tokens_per_expert[i-1]
|
||||
if start_idx == end_idx:
|
||||
continue
|
||||
expert = self.experts[i]
|
||||
exp_token_idx = token_idxs[start_idx:end_idx]
|
||||
expert_tokens = x[exp_token_idx]
|
||||
expert_out = expert(expert_tokens)
|
||||
expert_out.mul_(flat_expert_weights[idxs[start_idx:end_idx]])
|
||||
|
||||
# for fp16 and other dtype
|
||||
expert_cache = expert_cache.to(expert_out.dtype)
|
||||
expert_cache.scatter_reduce_(0, exp_token_idx.view(-1, 1).repeat(1, x.shape[-1]), expert_out, reduce='sum')
|
||||
return expert_cache
|
||||
@@ -0,0 +1,506 @@
|
||||
from typing import Any, Callable, Dict, Optional, Tuple, List
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import einops
|
||||
from einops import repeat
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.utils import USE_PEFT_BACKEND, is_torch_version, logging, scale_lora_layers, unscale_lora_layers
|
||||
from diffusers.utils.torch_utils import maybe_allow_in_graph
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from ..embeddings import PatchEmbed, PooledEmbed, TimestepEmbed, EmbedND, OutEmbed
|
||||
from ..attention import HiDreamAttention, FeedForwardSwiGLU
|
||||
from ..attention_processor import HiDreamAttnProcessor_flashattn
|
||||
from ..moe import MOEFeedForwardSwiGLU
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
class TextProjection(nn.Module):
|
||||
def __init__(self, in_features, hidden_size):
|
||||
super().__init__()
|
||||
self.linear = nn.Linear(in_features=in_features, out_features=hidden_size, bias=False)
|
||||
|
||||
def forward(self, caption):
|
||||
hidden_states = self.linear(caption)
|
||||
return hidden_states
|
||||
|
||||
class BlockType:
|
||||
TransformerBlock = 1
|
||||
SingleTransformerBlock = 2
|
||||
|
||||
@maybe_allow_in_graph
|
||||
class HiDreamImageSingleTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
num_routed_experts: int = 4,
|
||||
num_activated_experts: int = 2
|
||||
):
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(dim, 6 * dim, bias=True)
|
||||
)
|
||||
nn.init.zeros_(self.adaLN_modulation[1].weight)
|
||||
nn.init.zeros_(self.adaLN_modulation[1].bias)
|
||||
|
||||
# 1. Attention
|
||||
self.norm1_i = nn.LayerNorm(dim, eps = 1e-06, elementwise_affine = False)
|
||||
self.attn1 = HiDreamAttention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
processor = HiDreamAttnProcessor_flashattn(),
|
||||
single = True
|
||||
)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.norm3_i = nn.LayerNorm(dim, eps = 1e-06, elementwise_affine = False)
|
||||
if num_routed_experts > 0:
|
||||
self.ff_i = MOEFeedForwardSwiGLU(
|
||||
dim = dim,
|
||||
hidden_dim = 4 * dim,
|
||||
num_routed_experts = num_routed_experts,
|
||||
num_activated_experts = num_activated_experts,
|
||||
)
|
||||
else:
|
||||
self.ff_i = FeedForwardSwiGLU(dim = dim, hidden_dim = 4 * dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image_tokens: torch.FloatTensor,
|
||||
image_tokens_masks: Optional[torch.FloatTensor] = None,
|
||||
text_tokens: Optional[torch.FloatTensor] = None,
|
||||
adaln_input: Optional[torch.FloatTensor] = None,
|
||||
rope: torch.FloatTensor = None,
|
||||
|
||||
) -> torch.FloatTensor:
|
||||
wtype = image_tokens.dtype
|
||||
shift_msa_i, scale_msa_i, gate_msa_i, shift_mlp_i, scale_mlp_i, gate_mlp_i = \
|
||||
self.adaLN_modulation(adaln_input)[:,None].chunk(6, dim=-1)
|
||||
|
||||
# 1. MM-Attention
|
||||
norm_image_tokens = self.norm1_i(image_tokens).to(dtype=wtype)
|
||||
norm_image_tokens = norm_image_tokens * (1 + scale_msa_i) + shift_msa_i
|
||||
attn_output_i = self.attn1(
|
||||
norm_image_tokens,
|
||||
image_tokens_masks,
|
||||
rope = rope,
|
||||
)
|
||||
image_tokens = gate_msa_i * attn_output_i + image_tokens
|
||||
|
||||
# 2. Feed-forward
|
||||
norm_image_tokens = self.norm3_i(image_tokens).to(dtype=wtype)
|
||||
norm_image_tokens = norm_image_tokens * (1 + scale_mlp_i) + shift_mlp_i
|
||||
ff_output_i = gate_mlp_i * self.ff_i(norm_image_tokens.to(dtype=wtype))
|
||||
image_tokens = ff_output_i + image_tokens
|
||||
return image_tokens
|
||||
|
||||
@maybe_allow_in_graph
|
||||
class HiDreamImageTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
num_routed_experts: int = 4,
|
||||
num_activated_experts: int = 2
|
||||
):
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(dim, 12 * dim, bias=True)
|
||||
)
|
||||
nn.init.zeros_(self.adaLN_modulation[1].weight)
|
||||
nn.init.zeros_(self.adaLN_modulation[1].bias)
|
||||
|
||||
# 1. Attention
|
||||
self.norm1_i = nn.LayerNorm(dim, eps = 1e-06, elementwise_affine = False)
|
||||
self.norm1_t = nn.LayerNorm(dim, eps = 1e-06, elementwise_affine = False)
|
||||
self.attn1 = HiDreamAttention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
processor = HiDreamAttnProcessor_flashattn(),
|
||||
single = False
|
||||
)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.norm3_i = nn.LayerNorm(dim, eps = 1e-06, elementwise_affine = False)
|
||||
if num_routed_experts > 0:
|
||||
self.ff_i = MOEFeedForwardSwiGLU(
|
||||
dim = dim,
|
||||
hidden_dim = 4 * dim,
|
||||
num_routed_experts = num_routed_experts,
|
||||
num_activated_experts = num_activated_experts,
|
||||
)
|
||||
else:
|
||||
self.ff_i = FeedForwardSwiGLU(dim = dim, hidden_dim = 4 * dim)
|
||||
self.norm3_t = nn.LayerNorm(dim, eps = 1e-06, elementwise_affine = False)
|
||||
self.ff_t = FeedForwardSwiGLU(dim = dim, hidden_dim = 4 * dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image_tokens: torch.FloatTensor,
|
||||
image_tokens_masks: Optional[torch.FloatTensor] = None,
|
||||
text_tokens: Optional[torch.FloatTensor] = None,
|
||||
adaln_input: Optional[torch.FloatTensor] = None,
|
||||
rope: torch.FloatTensor = None,
|
||||
) -> torch.FloatTensor:
|
||||
wtype = image_tokens.dtype
|
||||
shift_msa_i, scale_msa_i, gate_msa_i, shift_mlp_i, scale_mlp_i, gate_mlp_i, \
|
||||
shift_msa_t, scale_msa_t, gate_msa_t, shift_mlp_t, scale_mlp_t, gate_mlp_t = \
|
||||
self.adaLN_modulation(adaln_input)[:,None].chunk(12, dim=-1)
|
||||
|
||||
# 1. MM-Attention
|
||||
norm_image_tokens = self.norm1_i(image_tokens).to(dtype=wtype)
|
||||
norm_image_tokens = norm_image_tokens * (1 + scale_msa_i) + shift_msa_i
|
||||
norm_text_tokens = self.norm1_t(text_tokens).to(dtype=wtype)
|
||||
norm_text_tokens = norm_text_tokens * (1 + scale_msa_t) + shift_msa_t
|
||||
|
||||
attn_output_i, attn_output_t = self.attn1(
|
||||
norm_image_tokens,
|
||||
image_tokens_masks,
|
||||
norm_text_tokens,
|
||||
rope = rope,
|
||||
)
|
||||
|
||||
image_tokens = gate_msa_i * attn_output_i + image_tokens
|
||||
text_tokens = gate_msa_t * attn_output_t + text_tokens
|
||||
|
||||
# 2. Feed-forward
|
||||
norm_image_tokens = self.norm3_i(image_tokens).to(dtype=wtype)
|
||||
norm_image_tokens = norm_image_tokens * (1 + scale_mlp_i) + shift_mlp_i
|
||||
norm_text_tokens = self.norm3_t(text_tokens).to(dtype=wtype)
|
||||
norm_text_tokens = norm_text_tokens * (1 + scale_mlp_t) + shift_mlp_t
|
||||
|
||||
ff_output_i = gate_mlp_i * self.ff_i(norm_image_tokens)
|
||||
ff_output_t = gate_mlp_t * self.ff_t(norm_text_tokens)
|
||||
image_tokens = ff_output_i + image_tokens
|
||||
text_tokens = ff_output_t + text_tokens
|
||||
return image_tokens, text_tokens
|
||||
|
||||
@maybe_allow_in_graph
|
||||
class HiDreamImageBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
num_routed_experts: int = 4,
|
||||
num_activated_experts: int = 2,
|
||||
block_type: BlockType = BlockType.TransformerBlock,
|
||||
):
|
||||
super().__init__()
|
||||
block_classes = {
|
||||
BlockType.TransformerBlock: HiDreamImageTransformerBlock,
|
||||
BlockType.SingleTransformerBlock: HiDreamImageSingleTransformerBlock,
|
||||
}
|
||||
self.block = block_classes[block_type](
|
||||
dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
num_routed_experts,
|
||||
num_activated_experts
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image_tokens: torch.FloatTensor,
|
||||
image_tokens_masks: Optional[torch.FloatTensor] = None,
|
||||
text_tokens: Optional[torch.FloatTensor] = None,
|
||||
adaln_input: torch.FloatTensor = None,
|
||||
rope: torch.FloatTensor = None,
|
||||
) -> torch.FloatTensor:
|
||||
return self.block(
|
||||
image_tokens,
|
||||
image_tokens_masks,
|
||||
text_tokens,
|
||||
adaln_input,
|
||||
rope,
|
||||
)
|
||||
|
||||
class HiDreamImageTransformer2DModel(
|
||||
ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin
|
||||
):
|
||||
_supports_gradient_checkpointing = True
|
||||
_no_split_modules = ["HiDreamImageBlock"]
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: Optional[int] = None,
|
||||
in_channels: int = 64,
|
||||
out_channels: Optional[int] = None,
|
||||
num_layers: int = 16,
|
||||
num_single_layers: int = 32,
|
||||
attention_head_dim: int = 128,
|
||||
num_attention_heads: int = 20,
|
||||
caption_channels: List[int] = None,
|
||||
text_emb_dim: int = 2048,
|
||||
num_routed_experts: int = 4,
|
||||
num_activated_experts: int = 2,
|
||||
axes_dims_rope: Tuple[int, int] = (32, 32),
|
||||
max_resolution: Tuple[int, int] = (128, 128),
|
||||
llama_layers: List[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.out_channels = out_channels or in_channels
|
||||
self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim
|
||||
self.llama_layers = llama_layers
|
||||
|
||||
self.t_embedder = TimestepEmbed(self.inner_dim)
|
||||
self.p_embedder = PooledEmbed(text_emb_dim, self.inner_dim)
|
||||
self.x_embedder = PatchEmbed(
|
||||
patch_size = patch_size,
|
||||
in_channels = in_channels,
|
||||
out_channels = self.inner_dim,
|
||||
)
|
||||
self.pe_embedder = EmbedND(theta=10000, axes_dim=axes_dims_rope)
|
||||
|
||||
self.double_stream_blocks = nn.ModuleList(
|
||||
[
|
||||
HiDreamImageBlock(
|
||||
dim = self.inner_dim,
|
||||
num_attention_heads = self.config.num_attention_heads,
|
||||
attention_head_dim = self.config.attention_head_dim,
|
||||
num_routed_experts = num_routed_experts,
|
||||
num_activated_experts = num_activated_experts,
|
||||
block_type = BlockType.TransformerBlock
|
||||
)
|
||||
for i in range(self.config.num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.single_stream_blocks = nn.ModuleList(
|
||||
[
|
||||
HiDreamImageBlock(
|
||||
dim = self.inner_dim,
|
||||
num_attention_heads = self.config.num_attention_heads,
|
||||
attention_head_dim = self.config.attention_head_dim,
|
||||
num_routed_experts = num_routed_experts,
|
||||
num_activated_experts = num_activated_experts,
|
||||
block_type = BlockType.SingleTransformerBlock
|
||||
)
|
||||
for i in range(self.config.num_single_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.final_layer = OutEmbed(self.inner_dim, patch_size, self.out_channels)
|
||||
|
||||
caption_channels = [caption_channels[1], ] * (num_layers + num_single_layers) + [caption_channels[0], ]
|
||||
caption_projection = []
|
||||
for caption_channel in caption_channels:
|
||||
caption_projection.append(TextProjection(in_features = caption_channel, hidden_size = self.inner_dim))
|
||||
self.caption_projection = nn.ModuleList(caption_projection)
|
||||
self.max_seq = max_resolution[0] * max_resolution[1] // (patch_size * patch_size)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
|
||||
def expand_timesteps(self, timesteps, batch_size, device):
|
||||
if not torch.is_tensor(timesteps):
|
||||
is_mps = device.type == "mps"
|
||||
if isinstance(timesteps, float):
|
||||
dtype = torch.float32 if is_mps else torch.float64
|
||||
else:
|
||||
dtype = torch.int32 if is_mps else torch.int64
|
||||
timesteps = torch.tensor([timesteps], dtype=dtype, device=device)
|
||||
elif len(timesteps.shape) == 0:
|
||||
timesteps = timesteps[None].to(device)
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timesteps = timesteps.expand(batch_size)
|
||||
return timesteps
|
||||
|
||||
# the implementation on hidream during train was wrong, just use the inference one.
|
||||
def unpatchify(self, x: torch.Tensor, img_sizes: List[Tuple[int, int]], is_training: bool) -> List[torch.Tensor]:
|
||||
# Process all images in the batch according to their specific dimensions
|
||||
x_arr = []
|
||||
for i, img_size in enumerate(img_sizes):
|
||||
pH, pW = img_size
|
||||
x_arr.append(
|
||||
einops.rearrange(
|
||||
x[i, :pH*pW].reshape(1, pH, pW, -1),
|
||||
'B H W (p1 p2 C) -> B C (H p1) (W p2)',
|
||||
p1=self.config.patch_size, p2=self.config.patch_size
|
||||
)
|
||||
)
|
||||
x = torch.cat(x_arr, dim=0)
|
||||
return x
|
||||
|
||||
def patchify(self, x, max_seq, img_sizes=None):
|
||||
pz2 = self.config.patch_size * self.config.patch_size
|
||||
if isinstance(x, torch.Tensor):
|
||||
B, C = x.shape[0], x.shape[1]
|
||||
device = x.device
|
||||
dtype = x.dtype
|
||||
else:
|
||||
B, C = len(x), x[0].shape[0]
|
||||
device = x[0].device
|
||||
dtype = x[0].dtype
|
||||
x_masks = torch.zeros((B, max_seq), dtype=dtype, device=device)
|
||||
|
||||
if img_sizes is not None:
|
||||
for i, img_size in enumerate(img_sizes):
|
||||
x_masks[i, 0:img_size[0] * img_size[1]] = 1
|
||||
x = einops.rearrange(x, 'B C S p -> B S (p C)', p=pz2)
|
||||
elif isinstance(x, torch.Tensor):
|
||||
pH, pW = x.shape[-2] // self.config.patch_size, x.shape[-1] // self.config.patch_size
|
||||
x = einops.rearrange(x, 'B C (H p1) (W p2) -> B (H W) (p1 p2 C)', p1=self.config.patch_size, p2=self.config.patch_size)
|
||||
img_sizes = [[pH, pW]] * B
|
||||
x_masks = None
|
||||
else:
|
||||
raise NotImplementedError
|
||||
return x, x_masks, img_sizes
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timesteps: torch.LongTensor = None,
|
||||
encoder_hidden_states: torch.Tensor = None,
|
||||
pooled_embeds: torch.Tensor = None,
|
||||
img_sizes: Optional[List[Tuple[int, int]]] = None,
|
||||
img_ids: Optional[torch.Tensor] = None,
|
||||
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = True,
|
||||
):
|
||||
if joint_attention_kwargs is not None:
|
||||
joint_attention_kwargs = joint_attention_kwargs.copy()
|
||||
lora_scale = joint_attention_kwargs.pop("scale", 1.0)
|
||||
else:
|
||||
lora_scale = 1.0
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# weight the lora layers by setting `lora_scale` for each PEFT layer
|
||||
scale_lora_layers(self, lora_scale)
|
||||
else:
|
||||
if joint_attention_kwargs is not None and joint_attention_kwargs.get("scale", None) is not None:
|
||||
logger.warning(
|
||||
"Passing `scale` via `joint_attention_kwargs` when not using the PEFT backend is ineffective."
|
||||
)
|
||||
|
||||
# spatial forward
|
||||
batch_size = hidden_states.shape[0]
|
||||
hidden_states_type = hidden_states.dtype
|
||||
|
||||
# 0. time
|
||||
timesteps = self.expand_timesteps(timesteps, batch_size, hidden_states.device)
|
||||
timesteps = self.t_embedder(timesteps, hidden_states_type)
|
||||
p_embedder = self.p_embedder(pooled_embeds)
|
||||
adaln_input = timesteps + p_embedder
|
||||
|
||||
hidden_states, image_tokens_masks, img_sizes = self.patchify(hidden_states, self.max_seq, img_sizes)
|
||||
if image_tokens_masks is None:
|
||||
pH, pW = img_sizes[0]
|
||||
img_ids = torch.zeros(pH, pW, 3, device=hidden_states.device)
|
||||
img_ids[..., 1] = img_ids[..., 1] + torch.arange(pH, device=hidden_states.device)[:, None]
|
||||
img_ids[..., 2] = img_ids[..., 2] + torch.arange(pW, device=hidden_states.device)[None, :]
|
||||
img_ids = repeat(img_ids, "h w c -> b (h w) c", b=batch_size)
|
||||
hidden_states = self.x_embedder(hidden_states)
|
||||
|
||||
T5_encoder_hidden_states = encoder_hidden_states[0]
|
||||
encoder_hidden_states = encoder_hidden_states[-1]
|
||||
encoder_hidden_states = [encoder_hidden_states[k] for k in self.llama_layers]
|
||||
|
||||
if self.caption_projection is not None:
|
||||
new_encoder_hidden_states = []
|
||||
for i, enc_hidden_state in enumerate(encoder_hidden_states):
|
||||
enc_hidden_state = self.caption_projection[i](enc_hidden_state)
|
||||
enc_hidden_state = enc_hidden_state.view(batch_size, -1, hidden_states.shape[-1])
|
||||
new_encoder_hidden_states.append(enc_hidden_state)
|
||||
encoder_hidden_states = new_encoder_hidden_states
|
||||
T5_encoder_hidden_states = self.caption_projection[-1](T5_encoder_hidden_states)
|
||||
T5_encoder_hidden_states = T5_encoder_hidden_states.view(batch_size, -1, hidden_states.shape[-1])
|
||||
encoder_hidden_states.append(T5_encoder_hidden_states)
|
||||
|
||||
txt_ids = torch.zeros(
|
||||
batch_size,
|
||||
encoder_hidden_states[-1].shape[1] + encoder_hidden_states[-2].shape[1] + encoder_hidden_states[0].shape[1],
|
||||
3,
|
||||
device=img_ids.device, dtype=img_ids.dtype
|
||||
)
|
||||
ids = torch.cat((img_ids, txt_ids), dim=1)
|
||||
rope = self.pe_embedder(ids)
|
||||
|
||||
# 2. Blocks
|
||||
block_id = 0
|
||||
initial_encoder_hidden_states = torch.cat([encoder_hidden_states[-1], encoder_hidden_states[-2]], dim=1)
|
||||
initial_encoder_hidden_states_seq_len = initial_encoder_hidden_states.shape[1]
|
||||
for bid, block in enumerate(self.double_stream_blocks):
|
||||
cur_llama31_encoder_hidden_states = encoder_hidden_states[block_id].detach()
|
||||
cur_encoder_hidden_states = torch.cat([initial_encoder_hidden_states, cur_llama31_encoder_hidden_states], dim=1)
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
hidden_states, initial_encoder_hidden_states = self._gradient_checkpointing_func(
|
||||
block,
|
||||
hidden_states,
|
||||
image_tokens_masks,
|
||||
cur_encoder_hidden_states,
|
||||
adaln_input.clone(),
|
||||
rope.clone(),
|
||||
)
|
||||
|
||||
else:
|
||||
hidden_states, initial_encoder_hidden_states = block(
|
||||
image_tokens = hidden_states,
|
||||
image_tokens_masks = image_tokens_masks,
|
||||
text_tokens = cur_encoder_hidden_states,
|
||||
adaln_input = adaln_input,
|
||||
rope = rope,
|
||||
)
|
||||
initial_encoder_hidden_states = initial_encoder_hidden_states[:, :initial_encoder_hidden_states_seq_len]
|
||||
block_id += 1
|
||||
|
||||
image_tokens_seq_len = hidden_states.shape[1]
|
||||
hidden_states = torch.cat([hidden_states, initial_encoder_hidden_states], dim=1)
|
||||
hidden_states_seq_len = hidden_states.shape[1]
|
||||
if image_tokens_masks is not None:
|
||||
encoder_attention_mask_ones = torch.ones(
|
||||
(batch_size, initial_encoder_hidden_states.shape[1] + cur_llama31_encoder_hidden_states.shape[1]),
|
||||
device=image_tokens_masks.device, dtype=image_tokens_masks.dtype
|
||||
)
|
||||
image_tokens_masks = torch.cat([image_tokens_masks, encoder_attention_mask_ones], dim=1)
|
||||
|
||||
for bid, block in enumerate(self.single_stream_blocks):
|
||||
cur_llama31_encoder_hidden_states = encoder_hidden_states[block_id].detach()
|
||||
hidden_states = torch.cat([hidden_states, cur_llama31_encoder_hidden_states], dim=1)
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block,
|
||||
hidden_states,
|
||||
image_tokens_masks,
|
||||
None,
|
||||
adaln_input.clone(),
|
||||
rope.clone(),
|
||||
)
|
||||
else:
|
||||
hidden_states = block(
|
||||
image_tokens = hidden_states,
|
||||
image_tokens_masks = image_tokens_masks,
|
||||
text_tokens = None,
|
||||
adaln_input = adaln_input,
|
||||
rope = rope,
|
||||
)
|
||||
hidden_states = hidden_states[:, :hidden_states_seq_len]
|
||||
block_id += 1
|
||||
|
||||
hidden_states = hidden_states[:, :image_tokens_seq_len, ...]
|
||||
output = self.final_layer(hidden_states, adaln_input)
|
||||
output = self.unpatchify(output, img_sizes, self.training)
|
||||
if image_tokens_masks is not None:
|
||||
image_tokens_masks = image_tokens_masks[:, :image_tokens_seq_len]
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not return_dict:
|
||||
return (output, image_tokens_masks)
|
||||
return Transformer2DModelOutput(sample=output, mask=image_tokens_masks)
|
||||
|
||||
@@ -0,0 +1,737 @@
|
||||
import inspect
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
import math
|
||||
import einops
|
||||
import torch
|
||||
from transformers import (
|
||||
CLIPTextModelWithProjection,
|
||||
CLIPTokenizer,
|
||||
T5EncoderModel,
|
||||
T5Tokenizer,
|
||||
LlamaForCausalLM,
|
||||
PreTrainedTokenizerFast
|
||||
)
|
||||
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
from diffusers.loaders import FromSingleFileMixin
|
||||
from diffusers.models.autoencoders import AutoencoderKL
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import (
|
||||
USE_PEFT_BACKEND,
|
||||
is_torch_xla_available,
|
||||
logging,
|
||||
)
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from .pipeline_output import HiDreamImagePipelineOutput
|
||||
from ...models.transformers.transformer_hidream_image import HiDreamImageTransformer2DModel
|
||||
from ...schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
|
||||
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
|
||||
|
||||
# Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shift
|
||||
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.15,
|
||||
):
|
||||
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 HiDreamImagePipeline(DiffusionPipeline, FromSingleFileMixin):
|
||||
model_cpu_offload_seq = "text_encoder->text_encoder_2->text_encoder_3->text_encoder_4->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: CLIPTextModelWithProjection,
|
||||
tokenizer: CLIPTokenizer,
|
||||
text_encoder_2: CLIPTextModelWithProjection,
|
||||
tokenizer_2: CLIPTokenizer,
|
||||
text_encoder_3: T5EncoderModel,
|
||||
tokenizer_3: T5Tokenizer,
|
||||
text_encoder_4: LlamaForCausalLM,
|
||||
tokenizer_4: PreTrainedTokenizerFast,
|
||||
transformer: HiDreamImageTransformer2DModel,
|
||||
aggressive_unloading: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
text_encoder_2=text_encoder_2,
|
||||
text_encoder_3=text_encoder_3,
|
||||
text_encoder_4=text_encoder_4,
|
||||
tokenizer=tokenizer,
|
||||
tokenizer_2=tokenizer_2,
|
||||
tokenizer_3=tokenizer_3,
|
||||
tokenizer_4=tokenizer_4,
|
||||
scheduler=scheduler,
|
||||
transformer=transformer,
|
||||
)
|
||||
self.vae_scale_factor = (
|
||||
2 ** (len(self.vae.config.block_out_channels) - 1) if hasattr(self, "vae") and self.vae is not None else 8
|
||||
)
|
||||
# HiDreamImage 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.default_sample_size = 128
|
||||
self.tokenizer_4.pad_token = self.tokenizer_4.eos_token
|
||||
self.aggressive_unloading = aggressive_unloading
|
||||
|
||||
def _get_t5_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
num_images_per_prompt: int = 1,
|
||||
max_sequence_length: int = 128,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder_3.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer_3(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=min(max_sequence_length, self.tokenizer_3.model_max_length),
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
attention_mask = text_inputs.attention_mask
|
||||
untruncated_ids = self.tokenizer_3(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_3.batch_decode(untruncated_ids[:, min(max_sequence_length, self.tokenizer_3.model_max_length) - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because `max_sequence_length` is set to "
|
||||
f" {min(max_sequence_length, self.tokenizer_3.model_max_length)} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
prompt_embeds = self.text_encoder_3(text_input_ids.to(device), attention_mask=attention_mask.to(device))[0]
|
||||
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 _get_clip_prompt_embeds(
|
||||
self,
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
prompt: Union[str, List[str]],
|
||||
num_images_per_prompt: int = 1,
|
||||
max_sequence_length: int = 128,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=min(max_sequence_length, 218),
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = 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 = tokenizer.batch_decode(untruncated_ids[:, 218 - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because CLIP can only handle sequences up to"
|
||||
f" {218} tokens: {removed_text}"
|
||||
)
|
||||
prompt_embeds = text_encoder(text_input_ids.to(device), output_hidden_states=True)
|
||||
|
||||
# Use pooled output of CLIPTextModel
|
||||
prompt_embeds = prompt_embeds[0]
|
||||
prompt_embeds = prompt_embeds.to(dtype=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_llama3_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
num_images_per_prompt: int = 1,
|
||||
max_sequence_length: int = 128,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder_4.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer_4(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=min(max_sequence_length, self.tokenizer_4.model_max_length),
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
attention_mask = text_inputs.attention_mask
|
||||
untruncated_ids = self.tokenizer_4(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_4.batch_decode(untruncated_ids[:, min(max_sequence_length, self.tokenizer_4.model_max_length) - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because `max_sequence_length` is set to "
|
||||
f" {min(max_sequence_length, self.tokenizer_4.model_max_length)} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
outputs = self.text_encoder_4(
|
||||
text_input_ids.to(device),
|
||||
attention_mask=attention_mask.to(device),
|
||||
output_hidden_states=True,
|
||||
output_attentions=True
|
||||
)
|
||||
|
||||
prompt_embeds = outputs.hidden_states[1:]
|
||||
prompt_embeds = torch.stack(prompt_embeds, dim=0)
|
||||
_, _, seq_len, dim = prompt_embeds.shape
|
||||
|
||||
# duplicate text embeddings and attention mask for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, 1, num_images_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(-1, batch_size * num_images_per_prompt, seq_len, dim)
|
||||
return prompt_embeds
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
prompt_2: Union[str, List[str]],
|
||||
prompt_3: Union[str, List[str]],
|
||||
prompt_4: Union[str, List[str]],
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
num_images_per_prompt: int = 1,
|
||||
do_classifier_free_guidance: bool = True,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_3: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_4: Optional[Union[str, List[str]]] = None,
|
||||
prompt_embeds: Optional[List[torch.FloatTensor]] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
max_sequence_length: int = 128,
|
||||
lora_scale: Optional[float] = None,
|
||||
):
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
if prompt is not None:
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds[0].shape[0]
|
||||
|
||||
prompt_embeds, pooled_prompt_embeds = self._encode_prompt(
|
||||
prompt = prompt,
|
||||
prompt_2 = prompt_2,
|
||||
prompt_3 = prompt_3,
|
||||
prompt_4 = prompt_4,
|
||||
device = device,
|
||||
dtype = dtype,
|
||||
num_images_per_prompt = num_images_per_prompt,
|
||||
prompt_embeds = prompt_embeds,
|
||||
pooled_prompt_embeds = pooled_prompt_embeds,
|
||||
max_sequence_length = max_sequence_length,
|
||||
)
|
||||
|
||||
if do_classifier_free_guidance and negative_prompt_embeds is None:
|
||||
negative_prompt = negative_prompt or ""
|
||||
negative_prompt_2 = negative_prompt_2 or negative_prompt
|
||||
negative_prompt_3 = negative_prompt_3 or negative_prompt
|
||||
negative_prompt_4 = negative_prompt_4 or negative_prompt
|
||||
|
||||
# normalize str to list
|
||||
negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt
|
||||
negative_prompt_2 = (
|
||||
batch_size * [negative_prompt_2] if isinstance(negative_prompt_2, str) else negative_prompt_2
|
||||
)
|
||||
negative_prompt_3 = (
|
||||
batch_size * [negative_prompt_3] if isinstance(negative_prompt_3, str) else negative_prompt_3
|
||||
)
|
||||
negative_prompt_4 = (
|
||||
batch_size * [negative_prompt_4] if isinstance(negative_prompt_4, str) else negative_prompt_4
|
||||
)
|
||||
|
||||
if prompt is not None and type(prompt) is not type(negative_prompt):
|
||||
raise TypeError(
|
||||
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||
f" {type(prompt)}."
|
||||
)
|
||||
elif batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
|
||||
negative_prompt_embeds, negative_pooled_prompt_embeds = self._encode_prompt(
|
||||
prompt = negative_prompt,
|
||||
prompt_2 = negative_prompt_2,
|
||||
prompt_3 = negative_prompt_3,
|
||||
prompt_4 = negative_prompt_4,
|
||||
device = device,
|
||||
dtype = dtype,
|
||||
num_images_per_prompt = num_images_per_prompt,
|
||||
prompt_embeds = negative_prompt_embeds,
|
||||
pooled_prompt_embeds = negative_pooled_prompt_embeds,
|
||||
max_sequence_length = max_sequence_length,
|
||||
)
|
||||
return prompt_embeds, negative_prompt_embeds, pooled_prompt_embeds, negative_pooled_prompt_embeds
|
||||
|
||||
def _encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
prompt_2: Union[str, List[str]],
|
||||
prompt_3: Union[str, List[str]],
|
||||
prompt_4: Union[str, List[str]],
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
num_images_per_prompt: int = 1,
|
||||
prompt_embeds: Optional[List[torch.FloatTensor]] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
max_sequence_length: int = 128,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_2 = prompt_2 or prompt
|
||||
prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2
|
||||
|
||||
prompt_3 = prompt_3 or prompt
|
||||
prompt_3 = [prompt_3] if isinstance(prompt_3, str) else prompt_3
|
||||
|
||||
prompt_4 = prompt_4 or prompt
|
||||
prompt_4 = [prompt_4] if isinstance(prompt_4, str) else prompt_4
|
||||
|
||||
pooled_prompt_embeds_1 = self._get_clip_prompt_embeds(
|
||||
self.tokenizer,
|
||||
self.text_encoder,
|
||||
prompt = prompt,
|
||||
num_images_per_prompt = num_images_per_prompt,
|
||||
max_sequence_length = max_sequence_length,
|
||||
device = device,
|
||||
dtype = dtype,
|
||||
)
|
||||
|
||||
pooled_prompt_embeds_2 = self._get_clip_prompt_embeds(
|
||||
self.tokenizer_2,
|
||||
self.text_encoder_2,
|
||||
prompt = prompt_2,
|
||||
num_images_per_prompt = num_images_per_prompt,
|
||||
max_sequence_length = max_sequence_length,
|
||||
device = device,
|
||||
dtype = dtype,
|
||||
)
|
||||
|
||||
pooled_prompt_embeds = torch.cat([pooled_prompt_embeds_1, pooled_prompt_embeds_2], dim=-1)
|
||||
|
||||
t5_prompt_embeds = self._get_t5_prompt_embeds(
|
||||
prompt = prompt_3,
|
||||
num_images_per_prompt = num_images_per_prompt,
|
||||
max_sequence_length = max_sequence_length,
|
||||
device = device,
|
||||
dtype = dtype
|
||||
)
|
||||
llama3_prompt_embeds = self._get_llama3_prompt_embeds(
|
||||
prompt = prompt_4,
|
||||
num_images_per_prompt = num_images_per_prompt,
|
||||
max_sequence_length = max_sequence_length,
|
||||
device = device,
|
||||
dtype = dtype
|
||||
)
|
||||
prompt_embeds = [t5_prompt_embeds, llama3_prompt_embeds]
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds
|
||||
|
||||
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 None:
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
else:
|
||||
if latents.shape != shape:
|
||||
raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {shape}")
|
||||
latents = latents.to(device)
|
||||
return latents
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
@property
|
||||
def do_classifier_free_guidance(self):
|
||||
return self._guidance_scale > 1
|
||||
|
||||
@property
|
||||
def joint_attention_kwargs(self):
|
||||
return self._joint_attention_kwargs
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
prompt_3: Optional[Union[str, List[str]]] = None,
|
||||
prompt_4: Optional[Union[str, List[str]]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 50,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
guidance_scale: float = 5.0,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_3: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_4: Optional[Union[str, List[str]]] = None,
|
||||
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,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_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 = 128,
|
||||
):
|
||||
height = height or self.default_sample_size * self.vae_scale_factor
|
||||
width = width or self.default_sample_size * self.vae_scale_factor
|
||||
|
||||
division = self.vae_scale_factor * 2
|
||||
S_max = (self.default_sample_size * self.vae_scale_factor) ** 2
|
||||
scale = S_max / (width * height)
|
||||
scale = math.sqrt(scale)
|
||||
width, height = int(width * scale // division * division), int(height * scale // division * division)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._joint_attention_kwargs = joint_attention_kwargs
|
||||
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[0].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
|
||||
)
|
||||
(
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
) = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
prompt_2=prompt_2,
|
||||
prompt_3=prompt_3,
|
||||
prompt_4=prompt_4,
|
||||
negative_prompt=negative_prompt,
|
||||
negative_prompt_2=negative_prompt_2,
|
||||
negative_prompt_3=negative_prompt_3,
|
||||
negative_prompt_4=negative_prompt_4,
|
||||
do_classifier_free_guidance=self.do_classifier_free_guidance,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
negative_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,
|
||||
)
|
||||
|
||||
if self.do_classifier_free_guidance:
|
||||
prompt_embeds_arr = []
|
||||
for n, p in zip(negative_prompt_embeds, prompt_embeds):
|
||||
if len(n.shape) == 3:
|
||||
prompt_embeds_arr.append(torch.cat([n, p], dim=0))
|
||||
else:
|
||||
prompt_embeds_arr.append(torch.cat([n, p], dim=1))
|
||||
prompt_embeds = prompt_embeds_arr
|
||||
pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, pooled_prompt_embeds], dim=0)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_images_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
pooled_prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
if latents.shape[-2] != latents.shape[-1]:
|
||||
B, C, H, W = latents.shape
|
||||
pH, pW = H // self.transformer.config.patch_size, W // self.transformer.config.patch_size
|
||||
|
||||
img_sizes = torch.tensor([pH, pW], dtype=torch.int64).reshape(-1)
|
||||
img_ids = torch.zeros(pH, pW, 3)
|
||||
img_ids[..., 1] = img_ids[..., 1] + torch.arange(pH)[:, None]
|
||||
img_ids[..., 2] = img_ids[..., 2] + torch.arange(pW)[None, :]
|
||||
img_ids = img_ids.reshape(pH * pW, -1)
|
||||
img_ids_pad = torch.zeros(self.transformer.max_seq, 3)
|
||||
img_ids_pad[:pH*pW, :] = img_ids
|
||||
|
||||
img_sizes = img_sizes.unsqueeze(0).to(latents.device)
|
||||
img_ids = img_ids_pad.unsqueeze(0).to(latents.device)
|
||||
if self.do_classifier_free_guidance:
|
||||
img_sizes = img_sizes.repeat(2 * B, 1)
|
||||
img_ids = img_ids.repeat(2 * B, 1, 1)
|
||||
else:
|
||||
img_sizes = img_ids = None
|
||||
|
||||
# 5. Prepare timesteps
|
||||
mu = calculate_shift(self.transformer.max_seq)
|
||||
scheduler_kwargs = {"mu": mu}
|
||||
if isinstance(self.scheduler, FlowUniPCMultistepScheduler):
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device, shift=math.exp(mu))
|
||||
timesteps = self.scheduler.timesteps
|
||||
else:
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
sigmas=sigmas,
|
||||
**scheduler_kwargs,
|
||||
)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# 6. Denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latent_model_input.shape[0])
|
||||
|
||||
if latent_model_input.shape[-2] != latent_model_input.shape[-1]:
|
||||
B, C, H, W = latent_model_input.shape
|
||||
patch_size = self.transformer.config.patch_size
|
||||
pH, pW = H // patch_size, W // patch_size
|
||||
out = torch.zeros(
|
||||
(B, C, self.transformer.max_seq, patch_size * patch_size),
|
||||
dtype=latent_model_input.dtype,
|
||||
device=latent_model_input.device
|
||||
)
|
||||
latent_model_input = einops.rearrange(latent_model_input, 'B C (H p1) (W p2) -> B C (H W) (p1 p2)', p1=patch_size, p2=patch_size)
|
||||
out[:, :, 0:pH*pW] = latent_model_input
|
||||
latent_model_input = out
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states = latent_model_input,
|
||||
timesteps = timestep,
|
||||
encoder_hidden_states = prompt_embeds,
|
||||
pooled_embeds = pooled_prompt_embeds,
|
||||
img_sizes = img_sizes,
|
||||
img_ids = img_ids,
|
||||
return_dict = False,
|
||||
)[0]
|
||||
noise_pred = -noise_pred
|
||||
|
||||
# perform guidance
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
# 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)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_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()
|
||||
|
||||
if output_type == "latent":
|
||||
image = latents
|
||||
|
||||
else:
|
||||
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 HiDreamImagePipelineOutput(images=image)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,21 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Union
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
|
||||
from diffusers.utils import BaseOutput
|
||||
|
||||
|
||||
@dataclass
|
||||
class HiDreamImagePipelineOutput(BaseOutput):
|
||||
"""
|
||||
Output class for HiDreamImage pipelines.
|
||||
|
||||
Args:
|
||||
images (`List[PIL.Image.Image]` or `np.ndarray`)
|
||||
List of denoised PIL images of length `batch_size` or numpy array of shape `(batch_size, height, width,
|
||||
num_channels)`. PIL images or numpy array present the denoised images of the diffusion pipeline.
|
||||
"""
|
||||
|
||||
images: Union[List[PIL.Image.Image], np.ndarray]
|
||||
@@ -0,0 +1,428 @@
|
||||
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, is_scipy_available, logging
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
if is_scipy_available():
|
||||
import scipy.stats
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlashFlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
Args:
|
||||
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
|
||||
denoising loop.
|
||||
"""
|
||||
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
|
||||
class FlashFlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
"""
|
||||
Euler scheduler.
|
||||
|
||||
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
|
||||
methods the library implements for all schedulers such as loading and saving.
|
||||
|
||||
Args:
|
||||
num_train_timesteps (`int`, defaults to 1000):
|
||||
The number of diffusion steps to train the model.
|
||||
timestep_spacing (`str`, defaults to `"linspace"`):
|
||||
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||
shift (`float`, defaults to 1.0):
|
||||
The shift value for the timestep schedule.
|
||||
"""
|
||||
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
shift: float = 1.0,
|
||||
use_dynamic_shifting=False,
|
||||
base_shift: Optional[float] = 0.5,
|
||||
max_shift: Optional[float] = 1.15,
|
||||
base_image_seq_len: Optional[int] = 256,
|
||||
max_image_seq_len: Optional[int] = 4096,
|
||||
invert_sigmas: bool = False,
|
||||
use_karras_sigmas: Optional[bool] = False,
|
||||
use_exponential_sigmas: Optional[bool] = False,
|
||||
use_beta_sigmas: Optional[bool] = False,
|
||||
):
|
||||
if self.config.use_beta_sigmas and not is_scipy_available():
|
||||
raise ImportError("Make sure to install scipy if you want to use beta sigmas.")
|
||||
if sum([self.config.use_beta_sigmas, self.config.use_exponential_sigmas, self.config.use_karras_sigmas]) > 1:
|
||||
raise ValueError(
|
||||
"Only one of `config.use_beta_sigmas`, `config.use_exponential_sigmas`, `config.use_karras_sigmas` can be used."
|
||||
)
|
||||
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
|
||||
sigmas = timesteps / num_train_timesteps
|
||||
if not use_dynamic_shifting:
|
||||
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
|
||||
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||
|
||||
self.timesteps = sigmas * num_train_timesteps
|
||||
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
self.sigmas = sigmas.to("cpu") # to avoid too much CPU/GPU communication
|
||||
self.sigma_min = self.sigmas[-1].item()
|
||||
self.sigma_max = self.sigmas[0].item()
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
"""
|
||||
The index counter for current timestep. It will increase 1 after each scheduler step.
|
||||
"""
|
||||
return self._step_index
|
||||
|
||||
@property
|
||||
def begin_index(self):
|
||||
"""
|
||||
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
|
||||
"""
|
||||
return self._begin_index
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
|
||||
def set_begin_index(self, begin_index: int = 0):
|
||||
"""
|
||||
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
|
||||
|
||||
Args:
|
||||
begin_index (`int`):
|
||||
The begin index for the scheduler.
|
||||
"""
|
||||
self._begin_index = begin_index
|
||||
|
||||
def scale_noise(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
noise: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
"""
|
||||
Forward process in flow-matching
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The input sample.
|
||||
timestep (`int`, *optional*):
|
||||
The current timestep in the diffusion chain.
|
||||
|
||||
Returns:
|
||||
`torch.FloatTensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
# Make sure sigmas and timesteps have the same device and dtype as original_samples
|
||||
sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype)
|
||||
|
||||
if sample.device.type == "mps" and torch.is_floating_point(timestep):
|
||||
# mps does not support float64
|
||||
schedule_timesteps = self.timesteps.to(sample.device, dtype=torch.float32)
|
||||
timestep = timestep.to(sample.device, dtype=torch.float32)
|
||||
else:
|
||||
schedule_timesteps = self.timesteps.to(sample.device)
|
||||
timestep = timestep.to(sample.device)
|
||||
|
||||
# self.begin_index is None when scheduler is used for training, or pipeline does not implement set_begin_index
|
||||
if self.begin_index is None:
|
||||
step_indices = [self.index_for_timestep(t, schedule_timesteps) for t in timestep]
|
||||
elif self.step_index is not None:
|
||||
# add_noise is called after first denoising step (for inpainting)
|
||||
step_indices = [self.step_index] * timestep.shape[0]
|
||||
else:
|
||||
# add noise is called before first denoising step to create initial latent(img2img)
|
||||
step_indices = [self.begin_index] * timestep.shape[0]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < len(sample.shape):
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
|
||||
sample = sigma * noise + (1.0 - sigma) * sample
|
||||
|
||||
return sample
|
||||
|
||||
def _sigma_to_t(self, sigma):
|
||||
return sigma * self.config.num_train_timesteps
|
||||
|
||||
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
|
||||
|
||||
def set_timesteps(
|
||||
self,
|
||||
num_inference_steps: int = None,
|
||||
device: Union[str, torch.device] = None,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
mu: Optional[float] = None,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
"""
|
||||
if self.config.use_dynamic_shifting and mu is None:
|
||||
raise ValueError(" you have a pass a value for `mu` when `use_dynamic_shifting` is set to be `True`")
|
||||
|
||||
if sigmas is None:
|
||||
timesteps = np.linspace(
|
||||
self._sigma_to_t(self.sigma_max), self._sigma_to_t(self.sigma_min), num_inference_steps
|
||||
)
|
||||
|
||||
sigmas = timesteps / self.config.num_train_timesteps
|
||||
else:
|
||||
sigmas = np.array(sigmas).astype(np.float32)
|
||||
num_inference_steps = len(sigmas)
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
if self.config.use_dynamic_shifting:
|
||||
sigmas = self.time_shift(mu, 1.0, sigmas)
|
||||
else:
|
||||
sigmas = self.config.shift * sigmas / (1 + (self.config.shift - 1) * sigmas)
|
||||
|
||||
if self.config.use_karras_sigmas:
|
||||
sigmas = self._convert_to_karras(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
|
||||
|
||||
elif self.config.use_exponential_sigmas:
|
||||
sigmas = self._convert_to_exponential(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
|
||||
|
||||
elif self.config.use_beta_sigmas:
|
||||
sigmas = self._convert_to_beta(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
|
||||
|
||||
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
|
||||
if self.config.invert_sigmas:
|
||||
sigmas = 1.0 - sigmas
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
sigmas = torch.cat([sigmas, torch.ones(1, device=sigmas.device)])
|
||||
else:
|
||||
sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
|
||||
|
||||
self.timesteps = timesteps.to(device=device)
|
||||
self.sigmas = sigmas
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||
if schedule_timesteps is None:
|
||||
schedule_timesteps = self.timesteps
|
||||
|
||||
indices = (schedule_timesteps == timestep).nonzero()
|
||||
|
||||
# The sigma index that is taken for the **very** first `step`
|
||||
# is always the second index (or the last index if there is only 1)
|
||||
# This way we can ensure we don't accidentally skip a sigma in
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
return indices[pos].item()
|
||||
|
||||
def _init_step_index(self, timestep):
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
self._step_index = self.index_for_timestep(timestep)
|
||||
else:
|
||||
self._step_index = self._begin_index
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
s_churn: float = 0.0,
|
||||
s_tmin: float = 0.0,
|
||||
s_tmax: float = float("inf"),
|
||||
s_noise: float = 1.0,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
return_dict: bool = True,
|
||||
) -> Union[FlashFlowMatchEulerDiscreteSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
process from the learned model outputs (most often the predicted noise).
|
||||
|
||||
Args:
|
||||
model_output (`torch.FloatTensor`):
|
||||
The direct output from learned diffusion model.
|
||||
timestep (`float`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.FloatTensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
s_churn (`float`):
|
||||
s_tmin (`float`):
|
||||
s_tmax (`float`):
|
||||
s_noise (`float`, defaults to 1.0):
|
||||
Scaling factor for noise added to the sample.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
return_dict (`bool`):
|
||||
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
|
||||
tuple.
|
||||
|
||||
Returns:
|
||||
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
|
||||
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
|
||||
if (
|
||||
isinstance(timestep, int)
|
||||
or isinstance(timestep, torch.IntTensor)
|
||||
or isinstance(timestep, torch.LongTensor)
|
||||
):
|
||||
raise ValueError(
|
||||
(
|
||||
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."
|
||||
),
|
||||
)
|
||||
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
|
||||
# Upcast to avoid precision issues when computing prev_sample
|
||||
|
||||
sigma = self.sigmas[self.step_index]
|
||||
|
||||
# Upcast to avoid precision issues when computing prev_sample
|
||||
sample = sample.to(torch.float32)
|
||||
|
||||
denoised = sample - model_output * sigma
|
||||
|
||||
if self.step_index < self.num_inference_steps - 1:
|
||||
sigma_next = self.sigmas[self.step_index + 1]
|
||||
noise = randn_tensor(
|
||||
model_output.shape,
|
||||
generator=generator,
|
||||
device=model_output.device,
|
||||
dtype=denoised.dtype,
|
||||
)
|
||||
sample = sigma_next * noise + (1.0 - sigma_next) * denoised
|
||||
|
||||
self._step_index += 1
|
||||
sample = sample.to(model_output.dtype)
|
||||
|
||||
if not return_dict:
|
||||
return (sample,)
|
||||
|
||||
return FlashFlowMatchEulerDiscreteSchedulerOutput(prev_sample=sample)
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras
|
||||
def _convert_to_karras(self, in_sigmas: torch.Tensor, num_inference_steps) -> torch.Tensor:
|
||||
"""Constructs the noise schedule of Karras et al. (2022)."""
|
||||
|
||||
# Hack to make sure that other schedulers which copy this function don't break
|
||||
# TODO: Add this logic to the other schedulers
|
||||
if hasattr(self.config, "sigma_min"):
|
||||
sigma_min = self.config.sigma_min
|
||||
else:
|
||||
sigma_min = None
|
||||
|
||||
if hasattr(self.config, "sigma_max"):
|
||||
sigma_max = self.config.sigma_max
|
||||
else:
|
||||
sigma_max = None
|
||||
|
||||
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||
|
||||
rho = 7.0 # 7.0 is the value used in the paper
|
||||
ramp = np.linspace(0, 1, num_inference_steps)
|
||||
min_inv_rho = sigma_min ** (1 / rho)
|
||||
max_inv_rho = sigma_max ** (1 / rho)
|
||||
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
|
||||
return sigmas
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential
|
||||
def _convert_to_exponential(self, in_sigmas: torch.Tensor, num_inference_steps: int) -> torch.Tensor:
|
||||
"""Constructs an exponential noise schedule."""
|
||||
|
||||
# Hack to make sure that other schedulers which copy this function don't break
|
||||
# TODO: Add this logic to the other schedulers
|
||||
if hasattr(self.config, "sigma_min"):
|
||||
sigma_min = self.config.sigma_min
|
||||
else:
|
||||
sigma_min = None
|
||||
|
||||
if hasattr(self.config, "sigma_max"):
|
||||
sigma_max = self.config.sigma_max
|
||||
else:
|
||||
sigma_max = None
|
||||
|
||||
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||
|
||||
sigmas = np.exp(np.linspace(math.log(sigma_max), math.log(sigma_min), num_inference_steps))
|
||||
return sigmas
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta
|
||||
def _convert_to_beta(
|
||||
self, in_sigmas: torch.Tensor, num_inference_steps: int, alpha: float = 0.6, beta: float = 0.6
|
||||
) -> torch.Tensor:
|
||||
"""From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)"""
|
||||
|
||||
# Hack to make sure that other schedulers which copy this function don't break
|
||||
# TODO: Add this logic to the other schedulers
|
||||
if hasattr(self.config, "sigma_min"):
|
||||
sigma_min = self.config.sigma_min
|
||||
else:
|
||||
sigma_min = None
|
||||
|
||||
if hasattr(self.config, "sigma_max"):
|
||||
sigma_max = self.config.sigma_max
|
||||
else:
|
||||
sigma_max = None
|
||||
|
||||
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||
|
||||
sigmas = np.array(
|
||||
[
|
||||
sigma_min + (ppf * (sigma_max - sigma_min))
|
||||
for ppf in [
|
||||
scipy.stats.beta.ppf(timestep, alpha, beta)
|
||||
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
|
||||
]
|
||||
]
|
||||
)
|
||||
return sigmas
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
@@ -0,0 +1,800 @@
|
||||
# Copied from https://github.com/huggingface/diffusers/blob/v0.31.0/src/diffusers/schedulers/scheduling_unipc_multistep.py
|
||||
# Convert unipc for flow matching
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
|
||||
import math
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import (KarrasDiffusionSchedulers,
|
||||
SchedulerMixin,
|
||||
SchedulerOutput)
|
||||
from diffusers.utils import deprecate, is_scipy_available
|
||||
|
||||
if is_scipy_available():
|
||||
import scipy.stats
|
||||
|
||||
|
||||
class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin):
|
||||
"""
|
||||
`UniPCMultistepScheduler` is a training-free framework designed for the fast sampling of diffusion models.
|
||||
|
||||
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
|
||||
methods the library implements for all schedulers such as loading and saving.
|
||||
|
||||
Args:
|
||||
num_train_timesteps (`int`, defaults to 1000):
|
||||
The number of diffusion steps to train the model.
|
||||
solver_order (`int`, default `2`):
|
||||
The UniPC order which can be any positive integer. The effective order of accuracy is `solver_order + 1`
|
||||
due to the UniC. It is recommended to use `solver_order=2` for guided sampling, and `solver_order=3` for
|
||||
unconditional sampling.
|
||||
prediction_type (`str`, defaults to "flow_prediction"):
|
||||
Prediction type of the scheduler function; must be `flow_prediction` for this scheduler, which predicts
|
||||
the flow of the diffusion process.
|
||||
thresholding (`bool`, defaults to `False`):
|
||||
Whether to use the "dynamic thresholding" method. This is unsuitable for latent-space diffusion models such
|
||||
as Stable Diffusion.
|
||||
dynamic_thresholding_ratio (`float`, defaults to 0.995):
|
||||
The ratio for the dynamic thresholding method. Valid only when `thresholding=True`.
|
||||
sample_max_value (`float`, defaults to 1.0):
|
||||
The threshold value for dynamic thresholding. Valid only when `thresholding=True` and `predict_x0=True`.
|
||||
predict_x0 (`bool`, defaults to `True`):
|
||||
Whether to use the updating algorithm on the predicted x0.
|
||||
solver_type (`str`, default `bh2`):
|
||||
Solver type for UniPC. It is recommended to use `bh1` for unconditional sampling when steps < 10, and `bh2`
|
||||
otherwise.
|
||||
lower_order_final (`bool`, default `True`):
|
||||
Whether to use lower-order solvers in the final steps. Only valid for < 15 inference steps. This can
|
||||
stabilize the sampling of DPMSolver for steps < 15, especially for steps <= 10.
|
||||
disable_corrector (`list`, default `[]`):
|
||||
Decides which step to disable the corrector to mitigate the misalignment between `epsilon_theta(x_t, c)`
|
||||
and `epsilon_theta(x_t^c, c)` which can influence convergence for a large guidance scale. Corrector is
|
||||
usually disabled during the first few steps.
|
||||
solver_p (`SchedulerMixin`, default `None`):
|
||||
Any other scheduler that if specified, the algorithm becomes `solver_p + UniC`.
|
||||
use_karras_sigmas (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use Karras sigmas for step sizes in the noise schedule during the sampling process. If `True`,
|
||||
the sigmas are determined according to a sequence of noise levels {σi}.
|
||||
use_exponential_sigmas (`bool`, *optional*, defaults to `False`):
|
||||
Whether to use exponential sigmas for step sizes in the noise schedule during the sampling process.
|
||||
timestep_spacing (`str`, defaults to `"linspace"`):
|
||||
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||
steps_offset (`int`, defaults to 0):
|
||||
An offset added to the inference steps, as required by some model families.
|
||||
final_sigmas_type (`str`, defaults to `"zero"`):
|
||||
The final `sigma` value for the noise schedule during the sampling process. If `"sigma_min"`, the final
|
||||
sigma is the same as the last sigma in the training schedule. If `zero`, the final sigma is set to 0.
|
||||
"""
|
||||
|
||||
_compatibles = [e.name for e in KarrasDiffusionSchedulers]
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
solver_order: int = 2,
|
||||
prediction_type: str = "flow_prediction",
|
||||
shift: Optional[float] = 1.0,
|
||||
use_dynamic_shifting=False,
|
||||
thresholding: bool = False,
|
||||
dynamic_thresholding_ratio: float = 0.995,
|
||||
sample_max_value: float = 1.0,
|
||||
predict_x0: bool = True,
|
||||
solver_type: str = "bh2",
|
||||
lower_order_final: bool = True,
|
||||
disable_corrector: List[int] = [],
|
||||
solver_p: SchedulerMixin = None,
|
||||
timestep_spacing: str = "linspace",
|
||||
steps_offset: int = 0,
|
||||
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min"
|
||||
):
|
||||
|
||||
if solver_type not in ["bh1", "bh2"]:
|
||||
if solver_type in ["midpoint", "heun", "logrho"]:
|
||||
self.register_to_config(solver_type="bh2")
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"{solver_type} is not implemented for {self.__class__}")
|
||||
|
||||
self.predict_x0 = predict_x0
|
||||
# setable values
|
||||
self.num_inference_steps = None
|
||||
alphas = np.linspace(1, 1 / num_train_timesteps,
|
||||
num_train_timesteps)[::-1].copy()
|
||||
sigmas = 1.0 - alphas
|
||||
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32)
|
||||
|
||||
if not use_dynamic_shifting:
|
||||
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
|
||||
sigmas = shift * sigmas / (1 +
|
||||
(shift - 1) * sigmas) # pyright: ignore
|
||||
|
||||
self.sigmas = sigmas
|
||||
self.timesteps = sigmas * num_train_timesteps
|
||||
|
||||
self.model_outputs = [None] * solver_order
|
||||
self.timestep_list = [None] * solver_order
|
||||
self.lower_order_nums = 0
|
||||
self.disable_corrector = disable_corrector
|
||||
self.solver_p = solver_p
|
||||
self.last_sample = None
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
self.sigmas = self.sigmas.to(
|
||||
"cpu") # to avoid too much CPU/GPU communication
|
||||
self.sigma_min = self.sigmas[-1].item()
|
||||
self.sigma_max = self.sigmas[0].item()
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
"""
|
||||
The index counter for current timestep. It will increase 1 after each scheduler step.
|
||||
"""
|
||||
return self._step_index
|
||||
|
||||
@property
|
||||
def begin_index(self):
|
||||
"""
|
||||
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
|
||||
"""
|
||||
return self._begin_index
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
|
||||
def set_begin_index(self, begin_index: int = 0):
|
||||
"""
|
||||
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
|
||||
|
||||
Args:
|
||||
begin_index (`int`):
|
||||
The begin index for the scheduler.
|
||||
"""
|
||||
self._begin_index = begin_index
|
||||
|
||||
# Modified from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler.set_timesteps
|
||||
def set_timesteps(
|
||||
self,
|
||||
num_inference_steps: Union[int, None] = None,
|
||||
device: Union[str, torch.device] = None,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
mu: Optional[Union[float, None]] = None,
|
||||
shift: Optional[Union[float, None]] = None,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
Total number of the spacing of the time steps.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
"""
|
||||
|
||||
if self.config.use_dynamic_shifting and mu is None:
|
||||
raise ValueError(
|
||||
" you have to pass a value for `mu` when `use_dynamic_shifting` is set to be `True`"
|
||||
)
|
||||
|
||||
if sigmas is None:
|
||||
sigmas = np.linspace(self.sigma_max, self.sigma_min,
|
||||
num_inference_steps +
|
||||
1).copy()[:-1] # pyright: ignore
|
||||
|
||||
if self.config.use_dynamic_shifting:
|
||||
sigmas = self.time_shift(mu, 1.0, sigmas) # pyright: ignore
|
||||
else:
|
||||
if shift is None:
|
||||
shift = self.config.shift
|
||||
sigmas = shift * sigmas / (1 +
|
||||
(shift - 1) * sigmas) # pyright: ignore
|
||||
|
||||
if self.config.final_sigmas_type == "sigma_min":
|
||||
sigma_last = ((1 - self.alphas_cumprod[0]) /
|
||||
self.alphas_cumprod[0])**0.5
|
||||
elif self.config.final_sigmas_type == "zero":
|
||||
sigma_last = 0
|
||||
else:
|
||||
raise ValueError(
|
||||
f"`final_sigmas_type` must be one of 'zero', or 'sigma_min', but got {self.config.final_sigmas_type}"
|
||||
)
|
||||
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
sigmas = np.concatenate([sigmas, [sigma_last]
|
||||
]).astype(np.float32) # pyright: ignore
|
||||
|
||||
self.sigmas = torch.from_numpy(sigmas)
|
||||
self.timesteps = torch.from_numpy(timesteps).to(
|
||||
device=device, dtype=torch.int64)
|
||||
|
||||
self.num_inference_steps = len(timesteps)
|
||||
|
||||
self.model_outputs = [
|
||||
None,
|
||||
] * self.config.solver_order
|
||||
self.lower_order_nums = 0
|
||||
self.last_sample = None
|
||||
if self.solver_p:
|
||||
self.solver_p.set_timesteps(self.num_inference_steps, device=device)
|
||||
|
||||
# add an index counter for schedulers that allow duplicated timesteps
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
self.sigmas = self.sigmas.to(
|
||||
"cpu") # to avoid too much CPU/GPU communication
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample
|
||||
def _threshold_sample(self, sample: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
"Dynamic thresholding: At each sampling step we set s to a certain percentile absolute pixel value in xt0 (the
|
||||
prediction of x_0 at timestep t), and if s > 1, then we threshold xt0 to the range [-s, s] and then divide by
|
||||
s. Dynamic thresholding pushes saturated pixels (those near -1 and 1) inwards, thereby actively preventing
|
||||
pixels from saturation at each step. We find that dynamic thresholding results in significantly better
|
||||
photorealism as well as better image-text alignment, especially when using very large guidance weights."
|
||||
|
||||
https://arxiv.org/abs/2205.11487
|
||||
"""
|
||||
dtype = sample.dtype
|
||||
batch_size, channels, *remaining_dims = sample.shape
|
||||
|
||||
if dtype not in (torch.float32, torch.float64):
|
||||
sample = sample.float(
|
||||
) # upcast for quantile calculation, and clamp not implemented for cpu half
|
||||
|
||||
# Flatten sample for doing quantile calculation along each image
|
||||
sample = sample.reshape(batch_size, channels * np.prod(remaining_dims))
|
||||
|
||||
abs_sample = sample.abs() # "a certain percentile absolute pixel value"
|
||||
|
||||
s = torch.quantile(
|
||||
abs_sample, self.config.dynamic_thresholding_ratio, dim=1)
|
||||
s = torch.clamp(
|
||||
s, min=1, max=self.config.sample_max_value
|
||||
) # When clamped to min=1, equivalent to standard clipping to [-1, 1]
|
||||
s = s.unsqueeze(
|
||||
1) # (batch_size, 1) because clamp will broadcast along dim=0
|
||||
sample = torch.clamp(
|
||||
sample, -s, s
|
||||
) / s # "we threshold xt0 to the range [-s, s] and then divide by s"
|
||||
|
||||
sample = sample.reshape(batch_size, channels, *remaining_dims)
|
||||
sample = sample.to(dtype)
|
||||
|
||||
return sample
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler._sigma_to_t
|
||||
def _sigma_to_t(self, sigma):
|
||||
return sigma * self.config.num_train_timesteps
|
||||
|
||||
def _sigma_to_alpha_sigma_t(self, sigma):
|
||||
return 1 - sigma, sigma
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.set_timesteps
|
||||
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma)
|
||||
|
||||
def convert_model_output(
|
||||
self,
|
||||
model_output: torch.Tensor,
|
||||
*args,
|
||||
sample: torch.Tensor = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Convert the model output to the corresponding type the UniPC algorithm needs.
|
||||
|
||||
Args:
|
||||
model_output (`torch.Tensor`):
|
||||
The direct output from the learned diffusion model.
|
||||
timestep (`int`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.Tensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
The converted model output.
|
||||
"""
|
||||
timestep = args[0] if len(args) > 0 else kwargs.pop("timestep", None)
|
||||
if sample is None:
|
||||
if len(args) > 1:
|
||||
sample = args[1]
|
||||
else:
|
||||
raise ValueError(
|
||||
"missing `sample` as a required keyward argument")
|
||||
if timestep is not None:
|
||||
deprecate(
|
||||
"timesteps",
|
||||
"1.0.0",
|
||||
"Passing `timesteps` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`",
|
||||
)
|
||||
|
||||
sigma = self.sigmas[self.step_index]
|
||||
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
|
||||
|
||||
if self.predict_x0:
|
||||
if self.config.prediction_type == "flow_prediction":
|
||||
sigma_t = self.sigmas[self.step_index]
|
||||
x0_pred = sample - sigma_t * model_output
|
||||
else:
|
||||
raise ValueError(
|
||||
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`,"
|
||||
" `v_prediction` or `flow_prediction` for the UniPCMultistepScheduler."
|
||||
)
|
||||
|
||||
if self.config.thresholding:
|
||||
x0_pred = self._threshold_sample(x0_pred)
|
||||
|
||||
return x0_pred
|
||||
else:
|
||||
if self.config.prediction_type == "flow_prediction":
|
||||
sigma_t = self.sigmas[self.step_index]
|
||||
epsilon = sample - (1 - sigma_t) * model_output
|
||||
else:
|
||||
raise ValueError(
|
||||
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`,"
|
||||
" `v_prediction` or `flow_prediction` for the UniPCMultistepScheduler."
|
||||
)
|
||||
|
||||
if self.config.thresholding:
|
||||
sigma_t = self.sigmas[self.step_index]
|
||||
x0_pred = sample - sigma_t * model_output
|
||||
x0_pred = self._threshold_sample(x0_pred)
|
||||
epsilon = model_output + x0_pred
|
||||
|
||||
return epsilon
|
||||
|
||||
def multistep_uni_p_bh_update(
|
||||
self,
|
||||
model_output: torch.Tensor,
|
||||
*args,
|
||||
sample: torch.Tensor = None,
|
||||
order: int = None, # pyright: ignore
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
One step for the UniP (B(h) version). Alternatively, `self.solver_p` is used if is specified.
|
||||
|
||||
Args:
|
||||
model_output (`torch.Tensor`):
|
||||
The direct output from the learned diffusion model at the current timestep.
|
||||
prev_timestep (`int`):
|
||||
The previous discrete timestep in the diffusion chain.
|
||||
sample (`torch.Tensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
order (`int`):
|
||||
The order of UniP at this timestep (corresponds to the *p* in UniPC-p).
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
The sample tensor at the previous timestep.
|
||||
"""
|
||||
prev_timestep = args[0] if len(args) > 0 else kwargs.pop(
|
||||
"prev_timestep", None)
|
||||
if sample is None:
|
||||
if len(args) > 1:
|
||||
sample = args[1]
|
||||
else:
|
||||
raise ValueError(
|
||||
" missing `sample` as a required keyward argument")
|
||||
if order is None:
|
||||
if len(args) > 2:
|
||||
order = args[2]
|
||||
else:
|
||||
raise ValueError(
|
||||
" missing `order` as a required keyward argument")
|
||||
if prev_timestep is not None:
|
||||
deprecate(
|
||||
"prev_timestep",
|
||||
"1.0.0",
|
||||
"Passing `prev_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`",
|
||||
)
|
||||
model_output_list = self.model_outputs
|
||||
|
||||
s0 = self.timestep_list[-1]
|
||||
m0 = model_output_list[-1]
|
||||
x = sample
|
||||
|
||||
if self.solver_p:
|
||||
x_t = self.solver_p.step(model_output, s0, x).prev_sample
|
||||
return x_t
|
||||
|
||||
sigma_t, sigma_s0 = self.sigmas[self.step_index + 1], self.sigmas[
|
||||
self.step_index] # pyright: ignore
|
||||
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
|
||||
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
|
||||
|
||||
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
|
||||
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
|
||||
|
||||
h = lambda_t - lambda_s0
|
||||
device = sample.device
|
||||
|
||||
rks = []
|
||||
D1s = []
|
||||
for i in range(1, order):
|
||||
si = self.step_index - i # pyright: ignore
|
||||
mi = model_output_list[-(i + 1)]
|
||||
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
|
||||
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
|
||||
rk = (lambda_si - lambda_s0) / h
|
||||
rks.append(rk)
|
||||
D1s.append((mi - m0) / rk) # pyright: ignore
|
||||
|
||||
rks.append(1.0)
|
||||
rks = torch.tensor(rks, device=device)
|
||||
|
||||
R = []
|
||||
b = []
|
||||
|
||||
hh = -h if self.predict_x0 else h
|
||||
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
|
||||
h_phi_k = h_phi_1 / hh - 1
|
||||
|
||||
factorial_i = 1
|
||||
|
||||
if self.config.solver_type == "bh1":
|
||||
B_h = hh
|
||||
elif self.config.solver_type == "bh2":
|
||||
B_h = torch.expm1(hh)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
for i in range(1, order + 1):
|
||||
R.append(torch.pow(rks, i - 1))
|
||||
b.append(h_phi_k * factorial_i / B_h)
|
||||
factorial_i *= i + 1
|
||||
h_phi_k = h_phi_k / hh - 1 / factorial_i
|
||||
|
||||
R = torch.stack(R)
|
||||
b = torch.tensor(b, device=device)
|
||||
|
||||
if len(D1s) > 0:
|
||||
D1s = torch.stack(D1s, dim=1) # (B, K)
|
||||
# for order 2, we use a simplified version
|
||||
if order == 2:
|
||||
rhos_p = torch.tensor([0.5], dtype=x.dtype, device=device)
|
||||
else:
|
||||
rhos_p = torch.linalg.solve(R[:-1, :-1],
|
||||
b[:-1]).to(device).to(x.dtype)
|
||||
else:
|
||||
D1s = None
|
||||
|
||||
if self.predict_x0:
|
||||
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
|
||||
if D1s is not None:
|
||||
pred_res = torch.einsum("k,bkc...->bc...", rhos_p,
|
||||
D1s) # pyright: ignore
|
||||
else:
|
||||
pred_res = 0
|
||||
x_t = x_t_ - alpha_t * B_h * pred_res
|
||||
else:
|
||||
x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0
|
||||
if D1s is not None:
|
||||
pred_res = torch.einsum("k,bkc...->bc...", rhos_p,
|
||||
D1s) # pyright: ignore
|
||||
else:
|
||||
pred_res = 0
|
||||
x_t = x_t_ - sigma_t * B_h * pred_res
|
||||
|
||||
x_t = x_t.to(x.dtype)
|
||||
return x_t
|
||||
|
||||
def multistep_uni_c_bh_update(
|
||||
self,
|
||||
this_model_output: torch.Tensor,
|
||||
*args,
|
||||
last_sample: torch.Tensor = None,
|
||||
this_sample: torch.Tensor = None,
|
||||
order: int = None, # pyright: ignore
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
One step for the UniC (B(h) version).
|
||||
|
||||
Args:
|
||||
this_model_output (`torch.Tensor`):
|
||||
The model outputs at `x_t`.
|
||||
this_timestep (`int`):
|
||||
The current timestep `t`.
|
||||
last_sample (`torch.Tensor`):
|
||||
The generated sample before the last predictor `x_{t-1}`.
|
||||
this_sample (`torch.Tensor`):
|
||||
The generated sample after the last predictor `x_{t}`.
|
||||
order (`int`):
|
||||
The `p` of UniC-p at this step. The effective order of accuracy should be `order + 1`.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
The corrected sample tensor at the current timestep.
|
||||
"""
|
||||
this_timestep = args[0] if len(args) > 0 else kwargs.pop(
|
||||
"this_timestep", None)
|
||||
if last_sample is None:
|
||||
if len(args) > 1:
|
||||
last_sample = args[1]
|
||||
else:
|
||||
raise ValueError(
|
||||
" missing`last_sample` as a required keyward argument")
|
||||
if this_sample is None:
|
||||
if len(args) > 2:
|
||||
this_sample = args[2]
|
||||
else:
|
||||
raise ValueError(
|
||||
" missing`this_sample` as a required keyward argument")
|
||||
if order is None:
|
||||
if len(args) > 3:
|
||||
order = args[3]
|
||||
else:
|
||||
raise ValueError(
|
||||
" missing`order` as a required keyward argument")
|
||||
if this_timestep is not None:
|
||||
deprecate(
|
||||
"this_timestep",
|
||||
"1.0.0",
|
||||
"Passing `this_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`",
|
||||
)
|
||||
|
||||
model_output_list = self.model_outputs
|
||||
|
||||
m0 = model_output_list[-1]
|
||||
x = last_sample
|
||||
x_t = this_sample
|
||||
model_t = this_model_output
|
||||
|
||||
sigma_t, sigma_s0 = self.sigmas[self.step_index], self.sigmas[
|
||||
self.step_index - 1] # pyright: ignore
|
||||
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
|
||||
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
|
||||
|
||||
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
|
||||
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
|
||||
|
||||
h = lambda_t - lambda_s0
|
||||
device = this_sample.device
|
||||
|
||||
rks = []
|
||||
D1s = []
|
||||
for i in range(1, order):
|
||||
si = self.step_index - (i + 1) # pyright: ignore
|
||||
mi = model_output_list[-(i + 1)]
|
||||
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
|
||||
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
|
||||
rk = (lambda_si - lambda_s0) / h
|
||||
rks.append(rk)
|
||||
D1s.append((mi - m0) / rk) # pyright: ignore
|
||||
|
||||
rks.append(1.0)
|
||||
rks = torch.tensor(rks, device=device)
|
||||
|
||||
R = []
|
||||
b = []
|
||||
|
||||
hh = -h if self.predict_x0 else h
|
||||
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
|
||||
h_phi_k = h_phi_1 / hh - 1
|
||||
|
||||
factorial_i = 1
|
||||
|
||||
if self.config.solver_type == "bh1":
|
||||
B_h = hh
|
||||
elif self.config.solver_type == "bh2":
|
||||
B_h = torch.expm1(hh)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
for i in range(1, order + 1):
|
||||
R.append(torch.pow(rks, i - 1))
|
||||
b.append(h_phi_k * factorial_i / B_h)
|
||||
factorial_i *= i + 1
|
||||
h_phi_k = h_phi_k / hh - 1 / factorial_i
|
||||
|
||||
R = torch.stack(R)
|
||||
b = torch.tensor(b, device=device)
|
||||
|
||||
if len(D1s) > 0:
|
||||
D1s = torch.stack(D1s, dim=1)
|
||||
else:
|
||||
D1s = None
|
||||
|
||||
# for order 1, we use a simplified version
|
||||
if order == 1:
|
||||
rhos_c = torch.tensor([0.5], dtype=x.dtype, device=device)
|
||||
else:
|
||||
rhos_c = torch.linalg.solve(R, b).to(device).to(x.dtype)
|
||||
|
||||
if self.predict_x0:
|
||||
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
|
||||
if D1s is not None:
|
||||
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s)
|
||||
else:
|
||||
corr_res = 0
|
||||
D1_t = model_t - m0
|
||||
x_t = x_t_ - alpha_t * B_h * (corr_res + rhos_c[-1] * D1_t)
|
||||
else:
|
||||
x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0
|
||||
if D1s is not None:
|
||||
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s)
|
||||
else:
|
||||
corr_res = 0
|
||||
D1_t = model_t - m0
|
||||
x_t = x_t_ - sigma_t * B_h * (corr_res + rhos_c[-1] * D1_t)
|
||||
x_t = x_t.to(x.dtype)
|
||||
return x_t
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||
if schedule_timesteps is None:
|
||||
schedule_timesteps = self.timesteps
|
||||
|
||||
indices = (schedule_timesteps == timestep).nonzero()
|
||||
|
||||
# The sigma index that is taken for the **very** first `step`
|
||||
# is always the second index (or the last index if there is only 1)
|
||||
# This way we can ensure we don't accidentally skip a sigma in
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
return indices[pos].item()
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._init_step_index
|
||||
def _init_step_index(self, timestep):
|
||||
"""
|
||||
Initialize the step_index counter for the scheduler.
|
||||
"""
|
||||
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
self._step_index = self.index_for_timestep(timestep)
|
||||
else:
|
||||
self._step_index = self._begin_index
|
||||
|
||||
def step(self,
|
||||
model_output: torch.Tensor,
|
||||
timestep: Union[int, torch.Tensor],
|
||||
sample: torch.Tensor,
|
||||
return_dict: bool = True,
|
||||
generator=None) -> Union[SchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
|
||||
the multistep UniPC.
|
||||
|
||||
Args:
|
||||
model_output (`torch.Tensor`):
|
||||
The direct output from learned diffusion model.
|
||||
timestep (`int`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.Tensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
return_dict (`bool`):
|
||||
Whether or not to return a [`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`.
|
||||
|
||||
Returns:
|
||||
[`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_utils.SchedulerOutput`] is returned, otherwise a
|
||||
tuple is returned where the first element is the sample tensor.
|
||||
|
||||
"""
|
||||
if self.num_inference_steps is None:
|
||||
raise ValueError(
|
||||
"Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler"
|
||||
)
|
||||
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
|
||||
use_corrector = (
|
||||
self.step_index > 0 and
|
||||
self.step_index - 1 not in self.disable_corrector and
|
||||
self.last_sample is not None # pyright: ignore
|
||||
)
|
||||
|
||||
model_output_convert = self.convert_model_output(
|
||||
model_output, sample=sample)
|
||||
if use_corrector:
|
||||
sample = self.multistep_uni_c_bh_update(
|
||||
this_model_output=model_output_convert,
|
||||
last_sample=self.last_sample,
|
||||
this_sample=sample,
|
||||
order=self.this_order,
|
||||
)
|
||||
|
||||
for i in range(self.config.solver_order - 1):
|
||||
self.model_outputs[i] = self.model_outputs[i + 1]
|
||||
self.timestep_list[i] = self.timestep_list[i + 1]
|
||||
|
||||
self.model_outputs[-1] = model_output_convert
|
||||
self.timestep_list[-1] = timestep # pyright: ignore
|
||||
|
||||
if self.config.lower_order_final:
|
||||
this_order = min(self.config.solver_order,
|
||||
len(self.timesteps) -
|
||||
self.step_index) # pyright: ignore
|
||||
else:
|
||||
this_order = self.config.solver_order
|
||||
|
||||
self.this_order = min(this_order,
|
||||
self.lower_order_nums + 1) # warmup for multistep
|
||||
assert self.this_order > 0
|
||||
|
||||
self.last_sample = sample
|
||||
prev_sample = self.multistep_uni_p_bh_update(
|
||||
model_output=model_output, # pass the original non-converted model output, in case solver-p is used
|
||||
sample=sample,
|
||||
order=self.this_order,
|
||||
)
|
||||
|
||||
if self.lower_order_nums < self.config.solver_order:
|
||||
self.lower_order_nums += 1
|
||||
|
||||
# upon completion increase step index by one
|
||||
self._step_index += 1 # pyright: ignore
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample,)
|
||||
|
||||
return SchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
def scale_model_input(self, sample: torch.Tensor, *args,
|
||||
**kwargs) -> torch.Tensor:
|
||||
"""
|
||||
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
|
||||
current timestep.
|
||||
|
||||
Args:
|
||||
sample (`torch.Tensor`):
|
||||
The input sample.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
return sample
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.add_noise
|
||||
def add_noise(
|
||||
self,
|
||||
original_samples: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
timesteps: torch.IntTensor,
|
||||
) -> torch.Tensor:
|
||||
# Make sure sigmas and timesteps have the same device and dtype as original_samples
|
||||
sigmas = self.sigmas.to(
|
||||
device=original_samples.device, dtype=original_samples.dtype)
|
||||
if original_samples.device.type == "mps" and torch.is_floating_point(
|
||||
timesteps):
|
||||
# mps does not support float64
|
||||
schedule_timesteps = self.timesteps.to(
|
||||
original_samples.device, dtype=torch.float32)
|
||||
timesteps = timesteps.to(
|
||||
original_samples.device, dtype=torch.float32)
|
||||
else:
|
||||
schedule_timesteps = self.timesteps.to(original_samples.device)
|
||||
timesteps = timesteps.to(original_samples.device)
|
||||
|
||||
# begin_index is None when the scheduler is used for training or pipeline does not implement set_begin_index
|
||||
if self.begin_index is None:
|
||||
step_indices = [
|
||||
self.index_for_timestep(t, schedule_timesteps)
|
||||
for t in timesteps
|
||||
]
|
||||
elif self.step_index is not None:
|
||||
# add_noise is called after first denoising step (for inpainting)
|
||||
step_indices = [self.step_index] * timesteps.shape[0]
|
||||
else:
|
||||
# add noise is called before first denoising step to create initial latent(img2img)
|
||||
step_indices = [self.begin_index] * timesteps.shape[0]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < len(original_samples.shape):
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
|
||||
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
|
||||
noisy_samples = alpha_t * original_samples + sigma_t * noise
|
||||
return noisy_samples
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
363
extensions_built_in/diffusion_models/omnigen2/__init__.py
Normal file
363
extensions_built_in/diffusion_models/omnigen2/__init__.py
Normal file
@@ -0,0 +1,363 @@
|
||||
import os
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
import torch
|
||||
import yaml
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||
from toolkit.models.base_model import BaseModel
|
||||
from diffusers import AutoencoderKL
|
||||
from toolkit.basic import flush
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.samplers.custom_flowmatch_sampler import (
|
||||
CustomFlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from toolkit.accelerator import unwrap_model
|
||||
from optimum.quanto import freeze
|
||||
from toolkit.util.quantize import quantize, get_qtype
|
||||
from .src.pipelines.omnigen2.pipeline_omnigen2 import OmniGen2Pipeline
|
||||
from .src.models.transformers import OmniGen2Transformer2DModel
|
||||
from .src.models.transformers.repo import OmniGen2RotaryPosEmbed
|
||||
from .src.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler as OmniFlowMatchEuler,
|
||||
)
|
||||
from PIL import Image
|
||||
from transformers import (
|
||||
CLIPProcessor,
|
||||
Qwen2_5_VLForConditionalGeneration,
|
||||
)
|
||||
import torch.nn.functional as F
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
|
||||
scheduler_config = {"num_train_timesteps": 1000}
|
||||
|
||||
BASE_MODEL_PATH = "OmniGen2/OmniGen2"
|
||||
|
||||
|
||||
class OmniGen2Model(BaseModel):
|
||||
arch = "omnigen2"
|
||||
|
||||
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 = ["OmniGen2Transformer2DModel"]
|
||||
self._control_latent = None
|
||||
|
||||
# static method to get the noise scheduler
|
||||
@staticmethod
|
||||
def get_train_scheduler():
|
||||
return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
|
||||
|
||||
def get_bucket_divisibility(self):
|
||||
return 16
|
||||
|
||||
def load_model(self):
|
||||
dtype = self.torch_dtype
|
||||
# HiDream-ai/HiDream-I1-Full
|
||||
self.print_and_status_update("Loading OmniGen2 model")
|
||||
# will be updated if we detect a existing checkpoint in training folder
|
||||
model_path = self.model_config.name_or_path
|
||||
extras_path = self.model_config.extras_name_or_path
|
||||
|
||||
scheduler = OmniGen2Model.get_train_scheduler()
|
||||
|
||||
self.print_and_status_update("Loading Qwen2.5 VL")
|
||||
processor = CLIPProcessor.from_pretrained(
|
||||
extras_path, subfolder="processor", use_fast=True
|
||||
)
|
||||
|
||||
mllm = Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
||||
extras_path, subfolder="mllm", torch_dtype=torch.bfloat16
|
||||
)
|
||||
mllm.to(self.device_torch, dtype=dtype)
|
||||
if self.model_config.quantize_te:
|
||||
self.print_and_status_update("Quantizing Qwen2.5 VL model")
|
||||
quantization_type = get_qtype(self.model_config.qtype_te)
|
||||
quantize(mllm, weights=quantization_type)
|
||||
freeze(mllm)
|
||||
|
||||
if self.low_vram:
|
||||
# unload it for now
|
||||
mllm.to("cpu")
|
||||
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Loading transformer")
|
||||
|
||||
transformer = OmniGen2Transformer2DModel.from_pretrained(
|
||||
model_path, subfolder="transformer", torch_dtype=torch.bfloat16
|
||||
)
|
||||
|
||||
if not self.low_vram:
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
|
||||
if self.model_config.quantize:
|
||||
self.print_and_status_update("Quantizing transformer")
|
||||
quantization_type = get_qtype(self.model_config.qtype)
|
||||
quantize(transformer, weights=quantization_type)
|
||||
freeze(transformer)
|
||||
|
||||
if self.low_vram:
|
||||
# unload it for now
|
||||
transformer.to("cpu")
|
||||
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Loading vae")
|
||||
|
||||
vae = AutoencoderKL.from_pretrained(
|
||||
extras_path, subfolder="vae", torch_dtype=torch.bfloat16
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
flush()
|
||||
self.print_and_status_update("Loading Qwen2.5 VLProcessor")
|
||||
|
||||
flush()
|
||||
|
||||
if self.low_vram:
|
||||
self.print_and_status_update("Moving everything to device")
|
||||
# move it all back
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
vae.to(self.device_torch, dtype=dtype)
|
||||
mllm.to(self.device_torch, dtype=dtype)
|
||||
|
||||
# set to eval mode
|
||||
# transformer.eval()
|
||||
vae.eval()
|
||||
mllm.eval()
|
||||
mllm.requires_grad_(False)
|
||||
|
||||
pipe: OmniGen2Pipeline = OmniGen2Pipeline(
|
||||
transformer=transformer,
|
||||
vae=vae,
|
||||
scheduler=scheduler,
|
||||
mllm=mllm,
|
||||
processor=processor,
|
||||
)
|
||||
|
||||
flush()
|
||||
|
||||
text_encoder_list = [mllm]
|
||||
tokenizer_list = [processor]
|
||||
|
||||
flush()
|
||||
|
||||
# save it to the model class
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder_list # list of text encoders
|
||||
self.tokenizer = tokenizer_list # list of tokenizers
|
||||
self.model = pipe.transformer
|
||||
self.pipeline = pipe
|
||||
|
||||
self.freqs_cis = OmniGen2RotaryPosEmbed.get_freqs_cis(
|
||||
transformer.config.axes_dim_rope,
|
||||
transformer.config.axes_lens,
|
||||
theta=10000,
|
||||
)
|
||||
|
||||
self.print_and_status_update("Model Loaded")
|
||||
|
||||
def get_generation_pipeline(self):
|
||||
scheduler = OmniFlowMatchEuler(
|
||||
dynamic_time_shift=True, num_train_timesteps=1000
|
||||
)
|
||||
|
||||
pipeline: OmniGen2Pipeline = OmniGen2Pipeline(
|
||||
transformer=self.model,
|
||||
vae=self.vae,
|
||||
scheduler=scheduler,
|
||||
mllm=self.text_encoder[0],
|
||||
processor=self.tokenizer[0],
|
||||
)
|
||||
|
||||
pipeline = pipeline.to(self.device_torch)
|
||||
|
||||
return pipeline
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: OmniGen2Pipeline,
|
||||
gen_config: GenerateImageConfig,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
unconditional_embeds: PromptEmbeds,
|
||||
generator: torch.Generator,
|
||||
extra: dict,
|
||||
):
|
||||
input_images = []
|
||||
if gen_config.ctrl_img is not None:
|
||||
control_img = Image.open(gen_config.ctrl_img)
|
||||
control_img = control_img.convert("RGB")
|
||||
# resize to width and height
|
||||
if control_img.size != (gen_config.width, gen_config.height):
|
||||
control_img = control_img.resize(
|
||||
(gen_config.width, gen_config.height), Image.BILINEAR
|
||||
)
|
||||
input_images = [control_img]
|
||||
|
||||
img = pipeline(
|
||||
prompt_embeds=conditional_embeds.text_embeds,
|
||||
prompt_attention_mask=conditional_embeds.attention_mask,
|
||||
negative_prompt_embeds=unconditional_embeds.text_embeds,
|
||||
negative_prompt_attention_mask=unconditional_embeds.attention_mask,
|
||||
height=gen_config.height,
|
||||
width=gen_config.width,
|
||||
num_inference_steps=gen_config.num_inference_steps,
|
||||
text_guidance_scale=gen_config.guidance_scale,
|
||||
image_guidance_scale=1.0, # reference image guidance scale. Add this for controls
|
||||
latents=gen_config.latents,
|
||||
align_res=False,
|
||||
generator=generator,
|
||||
input_images=input_images,
|
||||
**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,
|
||||
):
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
try:
|
||||
timestep = timestep.expand(latent_model_input.shape[0]).to(
|
||||
latent_model_input.dtype
|
||||
)
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
timesteps = timestep / 1000 # convert to 0 to 1 scale
|
||||
# timestep for model starts at 0 instead of 1. So we need to reverse them
|
||||
timestep = 1 - timesteps
|
||||
model_pred = self.model(
|
||||
latent_model_input,
|
||||
timestep,
|
||||
text_embeddings.text_embeds,
|
||||
self.freqs_cis,
|
||||
text_embeddings.attention_mask,
|
||||
ref_image_hidden_states=self._control_latent,
|
||||
)
|
||||
|
||||
return model_pred
|
||||
|
||||
def condition_noisy_latents(
|
||||
self, latents: torch.Tensor, batch: "DataLoaderBatchDTO"
|
||||
):
|
||||
# reset the control latent
|
||||
self._control_latent = None
|
||||
with torch.no_grad():
|
||||
control_tensor = batch.control_tensor
|
||||
if control_tensor is not None:
|
||||
self.vae.to(self.device_torch)
|
||||
# we are not packed here, so we just need to pass them so we can pack them later
|
||||
control_tensor = control_tensor * 2 - 1
|
||||
control_tensor = control_tensor.to(
|
||||
self.vae_device_torch, dtype=self.torch_dtype
|
||||
)
|
||||
|
||||
# if it is not the size of batch.tensor, (bs,ch,h,w) then we need to resize it
|
||||
# todo, we may not need to do this, check
|
||||
if batch.tensor is not None:
|
||||
target_h, target_w = batch.tensor.shape[2], batch.tensor.shape[3]
|
||||
else:
|
||||
# When caching latents, batch.tensor is None. We get the size from the file_items instead.
|
||||
target_h = batch.file_items[0].crop_height
|
||||
target_w = batch.file_items[0].crop_width
|
||||
|
||||
if (
|
||||
control_tensor.shape[2] != target_h
|
||||
or control_tensor.shape[3] != target_w
|
||||
):
|
||||
control_tensor = F.interpolate(
|
||||
control_tensor, size=(target_h, target_w), mode="bilinear"
|
||||
)
|
||||
|
||||
control_latent = self.encode_images(control_tensor).to(
|
||||
latents.device, latents.dtype
|
||||
)
|
||||
self._control_latent = [
|
||||
[x.squeeze(0)]
|
||||
for x in torch.chunk(control_latent, control_latent.shape[0], dim=0)
|
||||
]
|
||||
|
||||
return latents.detach()
|
||||
|
||||
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
prompt = [self.pipeline._apply_chat_template(_prompt) for _prompt in prompt]
|
||||
self.text_encoder_to(self.device_torch, dtype=self.torch_dtype)
|
||||
max_sequence_length = 256
|
||||
prompt_embeds, prompt_attention_mask, _, _ = self.pipeline.encode_prompt(
|
||||
prompt=prompt,
|
||||
do_classifier_free_guidance=False,
|
||||
device=self.device_torch,
|
||||
max_sequence_length=max_sequence_length,
|
||||
)
|
||||
pe = PromptEmbeds(prompt_embeds)
|
||||
pe.attention_mask = prompt_attention_mask
|
||||
return pe
|
||||
|
||||
def get_model_has_grad(self):
|
||||
# return from a weight if it has grad
|
||||
return False
|
||||
|
||||
def get_te_has_grad(self):
|
||||
# assume no one wants to finetune 4 text encoders.
|
||||
return False
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
# only save the transformer
|
||||
transformer: OmniGen2Transformer2DModel = 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")
|
||||
batch = kwargs.get("batch")
|
||||
# return (noise - batch.latents).detach()
|
||||
return (batch.latents - noise).detach()
|
||||
|
||||
def get_transformer_block_names(self) -> Optional[List[str]]:
|
||||
# omnigen2 had a few blocks for things like noise_refiner, ref_image_refiner, context_refiner, and layers.
|
||||
# lets do all but image refiner until we add it
|
||||
if self.model_config.model_kwargs.get("use_image_refiner", False):
|
||||
return ["noise_refiner", "context_refiner", "ref_image_refiner", "layers"]
|
||||
return ["noise_refiner", "context_refiner", "layers"]
|
||||
|
||||
def convert_lora_weights_before_save(self, state_dict):
|
||||
# currently starte with transformer. but needs to start with diffusion_model. for comfyui
|
||||
new_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
new_key = key.replace("transformer.", "diffusion_model.")
|
||||
new_sd[new_key] = value
|
||||
return new_sd
|
||||
|
||||
def convert_lora_weights_before_load(self, state_dict):
|
||||
# saved as diffusion_model. but needs to be transformer. for ai-toolkit
|
||||
new_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
new_key = key.replace("diffusion_model.", "transformer.")
|
||||
new_sd[new_key] = value
|
||||
return new_sd
|
||||
|
||||
def get_base_model_version(self):
|
||||
return "omnigen2"
|
||||
@@ -0,0 +1,357 @@
|
||||
"""
|
||||
OmniGen2 Attention Processor Module
|
||||
|
||||
Copyright 2025 BAAI, The OmniGen2 Team and The HuggingFace Team. All rights reserved.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
|
||||
import warnings
|
||||
import math
|
||||
from typing import Optional, Tuple, Dict, Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import repeat
|
||||
|
||||
from ..utils.import_utils import is_flash_attn_available
|
||||
|
||||
if is_flash_attn_available():
|
||||
from flash_attn import flash_attn_varlen_func
|
||||
from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input
|
||||
else:
|
||||
warnings.warn("Cannot import flash_attn, install flash_attn to use Flash2Varlen attention for better performance")
|
||||
|
||||
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from .embeddings import apply_rotary_emb
|
||||
|
||||
|
||||
class OmniGen2AttnProcessorFlash2Varlen:
|
||||
"""
|
||||
Processor for implementing scaled dot-product attention with flash attention and variable length sequences.
|
||||
|
||||
This processor implements:
|
||||
- Flash attention with variable length sequences
|
||||
- Rotary position embeddings (RoPE)
|
||||
- Query-Key normalization
|
||||
- Proportional attention scaling
|
||||
|
||||
Args:
|
||||
None
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize the attention processor."""
|
||||
if not is_flash_attn_available():
|
||||
raise ImportError(
|
||||
"OmniGen2AttnProcessorFlash2Varlen requires flash_attn. "
|
||||
"Please install flash_attn."
|
||||
)
|
||||
|
||||
def _upad_input(
|
||||
self,
|
||||
query_layer: torch.Tensor,
|
||||
key_layer: torch.Tensor,
|
||||
value_layer: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
query_length: int,
|
||||
num_heads: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, Tuple[torch.Tensor, torch.Tensor], Tuple[int, int]]:
|
||||
"""
|
||||
Unpad the input tensors for flash attention.
|
||||
|
||||
Args:
|
||||
query_layer: Query tensor of shape (batch_size, seq_len, num_heads, head_dim)
|
||||
key_layer: Key tensor of shape (batch_size, seq_len, num_kv_heads, head_dim)
|
||||
value_layer: Value tensor of shape (batch_size, seq_len, num_kv_heads, head_dim)
|
||||
attention_mask: Attention mask tensor of shape (batch_size, seq_len)
|
||||
query_length: Length of the query sequence
|
||||
num_heads: Number of attention heads
|
||||
|
||||
Returns:
|
||||
Tuple containing:
|
||||
- Unpadded query tensor
|
||||
- Unpadded key tensor
|
||||
- Unpadded value tensor
|
||||
- Query indices
|
||||
- Tuple of cumulative sequence lengths for query and key
|
||||
- Tuple of maximum sequence lengths for query and key
|
||||
"""
|
||||
def _get_unpad_data(attention_mask: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, int]:
|
||||
"""Helper function to get unpadding data from attention mask."""
|
||||
seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)
|
||||
indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()
|
||||
max_seqlen_in_batch = seqlens_in_batch.max().item()
|
||||
cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))
|
||||
return indices, cu_seqlens, max_seqlen_in_batch
|
||||
|
||||
indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)
|
||||
batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape
|
||||
|
||||
# Unpad key and value layers
|
||||
key_layer = index_first_axis(
|
||||
key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim),
|
||||
indices_k,
|
||||
)
|
||||
value_layer = index_first_axis(
|
||||
value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim),
|
||||
indices_k,
|
||||
)
|
||||
|
||||
# Handle different query length cases
|
||||
if query_length == kv_seq_len:
|
||||
query_layer = index_first_axis(
|
||||
query_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim),
|
||||
indices_k,
|
||||
)
|
||||
cu_seqlens_q = cu_seqlens_k
|
||||
max_seqlen_in_batch_q = max_seqlen_in_batch_k
|
||||
indices_q = indices_k
|
||||
elif query_length == 1:
|
||||
max_seqlen_in_batch_q = 1
|
||||
cu_seqlens_q = torch.arange(
|
||||
batch_size + 1, dtype=torch.int32, device=query_layer.device
|
||||
)
|
||||
indices_q = cu_seqlens_q[:-1]
|
||||
query_layer = query_layer.squeeze(1)
|
||||
else:
|
||||
attention_mask = attention_mask[:, -query_length:]
|
||||
query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)
|
||||
|
||||
return (
|
||||
query_layer,
|
||||
key_layer,
|
||||
value_layer,
|
||||
indices_q,
|
||||
(cu_seqlens_q, cu_seqlens_k),
|
||||
(max_seqlen_in_batch_q, max_seqlen_in_batch_k),
|
||||
)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
base_sequence_length: Optional[int] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Process attention computation with flash attention.
|
||||
|
||||
Args:
|
||||
attn: Attention module
|
||||
hidden_states: Hidden states tensor of shape (batch_size, seq_len, hidden_dim)
|
||||
encoder_hidden_states: Encoder hidden states tensor
|
||||
attention_mask: Optional attention mask tensor
|
||||
image_rotary_emb: Optional rotary embeddings for image tokens
|
||||
base_sequence_length: Optional base sequence length for proportional attention
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Processed hidden states after attention computation
|
||||
"""
|
||||
batch_size, sequence_length, _ = hidden_states.shape
|
||||
|
||||
# Get Query-Key-Value Pair
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
query_dim = query.shape[-1]
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = query_dim // attn.heads
|
||||
dtype = query.dtype
|
||||
|
||||
# Get key-value heads
|
||||
kv_heads = inner_dim // head_dim
|
||||
|
||||
# Reshape tensors for attention computation
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim)
|
||||
key = key.view(batch_size, -1, kv_heads, head_dim)
|
||||
value = value.view(batch_size, -1, kv_heads, head_dim)
|
||||
|
||||
# Apply Query-Key normalization
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
# Apply Rotary Position Embeddings
|
||||
if image_rotary_emb is not None:
|
||||
query = apply_rotary_emb(query, image_rotary_emb, use_real=False)
|
||||
key = apply_rotary_emb(key, image_rotary_emb, use_real=False)
|
||||
|
||||
query, key = query.to(dtype), key.to(dtype)
|
||||
|
||||
# Calculate attention scale
|
||||
if base_sequence_length is not None:
|
||||
softmax_scale = math.sqrt(math.log(sequence_length, base_sequence_length)) * attn.scale
|
||||
else:
|
||||
softmax_scale = attn.scale
|
||||
|
||||
# Unpad input for flash attention
|
||||
(
|
||||
query_states,
|
||||
key_states,
|
||||
value_states,
|
||||
indices_q,
|
||||
cu_seq_lens,
|
||||
max_seq_lens,
|
||||
) = self._upad_input(query, key, value, attention_mask, sequence_length, attn.heads)
|
||||
|
||||
cu_seqlens_q, cu_seqlens_k = cu_seq_lens
|
||||
max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens
|
||||
|
||||
# Handle different number of heads
|
||||
if kv_heads < attn.heads:
|
||||
key_states = repeat(key_states, "l h c -> l (h k) c", k=attn.heads // kv_heads)
|
||||
value_states = repeat(value_states, "l h c -> l (h k) c", k=attn.heads // kv_heads)
|
||||
|
||||
# Apply flash attention
|
||||
attn_output_unpad = flash_attn_varlen_func(
|
||||
query_states,
|
||||
key_states,
|
||||
value_states,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_in_batch_q,
|
||||
max_seqlen_k=max_seqlen_in_batch_k,
|
||||
dropout_p=0.0,
|
||||
causal=False,
|
||||
softmax_scale=softmax_scale,
|
||||
)
|
||||
|
||||
# Pad output and apply final transformations
|
||||
hidden_states = pad_input(attn_output_unpad, indices_q, batch_size, sequence_length)
|
||||
hidden_states = hidden_states.flatten(-2)
|
||||
hidden_states = hidden_states.type_as(query)
|
||||
|
||||
# Apply output projection
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class OmniGen2AttnProcessor:
|
||||
"""
|
||||
Processor for implementing scaled dot-product attention with flash attention and variable length sequences.
|
||||
|
||||
This processor is optimized for PyTorch 2.0 and implements:
|
||||
- Flash attention with variable length sequences
|
||||
- Rotary position embeddings (RoPE)
|
||||
- Query-Key normalization
|
||||
- Proportional attention scaling
|
||||
|
||||
Args:
|
||||
None
|
||||
|
||||
Raises:
|
||||
ImportError: If PyTorch version is less than 2.0
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize the attention processor."""
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError(
|
||||
"OmniGen2AttnProcessorFlash2Varlen requires PyTorch 2.0. "
|
||||
"Please upgrade PyTorch to version 2.0 or later."
|
||||
)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
base_sequence_length: Optional[int] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Process attention computation with flash attention.
|
||||
|
||||
Args:
|
||||
attn: Attention module
|
||||
hidden_states: Hidden states tensor of shape (batch_size, seq_len, hidden_dim)
|
||||
encoder_hidden_states: Encoder hidden states tensor
|
||||
attention_mask: Optional attention mask tensor
|
||||
image_rotary_emb: Optional rotary embeddings for image tokens
|
||||
base_sequence_length: Optional base sequence length for proportional attention
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Processed hidden states after attention computation
|
||||
"""
|
||||
batch_size, sequence_length, _ = hidden_states.shape
|
||||
|
||||
# Get Query-Key-Value Pair
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
query_dim = query.shape[-1]
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = query_dim // attn.heads
|
||||
dtype = query.dtype
|
||||
|
||||
# Get key-value heads
|
||||
kv_heads = inner_dim // head_dim
|
||||
|
||||
# Reshape tensors for attention computation
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim)
|
||||
key = key.view(batch_size, -1, kv_heads, head_dim)
|
||||
value = value.view(batch_size, -1, kv_heads, head_dim)
|
||||
|
||||
# Apply Query-Key normalization
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
# Apply Rotary Position Embeddings
|
||||
if image_rotary_emb is not None:
|
||||
query = apply_rotary_emb(query, image_rotary_emb, use_real=False)
|
||||
key = apply_rotary_emb(key, image_rotary_emb, use_real=False)
|
||||
|
||||
query, key = query.to(dtype), key.to(dtype)
|
||||
|
||||
# Calculate attention scale
|
||||
if base_sequence_length is not None:
|
||||
softmax_scale = math.sqrt(math.log(sequence_length, base_sequence_length)) * attn.scale
|
||||
else:
|
||||
softmax_scale = attn.scale
|
||||
|
||||
# scaled_dot_product_attention expects attention_mask shape to be
|
||||
# (batch, heads, source_length, target_length)
|
||||
if attention_mask is not None:
|
||||
attention_mask = attention_mask.bool().view(batch_size, 1, 1, -1)
|
||||
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
|
||||
# explicitly repeat key and value to match query length, otherwise using enable_gqa=True results in MATH backend of sdpa in our test of pytorch2.6
|
||||
key = key.repeat_interleave(query.size(-3) // key.size(-3), -3)
|
||||
value = value.repeat_interleave(query.size(-3) // value.size(-3), -3)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, scale=softmax_scale
|
||||
)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.type_as(query)
|
||||
|
||||
# Apply output projection
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,126 @@
|
||||
# Copyright 2024 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
from diffusers.models.activations import get_activation
|
||||
|
||||
|
||||
class TimestepEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
time_embed_dim: int,
|
||||
act_fn: str = "silu",
|
||||
out_dim: int = None,
|
||||
post_act_fn: Optional[str] = None,
|
||||
cond_proj_dim=None,
|
||||
sample_proj_bias=True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.linear_1 = nn.Linear(in_channels, time_embed_dim, sample_proj_bias)
|
||||
|
||||
if cond_proj_dim is not None:
|
||||
self.cond_proj = nn.Linear(cond_proj_dim, in_channels, bias=False)
|
||||
else:
|
||||
self.cond_proj = None
|
||||
|
||||
self.act = get_activation(act_fn)
|
||||
|
||||
if out_dim is not None:
|
||||
time_embed_dim_out = out_dim
|
||||
else:
|
||||
time_embed_dim_out = time_embed_dim
|
||||
self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim_out, sample_proj_bias)
|
||||
|
||||
if post_act_fn is None:
|
||||
self.post_act = None
|
||||
else:
|
||||
self.post_act = get_activation(post_act_fn)
|
||||
|
||||
self.initialize_weights()
|
||||
|
||||
def initialize_weights(self):
|
||||
nn.init.normal_(self.linear_1.weight, std=0.02)
|
||||
nn.init.zeros_(self.linear_1.bias)
|
||||
nn.init.normal_(self.linear_2.weight, std=0.02)
|
||||
nn.init.zeros_(self.linear_2.bias)
|
||||
|
||||
def forward(self, sample, condition=None):
|
||||
if condition is not None:
|
||||
sample = sample + self.cond_proj(condition)
|
||||
sample = self.linear_1(sample)
|
||||
|
||||
if self.act is not None:
|
||||
sample = self.act(sample)
|
||||
|
||||
sample = self.linear_2(sample)
|
||||
|
||||
if self.post_act is not None:
|
||||
sample = self.post_act(sample)
|
||||
return sample
|
||||
|
||||
|
||||
def apply_rotary_emb(
|
||||
x: torch.Tensor,
|
||||
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
|
||||
use_real: bool = True,
|
||||
use_real_unbind_dim: int = -1,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings
|
||||
to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are
|
||||
reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting
|
||||
tensors contain rotary embeddings and are returned as real tensors.
|
||||
|
||||
Args:
|
||||
x (`torch.Tensor`):
|
||||
Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply
|
||||
freqs_cis (`Tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],)
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
|
||||
"""
|
||||
if use_real:
|
||||
cos, sin = freqs_cis # [S, D]
|
||||
cos = cos[None, None]
|
||||
sin = sin[None, None]
|
||||
cos, sin = cos.to(x.device), sin.to(x.device)
|
||||
|
||||
if use_real_unbind_dim == -1:
|
||||
# Used for flux, cogvideox, hunyuan-dit
|
||||
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
|
||||
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
||||
elif use_real_unbind_dim == -2:
|
||||
# Used for Stable Audio, OmniGen and CogView4
|
||||
x_real, x_imag = x.reshape(*x.shape[:-1], 2, -1).unbind(-2) # [B, S, H, D//2]
|
||||
x_rotated = torch.cat([-x_imag, x_real], dim=-1)
|
||||
else:
|
||||
raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.")
|
||||
|
||||
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
|
||||
|
||||
return out
|
||||
else:
|
||||
# used for lumina
|
||||
# x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
|
||||
x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], x.shape[-1] // 2, 2))
|
||||
freqs_cis = freqs_cis.unsqueeze(2)
|
||||
x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3)
|
||||
|
||||
return x_out.type_as(x)
|
||||
@@ -0,0 +1,3 @@
|
||||
from .transformer_omnigen2 import OmniGen2Transformer2DModel
|
||||
|
||||
__all__ = ["OmniGen2Transformer2DModel"]
|
||||
@@ -0,0 +1,218 @@
|
||||
|
||||
# Copyright 2024 Alpha-VLLM Authors and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import warnings
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from diffusers.models.embeddings import Timesteps
|
||||
from ..embeddings import TimestepEmbedding
|
||||
|
||||
from ...utils.import_utils import is_flash_attn_available, is_triton_available
|
||||
|
||||
if is_triton_available():
|
||||
from ...ops.triton.layer_norm import RMSNorm
|
||||
else:
|
||||
from torch.nn import RMSNorm
|
||||
warnings.warn("Cannot import triton, install triton to use fused RMSNorm for better performance")
|
||||
|
||||
if is_flash_attn_available():
|
||||
from flash_attn.ops.activations import swiglu
|
||||
else:
|
||||
from .components import swiglu
|
||||
warnings.warn("Cannot import flash_attn, install flash_attn to use fused SwiGLU for better performance")
|
||||
|
||||
# try:
|
||||
# from flash_attn.ops.activations import swiglu as fused_swiglu
|
||||
# FUSEDSWIGLU_AVALIBLE = True
|
||||
# except ImportError:
|
||||
|
||||
# FUSEDSWIGLU_AVALIBLE = False
|
||||
# warnings.warn("Cannot import apex RMSNorm, switch to vanilla implementation")
|
||||
|
||||
class LuminaRMSNormZero(nn.Module):
|
||||
"""
|
||||
Norm layer adaptive RMS normalization zero.
|
||||
|
||||
Parameters:
|
||||
embedding_dim (`int`): The size of each embedding vector.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
norm_eps: float,
|
||||
norm_elementwise_affine: bool,
|
||||
):
|
||||
super().__init__()
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = nn.Linear(
|
||||
min(embedding_dim, 1024),
|
||||
4 * embedding_dim,
|
||||
bias=True,
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(embedding_dim, eps=norm_eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
emb: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
emb = self.linear(self.silu(emb))
|
||||
scale_msa, gate_msa, scale_mlp, gate_mlp = emb.chunk(4, dim=1)
|
||||
x = self.norm(x) * (1 + scale_msa[:, None])
|
||||
return x, gate_msa, scale_mlp, gate_mlp
|
||||
|
||||
|
||||
class LuminaLayerNormContinuous(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
conditioning_embedding_dim: int,
|
||||
# NOTE: It is a bit weird that the norm layer can be configured to have scale and shift parameters
|
||||
# because the output is immediately scaled and shifted by the projected conditioning embeddings.
|
||||
# Note that AdaLayerNorm does not let the norm layer have scale and shift parameters.
|
||||
# However, this is how it was implemented in the original code, and it's rather likely you should
|
||||
# set `elementwise_affine` to False.
|
||||
elementwise_affine=True,
|
||||
eps=1e-5,
|
||||
bias=True,
|
||||
norm_type="layer_norm",
|
||||
out_dim: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# AdaLN
|
||||
self.silu = nn.SiLU()
|
||||
self.linear_1 = nn.Linear(conditioning_embedding_dim, embedding_dim, bias=bias)
|
||||
|
||||
if norm_type == "layer_norm":
|
||||
self.norm = nn.LayerNorm(embedding_dim, eps, elementwise_affine, bias)
|
||||
elif norm_type == "rms_norm":
|
||||
self.norm = RMSNorm(embedding_dim, eps=eps, elementwise_affine=elementwise_affine)
|
||||
else:
|
||||
raise ValueError(f"unknown norm_type {norm_type}")
|
||||
|
||||
self.linear_2 = None
|
||||
if out_dim is not None:
|
||||
self.linear_2 = nn.Linear(embedding_dim, out_dim, bias=bias)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
conditioning_embedding: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
# convert back to the original dtype in case `conditioning_embedding`` is upcasted to float32 (needed for hunyuanDiT)
|
||||
emb = self.linear_1(self.silu(conditioning_embedding).to(x.dtype))
|
||||
scale = emb
|
||||
x = self.norm(x) * (1 + scale)[:, None, :]
|
||||
|
||||
if self.linear_2 is not None:
|
||||
x = self.linear_2(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class LuminaFeedForward(nn.Module):
|
||||
r"""
|
||||
A feed-forward layer.
|
||||
|
||||
Parameters:
|
||||
hidden_size (`int`):
|
||||
The dimensionality of the hidden layers in the model. This parameter determines the width of the model's
|
||||
hidden representations.
|
||||
intermediate_size (`int`): The intermediate dimension of the feedforward layer.
|
||||
multiple_of (`int`, *optional*): Value to ensure hidden dimension is a multiple
|
||||
of this value.
|
||||
ffn_dim_multiplier (float, *optional*): Custom multiplier for hidden
|
||||
dimension. Defaults to None.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
inner_dim: int,
|
||||
multiple_of: Optional[int] = 256,
|
||||
ffn_dim_multiplier: Optional[float] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.swiglu = swiglu
|
||||
|
||||
# custom hidden_size factor multiplier
|
||||
if ffn_dim_multiplier is not None:
|
||||
inner_dim = int(ffn_dim_multiplier * inner_dim)
|
||||
inner_dim = multiple_of * ((inner_dim + multiple_of - 1) // multiple_of)
|
||||
|
||||
self.linear_1 = nn.Linear(
|
||||
dim,
|
||||
inner_dim,
|
||||
bias=False,
|
||||
)
|
||||
self.linear_2 = nn.Linear(
|
||||
inner_dim,
|
||||
dim,
|
||||
bias=False,
|
||||
)
|
||||
self.linear_3 = nn.Linear(
|
||||
dim,
|
||||
inner_dim,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
h1, h2 = self.linear_1(x), self.linear_3(x)
|
||||
return self.linear_2(self.swiglu(h1, h2))
|
||||
|
||||
|
||||
class Lumina2CombinedTimestepCaptionEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int = 4096,
|
||||
text_feat_dim: int = 2048,
|
||||
frequency_embedding_size: int = 256,
|
||||
norm_eps: float = 1e-5,
|
||||
timestep_scale: float = 1.0,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.time_proj = Timesteps(
|
||||
num_channels=frequency_embedding_size, flip_sin_to_cos=True, downscale_freq_shift=0.0, scale=timestep_scale
|
||||
)
|
||||
|
||||
self.timestep_embedder = TimestepEmbedding(
|
||||
in_channels=frequency_embedding_size, time_embed_dim=min(hidden_size, 1024)
|
||||
)
|
||||
|
||||
self.caption_embedder = nn.Sequential(
|
||||
RMSNorm(text_feat_dim, eps=norm_eps),
|
||||
nn.Linear(text_feat_dim, hidden_size, bias=True),
|
||||
)
|
||||
|
||||
self._initialize_weights()
|
||||
|
||||
def _initialize_weights(self):
|
||||
nn.init.trunc_normal_(self.caption_embedder[1].weight, std=0.02)
|
||||
nn.init.zeros_(self.caption_embedder[1].bias)
|
||||
|
||||
def forward(
|
||||
self, timestep: torch.Tensor, text_hidden_states: torch.Tensor, dtype: torch.dtype
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
timestep_proj = self.time_proj(timestep).to(dtype=dtype)
|
||||
time_embed = self.timestep_embedder(timestep_proj)
|
||||
caption_embed = self.caption_embedder(text_hidden_states)
|
||||
return time_embed, caption_embed
|
||||
@@ -0,0 +1,4 @@
|
||||
import torch.nn.functional as F
|
||||
|
||||
def swiglu(x, y):
|
||||
return F.silu(x.float(), inplace=False).to(x.dtype) * y
|
||||
@@ -0,0 +1,135 @@
|
||||
from typing import List, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from einops import repeat
|
||||
from diffusers.models.embeddings import get_1d_rotary_pos_embed
|
||||
|
||||
class OmniGen2RotaryPosEmbed(nn.Module):
|
||||
def __init__(self, theta: int,
|
||||
axes_dim: Tuple[int, int, int],
|
||||
axes_lens: Tuple[int, int, int] = (300, 512, 512),
|
||||
patch_size: int = 2):
|
||||
super().__init__()
|
||||
self.theta = theta
|
||||
self.axes_dim = axes_dim
|
||||
self.axes_lens = axes_lens
|
||||
self.patch_size = patch_size
|
||||
|
||||
@staticmethod
|
||||
def get_freqs_cis(axes_dim: Tuple[int, int, int],
|
||||
axes_lens: Tuple[int, int, int],
|
||||
theta: int) -> List[torch.Tensor]:
|
||||
freqs_cis = []
|
||||
freqs_dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64
|
||||
for i, (d, e) in enumerate(zip(axes_dim, axes_lens)):
|
||||
emb = get_1d_rotary_pos_embed(d, e, theta=theta, freqs_dtype=freqs_dtype)
|
||||
freqs_cis.append(emb)
|
||||
return freqs_cis
|
||||
|
||||
def _get_freqs_cis(self, freqs_cis, ids: torch.Tensor) -> torch.Tensor:
|
||||
device = ids.device
|
||||
if ids.device.type == "mps":
|
||||
ids = ids.to("cpu")
|
||||
|
||||
result = []
|
||||
for i in range(len(self.axes_dim)):
|
||||
freqs = freqs_cis[i].to(ids.device)
|
||||
index = ids[:, :, i : i + 1].repeat(1, 1, freqs.shape[-1]).to(torch.int64)
|
||||
result.append(torch.gather(freqs.unsqueeze(0).repeat(index.shape[0], 1, 1), dim=1, index=index))
|
||||
return torch.cat(result, dim=-1).to(device)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
freqs_cis,
|
||||
attention_mask,
|
||||
l_effective_ref_img_len,
|
||||
l_effective_img_len,
|
||||
ref_img_sizes,
|
||||
img_sizes,
|
||||
device
|
||||
):
|
||||
batch_size = len(attention_mask)
|
||||
p = self.patch_size
|
||||
|
||||
encoder_seq_len = attention_mask.shape[1]
|
||||
l_effective_cap_len = attention_mask.sum(dim=1).tolist()
|
||||
|
||||
seq_lengths = [cap_len + sum(ref_img_len) + img_len for cap_len, ref_img_len, img_len in zip(l_effective_cap_len, l_effective_ref_img_len, l_effective_img_len)]
|
||||
|
||||
max_seq_len = int(max(seq_lengths))
|
||||
max_ref_img_len = max([sum(ref_img_len) for ref_img_len in l_effective_ref_img_len])
|
||||
max_img_len = max(l_effective_img_len)
|
||||
|
||||
# Create position IDs
|
||||
position_ids = torch.zeros(batch_size, max_seq_len, 3, dtype=torch.int32, device=device)
|
||||
|
||||
for i, (cap_seq_len, seq_len) in enumerate(zip(l_effective_cap_len, seq_lengths)):
|
||||
cap_seq_len = int(cap_seq_len)
|
||||
seq_len = int(seq_len)
|
||||
# add text position ids
|
||||
position_ids[i, :cap_seq_len] = repeat(torch.arange(cap_seq_len, dtype=torch.int32, device=device), "l -> l 3")
|
||||
|
||||
pe_shift = cap_seq_len
|
||||
pe_shift_len = cap_seq_len
|
||||
|
||||
if ref_img_sizes[i] is not None:
|
||||
for ref_img_size, ref_img_len in zip(ref_img_sizes[i], l_effective_ref_img_len[i]):
|
||||
H, W = ref_img_size
|
||||
ref_H_tokens, ref_W_tokens = H // p, W // p
|
||||
assert ref_H_tokens * ref_W_tokens == ref_img_len
|
||||
# add image position ids
|
||||
|
||||
row_ids = repeat(torch.arange(ref_H_tokens, dtype=torch.int32, device=device), "h -> h w", w=ref_W_tokens).flatten()
|
||||
col_ids = repeat(torch.arange(ref_W_tokens, dtype=torch.int32, device=device), "w -> h w", h=ref_H_tokens).flatten()
|
||||
position_ids[i, pe_shift_len:pe_shift_len + ref_img_len, 0] = pe_shift
|
||||
position_ids[i, pe_shift_len:pe_shift_len + ref_img_len, 1] = row_ids
|
||||
position_ids[i, pe_shift_len:pe_shift_len + ref_img_len, 2] = col_ids
|
||||
|
||||
pe_shift += max(ref_H_tokens, ref_W_tokens)
|
||||
pe_shift_len += ref_img_len
|
||||
|
||||
H, W = img_sizes[i]
|
||||
H_tokens, W_tokens = H // p, W // p
|
||||
assert H_tokens * W_tokens == l_effective_img_len[i]
|
||||
|
||||
row_ids = repeat(torch.arange(H_tokens, dtype=torch.int32, device=device), "h -> h w", w=W_tokens).flatten()
|
||||
col_ids = repeat(torch.arange(W_tokens, dtype=torch.int32, device=device), "w -> h w", h=H_tokens).flatten()
|
||||
|
||||
assert pe_shift_len + l_effective_img_len[i] == seq_len
|
||||
position_ids[i, pe_shift_len: seq_len, 0] = pe_shift
|
||||
position_ids[i, pe_shift_len: seq_len, 1] = row_ids
|
||||
position_ids[i, pe_shift_len: seq_len, 2] = col_ids
|
||||
|
||||
# Get combined rotary embeddings
|
||||
freqs_cis = self._get_freqs_cis(freqs_cis, position_ids)
|
||||
|
||||
# create separate rotary embeddings for captions and images
|
||||
cap_freqs_cis = torch.zeros(
|
||||
batch_size, encoder_seq_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype
|
||||
)
|
||||
ref_img_freqs_cis = torch.zeros(
|
||||
batch_size, max_ref_img_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype
|
||||
)
|
||||
img_freqs_cis = torch.zeros(
|
||||
batch_size, max_img_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype
|
||||
)
|
||||
|
||||
for i, (cap_seq_len, ref_img_len, img_len, seq_len) in enumerate(zip(l_effective_cap_len, l_effective_ref_img_len, l_effective_img_len, seq_lengths)):
|
||||
cap_seq_len = int(cap_seq_len)
|
||||
sum_ref_img_len = int(sum(ref_img_len))
|
||||
img_len = int(img_len)
|
||||
seq_len = int(seq_len)
|
||||
cap_freqs_cis[i, :cap_seq_len] = freqs_cis[i, :cap_seq_len]
|
||||
ref_img_freqs_cis[i, :sum_ref_img_len] = freqs_cis[i, cap_seq_len:cap_seq_len + sum_ref_img_len]
|
||||
img_freqs_cis[i, :img_len] = freqs_cis[i, cap_seq_len + sum_ref_img_len:cap_seq_len + sum_ref_img_len + img_len]
|
||||
|
||||
return (
|
||||
cap_freqs_cis,
|
||||
ref_img_freqs_cis,
|
||||
img_freqs_cis,
|
||||
freqs_cis,
|
||||
l_effective_cap_len,
|
||||
seq_lengths,
|
||||
)
|
||||
@@ -0,0 +1,621 @@
|
||||
import warnings
|
||||
import itertools
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import PeftAdapterMixin
|
||||
from diffusers.loaders.single_file_model import FromOriginalModelMixin
|
||||
from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
|
||||
from ..attention_processor import OmniGen2AttnProcessorFlash2Varlen, OmniGen2AttnProcessor
|
||||
from .repo import OmniGen2RotaryPosEmbed
|
||||
from .block_lumina2 import LuminaLayerNormContinuous, LuminaRMSNormZero, LuminaFeedForward, Lumina2CombinedTimestepCaptionEmbedding
|
||||
|
||||
from ...utils.import_utils import is_triton_available, is_flash_attn_available
|
||||
|
||||
if is_triton_available():
|
||||
from ...ops.triton.layer_norm import RMSNorm
|
||||
else:
|
||||
from torch.nn import RMSNorm
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
class OmniGen2TransformerBlock(nn.Module):
|
||||
"""
|
||||
Transformer block for OmniGen2 model.
|
||||
|
||||
This block implements a transformer layer with:
|
||||
- Multi-head attention with flash attention
|
||||
- Feed-forward network with SwiGLU activation
|
||||
- RMS normalization
|
||||
- Optional modulation for conditional generation
|
||||
|
||||
Args:
|
||||
dim: Dimension of the input and output tensors
|
||||
num_attention_heads: Number of attention heads
|
||||
num_kv_heads: Number of key-value heads
|
||||
multiple_of: Multiple of which the hidden dimension should be
|
||||
ffn_dim_multiplier: Multiplier for the feed-forward network dimension
|
||||
norm_eps: Epsilon value for normalization layers
|
||||
modulation: Whether to use modulation for conditional generation
|
||||
use_fused_rms_norm: Whether to use fused RMS normalization
|
||||
use_fused_swiglu: Whether to use fused SwiGLU activation
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
num_kv_heads: int,
|
||||
multiple_of: int,
|
||||
ffn_dim_multiplier: float,
|
||||
norm_eps: float,
|
||||
modulation: bool = True,
|
||||
) -> None:
|
||||
"""Initialize the transformer block."""
|
||||
super().__init__()
|
||||
self.head_dim = dim // num_attention_heads
|
||||
self.modulation = modulation
|
||||
|
||||
try:
|
||||
processor = OmniGen2AttnProcessorFlash2Varlen()
|
||||
except ImportError:
|
||||
processor = OmniGen2AttnProcessor()
|
||||
|
||||
# Initialize attention layer
|
||||
self.attn = Attention(
|
||||
query_dim=dim,
|
||||
cross_attention_dim=None,
|
||||
dim_head=dim // num_attention_heads,
|
||||
qk_norm="rms_norm",
|
||||
heads=num_attention_heads,
|
||||
kv_heads=num_kv_heads,
|
||||
eps=1e-5,
|
||||
bias=False,
|
||||
out_bias=False,
|
||||
processor=processor,
|
||||
)
|
||||
|
||||
# Initialize feed-forward network
|
||||
self.feed_forward = LuminaFeedForward(
|
||||
dim=dim,
|
||||
inner_dim=4 * dim,
|
||||
multiple_of=multiple_of,
|
||||
ffn_dim_multiplier=ffn_dim_multiplier
|
||||
)
|
||||
|
||||
# Initialize normalization layers
|
||||
if modulation:
|
||||
self.norm1 = LuminaRMSNormZero(
|
||||
embedding_dim=dim,
|
||||
norm_eps=norm_eps,
|
||||
norm_elementwise_affine=True
|
||||
)
|
||||
else:
|
||||
self.norm1 = RMSNorm(dim, eps=norm_eps)
|
||||
|
||||
self.ffn_norm1 = RMSNorm(dim, eps=norm_eps)
|
||||
self.norm2 = RMSNorm(dim, eps=norm_eps)
|
||||
self.ffn_norm2 = RMSNorm(dim, eps=norm_eps)
|
||||
|
||||
self.initialize_weights()
|
||||
|
||||
def initialize_weights(self) -> None:
|
||||
"""
|
||||
Initialize the weights of the transformer block.
|
||||
|
||||
Uses Xavier uniform initialization for linear layers and zero initialization for biases.
|
||||
"""
|
||||
nn.init.xavier_uniform_(self.attn.to_q.weight)
|
||||
nn.init.xavier_uniform_(self.attn.to_k.weight)
|
||||
nn.init.xavier_uniform_(self.attn.to_v.weight)
|
||||
nn.init.xavier_uniform_(self.attn.to_out[0].weight)
|
||||
|
||||
nn.init.xavier_uniform_(self.feed_forward.linear_1.weight)
|
||||
nn.init.xavier_uniform_(self.feed_forward.linear_2.weight)
|
||||
nn.init.xavier_uniform_(self.feed_forward.linear_3.weight)
|
||||
|
||||
if self.modulation:
|
||||
nn.init.zeros_(self.norm1.linear.weight)
|
||||
nn.init.zeros_(self.norm1.linear.bias)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
image_rotary_emb: torch.Tensor,
|
||||
temb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass of the transformer block.
|
||||
|
||||
Args:
|
||||
hidden_states: Input hidden states tensor
|
||||
attention_mask: Attention mask tensor
|
||||
image_rotary_emb: Rotary embeddings for image tokens
|
||||
temb: Optional timestep embedding tensor
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output hidden states after transformer block processing
|
||||
"""
|
||||
import time
|
||||
if self.modulation:
|
||||
if temb is None:
|
||||
raise ValueError("temb must be provided when modulation is enabled")
|
||||
|
||||
norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb)
|
||||
attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
hidden_states = hidden_states + gate_msa.unsqueeze(1).tanh() * self.norm2(attn_output)
|
||||
mlp_output = self.feed_forward(self.ffn_norm1(hidden_states) * (1 + scale_mlp.unsqueeze(1)))
|
||||
hidden_states = hidden_states + gate_mlp.unsqueeze(1).tanh() * self.ffn_norm2(mlp_output)
|
||||
else:
|
||||
norm_hidden_states = self.norm1(hidden_states)
|
||||
attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
hidden_states = hidden_states + self.norm2(attn_output)
|
||||
mlp_output = self.feed_forward(self.ffn_norm1(hidden_states))
|
||||
hidden_states = hidden_states + self.ffn_norm2(mlp_output)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class OmniGen2Transformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
|
||||
"""
|
||||
OmniGen2 Transformer 2D Model.
|
||||
|
||||
A transformer-based diffusion model for image generation with:
|
||||
- Patch-based image processing
|
||||
- Rotary position embeddings
|
||||
- Multi-head attention
|
||||
- Conditional generation support
|
||||
|
||||
Args:
|
||||
patch_size: Size of image patches
|
||||
in_channels: Number of input channels
|
||||
out_channels: Number of output channels (defaults to in_channels)
|
||||
hidden_size: Size of hidden layers
|
||||
num_layers: Number of transformer layers
|
||||
num_refiner_layers: Number of refiner layers
|
||||
num_attention_heads: Number of attention heads
|
||||
num_kv_heads: Number of key-value heads
|
||||
multiple_of: Multiple of which the hidden dimension should be
|
||||
ffn_dim_multiplier: Multiplier for feed-forward network dimension
|
||||
norm_eps: Epsilon value for normalization layers
|
||||
axes_dim_rope: Dimensions for rotary position embeddings
|
||||
axes_lens: Lengths for rotary position embeddings
|
||||
text_feat_dim: Dimension of text features
|
||||
timestep_scale: Scale factor for timestep embeddings
|
||||
use_fused_rms_norm: Whether to use fused RMS normalization
|
||||
use_fused_swiglu: Whether to use fused SwiGLU activation
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
_no_split_modules = ["Omnigen2TransformerBlock"]
|
||||
_skip_layerwise_casting_patterns = ["x_embedder", "norm"]
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int = 2,
|
||||
in_channels: int = 16,
|
||||
out_channels: Optional[int] = None,
|
||||
hidden_size: int = 2304,
|
||||
num_layers: int = 26,
|
||||
num_refiner_layers: int = 2,
|
||||
num_attention_heads: int = 24,
|
||||
num_kv_heads: int = 8,
|
||||
multiple_of: int = 256,
|
||||
ffn_dim_multiplier: Optional[float] = None,
|
||||
norm_eps: float = 1e-5,
|
||||
axes_dim_rope: Tuple[int, int, int] = (32, 32, 32),
|
||||
axes_lens: Tuple[int, int, int] = (300, 512, 512),
|
||||
text_feat_dim: int = 1024,
|
||||
timestep_scale: float = 1.0
|
||||
) -> None:
|
||||
"""Initialize the OmniGen2 transformer model."""
|
||||
super().__init__()
|
||||
|
||||
# Validate configuration
|
||||
if (hidden_size // num_attention_heads) != sum(axes_dim_rope):
|
||||
raise ValueError(
|
||||
f"hidden_size // num_attention_heads ({hidden_size // num_attention_heads}) "
|
||||
f"must equal sum(axes_dim_rope) ({sum(axes_dim_rope)})"
|
||||
)
|
||||
|
||||
self.out_channels = out_channels or in_channels
|
||||
|
||||
# Initialize embeddings
|
||||
self.rope_embedder = OmniGen2RotaryPosEmbed(
|
||||
theta=10000,
|
||||
axes_dim=axes_dim_rope,
|
||||
axes_lens=axes_lens,
|
||||
patch_size=patch_size,
|
||||
)
|
||||
|
||||
self.x_embedder = nn.Linear(
|
||||
in_features=patch_size * patch_size * in_channels,
|
||||
out_features=hidden_size,
|
||||
)
|
||||
|
||||
self.ref_image_patch_embedder = nn.Linear(
|
||||
in_features=patch_size * patch_size * in_channels,
|
||||
out_features=hidden_size,
|
||||
)
|
||||
|
||||
self.time_caption_embed = Lumina2CombinedTimestepCaptionEmbedding(
|
||||
hidden_size=hidden_size,
|
||||
text_feat_dim=text_feat_dim,
|
||||
norm_eps=norm_eps,
|
||||
timestep_scale=timestep_scale
|
||||
)
|
||||
|
||||
# Initialize transformer blocks
|
||||
self.noise_refiner = nn.ModuleList([
|
||||
OmniGen2TransformerBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
num_kv_heads,
|
||||
multiple_of,
|
||||
ffn_dim_multiplier,
|
||||
norm_eps,
|
||||
modulation=True
|
||||
)
|
||||
for _ in range(num_refiner_layers)
|
||||
])
|
||||
|
||||
self.ref_image_refiner = nn.ModuleList([
|
||||
OmniGen2TransformerBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
num_kv_heads,
|
||||
multiple_of,
|
||||
ffn_dim_multiplier,
|
||||
norm_eps,
|
||||
modulation=True
|
||||
)
|
||||
for _ in range(num_refiner_layers)
|
||||
])
|
||||
|
||||
self.context_refiner = nn.ModuleList(
|
||||
[
|
||||
OmniGen2TransformerBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
num_kv_heads,
|
||||
multiple_of,
|
||||
ffn_dim_multiplier,
|
||||
norm_eps,
|
||||
modulation=False
|
||||
)
|
||||
for _ in range(num_refiner_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
OmniGen2TransformerBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
num_kv_heads,
|
||||
multiple_of,
|
||||
ffn_dim_multiplier,
|
||||
norm_eps,
|
||||
modulation=True
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Output norm & projection
|
||||
self.norm_out = LuminaLayerNormContinuous(
|
||||
embedding_dim=hidden_size,
|
||||
conditioning_embedding_dim=min(hidden_size, 1024),
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
bias=True,
|
||||
out_dim=patch_size * patch_size * self.out_channels
|
||||
)
|
||||
|
||||
# Add learnable embeddings to distinguish different images
|
||||
self.image_index_embedding = nn.Parameter(torch.randn(5, hidden_size)) # support max 5 ref images
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
self.initialize_weights()
|
||||
|
||||
def initialize_weights(self) -> None:
|
||||
"""
|
||||
Initialize the weights of the model.
|
||||
|
||||
Uses Xavier uniform initialization for linear layers.
|
||||
"""
|
||||
nn.init.xavier_uniform_(self.x_embedder.weight)
|
||||
nn.init.constant_(self.x_embedder.bias, 0.0)
|
||||
|
||||
nn.init.xavier_uniform_(self.ref_image_patch_embedder.weight)
|
||||
nn.init.constant_(self.ref_image_patch_embedder.bias, 0.0)
|
||||
|
||||
nn.init.zeros_(self.norm_out.linear_1.weight)
|
||||
nn.init.zeros_(self.norm_out.linear_1.bias)
|
||||
nn.init.zeros_(self.norm_out.linear_2.weight)
|
||||
nn.init.zeros_(self.norm_out.linear_2.bias)
|
||||
|
||||
nn.init.normal_(self.image_index_embedding, std=0.02)
|
||||
|
||||
def img_patch_embed_and_refine(
|
||||
self,
|
||||
hidden_states,
|
||||
ref_image_hidden_states,
|
||||
padded_img_mask,
|
||||
padded_ref_img_mask,
|
||||
noise_rotary_emb,
|
||||
ref_img_rotary_emb,
|
||||
l_effective_ref_img_len,
|
||||
l_effective_img_len,
|
||||
temb
|
||||
):
|
||||
batch_size = len(hidden_states)
|
||||
max_combined_img_len = max([img_len + sum(ref_img_len) for img_len, ref_img_len in zip(l_effective_img_len, l_effective_ref_img_len)])
|
||||
|
||||
hidden_states = self.x_embedder(hidden_states)
|
||||
ref_image_hidden_states = self.ref_image_patch_embedder(ref_image_hidden_states)
|
||||
|
||||
for i in range(batch_size):
|
||||
shift = 0
|
||||
for j, ref_img_len in enumerate(l_effective_ref_img_len[i]):
|
||||
ref_image_hidden_states[i, shift:shift + ref_img_len, :] = ref_image_hidden_states[i, shift:shift + ref_img_len, :] + self.image_index_embedding[j]
|
||||
shift += ref_img_len
|
||||
|
||||
for layer in self.noise_refiner:
|
||||
hidden_states = layer(hidden_states, padded_img_mask, noise_rotary_emb, temb)
|
||||
|
||||
flat_l_effective_ref_img_len = list(itertools.chain(*l_effective_ref_img_len))
|
||||
num_ref_images = len(flat_l_effective_ref_img_len)
|
||||
max_ref_img_len = max(flat_l_effective_ref_img_len)
|
||||
|
||||
batch_ref_img_mask = ref_image_hidden_states.new_zeros(num_ref_images, max_ref_img_len, dtype=torch.bool)
|
||||
batch_ref_image_hidden_states = ref_image_hidden_states.new_zeros(num_ref_images, max_ref_img_len, self.config.hidden_size)
|
||||
batch_ref_img_rotary_emb = hidden_states.new_zeros(num_ref_images, max_ref_img_len, ref_img_rotary_emb.shape[-1], dtype=ref_img_rotary_emb.dtype)
|
||||
batch_temb = temb.new_zeros(num_ref_images, *temb.shape[1:], dtype=temb.dtype)
|
||||
|
||||
# sequence of ref imgs to batch
|
||||
idx = 0
|
||||
for i in range(batch_size):
|
||||
shift = 0
|
||||
for ref_img_len in l_effective_ref_img_len[i]:
|
||||
batch_ref_img_mask[idx, :ref_img_len] = True
|
||||
batch_ref_image_hidden_states[idx, :ref_img_len] = ref_image_hidden_states[i, shift:shift + ref_img_len]
|
||||
batch_ref_img_rotary_emb[idx, :ref_img_len] = ref_img_rotary_emb[i, shift:shift + ref_img_len]
|
||||
batch_temb[idx] = temb[i]
|
||||
shift += ref_img_len
|
||||
idx += 1
|
||||
|
||||
# refine ref imgs separately
|
||||
for layer in self.ref_image_refiner:
|
||||
batch_ref_image_hidden_states = layer(batch_ref_image_hidden_states, batch_ref_img_mask, batch_ref_img_rotary_emb, batch_temb)
|
||||
|
||||
# batch of ref imgs to sequence
|
||||
idx = 0
|
||||
for i in range(batch_size):
|
||||
shift = 0
|
||||
for ref_img_len in l_effective_ref_img_len[i]:
|
||||
ref_image_hidden_states[i, shift:shift + ref_img_len] = batch_ref_image_hidden_states[idx, :ref_img_len]
|
||||
shift += ref_img_len
|
||||
idx += 1
|
||||
|
||||
combined_img_hidden_states = hidden_states.new_zeros(batch_size, max_combined_img_len, self.config.hidden_size)
|
||||
for i, (ref_img_len, img_len) in enumerate(zip(l_effective_ref_img_len, l_effective_img_len)):
|
||||
combined_img_hidden_states[i, :sum(ref_img_len)] = ref_image_hidden_states[i, :sum(ref_img_len)]
|
||||
combined_img_hidden_states[i, sum(ref_img_len):sum(ref_img_len) + img_len] = hidden_states[i, :img_len]
|
||||
|
||||
return combined_img_hidden_states
|
||||
|
||||
def flat_and_pad_to_seq(self, hidden_states, ref_image_hidden_states):
|
||||
batch_size = len(hidden_states)
|
||||
p = self.config.patch_size
|
||||
device = hidden_states[0].device
|
||||
|
||||
img_sizes = [(img.size(1), img.size(2)) for img in hidden_states]
|
||||
l_effective_img_len = [(H // p) * (W // p) for (H, W) in img_sizes]
|
||||
|
||||
if ref_image_hidden_states is not None:
|
||||
ref_img_sizes = [[(img.size(1), img.size(2)) for img in imgs] if imgs is not None else None for imgs in ref_image_hidden_states]
|
||||
l_effective_ref_img_len = [[(ref_img_size[0] // p) * (ref_img_size[1] // p) for ref_img_size in _ref_img_sizes] if _ref_img_sizes is not None else [0] for _ref_img_sizes in ref_img_sizes]
|
||||
else:
|
||||
ref_img_sizes = [None for _ in range(batch_size)]
|
||||
l_effective_ref_img_len = [[0] for _ in range(batch_size)]
|
||||
|
||||
max_ref_img_len = max([sum(ref_img_len) for ref_img_len in l_effective_ref_img_len])
|
||||
max_img_len = max(l_effective_img_len)
|
||||
|
||||
# ref image patch embeddings
|
||||
flat_ref_img_hidden_states = []
|
||||
for i in range(batch_size):
|
||||
if ref_img_sizes[i] is not None:
|
||||
imgs = []
|
||||
for ref_img in ref_image_hidden_states[i]:
|
||||
C, H, W = ref_img.size()
|
||||
ref_img = rearrange(ref_img, 'c (h p1) (w p2) -> (h w) (p1 p2 c)', p1=p, p2=p)
|
||||
imgs.append(ref_img)
|
||||
|
||||
img = torch.cat(imgs, dim=0)
|
||||
flat_ref_img_hidden_states.append(img)
|
||||
else:
|
||||
flat_ref_img_hidden_states.append(None)
|
||||
|
||||
# image patch embeddings
|
||||
flat_hidden_states = []
|
||||
for i in range(batch_size):
|
||||
img = hidden_states[i]
|
||||
C, H, W = img.size()
|
||||
|
||||
img = rearrange(img, 'c (h p1) (w p2) -> (h w) (p1 p2 c)', p1=p, p2=p)
|
||||
flat_hidden_states.append(img)
|
||||
|
||||
padded_ref_img_hidden_states = torch.zeros(batch_size, max_ref_img_len, flat_hidden_states[0].shape[-1], device=device, dtype=flat_hidden_states[0].dtype)
|
||||
padded_ref_img_mask = torch.zeros(batch_size, max_ref_img_len, dtype=torch.bool, device=device)
|
||||
for i in range(batch_size):
|
||||
if ref_img_sizes[i] is not None:
|
||||
padded_ref_img_hidden_states[i, :sum(l_effective_ref_img_len[i])] = flat_ref_img_hidden_states[i]
|
||||
padded_ref_img_mask[i, :sum(l_effective_ref_img_len[i])] = True
|
||||
|
||||
padded_hidden_states = torch.zeros(batch_size, max_img_len, flat_hidden_states[0].shape[-1], device=device, dtype=flat_hidden_states[0].dtype)
|
||||
padded_img_mask = torch.zeros(batch_size, max_img_len, dtype=torch.bool, device=device)
|
||||
for i in range(batch_size):
|
||||
padded_hidden_states[i, :l_effective_img_len[i]] = flat_hidden_states[i]
|
||||
padded_img_mask[i, :l_effective_img_len[i]] = True
|
||||
|
||||
return (
|
||||
padded_hidden_states,
|
||||
padded_ref_img_hidden_states,
|
||||
padded_img_mask,
|
||||
padded_ref_img_mask,
|
||||
l_effective_ref_img_len,
|
||||
l_effective_img_len,
|
||||
ref_img_sizes,
|
||||
img_sizes,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.Tensor,
|
||||
text_hidden_states: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
text_attention_mask: torch.Tensor,
|
||||
ref_image_hidden_states: Optional[List[List[torch.Tensor]]] = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = False,
|
||||
) -> Union[torch.Tensor, Transformer2DModelOutput]:
|
||||
if attention_kwargs is not None:
|
||||
attention_kwargs = attention_kwargs.copy()
|
||||
lora_scale = attention_kwargs.pop("scale", 1.0)
|
||||
else:
|
||||
lora_scale = 1.0
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# weight the lora layers by setting `lora_scale` for each PEFT layer
|
||||
scale_lora_layers(self, lora_scale)
|
||||
else:
|
||||
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
|
||||
logger.warning(
|
||||
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
|
||||
)
|
||||
|
||||
# 1. Condition, positional & patch embedding
|
||||
batch_size = len(hidden_states)
|
||||
is_hidden_states_tensor = isinstance(hidden_states, torch.Tensor)
|
||||
|
||||
if is_hidden_states_tensor:
|
||||
assert hidden_states.ndim == 4
|
||||
hidden_states = [_hidden_states for _hidden_states in hidden_states]
|
||||
|
||||
device = hidden_states[0].device
|
||||
|
||||
temb, text_hidden_states = self.time_caption_embed(timestep, text_hidden_states, hidden_states[0].dtype)
|
||||
|
||||
(
|
||||
hidden_states,
|
||||
ref_image_hidden_states,
|
||||
img_mask,
|
||||
ref_img_mask,
|
||||
l_effective_ref_img_len,
|
||||
l_effective_img_len,
|
||||
ref_img_sizes,
|
||||
img_sizes,
|
||||
) = self.flat_and_pad_to_seq(hidden_states, ref_image_hidden_states)
|
||||
|
||||
(
|
||||
context_rotary_emb,
|
||||
ref_img_rotary_emb,
|
||||
noise_rotary_emb,
|
||||
rotary_emb,
|
||||
encoder_seq_lengths,
|
||||
seq_lengths,
|
||||
) = self.rope_embedder(
|
||||
freqs_cis,
|
||||
text_attention_mask,
|
||||
l_effective_ref_img_len,
|
||||
l_effective_img_len,
|
||||
ref_img_sizes,
|
||||
img_sizes,
|
||||
device,
|
||||
)
|
||||
|
||||
# 2. Context refinement
|
||||
for layer in self.context_refiner:
|
||||
text_hidden_states = layer(text_hidden_states, text_attention_mask, context_rotary_emb)
|
||||
|
||||
combined_img_hidden_states = self.img_patch_embed_and_refine(
|
||||
hidden_states,
|
||||
ref_image_hidden_states,
|
||||
img_mask,
|
||||
ref_img_mask,
|
||||
noise_rotary_emb,
|
||||
ref_img_rotary_emb,
|
||||
l_effective_ref_img_len,
|
||||
l_effective_img_len,
|
||||
temb,
|
||||
)
|
||||
|
||||
# 3. Joint Transformer blocks
|
||||
max_seq_len = int(max(seq_lengths))
|
||||
|
||||
attention_mask = hidden_states.new_zeros(batch_size, max_seq_len, dtype=torch.bool)
|
||||
joint_hidden_states = hidden_states.new_zeros(batch_size, max_seq_len, self.config.hidden_size)
|
||||
for i, (encoder_seq_len, seq_len) in enumerate(zip(encoder_seq_lengths, seq_lengths)):
|
||||
encoder_seq_len = int(encoder_seq_len)
|
||||
seq_len = int(seq_len)
|
||||
attention_mask[i, :seq_len] = True
|
||||
joint_hidden_states[i, :encoder_seq_len] = text_hidden_states[i, :encoder_seq_len]
|
||||
joint_hidden_states[i, encoder_seq_len:seq_len] = combined_img_hidden_states[i, :seq_len - encoder_seq_len]
|
||||
|
||||
hidden_states = joint_hidden_states
|
||||
|
||||
for layer_idx, layer in enumerate(self.layers):
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
layer, hidden_states, attention_mask, rotary_emb, temb
|
||||
)
|
||||
else:
|
||||
hidden_states = layer(hidden_states, attention_mask, rotary_emb, temb)
|
||||
|
||||
# 4. Output norm & projection
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
|
||||
p = self.config.patch_size
|
||||
output = []
|
||||
for i, (img_size, img_len, seq_len) in enumerate(zip(img_sizes, l_effective_img_len, seq_lengths)):
|
||||
img_len = int(img_len)
|
||||
seq_len = int(seq_len)
|
||||
height, width = img_size
|
||||
output.append(rearrange(hidden_states[i][seq_len - img_len:seq_len], '(h w) (p1 p2 c) -> c (h p1) (w p2)', h=height // p, w=width // p, p1=p, p2=p))
|
||||
if is_hidden_states_tensor:
|
||||
output = torch.stack(output, dim=0)
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not return_dict:
|
||||
return output
|
||||
return Transformer2DModelOutput(sample=output)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,266 @@
|
||||
# Copyright 2024 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import math
|
||||
import warnings
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from diffusers.image_processor import PipelineImageInput, VaeImageProcessor, is_valid_image_imagelist
|
||||
from diffusers.configuration_utils import register_to_config
|
||||
|
||||
class OmniGen2ImageProcessor(VaeImageProcessor):
|
||||
"""
|
||||
Image processor for PixArt image resize and crop.
|
||||
|
||||
Args:
|
||||
do_resize (`bool`, *optional*, defaults to `True`):
|
||||
Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`. Can accept
|
||||
`height` and `width` arguments from [`image_processor.VaeImageProcessor.preprocess`] method.
|
||||
vae_scale_factor (`int`, *optional*, defaults to `8`):
|
||||
VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor.
|
||||
resample (`str`, *optional*, defaults to `lanczos`):
|
||||
Resampling filter to use when resizing the image.
|
||||
do_normalize (`bool`, *optional*, defaults to `True`):
|
||||
Whether to normalize the image to [-1,1].
|
||||
do_binarize (`bool`, *optional*, defaults to `False`):
|
||||
Whether to binarize the image to 0/1.
|
||||
do_convert_rgb (`bool`, *optional*, defaults to be `False`):
|
||||
Whether to convert the images to RGB format.
|
||||
do_convert_grayscale (`bool`, *optional*, defaults to be `False`):
|
||||
Whether to convert the images to grayscale format.
|
||||
"""
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
do_resize: bool = True,
|
||||
vae_scale_factor: int = 16,
|
||||
resample: str = "lanczos",
|
||||
max_pixels: Optional[int] = None,
|
||||
max_side_length: Optional[int] = None,
|
||||
do_normalize: bool = True,
|
||||
do_binarize: bool = False,
|
||||
do_convert_grayscale: bool = False,
|
||||
):
|
||||
super().__init__(
|
||||
do_resize=do_resize,
|
||||
vae_scale_factor=vae_scale_factor,
|
||||
resample=resample,
|
||||
do_normalize=do_normalize,
|
||||
do_binarize=do_binarize,
|
||||
do_convert_grayscale=do_convert_grayscale,
|
||||
)
|
||||
|
||||
self.max_pixels = max_pixels
|
||||
self.max_side_length = max_side_length
|
||||
|
||||
def get_new_height_width(
|
||||
self,
|
||||
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
max_pixels: Optional[int] = None,
|
||||
max_side_length: Optional[int] = None,
|
||||
) -> Tuple[int, int]:
|
||||
r"""
|
||||
Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`.
|
||||
|
||||
Args:
|
||||
image (`Union[PIL.Image.Image, np.ndarray, torch.Tensor]`):
|
||||
The image input, which can be a PIL image, NumPy array, or PyTorch tensor. If it is a NumPy array, it
|
||||
should have shape `[batch, height, width]` or `[batch, height, width, channels]`. If it is a PyTorch
|
||||
tensor, it should have shape `[batch, channels, height, width]`.
|
||||
height (`Optional[int]`, *optional*, defaults to `None`):
|
||||
The height of the preprocessed image. If `None`, the height of the `image` input will be used.
|
||||
width (`Optional[int]`, *optional*, defaults to `None`):
|
||||
The width of the preprocessed image. If `None`, the width of the `image` input will be used.
|
||||
|
||||
Returns:
|
||||
`Tuple[int, int]`:
|
||||
A tuple containing the height and width, both resized to the nearest integer multiple of
|
||||
`vae_scale_factor`.
|
||||
"""
|
||||
|
||||
if height is None:
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
height = image.height
|
||||
elif isinstance(image, torch.Tensor):
|
||||
height = image.shape[2]
|
||||
else:
|
||||
height = image.shape[1]
|
||||
|
||||
if width is None:
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
width = image.width
|
||||
elif isinstance(image, torch.Tensor):
|
||||
width = image.shape[3]
|
||||
else:
|
||||
width = image.shape[2]
|
||||
|
||||
if max_side_length is None:
|
||||
max_side_length = self.max_side_length
|
||||
|
||||
if max_pixels is None:
|
||||
max_pixels = self.max_pixels
|
||||
|
||||
ratio = 1.0
|
||||
if max_side_length is not None:
|
||||
if height > width:
|
||||
max_side_length_ratio = max_side_length / height
|
||||
else:
|
||||
max_side_length_ratio = max_side_length / width
|
||||
|
||||
cur_pixels = height * width
|
||||
max_pixels_ratio = (max_pixels / cur_pixels) ** 0.5
|
||||
ratio = min(max_pixels_ratio, max_side_length_ratio, 1.0) # do not upscale input image
|
||||
|
||||
new_height, new_width = int(height * ratio) // self.config.vae_scale_factor * self.config.vae_scale_factor, int(width * ratio) // self.config.vae_scale_factor * self.config.vae_scale_factor
|
||||
return new_height, new_width
|
||||
|
||||
def preprocess(
|
||||
self,
|
||||
image: PipelineImageInput,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
max_pixels: Optional[int] = None,
|
||||
max_side_length: Optional[int] = None,
|
||||
resize_mode: str = "default", # "default", "fill", "crop"
|
||||
crops_coords: Optional[Tuple[int, int, int, int]] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Preprocess the image input.
|
||||
|
||||
Args:
|
||||
image (`PipelineImageInput`):
|
||||
The image input, accepted formats are PIL images, NumPy arrays, PyTorch tensors; Also accept list of
|
||||
supported formats.
|
||||
height (`int`, *optional*):
|
||||
The height in preprocessed image. If `None`, will use the `get_default_height_width()` to get default
|
||||
height.
|
||||
width (`int`, *optional*):
|
||||
The width in preprocessed. If `None`, will use get_default_height_width()` to get the default width.
|
||||
resize_mode (`str`, *optional*, defaults to `default`):
|
||||
The resize mode, can be one of `default` or `fill`. If `default`, will resize the image to fit within
|
||||
the specified width and height, and it may not maintaining the original aspect ratio. If `fill`, will
|
||||
resize the image to fit within the specified width and height, maintaining the aspect ratio, and then
|
||||
center the image within the dimensions, filling empty with data from image. If `crop`, will resize the
|
||||
image to fit within the specified width and height, maintaining the aspect ratio, and then center the
|
||||
image within the dimensions, cropping the excess. Note that resize_mode `fill` and `crop` are only
|
||||
supported for PIL image input.
|
||||
crops_coords (`List[Tuple[int, int, int, int]]`, *optional*, defaults to `None`):
|
||||
The crop coordinates for each image in the batch. If `None`, will not crop the image.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
The preprocessed image.
|
||||
"""
|
||||
supported_formats = (PIL.Image.Image, np.ndarray, torch.Tensor)
|
||||
|
||||
# Expand the missing dimension for 3-dimensional pytorch tensor or numpy array that represents grayscale image
|
||||
if self.config.do_convert_grayscale and isinstance(image, (torch.Tensor, np.ndarray)) and image.ndim == 3:
|
||||
if isinstance(image, torch.Tensor):
|
||||
# if image is a pytorch tensor could have 2 possible shapes:
|
||||
# 1. batch x height x width: we should insert the channel dimension at position 1
|
||||
# 2. channel x height x width: we should insert batch dimension at position 0,
|
||||
# however, since both channel and batch dimension has same size 1, it is same to insert at position 1
|
||||
# for simplicity, we insert a dimension of size 1 at position 1 for both cases
|
||||
image = image.unsqueeze(1)
|
||||
else:
|
||||
# if it is a numpy array, it could have 2 possible shapes:
|
||||
# 1. batch x height x width: insert channel dimension on last position
|
||||
# 2. height x width x channel: insert batch dimension on first position
|
||||
if image.shape[-1] == 1:
|
||||
image = np.expand_dims(image, axis=0)
|
||||
else:
|
||||
image = np.expand_dims(image, axis=-1)
|
||||
|
||||
if isinstance(image, list) and isinstance(image[0], np.ndarray) and image[0].ndim == 4:
|
||||
warnings.warn(
|
||||
"Passing `image` as a list of 4d np.ndarray is deprecated."
|
||||
"Please concatenate the list along the batch dimension and pass it as a single 4d np.ndarray",
|
||||
FutureWarning,
|
||||
)
|
||||
image = np.concatenate(image, axis=0)
|
||||
if isinstance(image, list) and isinstance(image[0], torch.Tensor) and image[0].ndim == 4:
|
||||
warnings.warn(
|
||||
"Passing `image` as a list of 4d torch.Tensor is deprecated."
|
||||
"Please concatenate the list along the batch dimension and pass it as a single 4d torch.Tensor",
|
||||
FutureWarning,
|
||||
)
|
||||
image = torch.cat(image, axis=0)
|
||||
|
||||
if not is_valid_image_imagelist(image):
|
||||
raise ValueError(
|
||||
f"Input is in incorrect format. Currently, we only support {', '.join(str(x) for x in supported_formats)}"
|
||||
)
|
||||
if not isinstance(image, list):
|
||||
image = [image]
|
||||
|
||||
if isinstance(image[0], PIL.Image.Image):
|
||||
if crops_coords is not None:
|
||||
image = [i.crop(crops_coords) for i in image]
|
||||
if self.config.do_resize:
|
||||
height, width = self.get_new_height_width(image[0], height, width, max_pixels, max_side_length)
|
||||
image = [self.resize(i, height, width, resize_mode=resize_mode) for i in image]
|
||||
if self.config.do_convert_rgb:
|
||||
image = [self.convert_to_rgb(i) for i in image]
|
||||
elif self.config.do_convert_grayscale:
|
||||
image = [self.convert_to_grayscale(i) for i in image]
|
||||
image = self.pil_to_numpy(image) # to np
|
||||
image = self.numpy_to_pt(image) # to pt
|
||||
|
||||
elif isinstance(image[0], np.ndarray):
|
||||
image = np.concatenate(image, axis=0) if image[0].ndim == 4 else np.stack(image, axis=0)
|
||||
|
||||
image = self.numpy_to_pt(image)
|
||||
|
||||
height, width = self.get_new_height_width(image, height, width, max_pixels, max_side_length)
|
||||
if self.config.do_resize:
|
||||
image = self.resize(image, height, width)
|
||||
|
||||
elif isinstance(image[0], torch.Tensor):
|
||||
image = torch.cat(image, axis=0) if image[0].ndim == 4 else torch.stack(image, axis=0)
|
||||
|
||||
if self.config.do_convert_grayscale and image.ndim == 3:
|
||||
image = image.unsqueeze(1)
|
||||
|
||||
channel = image.shape[1]
|
||||
# don't need any preprocess if the image is latents
|
||||
if channel == self.config.vae_latent_channels:
|
||||
return image
|
||||
|
||||
height, width = self.get_new_height_width(image, height, width, max_pixels, max_side_length)
|
||||
if self.config.do_resize:
|
||||
image = self.resize(image, height, width)
|
||||
|
||||
# expected range [0,1], normalize to [-1,1]
|
||||
do_normalize = self.config.do_normalize
|
||||
if do_normalize and image.min() < 0:
|
||||
warnings.warn(
|
||||
"Passing `image` as torch tensor with value range in [-1,1] is deprecated. The expected value range for image tensor is [0,1] "
|
||||
f"when passing as pytorch tensor or numpy Array. You passed `image` with value range [{image.min()},{image.max()}]",
|
||||
FutureWarning,
|
||||
)
|
||||
do_normalize = False
|
||||
if do_normalize:
|
||||
image = self.normalize(image)
|
||||
|
||||
if self.config.do_binarize:
|
||||
image = self.binarize(image)
|
||||
|
||||
return image
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user